diff --git a/README.md b/README.md index 8bb9a48fefd47472f4cc2c23a933f4bb70e9cfcb..5c7ef7c3125c1e0e0634703b447c23035dab901c 100644 --- a/README.md +++ b/README.md @@ -26,8 +26,8 @@ uses. ## Why split MiniMax-H3 is 195.9 GiB in bfloat16 and a ZeroGPU Space is evicted at **150 GB of storage**. An unquantized single -Space is therefore impossible. Cut the `MiniMaxH3Ref2VABlocks` sequence at its `text_encoder` step and both halves -fit unquantized: +Space is therefore impossible. Cut the `ref2va` branch of `MiniMaxH3Blocks` at its `text_encoder` step and both +halves fit unquantized: | Space | Subfolders | Download | Resident | |---|---|---|---| @@ -65,25 +65,29 @@ is 16 kHz mono on purpose: the audio VAE wants 32 kHz, so the example exercises ## How the split is expressed -`MiniMaxH3Ref2VABlocks` is a `SequentialPipelineBlocks` of eight steps: +`MiniMaxH3Blocks` is one `SequentialPipelineBlocks` whose branches are picked per request — and per `workflow=` — +from the inputs, `ref2va` being the branch `references` selects: ``` -setup -> text_encoder -> reference_encoder -> prepare_layout -> prepare_latents -> set_timesteps -> denoise -> decode +setup -> text_encoder -> reference_encoder -> denoise -> after_denoise -> decode ``` +where `denoise` is itself `prepare_layout -> prepare_latents -> set_timesteps -> denoise`, against the +`transformer_ref` partition. + `h3_split_blocks.py` subclasses it with the `text_encoder` step removed. Dropping the step drops the three components it declares, so `load_components` resolves `transformer_ref` / `vae` / `audio_vae` / the two schedulers out of the shared `modular_model_index.json` and never fetches the conditioner — and `prompt_embeds` and `text_token_tags` become ordinary required inputs of the pipeline call: ```py -pipe = MiniMaxH3Ref2VAGeneratorBlocks().init_pipeline("diffusers-internal-dev/MiniMax-H3") +pipe = MiniMaxH3Ref2VAGeneratorBlocks().init_pipeline("MiniMaxAI/MiniMax-H3") pipe.load_components(dtype=torch.bfloat16) state = pipe(prompt_embeds=..., text_token_tags=..., references=[...], height=544, width=960, num_frames=124, num_inference_steps=28) ``` -Only **text** encoding is remote. `reference_encoder` is the `ref2va` blockset's own encoder step — it runs the video +Only **text** encoding is remote. `reference_encoder` is the `ref2va` branch's own encoder step — it runs the video VAE over the image and video references and the audio VAE over the soundtracks, and it is where the references' latent geometry is resolved — so it stays on this side, next to the autoencoders the conditioner Space does not hold. @@ -169,6 +173,7 @@ on is cold and a cold one pays the lazy 72.16 GiB `PIPE.to("cuda")` inside its f | Variable | Default | Meaning | |---|---|---| | `H3_CONDITIONER` | `multimodalart/qwen3vl-conditioner` | The public Space this one asks for embeddings; the client passes no token, so the call runs on the caller's own quota. | +| `H3_MODEL_REPO` | `MiniMaxAI/MiniMax-H3` | The diffusers-layout checkpoint. Public. | | `H3_AOTI` | `0` | `1` loads the compiled block package. | | `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. | | `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. The two VAEs are pinned to torch SDPA instead: they are float32, which cuDNN has no kernel for. | @@ -176,15 +181,24 @@ on is cold and a cold one pays the lazy 72.16 GiB `PIPE.to("cuda")` inside its f | `H3_PLACEMENT_ALLOWANCE` | `90` | Seconds of the reservation set aside for a cold worker's placement. | | `H3_GPU_SIZE` | `xlarge` | ZeroGPU allocation size. `large` does not fit. | -## Required secret +## Secrets -`HF_TOKEN` — `diffusers-internal-dev/MiniMax-H3` is private. The conditioner is a public Space and is called without a token, so that round trip runs on the caller's own quota rather than this org's. +The weights are the public [`MiniMaxAI/MiniMax-H3`](https://huggingface.co/MiniMaxAI/MiniMax-H3) diffusers +checkpoint and the conditioner is a public Space called without a token, so that round trip runs on the caller's own +quota rather than this org's. `HF_TOKEN` is still read by `h3_aoti`, whose compiled-package repository is private — +without it, set `H3_AOTI=0`. ## Where diffusers comes from -MiniMax-H3 is modular-only and not in a released `diffusers`, so the integration branch's `src/diffusers` tree is -vendored here as a top-level `diffusers/` package; the working directory comes first on `sys.path`, so there is no -install step. `requirements.txt` only carries what that tree imports. +MiniMax-H3 is modular-only and not in a released `diffusers`, so `requirements.txt` installs it from the canonical +pull request, [huggingface/diffusers#14371](https://github.com/huggingface/diffusers/pull/14371), pinned to the +**commit** `665f5782` (`refs/pull/14371/head` at deploy time) rather than to the moving `minimax-h3-refactor` branch. + +That PR is a WIP: it needs **re-pinning whenever it updates**, and `h3_split_blocks.py` — which subclasses its block +classes to cut the pipeline in two — has to be re-checked against the new head at the same time. The PR refactored +the blocks into one workflow-selected pipeline (per-modality reference classes, `before_encode` / `after_denoise` +steps, no `packing` modules), so block names and the shape of the split are exactly what a new head is liable to +move. Two of those are `ref2va`-only and easy to miss. PyAV decodes a reference video or audio file as the reference is built, and **`torchaudio`** resamples a soundtrack that is not already at the audio VAE's 32 kHz — a 32 kHz @@ -194,4 +208,4 @@ reference skips the resample entirely, so the dependency only shows up once some ImportError: Resampling a MiniMax-H3 reference soundtrack from 24000 Hz to 32000 Hz needs `torchaudio`. ``` -The conditioner Space needs it as well: its `setup` step prepares the very same waveforms this one does. \ No newline at end of file +The conditioner Space needs it as well: its `setup` step normalizes the very same waveforms this one does. \ No newline at end of file diff --git a/app.py b/app.py index 4c041a9c768ee211882af3f761a0b121e9ea3537..048cbeadacca344d3012a48c15f1d89b95274af7 100644 --- a/app.py +++ b/app.py @@ -10,9 +10,9 @@ Why split at all: MiniMax-H3 is 195.9 GiB in bfloat16 and a ZeroGPU Space is evi unquantized single Space is impossible. Cut at the text-encoder step, this half pulls 77.3 GB (`transformer_ref/` 61.73 GiB + `vae/` 9.70 + `audio_vae/` 0.56) and the conditioner 66.7 GB, and neither is quantized. -The blockset is `MiniMaxH3Ref2VABlocks` with its `text_encoder` step removed — see `h3_split_blocks.py`. Only *text* -encoding is remote: `reference_encoder` is the `ref2va` blockset's own encoder step and runs here, next to the two -autoencoders it needs. +The blockset is the `ref2va` branch of `MiniMaxH3Blocks` with its `text_encoder` step removed — see +`h3_split_blocks.py`. Only *text* encoding is remote: `reference_encoder` is the `ref2va` branch's own encoder step +and runs here, next to the two autoencoders it needs. """ from __future__ import annotations @@ -27,7 +27,7 @@ import traceback import spaces import gradio as gr -MODEL_REPO = os.environ.get("H3_MODEL_REPO", "diffusers-internal-dev/MiniMax-H3") +MODEL_REPO = os.environ.get("H3_MODEL_REPO", "MiniMaxAI/MiniMax-H3") CONDITIONER_SPACE = os.environ.get("H3_CONDITIONER", "multimodalart/qwen3vl-conditioner") # `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. Neither puts anything on the card at *startup*, which is @@ -43,8 +43,9 @@ GPU_SIZE = os.environ.get("H3_GPU_SIZE", "xlarge") MIN_GPU_DURATION = int(os.environ.get("H3_GPU_DURATION_MIN", "120")) MAX_GPU_DURATION = int(os.environ.get("H3_GPU_DURATION_MAX", "1500")) -# MiniMax-H3's own canvases, i.e. `resolve_canvas_size` from `diffusers.modular_pipelines.minimax_h3.packing` -# evaluated for the six released aspect ratios. Hardcoded so the UI renders before `diffusers` is importable. +# MiniMax-H3's own canvases, i.e. `resolve_canvas_size` from +# `diffusers.modular_pipelines.minimax_h3.modular_pipeline` evaluated for the six released aspect ratios. Hardcoded +# so the UI renders before `diffusers` is importable. # Must stay identical to the conditioner's table: this Space forwards the *label* to the conditioner, so a canvas # that half does not know is rejected there and surfaces as a failure here. CANVASES = { @@ -137,7 +138,7 @@ def reference_rows(references: list[tuple[str, str]], num_frames: int) -> int: """ from PIL import Image - from diffusers.modular_pipelines.minimax_h3.packing import resolve_canvas_size + from diffusers.modular_pipelines.minimax_h3.modular_pipeline import resolve_canvas_size rows = 0 for kind, path in references: @@ -158,7 +159,7 @@ def reference_rows(references: list[tuple[str, str]], num_frames: int) -> int: with av.open(path) as container: stream = container.streams.video[0] source_height, source_width = stream.height, stream.width - canvas_height, canvas_width = resolve_canvas_size(source_width, source_height) + canvas_height, canvas_width = resolve_canvas_size(source_width, source_height, CANVAS_MULTIPLE) # Resampled onto 24 fps and capped at the generated length, then snapped down to `17 * n + 5`. frames = min(round(video_seconds * FPS), num_frames) snapped = max(1, (frames - LATENTS_PER_CHUNK) // FRAMES_PER_CHUNK) * FRAMES_PER_CHUNK + LATENTS_PER_CHUNK @@ -206,7 +207,8 @@ def load_models() -> str | None: """Load the denoising half. At **startup**, but *not* onto the card. `MiniMaxH3Ref2VAGeneratorBlocks` declares `transformer_ref`, `vae`, `audio_vae`, `scheduler`, `audio_scheduler` - and `video_processor`, so `load_components` fetches exactly those subfolders out of the shared + and `video_processor` as its pretrained components (plus an `image_processor` built from config), so + `load_components` fetches exactly those subfolders out of the shared `modular_model_index.json` — `text_encoder/` and the `transformer/` partition are never touched. Both autoencoders carry `_keep_in_fp32_modules` over every module, so the `dtype` below is refused for them and @@ -224,11 +226,6 @@ def load_models() -> str | None: if PIPE is not None or LOAD_ERROR is not None: return LOAD_ERROR - token = os.environ.get("HF_TOKEN") - if not token: - LOAD_ERROR = f"**`HF_TOKEN` secret is missing** and `{MODEL_REPO}` is private. Add it and restart." - return LOAD_ERROR - started = time.time() try: import torch @@ -240,7 +237,9 @@ def load_models() -> str | None: blocks = MiniMaxH3Ref2VAGeneratorBlocks() print(f"[ref2va] loading {[c.name for c in blocks.expected_components]} from {MODEL_REPO} ...", flush=True) pipe = blocks.init_pipeline(MODEL_REPO, components_manager=manager, collection="h3") - pipe.load_components(dtype=torch.bfloat16, token=token) + # `MiniMaxAI/MiniMax-H3` is public, so the weights are fetched without a token. `HF_TOKEN` is still read by + # `h3_aoti`, whose compiled-package repository is not. + pipe.load_components(dtype=torch.bfloat16) # Pin the two autoencoders to torch SDPA *before* the transformer takes cuDNN, and in that order. # @@ -348,6 +347,23 @@ def collect(image_paths, audio_path, video_path) -> list[tuple[str, str]]: return ordered +def build_references(references: list[tuple[str, str]]): + """The `(kind, path)` references of a request as decoded reference dataclasses, in packed order. + + One public class per modality since the blocks were refactored, each decoding its own file through `from_file` — + which is also what brings the rates along, a video its own frame rate and its soundtrack, a clip its sample rate. + The blocks themselves never open a media file. + """ + from diffusers.modular_pipelines.minimax_h3 import ( + MiniMaxH3AudioReference, + MiniMaxH3ImageReference, + MiniMaxH3VideoReference, + ) + + classes = {"image": MiniMaxH3ImageReference, "video": MiniMaxH3VideoReference, "audio": MiniMaxH3AudioReference} + return [classes[kind].from_file(path) for kind, path in references] + + def audio_bearing(references: list[tuple[str, str]]) -> list[tuple[str, float]]: """The references that carry a waveform, and how long it is. A video reference brings its own soundtrack.""" carried = [] @@ -396,12 +412,17 @@ def check(prompt: str, references: list[tuple[str, str]]) -> None: ) -def encode_remote(prompt, references, canvas, num_frames): +def encode_remote(prompt, references, canvas, num_frames, rewrite_prompt=False): """Ask the conditioner Space for `prompt_embeds` + `text_token_tags`. Off this Space's GPU time entirely. The references go over with the request: `ref2va`'s presentation puts a vision block in front of the prompt for every image and every merged video frame pair, so the conditioner has to see them. It decodes the very same files this Space does, which is what keeps the two `setup` runs in agreement. + + `rewrite_prompt` is the conditioner's prompt upsampling: it rewrites the request into MiniMax-H3's trained + reference format with its own Qwen3-VL — which is shown the references, so it can name what each one contributes — + and encodes that instead, handing the rewrite back under the plan's `refined_prompt`. It runs on the conditioner's + GPU booking, and this whole call happens before `_generate` books a card here, so `get_duration` is untouched. """ from gradio_client import handle_file from safetensors import safe_open @@ -412,6 +433,7 @@ def encode_remote(prompt, references, canvas, num_frames): kinds=",".join(kind for kind, _ in references), canvas=canvas, num_frames=num_frames, + rewrite_prompt=bool(rewrite_prompt), api_name="/encode_ref2va", ) with safe_open(path, framework="pt") as handle: @@ -431,8 +453,6 @@ def _generate(prompt_embeds, text_token_tags, references, height, width, num_fra """ import torch - from diffusers.modular_pipelines.minimax_h3 import MiniMaxH3Reference - if PLACEMENT == "lazy": # 72.16 GiB across PCIe on the first request of a worker, a no-op walk on every one after it. Startup # placement is not an option here — see `load_models` — and this is what buys the offload-free denoise loop. @@ -441,7 +461,7 @@ def _generate(prompt_embeds, text_token_tags, references, height, width, num_fra state = PIPE( prompt_embeds=prompt_embeds.to("cuda"), text_token_tags=text_token_tags, - references=[MiniMaxH3Reference(**{kind: path}) for kind, path in references], + references=build_references(references), height=height, width=width, num_frames=num_frames, @@ -472,8 +492,10 @@ def generate( duration=5, steps=28, seed=42, + upsample=False, progress=gr.Progress(track_tqdm=True), ): + """One request. `upsample` is appended last and defaults off, so an existing API client is untouched by it.""" if LOAD_ERROR: raise gr.Error(LOAD_ERROR) if PIPE is None: @@ -490,10 +512,12 @@ def generate( derivable = len(audio_bearing(references)) == 1 requested = 0 if (match and derivable) else snap_frames(duration) - progress(0.0, desc="Reading the prompt and references ...") + progress(0.0, desc="Upsampling the prompt ..." if upsample else "Reading the prompt and references ...") conditioned = time.time() try: - prompt_embeds, text_token_tags, metadata, plan = encode_remote(prompt, references, canvas, requested) + prompt_embeds, text_token_tags, metadata, plan = encode_remote( + prompt, references, canvas, requested, rewrite_prompt=upsample + ) except gr.Error: raise except Exception as error: @@ -506,6 +530,7 @@ def generate( ) from error condition_seconds = time.time() - conditioned height, width, num_frames = (int(metadata[key]) for key in ("height", "width", "num_frames")) + refined = plan.get("refined_prompt") or "" progress(0.1, desc=f"Generating {num_frames / FPS:.1f} s at {width}x{height} ...") started = time.time() @@ -522,11 +547,12 @@ def generate( print( f"[ref2va] {[kind for kind, _ in references]} · `{width}x{height}`, {num_frames} frames " f"({num_frames / FPS:.3f} s), {int(steps)} steps · conditioner {condition_seconds:.0f}s " - f"({plan['num_text_tokens']} tokens) · denoise + decode {generate_seconds:.0f}s " + f"({plan['num_text_tokens']} tokens{', upsampled' if refined else ''}) · " + f"denoise + decode {generate_seconds:.0f}s " f"({generate_seconds / int(steps):.1f} s/step) · seed {int(seed)}", flush=True, ) - return path + return path, refined load_models() @@ -592,9 +618,21 @@ with gr.Blocks(title="MiniMax-H3 Reference") as demo: ) steps = gr.Slider(label="Steps", minimum=10, maximum=40, step=1, value=28) seed = gr.Number(label="Seed", value=42, precision=0) + upsample = gr.Checkbox( + label="Upsample prompt", + value=False, + info="Rewrites the prompt into the model's trained format with the conditioner's Qwen3-VL before encoding.", + ) with gr.Column(): result = gr.Video(label="Video + soundtrack") + with gr.Accordion("Upsampled prompt", open=False): + upsampled = gr.Textbox( + show_label=False, + lines=8, + interactive=False, + placeholder="Turn on “Upsample prompt” to see the rewrite that was encoded.", + ) open_slots = gr.State(OPEN_IMAGE_SLOTS) @@ -613,8 +651,10 @@ with gr.Blocks(title="MiniMax-H3 Reference") as demo: duration_controls, [audio, video, match], [match, duration], show_progress="hidden", api_name=False ) - # Same order as `generate`'s signature: the exampled five first, then the remaining image slots. - request = [prompt, images[0], audio, video, canvas, *images[1:], match, duration, steps, seed] + # Same order as `generate`'s signature: the exampled five first, then the remaining image slots. `upsample` is + # appended after every input that was already here and every existing input keeps its position, so a positional + # API client that predates it keeps working and simply takes the default. + request = [prompt, images[0], audio, video, canvas, *images[1:], match, duration, steps, seed, upsample] gr.Examples( examples=[ @@ -641,13 +681,14 @@ with gr.Blocks(title="MiniMax-H3 Reference") as demo: ], ], inputs=[prompt, images[0], audio, video, canvas], - outputs=result, + outputs=[result, upsampled], fn=generate, cache_examples=True, cache_mode="lazy", ) - run.click(generate, request, result, api_name="generate") + # The video stays the first output and the upsampled prompt is appended last, so existing consumers are untouched. + run.click(generate, request, [result, upsampled], api_name="generate") if __name__ == "__main__": diff --git a/diffusers/__init__.py b/diffusers/__init__.py deleted file mode 100644 index 3c8a46426d2a04d57aa2a7b07b5ce02dc0bc3f4b..0000000000000000000000000000000000000000 --- a/diffusers/__init__.py +++ /dev/null @@ -1,1750 +0,0 @@ -__version__ = "0.40.0.dev0" - -from typing import TYPE_CHECKING - -from .utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - is_accelerate_available, - is_auto_round_available, - is_bitsandbytes_available, - is_gguf_available, - is_librosa_available, - is_note_seq_available, - is_nvidia_modelopt_available, - is_onnx_available, - is_opencv_available, - is_optimum_quanto_available, - is_scipy_available, - is_sdnq_available, - is_sentencepiece_available, - is_torch_available, - is_torchao_available, - is_torchsde_available, - is_transformers_available, - is_transformers_version, -) - - -# Lazy Import based on -# https://github.com/huggingface/transformers/blob/main/src/transformers/__init__.py - -# 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, -# and is used to defer the actual importing for when the objects are requested. -# This way `import diffusers` provides the names in the namespace without actually importing anything (and especially none of the backends). - -_import_structure = { - "configuration_utils": ["ConfigMixin"], - "guiders": [], - "hooks": [], - "loaders": ["FromOriginalModelMixin"], - "models": [], - "modular_pipelines": [], - "pipelines": [], - "quantizers.pipe_quant_config": ["PipelineQuantizationConfig"], - "quantizers.quantization_config": [], - "schedulers": [], - "utils": [ - "OptionalDependencyNotAvailable", - "is_inflect_available", - "is_invisible_watermark_available", - "is_librosa_available", - "is_note_seq_available", - "is_onnx_available", - "is_scipy_available", - "is_torch_available", - "is_torchsde_available", - "is_transformers_available", - "is_transformers_version", - "is_unidecode_available", - "logging", - ], -} - -try: - if not is_torch_available() and not is_accelerate_available() and not is_bitsandbytes_available(): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_bitsandbytes_objects - - _import_structure["utils.dummy_bitsandbytes_objects"] = [ - name for name in dir(dummy_bitsandbytes_objects) if not name.startswith("_") - ] -else: - _import_structure["quantizers.quantization_config"].append("BitsAndBytesConfig") - -try: - if not is_torch_available() and not is_accelerate_available() and not is_gguf_available(): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_gguf_objects - - _import_structure["utils.dummy_gguf_objects"] = [ - name for name in dir(dummy_gguf_objects) if not name.startswith("_") - ] -else: - _import_structure["quantizers.quantization_config"].append("GGUFQuantizationConfig") - -try: - if not is_torch_available() and not is_accelerate_available() and not is_torchao_available(): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_torchao_objects - - _import_structure["utils.dummy_torchao_objects"] = [ - name for name in dir(dummy_torchao_objects) if not name.startswith("_") - ] -else: - _import_structure["quantizers.quantization_config"].append("TorchAoConfig") - -try: - if not is_torch_available() and not is_accelerate_available() and not is_optimum_quanto_available(): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_optimum_quanto_objects - - _import_structure["utils.dummy_optimum_quanto_objects"] = [ - name for name in dir(dummy_optimum_quanto_objects) if not name.startswith("_") - ] -else: - _import_structure["quantizers.quantization_config"].append("QuantoConfig") - -try: - if not is_torch_available() and not is_accelerate_available() and not is_nvidia_modelopt_available(): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_nvidia_modelopt_objects - - _import_structure["utils.dummy_nvidia_modelopt_objects"] = [ - name for name in dir(dummy_nvidia_modelopt_objects) if not name.startswith("_") - ] -else: - _import_structure["quantizers.quantization_config"].append("NVIDIAModelOptConfig") - -try: - if not is_torch_available(): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_nunchaku_lite_objects - - _import_structure["utils.dummy_nunchaku_lite_objects"] = [ - name for name in dir(dummy_nunchaku_lite_objects) if not name.startswith("_") - ] -else: - _import_structure["quantizers.quantization_config"].append("NunchakuLiteQuantizationConfig") - -try: - if not is_auto_round_available(): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_auto_round_objects - - _import_structure["utils.dummy_auto_round_objects"] = [ - name for name in dir(dummy_auto_round_objects) if not name.startswith("_") - ] -else: - _import_structure["quantizers.quantization_config"].append("AutoRoundConfig") - -try: - if not is_torch_available() and not is_accelerate_available() and not is_sdnq_available(): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_sdnq_objects - - _import_structure["utils.dummy_sdnq_objects"] = [ - name for name in dir(dummy_sdnq_objects) if not name.startswith("_") - ] -else: - _import_structure["quantizers.quantization_config"].append("SDNQConfig") - -try: - if not is_onnx_available(): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_onnx_objects # noqa F403 - - _import_structure["utils.dummy_onnx_objects"] = [ - name for name in dir(dummy_onnx_objects) if not name.startswith("_") - ] - -else: - _import_structure["pipelines"].extend(["OnnxRuntimeModel"]) - -try: - if not is_torch_available(): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_pt_objects # noqa F403 - - _import_structure["utils.dummy_pt_objects"] = [name for name in dir(dummy_pt_objects) if not name.startswith("_")] - -else: - _import_structure["guiders"].extend( - [ - "AdaptiveProjectedGuidance", - "AdaptiveProjectedMixGuidance", - "AutoGuidance", - "BaseGuidance", - "ClassifierFreeGuidance", - "ClassifierFreeZeroStarGuidance", - "FrequencyDecoupledGuidance", - "PerturbedAttentionGuidance", - "SkipLayerGuidance", - "SmoothedEnergyGuidance", - "TangentialClassifierFreeGuidance", - ] - ) - _import_structure["hooks"].extend( - [ - "FasterCacheConfig", - "FirstBlockCacheConfig", - "HookRegistry", - "LayerSkipConfig", - "MagCacheConfig", - "PyramidAttentionBroadcastConfig", - "SmoothedEnergyGuidanceConfig", - "TaylorSeerCacheConfig", - "TextKVCacheConfig", - "apply_faster_cache", - "apply_first_block_cache", - "apply_layer_skip", - "apply_mag_cache", - "apply_pyramid_attention_broadcast", - "apply_taylorseer_cache", - "apply_text_kv_cache", - ] - ) - _import_structure["image_processor"] = [ - "InpaintProcessor", - "IPAdapterMaskProcessor", - "PixArtImageProcessor", - "VaeImageProcessor", - "VaeImageProcessorLDM3D", - ] - _import_structure["models"].extend( - [ - "AceStepTransformer1DModel", - "AllegroTransformer3DModel", - "AnimaTextConditioner", - "AnyFlowFARTransformer3DModel", - "AnyFlowTransformer3DModel", - "AsymmetricAutoencoderKL", - "AttentionBackendName", - "AuraFlowTransformer2DModel", - "AutoencoderDC", - "AutoencoderKL", - "AutoencoderKLAllegro", - "AutoencoderKLCogVideoX", - "AutoencoderKLCosmos", - "AutoencoderKLFlux2", - "AutoencoderKLHunyuanImage", - "AutoencoderKLHunyuanImageRefiner", - "AutoencoderKLHunyuanVideo", - "AutoencoderKLHunyuanVideo15", - "AutoencoderKLKVAE", - "AutoencoderKLKVAEVideo", - "AutoencoderKLLTX2Audio", - "AutoencoderKLLTX2Video", - "AutoencoderKLLTXVideo", - "AutoencoderKLMagvit", - "AutoencoderKLMiniMaxH3", - "AutoencoderKLMiniMaxH3Audio", - "AutoencoderKLMochi", - "AutoencoderKLQwenImage", - "AutoencoderKLTemporalDecoder", - "AutoencoderKLWan", - "AutoencoderOobleck", - "AutoencoderRAE", - "AutoencoderTiny", - "AutoencoderVidTok", - "AutoModel", - "BriaFiboTransformer2DModel", - "BriaTransformer2DModel", - "CacheMixin", - "ChromaTransformer2DModel", - "ChronoEditTransformer3DModel", - "CogVideoXTransformer3DModel", - "CogView3PlusTransformer2DModel", - "CogView4Transformer2DModel", - "ConsisIDTransformer3DModel", - "ConsistencyDecoderVAE", - "ContextParallelConfig", - "ControlNetModel", - "ControlNetUnionModel", - "ControlNetXSAdapter", - "Cosmos3AVAEAudioTokenizer", - "Cosmos3OmniTransformer", - "CosmosControlNetModel", - "CosmosTransformer3DModel", - "DiTTransformer2DModel", - "DreamLiteTransformer2DModel", - "DreamLiteUNetModel", - "EasyAnimateTransformer3DModel", - "ErnieImageTransformer2DModel", - "Flux2Transformer2DModel", - "FluxControlNetModel", - "FluxMultiControlNetModel", - "FluxTransformer2DModel", - "GlmImageTransformer2DModel", - "HeliosTransformer3DModel", - "HiDreamImageTransformer2DModel", - "HunyuanDiT2DControlNetModel", - "HunyuanDiT2DModel", - "HunyuanDiT2DMultiControlNetModel", - "HunyuanImageTransformer2DModel", - "HunyuanVideo15Transformer3DModel", - "HunyuanVideoFramepackTransformer3DModel", - "HunyuanVideoTransformer3DModel", - "I2VGenXLUNet", - "Ideogram4Transformer2DModel", - "JoyImageEditPlusTransformer3DModel", - "JoyImageEditTransformer3DModel", - "Kandinsky3UNet", - "Kandinsky5Transformer3DModel", - "Krea2Transformer2DModel", - "LatteTransformer3DModel", - "LongCatAudioDiTTransformer", - "LongCatAudioDiTVae", - "LongCatImageTransformer2DModel", - "LTX2VideoTransformer3DModel", - "LTXVideoTransformer3DModel", - "Lumina2Transformer2DModel", - "LuminaNextDiT2DModel", - "MiniMaxH3Transformer3DModel", - "MochiTransformer3DModel", - "ModelMixin", - "MotifVideoTransformer3DModel", - "MotionAdapter", - "MultiAdapter", - "MultiControlNetModel", - "NucleusMoEImageTransformer2DModel", - "OmniGenTransformer2DModel", - "OvisImageTransformer2DModel", - "ParallelConfig", - "PixArtTransformer2DModel", - "PriorTransformer", - "PRXTransformer2DModel", - "QwenImageControlNetModel", - "QwenImageMultiControlNetModel", - "QwenImageTransformer2DModel", - "SanaControlNetModel", - "SanaTransformer2DModel", - "SanaVideoTransformer3DModel", - "SD3ControlNetModel", - "SD3MultiControlNetModel", - "SD3Transformer2DModel", - "SkyReelsV2Transformer3DModel", - "SparseControlNetModel", - "StableAudioDiTModel", - "StableCascadeUNet", - "T2IAdapter", - "T5FilmDecoder", - "Transformer2DModel", - "TransformerTemporalModel", - "UNet1DModel", - "UNet2DConditionModel", - "UNet2DModel", - "UNet3DConditionModel", - "UNetControlNetXSModel", - "UNetMotionModel", - "UNetSpatioTemporalConditionModel", - "UVit2DModel", - "VQModel", - "WanAnimateTransformer3DModel", - "WanTransformer3DModel", - "WanVACETransformer3DModel", - "ZImageControlNetModel", - "ZImageTransformer2DModel", - "attention_backend", - ] - ) - _import_structure["modular_pipelines"].extend( - [ - "AutoPipelineBlocks", - "ComponentsManager", - "ComponentSpec", - "ConditionalPipelineBlocks", - "ConfigSpec", - "InputParam", - "LoopSequentialPipelineBlocks", - "ModularPipeline", - "ModularPipelineBlocks", - "OutputParam", - "SequentialPipelineBlocks", - ] - ) - _import_structure["optimization"] = [ - "get_constant_schedule", - "get_constant_schedule_with_warmup", - "get_cosine_schedule_with_warmup", - "get_cosine_with_hard_restarts_schedule_with_warmup", - "get_linear_schedule_with_warmup", - "get_polynomial_decay_schedule_with_warmup", - "get_scheduler", - ] - _import_structure["pipelines"].extend( - [ - "AudioPipelineOutput", - "AutoPipelineForImage2Image", - "AutoPipelineForInpainting", - "AutoPipelineForText2Audio", - "AutoPipelineForText2Image", - "ConsistencyModelPipeline", - "DanceDiffusionPipeline", - "DDIMPipeline", - "DDPMPipeline", - "DiffusionPipeline", - "DiTPipeline", - "ImagePipelineOutput", - "KarrasVePipeline", - "LDMPipeline", - "LDMSuperResolutionPipeline", - "PNDMPipeline", - "RePaintPipeline", - "ScoreSdeVePipeline", - "StableDiffusionMixin", - ] - ) - _import_structure["quantizers"] = ["DiffusersQuantizer"] - _import_structure["schedulers"].extend( - [ - "AmusedScheduler", - "BlockRefinementScheduler", - "BlockRefinementSchedulerOutput", - "CMStochasticIterativeScheduler", - "CogVideoXDDIMScheduler", - "CogVideoXDPMScheduler", - "DDIMInverseScheduler", - "DDIMParallelScheduler", - "DDIMScheduler", - "DDPMParallelScheduler", - "DDPMScheduler", - "DDPMWuerstchenScheduler", - "DEISMultistepScheduler", - "DiscreteDDIMScheduler", - "DiscreteDDIMSchedulerOutput", - "DPMSolverMultistepInverseScheduler", - "DPMSolverMultistepScheduler", - "DPMSolverSinglestepScheduler", - "EDMDPMSolverMultistepScheduler", - "EDMEulerScheduler", - "EntropyBoundScheduler", - "EntropyBoundSchedulerOutput", - "EulerAncestralDiscreteScheduler", - "EulerDiscreteScheduler", - "FlowMapEulerDiscreteScheduler", - "FlowMatchEulerDiscreteScheduler", - "FlowMatchHeunDiscreteScheduler", - "FlowMatchLCMScheduler", - "HeliosDMDScheduler", - "HeliosScheduler", - "HeunDiscreteScheduler", - "IPNDMScheduler", - "KarrasVeScheduler", - "KDPM2AncestralDiscreteScheduler", - "KDPM2DiscreteScheduler", - "LCMScheduler", - "LTXEulerAncestralRFScheduler", - "MiniMaxH3Scheduler", - "PNDMScheduler", - "RePaintScheduler", - "SASolverScheduler", - "SchedulerMixin", - "SCMScheduler", - "ScoreSdeVeScheduler", - "TCDScheduler", - "UnCLIPScheduler", - "UniPCMultistepScheduler", - "VQDiffusionScheduler", - ] - ) - _import_structure["training_utils"] = ["EMAModel"] - _import_structure["video_processor"] = ["VideoProcessor"] - -try: - if not (is_torch_available() and is_scipy_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_torch_and_scipy_objects # noqa F403 - - _import_structure["utils.dummy_torch_and_scipy_objects"] = [ - name for name in dir(dummy_torch_and_scipy_objects) if not name.startswith("_") - ] - -else: - _import_structure["schedulers"].extend(["LMSDiscreteScheduler"]) - -try: - if not (is_torch_available() and is_torchsde_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_torch_and_torchsde_objects # noqa F403 - - _import_structure["utils.dummy_torch_and_torchsde_objects"] = [ - name for name in dir(dummy_torch_and_torchsde_objects) if not name.startswith("_") - ] - -else: - _import_structure["schedulers"].extend(["CosineDPMSolverMultistepScheduler", "DPMSolverSDEScheduler"]) - -try: - if not (is_torch_available() and is_transformers_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_torch_and_transformers_objects # noqa F403 - - _import_structure["utils.dummy_torch_and_transformers_objects"] = [ - name for name in dir(dummy_torch_and_transformers_objects) if not name.startswith("_") - ] - -else: - _import_structure["modular_pipelines"].extend( - [ - "AnimaAutoBlocks", - "AnimaModularPipeline", - "Cosmos3DistilledBlocks", - "Cosmos3DistilledModularPipeline", - "Cosmos3OmniBlocks", - "Cosmos3OmniModularPipeline", - "ErnieImageAutoBlocks", - "ErnieImageModularPipeline", - "Flux2AutoBlocks", - "Flux2KleinAutoBlocks", - "Flux2KleinBaseAutoBlocks", - "Flux2KleinBaseModularPipeline", - "Flux2KleinModularPipeline", - "Flux2ModularPipeline", - "FluxAutoBlocks", - "FluxKontextAutoBlocks", - "FluxKontextModularPipeline", - "FluxModularPipeline", - "HeliosAutoBlocks", - "HeliosModularPipeline", - "HeliosPyramidAutoBlocks", - "HeliosPyramidDistilledAutoBlocks", - "HeliosPyramidDistilledModularPipeline", - "HeliosPyramidModularPipeline", - "HunyuanVideo15AutoBlocks", - "HunyuanVideo15ModularPipeline", - "Ideogram4AutoBlocks", - "Ideogram4ModularPipeline", - "Krea2AutoBlocks", - "Krea2ModularPipeline", - "Krea2TurboAutoBlocks", - "Krea2TurboModularPipeline", - "LTXAutoBlocks", - "LTXModularPipeline", - "MiniMaxH3Blocks", - "MiniMaxH3ModularPipeline", - "MiniMaxH3Ref2VABlocks", - "MiniMaxH3Ref2VAModularPipeline", - "QwenImageAutoBlocks", - "QwenImageEditAutoBlocks", - "QwenImageEditModularPipeline", - "QwenImageEditPlusAutoBlocks", - "QwenImageEditPlusModularPipeline", - "QwenImageLayeredAutoBlocks", - "QwenImageLayeredModularPipeline", - "QwenImageModularPipeline", - "StableDiffusion3AutoBlocks", - "StableDiffusion3ModularPipeline", - "StableDiffusionXLAutoBlocks", - "StableDiffusionXLModularPipeline", - "Wan22Blocks", - "Wan22Image2VideoBlocks", - "Wan22Image2VideoModularPipeline", - "Wan22ModularPipeline", - "WanBlocks", - "WanImage2VideoAutoBlocks", - "WanImage2VideoModularPipeline", - "WanModularPipeline", - "ZImageAutoBlocks", - "ZImageModularPipeline", - ] - ) - _import_structure["pipelines"].extend( - [ - "AceStepAudioTokenDetokenizer", - "AceStepAudioTokenizer", - "AceStepConditionEncoder", - "AceStepPipeline", - "AllegroPipeline", - "AltDiffusionImg2ImgPipeline", - "AltDiffusionPipeline", - "AmusedImg2ImgPipeline", - "AmusedInpaintPipeline", - "AmusedPipeline", - "AnimateDiffControlNetPipeline", - "AnimateDiffPAGPipeline", - "AnimateDiffPipeline", - "AnimateDiffSDXLPipeline", - "AnimateDiffSparseControlNetPipeline", - "AnimateDiffVideoToVideoControlNetPipeline", - "AnimateDiffVideoToVideoPipeline", - "AnyFlowFARPipeline", - "AnyFlowPipeline", - "AudioLDM2Pipeline", - "AudioLDM2ProjectionModel", - "AudioLDM2UNet2DConditionModel", - "AudioLDMPipeline", - "AuraFlowPipeline", - "BlipDiffusionControlNetPipeline", - "BlipDiffusionPipeline", - "BriaFiboEditPipeline", - "BriaFiboPipeline", - "BriaPipeline", - "ChromaImg2ImgPipeline", - "ChromaInpaintPipeline", - "ChromaPipeline", - "ChronoEditPipeline", - "CLIPImageProjection", - "CogVideoXFunControlPipeline", - "CogVideoXImageToVideoPipeline", - "CogVideoXPipeline", - "CogVideoXVideoToVideoPipeline", - "CogView3PlusPipeline", - "CogView4ControlPipeline", - "CogView4Pipeline", - "ConsisIDPipeline", - "Cosmos2_5_PredictBasePipeline", - "Cosmos2_5_TransferPipeline", - "Cosmos2TextToImagePipeline", - "Cosmos2VideoToWorldPipeline", - "Cosmos3OmniPipeline", - "CosmosActionCondition", - "CosmosTextToWorldPipeline", - "CosmosVideoToWorldPipeline", - "CycleDiffusionPipeline", - "DiffusionGemmaPipeline", - "DiffusionGemmaPipelineOutput", - "DreamLiteMobilePipeline", - "DreamLitePipeline", - "DreamLitePipelineOutput", - "EasyAnimateControlPipeline", - "EasyAnimateInpaintPipeline", - "EasyAnimatePipeline", - "ErnieImagePipeline", - "Flux2KleinInpaintPipeline", - "Flux2KleinKVPipeline", - "Flux2KleinPipeline", - "Flux2Pipeline", - "FluxControlImg2ImgPipeline", - "FluxControlInpaintPipeline", - "FluxControlNetImg2ImgPipeline", - "FluxControlNetInpaintPipeline", - "FluxControlNetPipeline", - "FluxControlPipeline", - "FluxFillPipeline", - "FluxImg2ImgPipeline", - "FluxInpaintPipeline", - "FluxKontextInpaintPipeline", - "FluxKontextPipeline", - "FluxPipeline", - "FluxPriorReduxPipeline", - "GlmImagePipeline", - "HeliosPipeline", - "HeliosPyramidPipeline", - "HiDreamImagePipeline", - "HunyuanDiTControlNetPipeline", - "HunyuanDiTPAGPipeline", - "HunyuanDiTPipeline", - "HunyuanImagePipeline", - "HunyuanImageRefinerPipeline", - "HunyuanSkyreelsImageToVideoPipeline", - "HunyuanVideo15ImageToVideoPipeline", - "HunyuanVideo15Pipeline", - "HunyuanVideoFramepackPipeline", - "HunyuanVideoImageToVideoPipeline", - "HunyuanVideoPipeline", - "I2VGenXLPipeline", - "Ideogram4Pipeline", - "Ideogram4PromptEnhancerHead", - "IFImg2ImgPipeline", - "IFImg2ImgSuperResolutionPipeline", - "IFInpaintingPipeline", - "IFInpaintingSuperResolutionPipeline", - "IFPipeline", - "IFSuperResolutionPipeline", - "ImageTextPipelineOutput", - "JoyImageEditPipeline", - "JoyImageEditPipelineOutput", - "JoyImageEditPlusPipeline", - "JoyImageEditPlusPipelineOutput", - "Kandinsky3Img2ImgPipeline", - "Kandinsky3Pipeline", - "Kandinsky5I2IPipeline", - "Kandinsky5I2VPipeline", - "Kandinsky5T2IPipeline", - "Kandinsky5T2VPipeline", - "KandinskyCombinedPipeline", - "KandinskyImg2ImgCombinedPipeline", - "KandinskyImg2ImgPipeline", - "KandinskyInpaintCombinedPipeline", - "KandinskyInpaintPipeline", - "KandinskyPipeline", - "KandinskyPriorPipeline", - "KandinskyV22CombinedPipeline", - "KandinskyV22ControlnetImg2ImgPipeline", - "KandinskyV22ControlnetPipeline", - "KandinskyV22Img2ImgCombinedPipeline", - "KandinskyV22Img2ImgPipeline", - "KandinskyV22InpaintCombinedPipeline", - "KandinskyV22InpaintPipeline", - "KandinskyV22Pipeline", - "KandinskyV22PriorEmb2EmbPipeline", - "KandinskyV22PriorPipeline", - "Krea2Pipeline", - "LatentConsistencyModelImg2ImgPipeline", - "LatentConsistencyModelPipeline", - "LattePipeline", - "LDMTextToImagePipeline", - "LEditsPPPipelineStableDiffusion", - "LEditsPPPipelineStableDiffusionXL", - "LLaDA2Pipeline", - "LLaDA2PipelineOutput", - "LongCatAudioDiTPipeline", - "LongCatImageEditPipeline", - "LongCatImagePipeline", - "LTX2ConditionPipeline", - "LTX2HDRPipeline", - "LTX2ImageToVideoPipeline", - "LTX2InContextPipeline", - "LTX2LatentUpsamplePipeline", - "LTX2Pipeline", - "LTXConditionPipeline", - "LTXI2VLongMultiPromptPipeline", - "LTXImageToVideoPipeline", - "LTXLatentUpsamplePipeline", - "LTXPipeline", - "LucyEditPipeline", - "Lumina2Pipeline", - "Lumina2Text2ImgPipeline", - "LuminaPipeline", - "LuminaText2ImgPipeline", - "MarigoldDepthPipeline", - "MarigoldIntrinsicsPipeline", - "MarigoldNormalsPipeline", - "MochiPipeline", - "MotifVideoImage2VideoPipeline", - "MotifVideoPipeline", - "MotifVideoPipelineOutput", - "MusicLDMPipeline", - "NucleusMoEImagePipeline", - "OmniGenPipeline", - "OvisImagePipeline", - "PaintByExamplePipeline", - "PIAPipeline", - "PixArtAlphaPipeline", - "PixArtSigmaPAGPipeline", - "PixArtSigmaPipeline", - "PRXPipeline", - "PRXPixelPipeline", - "QwenImageControlNetInpaintPipeline", - "QwenImageControlNetPipeline", - "QwenImageEditInpaintPipeline", - "QwenImageEditPipeline", - "QwenImageEditPlusPipeline", - "QwenImageImg2ImgPipeline", - "QwenImageInpaintPipeline", - "QwenImageLayeredPipeline", - "QwenImagePipeline", - "ReduxImageEncoder", - "SanaControlNetPipeline", - "SanaImageToVideoPipeline", - "SanaPAGPipeline", - "SanaPipeline", - "SanaSprintImg2ImgPipeline", - "SanaSprintPipeline", - "SanaVideoPipeline", - "SanaVideoPipeline", - "SemanticStableDiffusionPipeline", - "ShapEImg2ImgPipeline", - "ShapEPipeline", - "SkyReelsV2DiffusionForcingImageToVideoPipeline", - "SkyReelsV2DiffusionForcingPipeline", - "SkyReelsV2DiffusionForcingVideoToVideoPipeline", - "SkyReelsV2ImageToVideoPipeline", - "SkyReelsV2Pipeline", - "StableAudioPipeline", - "StableAudioProjectionModel", - "StableCascadeCombinedPipeline", - "StableCascadeDecoderPipeline", - "StableCascadePriorPipeline", - "StableDiffusion3ControlNetInpaintingPipeline", - "StableDiffusion3ControlNetPipeline", - "StableDiffusion3Img2ImgPipeline", - "StableDiffusion3InpaintPipeline", - "StableDiffusion3PAGImg2ImgPipeline", - "StableDiffusion3PAGImg2ImgPipeline", - "StableDiffusion3PAGPipeline", - "StableDiffusion3Pipeline", - "StableDiffusionAdapterPipeline", - "StableDiffusionAttendAndExcitePipeline", - "StableDiffusionControlNetImg2ImgPipeline", - "StableDiffusionControlNetInpaintPipeline", - "StableDiffusionControlNetPAGInpaintPipeline", - "StableDiffusionControlNetPAGPipeline", - "StableDiffusionControlNetPipeline", - "StableDiffusionControlNetXSPipeline", - "StableDiffusionDepth2ImgPipeline", - "StableDiffusionDiffEditPipeline", - "StableDiffusionGLIGENPipeline", - "StableDiffusionGLIGENTextImagePipeline", - "StableDiffusionImageVariationPipeline", - "StableDiffusionImg2ImgPipeline", - "StableDiffusionInpaintPipeline", - "StableDiffusionInpaintPipelineLegacy", - "StableDiffusionInstructPix2PixPipeline", - "StableDiffusionLatentUpscalePipeline", - "StableDiffusionLDM3DPipeline", - "StableDiffusionModelEditingPipeline", - "StableDiffusionPAGImg2ImgPipeline", - "StableDiffusionPAGInpaintPipeline", - "StableDiffusionPAGPipeline", - "StableDiffusionPanoramaPipeline", - "StableDiffusionParadigmsPipeline", - "StableDiffusionPipeline", - "StableDiffusionPipelineSafe", - "StableDiffusionPix2PixZeroPipeline", - "StableDiffusionSAGPipeline", - "StableDiffusionUpscalePipeline", - "StableDiffusionXLAdapterPipeline", - "StableDiffusionXLControlNetImg2ImgPipeline", - "StableDiffusionXLControlNetInpaintPipeline", - "StableDiffusionXLControlNetPAGImg2ImgPipeline", - "StableDiffusionXLControlNetPAGPipeline", - "StableDiffusionXLControlNetPipeline", - "StableDiffusionXLControlNetUnionImg2ImgPipeline", - "StableDiffusionXLControlNetUnionInpaintPipeline", - "StableDiffusionXLControlNetUnionPipeline", - "StableDiffusionXLControlNetXSPipeline", - "StableDiffusionXLImg2ImgPipeline", - "StableDiffusionXLInpaintPipeline", - "StableDiffusionXLInstructPix2PixPipeline", - "StableDiffusionXLPAGImg2ImgPipeline", - "StableDiffusionXLPAGInpaintPipeline", - "StableDiffusionXLPAGPipeline", - "StableDiffusionXLPipeline", - "StableUnCLIPImg2ImgPipeline", - "StableUnCLIPPipeline", - "StableVideoDiffusionPipeline", - "TextToVideoSDPipeline", - "TextToVideoZeroPipeline", - "TextToVideoZeroSDXLPipeline", - "UnCLIPImageVariationPipeline", - "UnCLIPPipeline", - "UniDiffuserModel", - "UniDiffuserPipeline", - "UniDiffuserTextDecoder", - "VersatileDiffusionDualGuidedPipeline", - "VersatileDiffusionImageVariationPipeline", - "VersatileDiffusionPipeline", - "VersatileDiffusionTextToImagePipeline", - "VideoToVideoSDPipeline", - "VisualClozeGenerationPipeline", - "VisualClozePipeline", - "VQDiffusionPipeline", - "WanAnimatePipeline", - "WanImageToVideoPipeline", - "WanPipeline", - "WanVACEPipeline", - "WanVideoToVideoPipeline", - "WuerstchenCombinedPipeline", - "WuerstchenDecoderPipeline", - "WuerstchenPriorPipeline", - "ZImageControlNetInpaintPipeline", - "ZImageControlNetPipeline", - "ZImageImg2ImgPipeline", - "ZImageInpaintPipeline", - "ZImageOmniPipeline", - "ZImagePipeline", - ] - ) - - -try: - if not (is_torch_available() and is_transformers_available() and is_opencv_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_torch_and_transformers_and_opencv_objects # noqa F403 - - _import_structure["utils.dummy_torch_and_transformers_and_opencv_objects"] = [ - name for name in dir(dummy_torch_and_transformers_and_opencv_objects) if not name.startswith("_") - ] - -else: - _import_structure["pipelines"].extend(["ConsisIDPipeline"]) - -try: - if not (is_torch_available() and is_transformers_available() and is_sentencepiece_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_torch_and_transformers_and_sentencepiece_objects # noqa F403 - - _import_structure["utils.dummy_torch_and_transformers_and_sentencepiece_objects"] = [ - name for name in dir(dummy_torch_and_transformers_and_sentencepiece_objects) if not name.startswith("_") - ] - -else: - _import_structure["pipelines"].extend(["KolorsImg2ImgPipeline", "KolorsPAGPipeline", "KolorsPipeline"]) - -try: - if not (is_torch_available() and is_transformers_available() and is_onnx_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_torch_and_transformers_and_onnx_objects # noqa F403 - - _import_structure["utils.dummy_torch_and_transformers_and_onnx_objects"] = [ - name for name in dir(dummy_torch_and_transformers_and_onnx_objects) if not name.startswith("_") - ] - -else: - _import_structure["pipelines"].extend( - [ - "OnnxStableDiffusionImg2ImgPipeline", - "OnnxStableDiffusionInpaintPipeline", - "OnnxStableDiffusionInpaintPipelineLegacy", - "OnnxStableDiffusionPipeline", - "OnnxStableDiffusionUpscalePipeline", - "StableDiffusionOnnxPipeline", - ] - ) - -try: - if not (is_torch_available() and is_librosa_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_torch_and_librosa_objects # noqa F403 - - _import_structure["utils.dummy_torch_and_librosa_objects"] = [ - name for name in dir(dummy_torch_and_librosa_objects) if not name.startswith("_") - ] - -else: - _import_structure["pipelines"].extend(["AudioDiffusionPipeline", "Mel"]) - -try: - if not (is_transformers_available() and is_torch_available() and is_note_seq_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_transformers_and_torch_and_note_seq_objects # noqa F403 - - _import_structure["utils.dummy_transformers_and_torch_and_note_seq_objects"] = [ - name for name in dir(dummy_transformers_and_torch_and_note_seq_objects) if not name.startswith("_") - ] - - -else: - _import_structure["pipelines"].extend(["SpectrogramDiffusionPipeline"]) - -try: - if not (is_note_seq_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from .utils import dummy_note_seq_objects # noqa F403 - - _import_structure["utils.dummy_note_seq_objects"] = [ - name for name in dir(dummy_note_seq_objects) if not name.startswith("_") - ] - - -else: - _import_structure["pipelines"].extend(["MidiProcessor"]) - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - from .configuration_utils import ConfigMixin - from .quantizers import PipelineQuantizationConfig - - try: - if not is_bitsandbytes_available(): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_bitsandbytes_objects import * - else: - from .quantizers.quantization_config import BitsAndBytesConfig - - try: - if not is_gguf_available(): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_gguf_objects import * - else: - from .quantizers.quantization_config import GGUFQuantizationConfig - - try: - if not is_torchao_available(): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_torchao_objects import * - else: - from .quantizers.quantization_config import TorchAoConfig - - try: - if not is_optimum_quanto_available(): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_optimum_quanto_objects import * - else: - from .quantizers.quantization_config import QuantoConfig - - try: - if not is_nvidia_modelopt_available(): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_nvidia_modelopt_objects import * - else: - from .quantizers.quantization_config import NVIDIAModelOptConfig - - try: - if not is_torch_available(): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_nunchaku_lite_objects import * - else: - from .quantizers.quantization_config import NunchakuLiteQuantizationConfig - - try: - if not is_auto_round_available(): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_auto_round_objects import * - else: - from .quantizers.quantization_config import AutoRoundConfig - - try: - if not is_sdnq_available(): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_sdnq_objects import * - else: - from .quantizers.quantization_config import SDNQConfig - - try: - if not is_onnx_available(): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_onnx_objects import * # noqa F403 - else: - from .pipelines import OnnxRuntimeModel - - try: - if not is_torch_available(): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_pt_objects import * # noqa F403 - else: - from .guiders import ( - AdaptiveProjectedGuidance, - AdaptiveProjectedMixGuidance, - AutoGuidance, - BaseGuidance, - ClassifierFreeGuidance, - ClassifierFreeZeroStarGuidance, - FrequencyDecoupledGuidance, - PerturbedAttentionGuidance, - SkipLayerGuidance, - SmoothedEnergyGuidance, - TangentialClassifierFreeGuidance, - ) - from .hooks import ( - FasterCacheConfig, - FirstBlockCacheConfig, - HookRegistry, - LayerSkipConfig, - MagCacheConfig, - PyramidAttentionBroadcastConfig, - SmoothedEnergyGuidanceConfig, - TaylorSeerCacheConfig, - TextKVCacheConfig, - apply_faster_cache, - apply_first_block_cache, - apply_layer_skip, - apply_mag_cache, - apply_pyramid_attention_broadcast, - apply_taylorseer_cache, - apply_text_kv_cache, - ) - from .image_processor import ( - InpaintProcessor, - IPAdapterMaskProcessor, - PixArtImageProcessor, - VaeImageProcessor, - VaeImageProcessorLDM3D, - ) - from .models import ( - AceStepTransformer1DModel, - AllegroTransformer3DModel, - AnimaTextConditioner, - AnyFlowFARTransformer3DModel, - AnyFlowTransformer3DModel, - AsymmetricAutoencoderKL, - AttentionBackendName, - AuraFlowTransformer2DModel, - AutoencoderDC, - AutoencoderKL, - AutoencoderKLAllegro, - AutoencoderKLCogVideoX, - AutoencoderKLCosmos, - AutoencoderKLFlux2, - AutoencoderKLHunyuanImage, - AutoencoderKLHunyuanImageRefiner, - AutoencoderKLHunyuanVideo, - AutoencoderKLHunyuanVideo15, - AutoencoderKLKVAE, - AutoencoderKLKVAEVideo, - AutoencoderKLLTX2Audio, - AutoencoderKLLTX2Video, - AutoencoderKLLTXVideo, - AutoencoderKLMagvit, - AutoencoderKLMiniMaxH3, - AutoencoderKLMiniMaxH3Audio, - AutoencoderKLMochi, - AutoencoderKLQwenImage, - AutoencoderKLTemporalDecoder, - AutoencoderKLWan, - AutoencoderOobleck, - AutoencoderRAE, - AutoencoderTiny, - AutoencoderVidTok, - AutoModel, - BriaFiboTransformer2DModel, - BriaTransformer2DModel, - CacheMixin, - ChromaTransformer2DModel, - ChronoEditTransformer3DModel, - CogVideoXTransformer3DModel, - CogView3PlusTransformer2DModel, - CogView4Transformer2DModel, - ConsisIDTransformer3DModel, - ConsistencyDecoderVAE, - ContextParallelConfig, - ControlNetModel, - ControlNetUnionModel, - ControlNetXSAdapter, - Cosmos3AVAEAudioTokenizer, - Cosmos3OmniTransformer, - CosmosControlNetModel, - CosmosTransformer3DModel, - DiTTransformer2DModel, - DreamLiteTransformer2DModel, - DreamLiteUNetModel, - EasyAnimateTransformer3DModel, - ErnieImageTransformer2DModel, - Flux2Transformer2DModel, - FluxControlNetModel, - FluxMultiControlNetModel, - FluxTransformer2DModel, - GlmImageTransformer2DModel, - HeliosTransformer3DModel, - HiDreamImageTransformer2DModel, - HunyuanDiT2DControlNetModel, - HunyuanDiT2DModel, - HunyuanDiT2DMultiControlNetModel, - HunyuanImageTransformer2DModel, - HunyuanVideo15Transformer3DModel, - HunyuanVideoFramepackTransformer3DModel, - HunyuanVideoTransformer3DModel, - I2VGenXLUNet, - Ideogram4Transformer2DModel, - JoyImageEditPlusTransformer3DModel, - JoyImageEditTransformer3DModel, - Kandinsky3UNet, - Kandinsky5Transformer3DModel, - Krea2Transformer2DModel, - LatteTransformer3DModel, - LongCatAudioDiTTransformer, - LongCatAudioDiTVae, - LongCatImageTransformer2DModel, - LTX2VideoTransformer3DModel, - LTXVideoTransformer3DModel, - Lumina2Transformer2DModel, - LuminaNextDiT2DModel, - MiniMaxH3Transformer3DModel, - MochiTransformer3DModel, - ModelMixin, - MotifVideoTransformer3DModel, - MotionAdapter, - MultiAdapter, - MultiControlNetModel, - NucleusMoEImageTransformer2DModel, - OmniGenTransformer2DModel, - OvisImageTransformer2DModel, - ParallelConfig, - PixArtTransformer2DModel, - PriorTransformer, - PRXTransformer2DModel, - QwenImageControlNetModel, - QwenImageMultiControlNetModel, - QwenImageTransformer2DModel, - SanaControlNetModel, - SanaTransformer2DModel, - SanaVideoTransformer3DModel, - SD3ControlNetModel, - SD3MultiControlNetModel, - SD3Transformer2DModel, - SkyReelsV2Transformer3DModel, - SparseControlNetModel, - StableAudioDiTModel, - T2IAdapter, - T5FilmDecoder, - Transformer2DModel, - TransformerTemporalModel, - UNet1DModel, - UNet2DConditionModel, - UNet2DModel, - UNet3DConditionModel, - UNetControlNetXSModel, - UNetMotionModel, - UNetSpatioTemporalConditionModel, - UVit2DModel, - VQModel, - WanAnimateTransformer3DModel, - WanTransformer3DModel, - WanVACETransformer3DModel, - ZImageControlNetModel, - ZImageTransformer2DModel, - attention_backend, - ) - from .modular_pipelines import ( - AutoPipelineBlocks, - ComponentsManager, - ComponentSpec, - ConditionalPipelineBlocks, - ConfigSpec, - InputParam, - LoopSequentialPipelineBlocks, - ModularPipeline, - ModularPipelineBlocks, - OutputParam, - SequentialPipelineBlocks, - ) - from .optimization import ( - get_constant_schedule, - get_constant_schedule_with_warmup, - get_cosine_schedule_with_warmup, - get_cosine_with_hard_restarts_schedule_with_warmup, - get_linear_schedule_with_warmup, - get_polynomial_decay_schedule_with_warmup, - get_scheduler, - ) - from .pipelines import ( - AudioPipelineOutput, - AutoPipelineForImage2Image, - AutoPipelineForInpainting, - AutoPipelineForText2Audio, - AutoPipelineForText2Image, - BlipDiffusionControlNetPipeline, - BlipDiffusionPipeline, - CLIPImageProjection, - ConsistencyModelPipeline, - DanceDiffusionPipeline, - DDIMPipeline, - DDPMPipeline, - DiffusionPipeline, - DiTPipeline, - ImagePipelineOutput, - KarrasVePipeline, - LDMPipeline, - LDMSuperResolutionPipeline, - PNDMPipeline, - RePaintPipeline, - ScoreSdeVePipeline, - StableDiffusionMixin, - ) - from .quantizers import DiffusersQuantizer - from .schedulers import ( - AmusedScheduler, - BlockRefinementScheduler, - BlockRefinementSchedulerOutput, - CMStochasticIterativeScheduler, - CogVideoXDDIMScheduler, - CogVideoXDPMScheduler, - DDIMInverseScheduler, - DDIMParallelScheduler, - DDIMScheduler, - DDPMParallelScheduler, - DDPMScheduler, - DDPMWuerstchenScheduler, - DEISMultistepScheduler, - DiscreteDDIMScheduler, - DiscreteDDIMSchedulerOutput, - DPMSolverMultistepInverseScheduler, - DPMSolverMultistepScheduler, - DPMSolverSinglestepScheduler, - EDMDPMSolverMultistepScheduler, - EDMEulerScheduler, - EntropyBoundScheduler, - EntropyBoundSchedulerOutput, - EulerAncestralDiscreteScheduler, - EulerDiscreteScheduler, - FlowMapEulerDiscreteScheduler, - FlowMatchEulerDiscreteScheduler, - FlowMatchHeunDiscreteScheduler, - FlowMatchLCMScheduler, - HeliosDMDScheduler, - HeliosScheduler, - HeunDiscreteScheduler, - IPNDMScheduler, - KarrasVeScheduler, - KDPM2AncestralDiscreteScheduler, - KDPM2DiscreteScheduler, - LCMScheduler, - LTXEulerAncestralRFScheduler, - MiniMaxH3Scheduler, - PNDMScheduler, - RePaintScheduler, - SASolverScheduler, - SchedulerMixin, - SCMScheduler, - ScoreSdeVeScheduler, - TCDScheduler, - UnCLIPScheduler, - UniPCMultistepScheduler, - VQDiffusionScheduler, - ) - from .training_utils import EMAModel - from .video_processor import VideoProcessor - - try: - if not (is_torch_available() and is_scipy_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_torch_and_scipy_objects import * # noqa F403 - else: - from .schedulers import LMSDiscreteScheduler - - try: - if not (is_torch_available() and is_torchsde_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_torch_and_torchsde_objects import * # noqa F403 - else: - from .schedulers import CosineDPMSolverMultistepScheduler, DPMSolverSDEScheduler - - try: - if not (is_torch_available() and is_transformers_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .modular_pipelines import ( - AnimaAutoBlocks, - AnimaModularPipeline, - Cosmos3DistilledBlocks, - Cosmos3DistilledModularPipeline, - Cosmos3OmniBlocks, - Cosmos3OmniModularPipeline, - ErnieImageAutoBlocks, - ErnieImageModularPipeline, - Flux2AutoBlocks, - Flux2KleinAutoBlocks, - Flux2KleinBaseAutoBlocks, - Flux2KleinBaseModularPipeline, - Flux2KleinModularPipeline, - Flux2ModularPipeline, - FluxAutoBlocks, - FluxKontextAutoBlocks, - FluxKontextModularPipeline, - FluxModularPipeline, - HeliosAutoBlocks, - HeliosModularPipeline, - HeliosPyramidAutoBlocks, - HeliosPyramidDistilledAutoBlocks, - HeliosPyramidDistilledModularPipeline, - HeliosPyramidModularPipeline, - HunyuanVideo15AutoBlocks, - HunyuanVideo15ModularPipeline, - Ideogram4AutoBlocks, - Ideogram4ModularPipeline, - Krea2AutoBlocks, - Krea2ModularPipeline, - Krea2TurboAutoBlocks, - Krea2TurboModularPipeline, - LTXAutoBlocks, - LTXModularPipeline, - MiniMaxH3Blocks, - MiniMaxH3ModularPipeline, - MiniMaxH3Ref2VABlocks, - MiniMaxH3Ref2VAModularPipeline, - QwenImageAutoBlocks, - QwenImageEditAutoBlocks, - QwenImageEditModularPipeline, - QwenImageEditPlusAutoBlocks, - QwenImageEditPlusModularPipeline, - QwenImageLayeredAutoBlocks, - QwenImageLayeredModularPipeline, - QwenImageModularPipeline, - StableDiffusion3AutoBlocks, - StableDiffusion3ModularPipeline, - StableDiffusionXLAutoBlocks, - StableDiffusionXLModularPipeline, - Wan22Blocks, - Wan22Image2VideoBlocks, - Wan22Image2VideoModularPipeline, - Wan22ModularPipeline, - WanBlocks, - WanImage2VideoAutoBlocks, - WanImage2VideoModularPipeline, - WanModularPipeline, - ZImageAutoBlocks, - ZImageModularPipeline, - ) - from .pipelines import ( - AceStepAudioTokenDetokenizer, - AceStepAudioTokenizer, - AceStepConditionEncoder, - AceStepPipeline, - AllegroPipeline, - AltDiffusionImg2ImgPipeline, - AltDiffusionPipeline, - AmusedImg2ImgPipeline, - AmusedInpaintPipeline, - AmusedPipeline, - AnimateDiffControlNetPipeline, - AnimateDiffPAGPipeline, - AnimateDiffPipeline, - AnimateDiffSDXLPipeline, - AnimateDiffSparseControlNetPipeline, - AnimateDiffVideoToVideoControlNetPipeline, - AnimateDiffVideoToVideoPipeline, - AnyFlowFARPipeline, - AnyFlowPipeline, - AudioLDM2Pipeline, - AudioLDM2ProjectionModel, - AudioLDM2UNet2DConditionModel, - AudioLDMPipeline, - AuraFlowPipeline, - BriaFiboEditPipeline, - BriaFiboPipeline, - BriaPipeline, - ChromaImg2ImgPipeline, - ChromaInpaintPipeline, - ChromaPipeline, - ChronoEditPipeline, - CLIPImageProjection, - CogVideoXFunControlPipeline, - CogVideoXImageToVideoPipeline, - CogVideoXPipeline, - CogVideoXVideoToVideoPipeline, - CogView3PlusPipeline, - CogView4ControlPipeline, - CogView4Pipeline, - ConsisIDPipeline, - Cosmos2_5_PredictBasePipeline, - Cosmos2_5_TransferPipeline, - Cosmos2TextToImagePipeline, - Cosmos2VideoToWorldPipeline, - Cosmos3OmniPipeline, - CosmosActionCondition, - CosmosTextToWorldPipeline, - CosmosVideoToWorldPipeline, - CycleDiffusionPipeline, - DiffusionGemmaPipeline, - DiffusionGemmaPipelineOutput, - DreamLiteMobilePipeline, - DreamLitePipeline, - DreamLitePipelineOutput, - EasyAnimateControlPipeline, - EasyAnimateInpaintPipeline, - EasyAnimatePipeline, - ErnieImagePipeline, - Flux2KleinInpaintPipeline, - Flux2KleinKVPipeline, - Flux2KleinPipeline, - Flux2Pipeline, - FluxControlImg2ImgPipeline, - FluxControlInpaintPipeline, - FluxControlNetImg2ImgPipeline, - FluxControlNetInpaintPipeline, - FluxControlNetPipeline, - FluxControlPipeline, - FluxFillPipeline, - FluxImg2ImgPipeline, - FluxInpaintPipeline, - FluxKontextInpaintPipeline, - FluxKontextPipeline, - FluxPipeline, - FluxPriorReduxPipeline, - GlmImagePipeline, - HeliosPipeline, - HeliosPyramidPipeline, - HiDreamImagePipeline, - HunyuanDiTControlNetPipeline, - HunyuanDiTPAGPipeline, - HunyuanDiTPipeline, - HunyuanImagePipeline, - HunyuanImageRefinerPipeline, - HunyuanSkyreelsImageToVideoPipeline, - HunyuanVideo15ImageToVideoPipeline, - HunyuanVideo15Pipeline, - HunyuanVideoFramepackPipeline, - HunyuanVideoImageToVideoPipeline, - HunyuanVideoPipeline, - I2VGenXLPipeline, - Ideogram4Pipeline, - Ideogram4PromptEnhancerHead, - IFImg2ImgPipeline, - IFImg2ImgSuperResolutionPipeline, - IFInpaintingPipeline, - IFInpaintingSuperResolutionPipeline, - IFPipeline, - IFSuperResolutionPipeline, - ImageTextPipelineOutput, - JoyImageEditPipeline, - JoyImageEditPipelineOutput, - JoyImageEditPlusPipeline, - JoyImageEditPlusPipelineOutput, - Kandinsky3Img2ImgPipeline, - Kandinsky3Pipeline, - Kandinsky5I2IPipeline, - Kandinsky5I2VPipeline, - Kandinsky5T2IPipeline, - Kandinsky5T2VPipeline, - KandinskyCombinedPipeline, - KandinskyImg2ImgCombinedPipeline, - KandinskyImg2ImgPipeline, - KandinskyInpaintCombinedPipeline, - KandinskyInpaintPipeline, - KandinskyPipeline, - KandinskyPriorPipeline, - KandinskyV22CombinedPipeline, - KandinskyV22ControlnetImg2ImgPipeline, - KandinskyV22ControlnetPipeline, - KandinskyV22Img2ImgCombinedPipeline, - KandinskyV22Img2ImgPipeline, - KandinskyV22InpaintCombinedPipeline, - KandinskyV22InpaintPipeline, - KandinskyV22Pipeline, - KandinskyV22PriorEmb2EmbPipeline, - KandinskyV22PriorPipeline, - Krea2Pipeline, - LatentConsistencyModelImg2ImgPipeline, - LatentConsistencyModelPipeline, - LattePipeline, - LDMTextToImagePipeline, - LEditsPPPipelineStableDiffusion, - LEditsPPPipelineStableDiffusionXL, - LLaDA2Pipeline, - LLaDA2PipelineOutput, - LongCatAudioDiTPipeline, - LongCatImageEditPipeline, - LongCatImagePipeline, - LTX2ConditionPipeline, - LTX2HDRPipeline, - LTX2ImageToVideoPipeline, - LTX2InContextPipeline, - LTX2LatentUpsamplePipeline, - LTX2Pipeline, - LTXConditionPipeline, - LTXI2VLongMultiPromptPipeline, - LTXImageToVideoPipeline, - LTXLatentUpsamplePipeline, - LTXPipeline, - LucyEditPipeline, - Lumina2Pipeline, - Lumina2Text2ImgPipeline, - LuminaPipeline, - LuminaText2ImgPipeline, - MarigoldDepthPipeline, - MarigoldIntrinsicsPipeline, - MarigoldNormalsPipeline, - MochiPipeline, - MotifVideoImage2VideoPipeline, - MotifVideoPipeline, - MotifVideoPipelineOutput, - MusicLDMPipeline, - NucleusMoEImagePipeline, - OmniGenPipeline, - OvisImagePipeline, - PaintByExamplePipeline, - PIAPipeline, - PixArtAlphaPipeline, - PixArtSigmaPAGPipeline, - PixArtSigmaPipeline, - PRXPipeline, - PRXPixelPipeline, - QwenImageControlNetInpaintPipeline, - QwenImageControlNetPipeline, - QwenImageEditInpaintPipeline, - QwenImageEditPipeline, - QwenImageEditPlusPipeline, - QwenImageImg2ImgPipeline, - QwenImageInpaintPipeline, - QwenImageLayeredPipeline, - QwenImagePipeline, - ReduxImageEncoder, - SanaControlNetPipeline, - SanaImageToVideoPipeline, - SanaPAGPipeline, - SanaPipeline, - SanaSprintImg2ImgPipeline, - SanaSprintPipeline, - SanaVideoPipeline, - SemanticStableDiffusionPipeline, - ShapEImg2ImgPipeline, - ShapEPipeline, - SkyReelsV2DiffusionForcingImageToVideoPipeline, - SkyReelsV2DiffusionForcingPipeline, - SkyReelsV2DiffusionForcingVideoToVideoPipeline, - SkyReelsV2ImageToVideoPipeline, - SkyReelsV2Pipeline, - StableAudioPipeline, - StableAudioProjectionModel, - StableCascadeCombinedPipeline, - StableCascadeDecoderPipeline, - StableCascadePriorPipeline, - StableDiffusion3ControlNetInpaintingPipeline, - StableDiffusion3ControlNetPipeline, - StableDiffusion3Img2ImgPipeline, - StableDiffusion3InpaintPipeline, - StableDiffusion3PAGImg2ImgPipeline, - StableDiffusion3PAGPipeline, - StableDiffusion3Pipeline, - StableDiffusionAdapterPipeline, - StableDiffusionAttendAndExcitePipeline, - StableDiffusionControlNetImg2ImgPipeline, - StableDiffusionControlNetInpaintPipeline, - StableDiffusionControlNetPAGInpaintPipeline, - StableDiffusionControlNetPAGPipeline, - StableDiffusionControlNetPipeline, - StableDiffusionControlNetXSPipeline, - StableDiffusionDepth2ImgPipeline, - StableDiffusionDiffEditPipeline, - StableDiffusionGLIGENPipeline, - StableDiffusionGLIGENTextImagePipeline, - StableDiffusionImageVariationPipeline, - StableDiffusionImg2ImgPipeline, - StableDiffusionInpaintPipeline, - StableDiffusionInpaintPipelineLegacy, - StableDiffusionInstructPix2PixPipeline, - StableDiffusionLatentUpscalePipeline, - StableDiffusionLDM3DPipeline, - StableDiffusionModelEditingPipeline, - StableDiffusionPAGImg2ImgPipeline, - StableDiffusionPAGInpaintPipeline, - StableDiffusionPAGPipeline, - StableDiffusionPanoramaPipeline, - StableDiffusionParadigmsPipeline, - StableDiffusionPipeline, - StableDiffusionPipelineSafe, - StableDiffusionPix2PixZeroPipeline, - StableDiffusionSAGPipeline, - StableDiffusionUpscalePipeline, - StableDiffusionXLAdapterPipeline, - StableDiffusionXLControlNetImg2ImgPipeline, - StableDiffusionXLControlNetInpaintPipeline, - StableDiffusionXLControlNetPAGImg2ImgPipeline, - StableDiffusionXLControlNetPAGPipeline, - StableDiffusionXLControlNetPipeline, - StableDiffusionXLControlNetUnionImg2ImgPipeline, - StableDiffusionXLControlNetUnionInpaintPipeline, - StableDiffusionXLControlNetUnionPipeline, - StableDiffusionXLControlNetXSPipeline, - StableDiffusionXLImg2ImgPipeline, - StableDiffusionXLInpaintPipeline, - StableDiffusionXLInstructPix2PixPipeline, - StableDiffusionXLPAGImg2ImgPipeline, - StableDiffusionXLPAGInpaintPipeline, - StableDiffusionXLPAGPipeline, - StableDiffusionXLPipeline, - StableUnCLIPImg2ImgPipeline, - StableUnCLIPPipeline, - StableVideoDiffusionPipeline, - TextToVideoSDPipeline, - TextToVideoZeroPipeline, - TextToVideoZeroSDXLPipeline, - UnCLIPImageVariationPipeline, - UnCLIPPipeline, - UniDiffuserModel, - UniDiffuserPipeline, - UniDiffuserTextDecoder, - VersatileDiffusionDualGuidedPipeline, - VersatileDiffusionImageVariationPipeline, - VersatileDiffusionPipeline, - VersatileDiffusionTextToImagePipeline, - VideoToVideoSDPipeline, - VisualClozeGenerationPipeline, - VisualClozePipeline, - VQDiffusionPipeline, - WanAnimatePipeline, - WanImageToVideoPipeline, - WanPipeline, - WanVACEPipeline, - WanVideoToVideoPipeline, - WuerstchenCombinedPipeline, - WuerstchenDecoderPipeline, - WuerstchenPriorPipeline, - ZImageControlNetInpaintPipeline, - ZImageControlNetPipeline, - ZImageImg2ImgPipeline, - ZImageInpaintPipeline, - ZImageOmniPipeline, - ZImagePipeline, - ) - - try: - if not (is_torch_available() and is_transformers_available() and is_sentencepiece_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_torch_and_transformers_and_sentencepiece_objects import * # noqa F403 - else: - from .pipelines import KolorsImg2ImgPipeline, KolorsPAGPipeline, KolorsPipeline - - try: - if not (is_torch_available() and is_transformers_available() and is_opencv_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_torch_and_transformers_and_opencv_objects import * # noqa F403 - else: - from .pipelines import ConsisIDPipeline - - try: - if not (is_torch_available() and is_transformers_available() and is_onnx_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_torch_and_transformers_and_onnx_objects import * # noqa F403 - else: - from .pipelines import ( - OnnxStableDiffusionImg2ImgPipeline, - OnnxStableDiffusionInpaintPipeline, - OnnxStableDiffusionInpaintPipelineLegacy, - OnnxStableDiffusionPipeline, - OnnxStableDiffusionUpscalePipeline, - StableDiffusionOnnxPipeline, - ) - - try: - if not (is_torch_available() and is_librosa_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_torch_and_librosa_objects import * # noqa F403 - else: - from .pipelines import AudioDiffusionPipeline, Mel - - try: - if not (is_transformers_available() and is_torch_available() and is_note_seq_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_transformers_and_torch_and_note_seq_objects import * # noqa F403 - else: - from .pipelines import SpectrogramDiffusionPipeline - - try: - if not (is_note_seq_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from .utils.dummy_note_seq_objects import * # noqa F403 - else: - from .pipelines import MidiProcessor - -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - extra_objects={"__version__": __version__}, - ) diff --git a/diffusers/callbacks.py b/diffusers/callbacks.py deleted file mode 100644 index 087a6b7fee565add21a99f98207d48efa8280d3c..0000000000000000000000000000000000000000 --- a/diffusers/callbacks.py +++ /dev/null @@ -1,244 +0,0 @@ -from typing import Any - -from .configuration_utils import ConfigMixin, register_to_config -from .utils import CONFIG_NAME - - -class PipelineCallback(ConfigMixin): - """ - Base class for all the official callbacks used in a pipeline. This class provides a structure for implementing - custom callbacks and ensures that all callbacks have a consistent interface. - - Please implement the following: - `tensor_inputs`: This should return a list of tensor inputs specific to your callback. You will only be able to - include - variables listed in the `._callback_tensor_inputs` attribute of your pipeline class. - `callback_fn`: This method defines the core functionality of your callback. - """ - - config_name = CONFIG_NAME - - @register_to_config - def __init__(self, cutoff_step_ratio=1.0, cutoff_step_index=None): - super().__init__() - - if (cutoff_step_ratio is None and cutoff_step_index is None) or ( - cutoff_step_ratio is not None and cutoff_step_index is not None - ): - raise ValueError("Either cutoff_step_ratio or cutoff_step_index should be provided, not both or none.") - - if cutoff_step_ratio is not None and ( - not isinstance(cutoff_step_ratio, float) or not (0.0 <= cutoff_step_ratio <= 1.0) - ): - raise ValueError("cutoff_step_ratio must be a float between 0.0 and 1.0.") - - @property - def tensor_inputs(self) -> list[str]: - raise NotImplementedError(f"You need to set the attribute `tensor_inputs` for {self.__class__}") - - def callback_fn(self, pipeline, step_index, timesteps, callback_kwargs) -> dict[str, Any]: - raise NotImplementedError(f"You need to implement the method `callback_fn` for {self.__class__}") - - def __call__(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]: - return self.callback_fn(pipeline, step_index, timestep, callback_kwargs) - - -class MultiPipelineCallbacks: - """ - This class is designed to handle multiple pipeline callbacks. It accepts a list of PipelineCallback objects and - provides a unified interface for calling all of them. - """ - - def __init__(self, callbacks: list[PipelineCallback]): - self.callbacks = callbacks - - @property - def tensor_inputs(self) -> list[str]: - return [input for callback in self.callbacks for input in callback.tensor_inputs] - - def __call__(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]: - """ - Calls all the callbacks in order with the given arguments and returns the final callback_kwargs. - """ - for callback in self.callbacks: - callback_kwargs = callback(pipeline, step_index, timestep, callback_kwargs) - - return callback_kwargs - - -class SDCFGCutoffCallback(PipelineCallback): - """ - Callback function for Stable Diffusion Pipelines. After certain number of steps (set by `cutoff_step_ratio` or - `cutoff_step_index`), this callback will disable the CFG. - - Note: This callback mutates the pipeline by changing the `_guidance_scale` attribute to 0.0 after the cutoff step. - """ - - tensor_inputs = ["prompt_embeds"] - - def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]: - cutoff_step_ratio = self.config.cutoff_step_ratio - cutoff_step_index = self.config.cutoff_step_index - - # Use cutoff_step_index if it's not None, otherwise use cutoff_step_ratio - cutoff_step = ( - cutoff_step_index if cutoff_step_index is not None else int(pipeline.num_timesteps * cutoff_step_ratio) - ) - - if step_index == cutoff_step: - prompt_embeds = callback_kwargs[self.tensor_inputs[0]] - prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens. - - pipeline._guidance_scale = 0.0 - - callback_kwargs[self.tensor_inputs[0]] = prompt_embeds - return callback_kwargs - - -class SDXLCFGCutoffCallback(PipelineCallback): - """ - Callback function for the base Stable Diffusion XL Pipelines. After certain number of steps (set by - `cutoff_step_ratio` or `cutoff_step_index`), this callback will disable the CFG. - - Note: This callback mutates the pipeline by changing the `_guidance_scale` attribute to 0.0 after the cutoff step. - """ - - tensor_inputs = [ - "prompt_embeds", - "add_text_embeds", - "add_time_ids", - ] - - def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]: - cutoff_step_ratio = self.config.cutoff_step_ratio - cutoff_step_index = self.config.cutoff_step_index - - # Use cutoff_step_index if it's not None, otherwise use cutoff_step_ratio - cutoff_step = ( - cutoff_step_index if cutoff_step_index is not None else int(pipeline.num_timesteps * cutoff_step_ratio) - ) - - if step_index == cutoff_step: - prompt_embeds = callback_kwargs[self.tensor_inputs[0]] - prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens. - - add_text_embeds = callback_kwargs[self.tensor_inputs[1]] - add_text_embeds = add_text_embeds[-1:] # "-1" denotes the embeddings for conditional pooled text tokens - - add_time_ids = callback_kwargs[self.tensor_inputs[2]] - add_time_ids = add_time_ids[-1:] # "-1" denotes the embeddings for conditional added time vector - - pipeline._guidance_scale = 0.0 - - callback_kwargs[self.tensor_inputs[0]] = prompt_embeds - callback_kwargs[self.tensor_inputs[1]] = add_text_embeds - callback_kwargs[self.tensor_inputs[2]] = add_time_ids - - return callback_kwargs - - -class SDXLControlnetCFGCutoffCallback(PipelineCallback): - """ - Callback function for the Controlnet Stable Diffusion XL Pipelines. After certain number of steps (set by - `cutoff_step_ratio` or `cutoff_step_index`), this callback will disable the CFG. - - Note: This callback mutates the pipeline by changing the `_guidance_scale` attribute to 0.0 after the cutoff step. - """ - - tensor_inputs = [ - "prompt_embeds", - "add_text_embeds", - "add_time_ids", - "image", - ] - - def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]: - cutoff_step_ratio = self.config.cutoff_step_ratio - cutoff_step_index = self.config.cutoff_step_index - - # Use cutoff_step_index if it's not None, otherwise use cutoff_step_ratio - cutoff_step = ( - cutoff_step_index if cutoff_step_index is not None else int(pipeline.num_timesteps * cutoff_step_ratio) - ) - - if step_index == cutoff_step: - prompt_embeds = callback_kwargs[self.tensor_inputs[0]] - prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens. - - add_text_embeds = callback_kwargs[self.tensor_inputs[1]] - add_text_embeds = add_text_embeds[-1:] # "-1" denotes the embeddings for conditional pooled text tokens - - add_time_ids = callback_kwargs[self.tensor_inputs[2]] - add_time_ids = add_time_ids[-1:] # "-1" denotes the embeddings for conditional added time vector - - # For Controlnet - image = callback_kwargs[self.tensor_inputs[3]] - image = image[-1:] - - pipeline._guidance_scale = 0.0 - - callback_kwargs[self.tensor_inputs[0]] = prompt_embeds - callback_kwargs[self.tensor_inputs[1]] = add_text_embeds - callback_kwargs[self.tensor_inputs[2]] = add_time_ids - callback_kwargs[self.tensor_inputs[3]] = image - - return callback_kwargs - - -class IPAdapterScaleCutoffCallback(PipelineCallback): - """ - Callback function for any pipeline that inherits `IPAdapterMixin`. After certain number of steps (set by - `cutoff_step_ratio` or `cutoff_step_index`), this callback will set the IP Adapter scale to `0.0`. - - Note: This callback mutates the IP Adapter attention processors by setting the scale to 0.0 after the cutoff step. - """ - - tensor_inputs = [] - - def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]: - cutoff_step_ratio = self.config.cutoff_step_ratio - cutoff_step_index = self.config.cutoff_step_index - - # Use cutoff_step_index if it's not None, otherwise use cutoff_step_ratio - cutoff_step = ( - cutoff_step_index if cutoff_step_index is not None else int(pipeline.num_timesteps * cutoff_step_ratio) - ) - - if step_index == cutoff_step: - pipeline.set_ip_adapter_scale(0.0) - return callback_kwargs - - -class SD3CFGCutoffCallback(PipelineCallback): - """ - Callback function for Stable Diffusion 3 Pipelines. After certain number of steps (set by `cutoff_step_ratio` or - `cutoff_step_index`), this callback will disable the CFG. - - Note: This callback mutates the pipeline by changing the `_guidance_scale` attribute to 0.0 after the cutoff step. - """ - - tensor_inputs = ["prompt_embeds", "pooled_prompt_embeds"] - - def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[str, Any]: - cutoff_step_ratio = self.config.cutoff_step_ratio - cutoff_step_index = self.config.cutoff_step_index - - # Use cutoff_step_index if it's not None, otherwise use cutoff_step_ratio - cutoff_step = ( - cutoff_step_index if cutoff_step_index is not None else int(pipeline.num_timesteps * cutoff_step_ratio) - ) - - if step_index == cutoff_step: - prompt_embeds = callback_kwargs[self.tensor_inputs[0]] - prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens. - - pooled_prompt_embeds = callback_kwargs[self.tensor_inputs[1]] - pooled_prompt_embeds = pooled_prompt_embeds[ - -1: - ] # "-1" denotes the embeddings for conditional pooled text tokens. - - pipeline._guidance_scale = 0.0 - - callback_kwargs[self.tensor_inputs[0]] = prompt_embeds - callback_kwargs[self.tensor_inputs[1]] = pooled_prompt_embeds - return callback_kwargs diff --git a/diffusers/commands/__init__.py b/diffusers/commands/__init__.py deleted file mode 100644 index 9f1a4e407bddd66c7ca3eb5657ef627f87368923..0000000000000000000000000000000000000000 --- a/diffusers/commands/__init__.py +++ /dev/null @@ -1,27 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from abc import ABC, abstractmethod -from argparse import ArgumentParser - - -class BaseDiffusersCLICommand(ABC): - @staticmethod - @abstractmethod - def register_subcommand(parser: ArgumentParser): - raise NotImplementedError() - - @abstractmethod - def run(self): - raise NotImplementedError() diff --git a/diffusers/commands/custom_blocks.py b/diffusers/commands/custom_blocks.py deleted file mode 100644 index 7ebaf785ba48669adb04b8952564c8e3eca12bdb..0000000000000000000000000000000000000000 --- a/diffusers/commands/custom_blocks.py +++ /dev/null @@ -1,140 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""`diffusers-cli custom_blocks` — package a local `ModularPipelineBlocks` subclass for the Hub. - -Parses `block.py` (or `--block_module_name`), instantiates the chosen block, and calls `save_pretrained` in the current -working directory. -""" - -import ast -import importlib.util -import os -from argparse import ArgumentParser, Namespace -from pathlib import Path - -from ..utils import logging -from . import BaseDiffusersCLICommand - - -EXPECTED_PARENT_CLASSES = ["ModularPipelineBlocks"] - - -def conversion_command_factory(args: Namespace): - return CustomBlocksCommand(args.block_module_name, args.block_class_name) - - -class CustomBlocksCommand(BaseDiffusersCLICommand): - @staticmethod - def register_subcommand(parser: ArgumentParser): - from argparse import RawDescriptionHelpFormatter - - epilog = ( - "Examples\n" - " $ diffusers-cli custom_blocks\n" - " $ diffusers-cli custom_blocks --block_module_name my_block.py\n" - " $ diffusers-cli custom_blocks --block_module_name my_block.py --block_class_name MyDenoiseBlock\n" - "\n" - "Learn more\n" - " Use `diffusers-cli --help` for more information about a command.\n" - " Read the documentation at https://huggingface.co/docs/diffusers\n" - ) - - conversion_parser = parser.add_parser( - "custom_blocks", - help="Package a local ModularPipelineBlocks subclass for the Hub.", - usage="\n diffusers-cli custom_blocks [options]", - epilog=epilog, - formatter_class=RawDescriptionHelpFormatter, - ) - conversion_parser._optionals.title = "Options" - conversion_parser.add_argument( - "--block_module_name", - type=str, - default="block.py", - help="Module filename in which the custom block will be implemented.", - ) - conversion_parser.add_argument( - "--block_class_name", - type=str, - default=None, - help="Name of the custom block. If provided None, we will try to infer it.", - ) - conversion_parser.set_defaults(func=conversion_command_factory) - - def __init__(self, block_module_name: str = "block.py", block_class_name: str = None): - self.logger = logging.get_logger("diffusers-cli/custom_blocks") - self.block_module_name = Path(block_module_name) - self.block_class_name = block_class_name - - def run(self): - # determine the block to be saved. - out = self._get_class_names(self.block_module_name) - classes_found = list({cls for cls, _ in out}) - - if self.block_class_name is not None: - child_class, parent_class = self._choose_block(out, self.block_class_name) - if child_class is None and parent_class is None: - raise ValueError( - "`block_class_name` could not be retrieved. Available classes from " - f"{self.block_module_name}:\n{classes_found}" - ) - else: - self.logger.info( - f"Found classes: {classes_found} will be using {classes_found[0]}. " - "If this needs to be changed, re-run the command specifying `block_class_name`." - ) - child_class, parent_class = out[0][0], out[0][1] - - # dynamically get the custom block and initialize it to call `save_pretrained` in the current directory. - # the user is responsible for running it, so I guess that is safe? - module_name = f"__dynamic__{self.block_module_name.stem}" - spec = importlib.util.spec_from_file_location(module_name, str(self.block_module_name)) - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - getattr(module, child_class)().save_pretrained(os.getcwd()) - - def _choose_block(self, candidates, chosen=None): - for cls, base in candidates: - if cls == chosen: - return cls, base - return None, None - - def _get_class_names(self, file_path): - source = file_path.read_text(encoding="utf-8") - try: - tree = ast.parse(source, filename=file_path) - except SyntaxError as e: - raise ValueError(f"Could not parse {file_path!r}: {e}") from e - - results: list[tuple[str, str]] = [] - for node in tree.body: - if not isinstance(node, ast.ClassDef): - continue - - base_names = [bname for b in node.bases if (bname := self._get_base_name(b)) is not None] - - for allowed in EXPECTED_PARENT_CLASSES: - if allowed in base_names: - results.append((node.name, allowed)) - - return results - - def _get_base_name(self, node: ast.expr): - if isinstance(node, ast.Name): - return node.id - elif isinstance(node, ast.Attribute): - val = self._get_base_name(node.value) - return f"{val}.{node.attr}" if val else node.attr - return None diff --git a/diffusers/commands/diffusers_cli.py b/diffusers/commands/diffusers_cli.py deleted file mode 100644 index 0e4f2c27fb64f5a829fdf0443e6e7021ed380607..0000000000000000000000000000000000000000 --- a/diffusers/commands/diffusers_cli.py +++ /dev/null @@ -1,69 +0,0 @@ -#!/usr/bin/env python -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from argparse import ArgumentParser - -from huggingface_hub.cli._output import OutputFormat, out - -from .custom_blocks import CustomBlocksCommand -from .env import EnvironmentCommand -from .fp16_safetensors import FP16SafetensorsCommand -from .run import RunCommand -from .schema import SchemaCommand -from .skills import SkillsCommand - - -def main(): - parser = ArgumentParser( - prog="diffusers-cli", - usage="\n diffusers-cli [--format ] [options]", - ) - parser._optionals.title = "Options" - parser.add_argument( - "--format", - choices=[m.value for m in OutputFormat], - default=OutputFormat.auto.value, - help=( - "Output format. 'auto' (default) picks 'agent' when an AI coding agent is detected " - "(via CLAUDECODE/CURSOR_AI/AIDER_AI_CONTEXT/... env vars) and 'human' otherwise. " - "Must appear before the subcommand." - ), - ) - commands_parser = parser.add_subparsers(title="Commands", metavar="") - - # Register commands - EnvironmentCommand.register_subcommand(commands_parser) - FP16SafetensorsCommand.register_subcommand(commands_parser) - CustomBlocksCommand.register_subcommand(commands_parser) - RunCommand.register_subcommand(commands_parser) - SchemaCommand.register_subcommand(commands_parser) - SkillsCommand.register_subcommand(commands_parser) - - # Let's go - args = parser.parse_args() - - out.set_mode(OutputFormat(args.format)) - - if not hasattr(args, "func"): - parser.print_help() - exit(1) - - # Run - service = args.func(args) - service.run() - - -if __name__ == "__main__": - main() diff --git a/diffusers/commands/env.py b/diffusers/commands/env.py deleted file mode 100644 index cbd2d111385be480f1ca1fbbb65a379291383d7f..0000000000000000000000000000000000000000 --- a/diffusers/commands/env.py +++ /dev/null @@ -1,185 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import importlib.metadata -import platform -import subprocess -from argparse import ArgumentParser - -import huggingface_hub - -from .. import __version__ as version -from ..utils import ( - is_accelerate_available, - is_bitsandbytes_available, - is_gguf_available, - is_google_colab, - is_nvidia_modelopt_available, - is_optimum_quanto_available, - is_peft_available, - is_safetensors_available, - is_torch_available, - is_torchao_available, - is_transformers_available, - is_xformers_available, -) -from . import BaseDiffusersCLICommand - - -# (display name, availability_fn, pypi distribution name for importlib.metadata.version) -_QUANTIZATION_BACKENDS = ( - ("bitsandbytes", is_bitsandbytes_available, "bitsandbytes"), - ("gguf", is_gguf_available, "gguf"), - ("optimum-quanto", is_optimum_quanto_available, "optimum-quanto"), - ("torchao", is_torchao_available, "torchao"), - ("nvidia-modelopt", is_nvidia_modelopt_available, "nvidia-modelopt"), -) - - -def info_command_factory(_): - return EnvironmentCommand() - - -class EnvironmentCommand(BaseDiffusersCLICommand): - @staticmethod - def register_subcommand(parser: ArgumentParser) -> None: - download_parser = parser.add_parser( - "env", - help="Print versions of diffusers and its dependencies (for bug reports).", - usage="\n diffusers-cli env", - ) - download_parser._optionals.title = "Options" - download_parser.set_defaults(func=info_command_factory) - - def run(self) -> dict: - hub_version = huggingface_hub.__version__ - - safetensors_version = "not installed" - if is_safetensors_available(): - import safetensors - - safetensors_version = safetensors.__version__ - - pt_version = "not installed" - pt_cuda_available = "NA" - if is_torch_available(): - import torch - - pt_version = torch.__version__ - pt_cuda_available = torch.cuda.is_available() - - transformers_version = "not installed" - if is_transformers_available(): - import transformers - - transformers_version = transformers.__version__ - - accelerate_version = "not installed" - if is_accelerate_available(): - import accelerate - - accelerate_version = accelerate.__version__ - - peft_version = "not installed" - if is_peft_available(): - import peft - - peft_version = peft.__version__ - - quantization_versions = {} - for backend_name, is_available_fn, dist_name in _QUANTIZATION_BACKENDS: - if not is_available_fn(): - continue - try: - quantization_versions[backend_name] = importlib.metadata.version(dist_name) - except importlib.metadata.PackageNotFoundError: - quantization_versions[backend_name] = "N/A" - - xformers_version = "not installed" - if is_xformers_available(): - import xformers - - xformers_version = xformers.__version__ - - platform_info = platform.platform() - - is_google_colab_str = "Yes" if is_google_colab() else "No" - - accelerator = "NA" - if platform.system() in {"Linux", "Windows"}: - try: - sp = subprocess.Popen( - ["nvidia-smi", "--query-gpu=gpu_name,memory.total", "--format=csv,noheader"], - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - ) - out_str, _ = sp.communicate() - out_str = out_str.decode("utf-8") - - if len(out_str) > 0: - accelerator = out_str.strip() - except FileNotFoundError: - pass - elif platform.system() == "Darwin": # Mac OS - try: - sp = subprocess.Popen( - ["system_profiler", "SPDisplaysDataType"], - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - ) - out_str, _ = sp.communicate() - out_str = out_str.decode("utf-8") - - start = out_str.find("Chipset Model:") - if start != -1: - start += len("Chipset Model:") - end = out_str.find("\n", start) - accelerator = out_str[start:end].strip() - - start = out_str.find("VRAM (Total):") - if start != -1: - start += len("VRAM (Total):") - end = out_str.find("\n", start) - accelerator += " VRAM: " + out_str[start:end].strip() - except FileNotFoundError: - pass - else: - print("It seems you are running an unusual OS. Could you fill in the accelerator manually?") - - info = { - "🤗 Diffusers version": version, - "Platform": platform_info, - "Running on Google Colab?": is_google_colab_str, - "Python version": platform.python_version(), - "PyTorch version (GPU?)": f"{pt_version} ({pt_cuda_available})", - "Huggingface_hub version": hub_version, - "Transformers version": transformers_version, - "Accelerate version": accelerate_version, - "PEFT version": peft_version, - **{f"{name} version": ver for name, ver in quantization_versions.items()}, - "Safetensors version": safetensors_version, - "xFormers version": xformers_version, - "Accelerator": accelerator, - "Using GPU in script?": "", - "Using distributed or parallel set-up in script?": "", - } - - print("\nCopy-and-paste the text below in your GitHub issue and FILL OUT the two last points.\n") - print(self.format_dict(info)) - - return info - - @staticmethod - def format_dict(d: dict) -> str: - return "\n".join([f"- {prop}: {val}" for prop, val in d.items()]) + "\n" diff --git a/diffusers/commands/fp16_safetensors.py b/diffusers/commands/fp16_safetensors.py deleted file mode 100644 index ec91bba357a036824aeba010fcabb963f8877ea0..0000000000000000000000000000000000000000 --- a/diffusers/commands/fp16_safetensors.py +++ /dev/null @@ -1,144 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -""" -Usage example: - diffusers-cli fp16_safetensors --ckpt_id=openai/shap-e --fp16 --use_safetensors -""" - -import glob -import json -import warnings -from argparse import ArgumentParser, Namespace -from importlib import import_module - -import huggingface_hub -import torch -from huggingface_hub import hf_hub_download -from packaging import version - -from ..utils import logging -from . import BaseDiffusersCLICommand - - -def conversion_command_factory(args: Namespace): - warnings.warn( - "`diffusers-cli fp16_safetensors` is deprecated and will be removed in a future version. " - "Convert weights to fp16 safetensors directly with `safetensors.torch.save_file` or via " - "`pipeline.save_pretrained(..., safe_serialization=True, variant='fp16')`.", - FutureWarning, - stacklevel=2, - ) - if args.use_auth_token: - warnings.warn( - "The `--use_auth_token` flag is deprecated and will be removed in a future version." - "Authentication is now handled automatically if the user is logged in." - ) - return FP16SafetensorsCommand(args.ckpt_id, args.fp16, args.use_safetensors) - - -class FP16SafetensorsCommand(BaseDiffusersCLICommand): - @staticmethod - def register_subcommand(parser: ArgumentParser): - conversion_parser = parser.add_parser( - "fp16_safetensors", - help="[DEPRECATED] Convert a Hub checkpoint's weights to fp16 safetensors and push back as a PR.", - usage="\n diffusers-cli fp16_safetensors [options]", - ) - conversion_parser._optionals.title = "Options" - conversion_parser.add_argument( - "--ckpt_id", - type=str, - help="Repo id of the checkpoints on which to run the conversion. Example: 'openai/shap-e'.", - ) - conversion_parser.add_argument( - "--fp16", action="store_true", help="If serializing the variables in FP16 precision." - ) - conversion_parser.add_argument( - "--use_safetensors", action="store_true", help="If serializing in the safetensors format." - ) - conversion_parser.add_argument( - "--use_auth_token", - action="store_true", - help="When working with checkpoints having private visibility. When used `hf auth login` needs to be run beforehand.", - ) - conversion_parser.set_defaults(func=conversion_command_factory) - - def __init__(self, ckpt_id: str, fp16: bool, use_safetensors: bool): - self.logger = logging.get_logger("diffusers-cli/fp16_safetensors") - self.ckpt_id = ckpt_id - self.local_ckpt_dir = f"/tmp/{ckpt_id}" - self.fp16 = fp16 - - self.use_safetensors = use_safetensors - - if not self.use_safetensors and not self.fp16: - raise NotImplementedError( - "When `use_safetensors` and `fp16` both are False, then this command is of no use." - ) - - def run(self): - if version.parse(huggingface_hub.__version__) < version.parse("0.9.0"): - raise ImportError( - "The huggingface_hub version must be >= 0.9.0 to use this command. Please update your huggingface_hub" - " installation." - ) - else: - from huggingface_hub import create_commit - from huggingface_hub._commit_api import CommitOperationAdd - - model_index = hf_hub_download(repo_id=self.ckpt_id, filename="model_index.json") - with open(model_index, "r") as f: - pipeline_class_name = json.load(f)["_class_name"] - pipeline_class = getattr(import_module("diffusers"), pipeline_class_name) - self.logger.info(f"Pipeline class imported: {pipeline_class_name}.") - - # Load the appropriate pipeline. We could have used `DiffusionPipeline` - # here, but just to avoid potential edge cases. - pipeline = pipeline_class.from_pretrained( - self.ckpt_id, torch_dtype=torch.float16 if self.fp16 else torch.float32 - ) - pipeline.save_pretrained( - self.local_ckpt_dir, - safe_serialization=True if self.use_safetensors else False, - variant="fp16" if self.fp16 else None, - ) - self.logger.info(f"Pipeline locally saved to {self.local_ckpt_dir}.") - - # Fetch all the paths. - if self.fp16: - modified_paths = glob.glob(f"{self.local_ckpt_dir}/*/*.fp16.*") - elif self.use_safetensors: - modified_paths = glob.glob(f"{self.local_ckpt_dir}/*/*.safetensors") - - # Prepare for the PR. - commit_message = f"Serialize variables with FP16: {self.fp16} and safetensors: {self.use_safetensors}." - operations = [] - for path in modified_paths: - operations.append(CommitOperationAdd(path_in_repo="/".join(path.split("/")[4:]), path_or_fileobj=path)) - - # Open the PR. - commit_description = ( - "Variables converted by the [`diffusers`' `fp16_safetensors`" - " CLI](https://github.com/huggingface/diffusers/blob/main/src/diffusers/commands/fp16_safetensors.py)." - ) - hub_pr_url = create_commit( - repo_id=self.ckpt_id, - operations=operations, - commit_message=commit_message, - commit_description=commit_description, - repo_type="model", - create_pr=True, - ).pr_url - self.logger.info(f"PR created here: {hub_pr_url}.") diff --git a/diffusers/commands/run.py b/diffusers/commands/run.py deleted file mode 100644 index 9cd63854783451d44d0384e9a75f4d0705e8c317..0000000000000000000000000000000000000000 --- a/diffusers/commands/run.py +++ /dev/null @@ -1,1227 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""`diffusers-cli run` — single agentic entry point. - -Runs any diffusers pipeline (standard or modular) by forwarding `--pipeline-kwargs` verbatim, saves the output by -detecting its runtime type, and can submit the same call to an HF Sandbox via `--remote`. -""" - -from __future__ import annotations - -import json -import os -import sys -from argparse import ArgumentParser, Namespace, _SubParsersAction -from pathlib import Path -from typing import Any - -from huggingface_hub.cli._output import out - -from diffusers.models.attention_dispatch import _HUB_KERNELS_REGISTRY -from diffusers.utils import load_image, load_video, logging - -from . import BaseDiffusersCLICommand - - -logger = logging.get_logger("diffusers-cli/run") - - -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- - -DEFAULT_OUTPUT_DIR = str(Path.home() / ".diffusers" / "cli" / "run" / "outputs") -DTYPE_CHOICES = ("auto", "float16", "fp16", "bfloat16", "bf16", "float32", "fp32") -CPU_OFFLOAD_CHOICES = ("model", "group") - - -ATTENTION_BACKEND_CHOICES = ("default", *sorted(b.value for b in _HUB_KERNELS_REGISTRY)) - -# Kwarg keys whose string value gets auto-loaded before being passed to the pipeline call. -# Images resolve via `diffusers.utils.load_image` → PIL.Image.Image; videos resolve via -# `diffusers.utils.load_video` → list[PIL.Image.Image]. -_IMAGE_INPUT_KEYS = ( - "image", - "mask_image", - "control_image", - "ip_adapter_image", - "image_2", -) -_VIDEO_INPUT_KEYS = ( - "video", - "control_video", -) -_AUDIO_INPUT_KEYS = ( - "initial_audio_waveforms", - "reference_audio", - "src_audio", -) - -# Pipeline attribute prefixes that identify a denoiser submodule. Matches base names -# (`transformer`, `unet`) and their numbered variants (`transformer_2`, etc.). -_DENOISER_COMPONENT_KEYS = ("transformer", "unet") - -_DEFAULT_REMOTE_DEPS = ( - "diffusers", - "accelerate", - "transformers", - "safetensors", - "sentencepiece", # required by several text-encoder tokenizers (T5, LLaMA, …) - "ftfy", # required by older CLIP text-encoder paths -) - -# Base sandbox image — provides torch + CUDA so `uv pip install --system` -# only has to add the small Python deps. cuda12.8 is the highest cuda12.x tag -# below the HF Jobs host driver's CUDA 12.9 max. -_DEFAULT_REMOTE_IMAGE = "pytorch/pytorch:2.10.0-cuda12.8-cudnn9-runtime" - -# Installed console-script name invoked inside the sandbox after the deps land. -_CONTAINER_CLI_BINARY = "diffusers-cli" - -# Working directories inside the sandbox: local media from `--pipeline-kwargs` is uploaded -# under _SANDBOX_INPUTS_DIR, and the sandbox CLI is told to write its outputs under -# _SANDBOX_OUTPUTS_DIR so we can download them back afterwards. -_SANDBOX_INPUTS_DIR = "/tmp/diffusers-cli/inputs" -_SANDBOX_OUTPUTS_DIR = "/tmp/diffusers-cli/outputs" - -RUN_ID_ENV = "DIFFUSERS_CLI_RUN_ID" - -# Namespace keys that control *how* a remote run is dispatched, not what the sandbox CLI -# runs. They are stripped when forwarding argv to the sandbox. -REMOTE_KEYS = frozenset( - { - "remote", - "flavor", - "timeout", - "dependencies", - "namespace", - "image", - "keep_alive", - "sandbox_id", - "idle_timeout", - "volume", - "func", - "format", # top-level --format is a local rendering flag; never forward to the sandbox - } -) - - -# --------------------------------------------------------------------------- -# Argparse helpers -# --------------------------------------------------------------------------- - - -def _add_loading_arguments(parser: ArgumentParser) -> None: - parser.add_argument("--model", "-m", required=True, help="Model id on the Hugging Face Hub or local path.") - parser.add_argument( - "--device-map", - default=None, - help=( - "Component placement. Accepts a torch device string (`cuda`, `cuda:0`, `cpu`, `mps`), " - "`balanced` for pipeline-level auto-split across visible GPUs, or a JSON dict of " - '`{"": }` for explicit per-component placement. Auto-detected if omitted.' - ), - ) - parser.add_argument("--dtype", default="auto", choices=DTYPE_CHOICES, help="Torch dtype for pipeline weights.") - parser.add_argument("--variant", default=None, help='Optional weight variant (e.g. "fp16").') - parser.add_argument("--revision", default=None, help="Model revision (branch, tag, or commit SHA).") - parser.add_argument("--token", default=None, help="Hugging Face token for gated/private models.") - parser.add_argument("--trust-remote-code", action="store_true", help="Allow custom code from the Hub.") - parser.add_argument( - "--lora", - action="append", - default=None, - metavar="JSON", - help=( - "JSON dict describing a LoRA adapter to attach after the pipeline loads. Repeat to stack " - 'multiple adapters. Format: \'{"lora_id": "", "lora_scale": }\'. `lora_scale` ' - "defaults to 1.0; `adapter_name` is optional (auto-generated as `lora_` when stacking)." - ), - ) - - -def _add_optimization_arguments(parser: ArgumentParser) -> None: - parser.add_argument( - "--cpu-offload", - choices=CPU_OFFLOAD_CHOICES, - default=None, - help=( - "Offload pipeline components to CPU during inference. " - "'model' uses enable_model_cpu_offload, " - "'group' uses pipeline.enable_group_offload(leaf_level, use_stream=True)." - ), - ) - parser.add_argument( - "--attention-backend", - choices=ATTENTION_BACKEND_CHOICES, - default="default", - help=( - "Override the attention backend on the transformer/UNet. " - "Only Hub-hosted kernels are exposed — they auto-download on first use." - ), - ) - parser.add_argument("--vae-tiling", action="store_true", help="Enable VAE tiling (lower peak VRAM).") - parser.add_argument("--vae-slicing", action="store_true", help="Enable VAE slicing (lower peak VRAM).") - parser.add_argument( - "--context-parallel", - action="store_true", - help=( - "Enable Ulysses-style context parallelism (ulysses_anything mode). " - "Requires a DiT-based pipeline and launching the CLI under torchrun with ≥2 GPUs." - ), - ) - parser.add_argument( - "--compile", - nargs="?", - const='{"fullgraph": true}', - default=None, - metavar="JSON", - help=( - "torch.compile every denoiser submodule on the pipeline. Accepts an optional JSON " - 'object of kwargs forwarded to `torch.compile`, e.g. \'{"mode": "max-autotune", ' - '"fullgraph": true}\'. Bare `--compile` uses `fullgraph=true`. Adds a one-time ' - "compilation cost on the first step but speeds up every subsequent step — worth it " - "for multi-step generation (50+ steps)." - ), - ) - - -def _add_output_arguments(parser: ArgumentParser) -> None: - parser.add_argument( - "--output", - "-o", - default=None, - help=( - "Output file or directory. Defaults to " - "~/.diffusers/cli/run/outputs/diffusers-run--/.." - ), - ) - parser.add_argument( - "--push-to", - default=None, - help=( - "Upload the generated files to this HF bucket after saving (created if missing). Accepts " - "an HF bucket id (`/`), an `hf://buckets//[/]` " - "URI, or a browser URL for the same — a subpath is used as a folder prefix. Under --remote " - "the upload runs inside the sandbox; without an explicit --output the bucket becomes the " - "sole destination and nothing is downloaded back." - ), - ) - - -def _add_remote_arguments(parser: ArgumentParser) -> None: - parser.add_argument( - "--remote", - action="store_true", - help="Run this command in a Hugging Face Sandbox instead of on the local machine.", - ) - parser.add_argument( - "--flavor", - default="a10g-small", - help="HF Sandbox hardware flavor for --remote (e.g. a10g-small, a100-large, cpu-basic).", - ) - parser.add_argument( - "--timeout", - default="10m", - help="Max wallclock for the run command inside the sandbox (e.g. 30m, 2h). Defaults to 10m.", - ) - parser.add_argument( - "--dependencies", - action="append", - default=None, - help="Extra pip dependencies to install in the sandbox. Repeat to add multiple.", - ) - parser.add_argument( - "--namespace", - default=None, - help="HF namespace to create the sandbox under (defaults to the current user).", - ) - parser.add_argument( - "--image", - default=None, - help=( - "Sandbox image for --remote (defaults to " - f"{_DEFAULT_REMOTE_IMAGE!r}). Must provide torch + CUDA; the CLI installs the " - "small Python deps on top via `uv pip install --system`." - ), - ) - parser.add_argument( - "--keep-alive", - action="store_true", - help=( - "Don't terminate the sandbox after the run. Its id is printed so a later --remote run " - "can reconnect with --sandbox-id and reuse the warm deps/weights/compile cache." - ), - ) - parser.add_argument( - "--sandbox-id", - default=None, - help=( - "Reconnect to an existing sandbox (from a prior --keep-alive run) instead of creating a new " - "one, reusing its warm deps/weights/compile cache. Implies --keep-alive; stop it with " - "`hf sandbox kill `." - ), - ) - parser.add_argument( - "--idle-timeout", - default="10m", - help=( - "Auto-shutdown the sandbox after this much inactivity (e.g. 30m, 1h). Defaults to 10m. " - "Only applied on new sandbox creation — ignored when reconnecting via --sandbox-id." - ), - ) - parser.add_argument( - "--volume", - action="append", - default=None, - metavar="BUCKET_ID[:MOUNT_PATH]", - help=( - "Mount an HF bucket into the sandbox as a read-write directory. Repeatable. Format: " - "`/` (mounts at `/mnt/buckets//`) or " - "`/:/some/path` for a custom path. Reference mounted files from " - "--pipeline-kwargs like any other local path. Applied only on new sandbox creation — " - "ignored when reconnecting via --sandbox-id." - ), - ) - - -# --------------------------------------------------------------------------- -# Pipeline loading + optimization -# --------------------------------------------------------------------------- - - -def _resolve_dtype(name: str | None): - if name in (None, "auto"): - return "auto" - import torch - - mapping = { - "fp32": torch.float32, - "float32": torch.float32, - "fp16": torch.float16, - "float16": torch.float16, - "bf16": torch.bfloat16, - "bfloat16": torch.bfloat16, - } - if name not in mapping: - raise ValueError(f"Unknown dtype: {name}") - return mapping[name] - - -def _resolve_device_map(raw: str | None) -> str | dict: - """Parse `--device-map` into a value acceptable by `from_pretrained(device_map=...)`. - - Returns a JSON dict if the value looks like one, `"balanced"` verbatim, or a single-device string (e.g. `"cuda"`, - `"cuda:1"`, `"cpu"`, `"mps"`). Auto-detects when `raw is None`, pinning to `cuda:$LOCAL_RANK` under torchrun. - """ - if raw is None: - from diffusers.utils.torch_utils import torch_device - - if torch_device == "cuda": - local_rank = os.environ.get("LOCAL_RANK") - if local_rank is not None: - import torch - - torch.cuda.set_device(int(local_rank)) - return f"cuda:{local_rank}" - return torch_device - - if raw.strip().startswith("{"): - try: - parsed = json.loads(raw) - except json.JSONDecodeError as e: - raise SystemExit(f"--device-map must be a device string or a JSON dict: {e}") from e - if not isinstance(parsed, dict): - raise SystemExit("--device-map JSON must decode to an object.") - return parsed - - return raw - - -def _apply_cpu_offload(pipeline: Any, mode: str, device_map: str | dict) -> None: - """Apply model or group CPU offload. Requires a single-device target (not balanced or dict).""" - if not isinstance(device_map, str) or device_map == "balanced": - raise SystemExit( - "--cpu-offload requires --device-map to be a single device string (e.g. 'cuda'); " - f"got {device_map!r}. balanced/dict placement is incompatible with CPU offload." - ) - - if mode == "model": - pipeline.enable_model_cpu_offload(device=device_map) - elif mode == "group": - import torch - - pipeline.enable_group_offload( - onload_device=torch.device(device_map), - offload_type="leaf_level", - use_stream=True, - ) - - -def _set_attention_backend(pipeline: Any, backend: str) -> None: - transformer = getattr(pipeline, "transformer", None) - if transformer is None or not hasattr(transformer, "set_attention_backend"): - logger.warning( - f"--attention-backend is only supported on transformer-based pipelines; " - f"{type(pipeline).__name__} uses the legacy UNet attention path." - ) - return - try: - transformer.set_attention_backend(backend) - except (ValueError, ImportError, RuntimeError) as e: - logger.warning( - f"Attention backend {backend!r} could not be set on {type(transformer).__name__}: " - f"{type(e).__name__}: {e}. Falling back to the model's default backend." - ) - - -def _enable_context_parallel(pipeline: Any) -> None: - import torch - - if not torch.distributed.is_available(): - raise SystemExit("--context-parallel requires a torch build with distributed support.") - - if not torch.distributed.is_initialized(): - # Hybrid backend: ulysses_anything's per-rank size coordination wants Gloo on CPU - # (avoids H2D/D2H for a tiny int tensor); the main attention all-to-all stays on NCCL. - torch.distributed.init_process_group(backend="cpu:gloo,cuda:nccl") - - transformer = getattr(pipeline, "transformer", None) - if transformer is None or not hasattr(transformer, "enable_parallelism"): - raise SystemExit( - "--context-parallel requires a DiT-based pipeline. " - f"{type(pipeline).__name__} does not expose a `transformer` with `enable_parallelism`." - ) - - from diffusers import ContextParallelConfig - - transformer.enable_parallelism( - config=ContextParallelConfig( - ulysses_degree=torch.distributed.get_world_size(), - ring_degree=1, - ulysses_anything=True, - ) - ) - - -def _apply_optimizations(pipeline: Any, args: Namespace) -> None: - """Apply VAE tiling/slicing, attention backend, context-parallel, and torch.compile toggles.""" - vae = getattr(pipeline, "vae", None) - if args.vae_tiling and vae is not None and hasattr(vae, "enable_tiling"): - vae.enable_tiling() - if args.vae_slicing and vae is not None and hasattr(vae, "enable_slicing"): - vae.enable_slicing() - if args.attention_backend != "default": - _set_attention_backend(pipeline, args.attention_backend) - if args.context_parallel: - _enable_context_parallel(pipeline) - if args.compile is not None: - if args.context_parallel: - logger.warning("--compile is currently not supported with --context-parallel; skipping compile.") - else: - _compile_denoiser(pipeline, args.compile) - - -def _compile_denoiser(pipeline: Any, compile_spec: str) -> None: - """Compile every `transformer*` and `unet*` submodule on the pipeline. - - `compile_spec` is the raw JSON string from `--compile` (`"{}"` for bare flag). Decoded into kwargs and forwarded - verbatim to the compile call. - - Prefers regional compilation via `module.compile_repeated_blocks(**kwargs)` — only compiles the repeated inner - blocks (the bulk of the compute), much faster first-step latency than compiling the whole module. Falls back to - full `torch.compile` if the model doesn't expose `_repeated_blocks`. - """ - import torch - - try: - compile_kwargs = json.loads(compile_spec) - except json.JSONDecodeError as e: - raise SystemExit(f"--compile must be valid JSON: {e}") from e - if not isinstance(compile_kwargs, dict): - raise SystemExit("--compile must decode to a JSON object.") - - for attr in dir(pipeline): - if not any(attr.startswith(key) for key in _DENOISER_COMPONENT_KEYS): - continue - module = getattr(pipeline, attr, None) - if not isinstance(module, torch.nn.Module): - continue - - if getattr(module, "_repeated_blocks", None): - # Regional compile — only the repeated blocks. Mutates `module` in place. - module.compile_repeated_blocks(**compile_kwargs) - else: - # No regional metadata declared; fall back to compiling the whole module. - setattr(pipeline, attr, torch.compile(module, **compile_kwargs)) - - -def _load_lora(pipeline: Any, args: Namespace) -> None: - """Attach one or more LoRA adapters. Each `--lora` value is a JSON dict. - - Per-entry fields: `lora_id` (required), `lora_scale` (optional float, default 1.0), `adapter_name` (optional; - auto-generated as `lora_` when stacking). Multiple `--lora` flags stack via a single `set_adapters(...)` call at - the end. - """ - if not args.lora: - return - specs = [] - for raw in args.lora: - try: - parsed = json.loads(raw) - except json.JSONDecodeError as e: - raise SystemExit(f"--lora must be valid JSON: {e}") from e - if not isinstance(parsed, dict): - raise SystemExit(f"--lora must decode to a JSON object; got {type(parsed).__name__}.") - specs.append(parsed) - if not hasattr(pipeline, "load_lora_weights"): - raise SystemExit(f"{type(pipeline).__name__} does not support LoRA loading.") - - names: list[str] = [] - scales: list[float] = [] - for i, spec in enumerate(specs): - lora_id = spec.get("lora_id") - if not lora_id: - raise SystemExit(f"--lora entry {i} is missing 'lora_id'.") - adapter_name = spec.get("adapter_name") or (f"lora_{i}" if len(specs) > 1 else "default") - pipeline.load_lora_weights(lora_id, adapter_name=adapter_name) - names.append(adapter_name) - scales.append(float(spec.get("lora_scale", 1.0))) - - if hasattr(pipeline, "set_adapters"): - pipeline.set_adapters(names, adapter_weights=scales) - - -def _load_pipeline(args: Namespace) -> Any: - import diffusers - - # Detect modular repos by trying the standard config; `ModularPipeline` repos ship - # `modular_model_index.json` instead of `model_index.json`, so `load_config` OSErrors. - try: - diffusers.DiffusionPipeline.load_config(args.model, token=args.token, revision=args.revision) - modular = False - except OSError: - modular = True - - dtype = _resolve_dtype(args.dtype) - device_map = _resolve_device_map(args.device_map) - common_kwargs: dict[str, Any] = { - "trust_remote_code": args.trust_remote_code, - } - if dtype != "auto": - common_kwargs["torch_dtype"] = dtype - if args.variant: - common_kwargs["variant"] = args.variant - if args.token: - common_kwargs["token"] = args.token - # CPU offload sets up its own placement hooks, so leave weights on CPU at load time. - if not args.cpu_offload: - common_kwargs["device_map"] = device_map - - if modular: - # ModularPipeline.from_pretrained fetches only the pipeline config; component - # weights come in via load_components(). `revision` scopes the config fetch, - # so it stays on from_pretrained — each ComponentSpec pins its own revision, - # and forwarding a global `revision` to load_components() would override those. - pipeline = diffusers.ModularPipeline.from_pretrained( - args.model, - trust_remote_code=args.trust_remote_code, - token=args.token, - revision=args.revision, - ) - pipeline.load_components(**common_kwargs) - else: - pipeline = diffusers.DiffusionPipeline.from_pretrained(args.model, revision=args.revision, **common_kwargs) - - _load_lora(pipeline, args) - if args.cpu_offload: - _apply_cpu_offload(pipeline, args.cpu_offload, device_map) - _apply_optimizations(pipeline, args) - - return pipeline - - -# --------------------------------------------------------------------------- -# Pipeline call helpers -# --------------------------------------------------------------------------- - - -def _parse_pipeline_kwargs(raw: str | None) -> dict[str, Any]: - if not raw: - return {} - try: - parsed = json.loads(raw) - except json.JSONDecodeError as e: - raise SystemExit(f"--pipeline-kwargs must be valid JSON: {e}") from e - if not isinstance(parsed, dict): - raise SystemExit("--pipeline-kwargs must decode to a JSON object.") - return parsed - - -def _load_audio(url_or_path: str) -> tuple[Any, int]: - """Load audio from a URL or local path via torchaudio. Returns `(waveform, sampling_rate)`.""" - import torchaudio - - if url_or_path.startswith(("http://", "https://")): - import io - - import httpx - - from ..utils.constants import DIFFUSERS_REQUEST_TIMEOUT - - resp = httpx.get(url_or_path, follow_redirects=True, timeout=DIFFUSERS_REQUEST_TIMEOUT) - resp.raise_for_status() - return torchaudio.load(io.BytesIO(resp.content)) - return torchaudio.load(url_or_path) - - -def _resolve_media_inputs(call_kwargs: dict[str, Any]) -> None: - """Replace string paths/URLs at known media-input keys with loaded tensors. - - Images resolve to `PIL.Image.Image` via `load_image`; videos to `list[PIL.Image.Image]` via `load_video`; audio to - a `torch.Tensor` via `_load_audio` (also auto-sets the paired sampling-rate kwarg for `initial_audio_waveforms` if - the user didn't supply it). A `list[str]` at any key is treated as a batch: each entry is loaded and the value - becomes a list of loaded objects. Non-string, non-list values pass through untouched. - """ - - def _is_string_list(v: Any) -> bool: - return isinstance(v, list) and bool(v) and all(isinstance(x, str) for x in v) - - for key in _IMAGE_INPUT_KEYS: - value = call_kwargs.get(key) - if isinstance(value, str): - call_kwargs[key] = load_image(value) - elif _is_string_list(value): - call_kwargs[key] = [load_image(v) for v in value] - for key in _VIDEO_INPUT_KEYS: - value = call_kwargs.get(key) - if isinstance(value, str): - call_kwargs[key] = load_video(value) - elif _is_string_list(value): - call_kwargs[key] = [load_video(v) for v in value] - for key in _AUDIO_INPUT_KEYS: - value = call_kwargs.get(key) - if isinstance(value, str): - waveform, sr = _load_audio(value) - call_kwargs[key] = waveform - if key == "initial_audio_waveforms" and "initial_audio_sampling_rate" not in call_kwargs: - call_kwargs["initial_audio_sampling_rate"] = sr - elif _is_string_list(value): - pairs = [_load_audio(v) for v in value] - call_kwargs[key] = [w for w, _ in pairs] - if key == "initial_audio_waveforms" and "initial_audio_sampling_rate" not in call_kwargs: - # All batched waveforms must share a sampling rate; use the first entry's. - call_kwargs["initial_audio_sampling_rate"] = pairs[0][1] - - -def _get_generator(seed: int | None, device: str): - if seed is None: - return None - import torch - - generator_device = "cpu" if device == "mps" else device - return torch.Generator(device=generator_device).manual_seed(seed) - - -def _unwrap_pipeline_output(result: Any) -> Any: - """Unwrap a pipeline-output object into the raw payload the saver can dispatch on.""" - if hasattr(result, "images"): - return result.images - if hasattr(result, "frames"): - return result.frames[0] - if hasattr(result, "audios"): - return result.audios - return result - - -# --------------------------------------------------------------------------- -# Output saving (dispatch by type) -# --------------------------------------------------------------------------- - - -def _get_or_create_run_id() -> str: - """Return the current run's id, creating one if not yet set. - - Format: `diffusers-run--<6-char-uuid>`. Same id is reused as the local output subdirectory, the - remote bucket prefix, and the container-side `RUN_ID_ENV` so a run's artifacts are traceable end-to-end. - """ - import uuid - from datetime import datetime - - existing = os.environ.get(RUN_ID_ENV) - if existing: - return existing - run_id = f"diffusers-run-{datetime.now().strftime('%Y%m%dT%H%M%S')}-{uuid.uuid4().hex[:6]}" - os.environ[RUN_ID_ENV] = run_id - return run_id - - -def _resolve_output_paths(task: str, num: int, explicit: str | None, ext: str) -> list[Path]: - if explicit is None: - base = Path(DEFAULT_OUTPUT_DIR) / _get_or_create_run_id() - base.mkdir(parents=True, exist_ok=True) - return [base / f"{i:04d}.{ext}" for i in range(num)] - - p = Path(explicit) - if explicit.endswith(os.sep) or p.is_dir(): - p.mkdir(parents=True, exist_ok=True) - return [p / f"{i:04d}.{ext}" for i in range(num)] - - p.parent.mkdir(parents=True, exist_ok=True) - if num == 1: - return [p] - stem, suffix = p.stem, p.suffix or f".{ext}" - return [p.with_name(f"{stem}-{i:04d}{suffix}") for i in range(num)] - - -def _as_pil_list(value: Any): - try: - from PIL.Image import Image as PILImage - except ImportError: - return None - if isinstance(value, PILImage): - return [value] - if isinstance(value, (list, tuple)) and value and all(isinstance(v, PILImage) for v in value): - return list(value) - return None - - -def _as_frame_sequence(value: Any): - try: - from PIL.Image import Image as PILImage - except ImportError: - PILImage = None # type: ignore[assignment] - - if isinstance(value, (list, tuple)) and len(value) >= 2: - first = value[0] - if PILImage is not None and isinstance(first, PILImage): - return list(value) - try: - import numpy as np - - if isinstance(first, np.ndarray): - return list(value) - except ImportError: - pass - return None - - -def _as_audio_arrays(value: Any): - try: - import numpy as np - except ImportError: - return None - if isinstance(value, np.ndarray) and value.ndim <= 2: - return [value] - if isinstance(value, (list, tuple)) and value and all(isinstance(v, np.ndarray) for v in value): - return list(value) - return None - - -def _save_audio_arrays(audios, sampling_rate: int, args: Namespace, task: str) -> list[str]: - """Write each numpy audio array to a 16-bit PCM WAV at `sampling_rate` Hz. - - Uses the stdlib `wave` module so no scipy dependency is required. - """ - import wave - - import numpy as np - - paths = _resolve_output_paths(task, len(audios), args.output, ext="wav") - saved: list[str] = [] - for audio, path in zip(audios, paths): - data = np.asarray(audio) - if data.dtype.kind == "f": - data = (np.clip(data, -1.0, 1.0) * 32767).astype(np.int16) - else: - data = data.astype(np.int16) - if data.ndim == 1: - n_channels = 1 - else: - # Heuristic: shorter axis is channels (interleaved layout for `wave` is - # samples × channels, so transpose if needed). - if data.shape[0] < data.shape[-1]: - data = data.T - n_channels = data.shape[1] - with wave.open(str(path), "wb") as w: - w.setnchannels(n_channels) - w.setsampwidth(2) # 16-bit PCM - w.setframerate(sampling_rate) - w.writeframes(data.tobytes()) - saved.append(str(path)) - return saved - - -def _save_output(value: Any, args: Namespace, task: str) -> list[str]: - """Save `value` by dispatching on its runtime type.""" - pil_images = _as_pil_list(value) - if pil_images is not None: - paths = _resolve_output_paths(task, len(pil_images), args.output, ext="png") - for img, path in zip(pil_images, paths): - img.save(path) - return [str(p) for p in paths] - - frames = _as_frame_sequence(value) - if frames is not None: - from diffusers.utils import export_to_video - - path = _resolve_output_paths(task, 1, args.output, ext="mp4")[0] - export_to_video(frames, str(path), fps=args.fps) - return [str(path)] - - audios = _as_audio_arrays(value) - if audios is not None: - return _save_audio_arrays(audios, args.sampling_rate or 16000, args, task) - - path = _resolve_output_paths(task, 1, args.output, ext="json")[0] - Path(path).write_text(json.dumps(value, default=str, indent=2)) - return [str(path)] - - -# --------------------------------------------------------------------------- -# Hub bucket upload (--push-to) -# --------------------------------------------------------------------------- - - -def _parse_push_to(spec: str) -> tuple[str, str]: - """Split `--push-to` into a bucket id and an optional subpath prefix. - - Accepts an HF bucket id (`/[/]`), a canonical - `hf://buckets//[/]` URI, or a Hub web URL for the same. Non-bucket URIs (models, - datasets, spaces) are rejected — `--push-to` targets storage buckets only. - """ - from huggingface_hub import parse_hf_uri - - # Bare shorthand → canonical URI so a single parser handles every accepted form. - if not spec.startswith(("hf://", "http://", "https://")): - spec = f"hf://buckets/{spec.strip('/')}" - uri = parse_hf_uri(spec) - if not uri.is_bucket: - raise SystemExit(f"--push-to must point at a bucket; got {uri.type!r} URI {spec!r}.") - return uri.id, uri.path_in_repo - - -def _push_outputs(args: Namespace, saved_paths: list[str], task: str) -> dict[str, Any] | None: - """Upload `saved_paths` to the `--push-to` bucket. Returns a summary or None.""" - if not args.push_to: - return None - - from huggingface_hub import HfApi - - bucket_id, subpath = _parse_push_to(args.push_to) - api = HfApi(token=args.token) - api.create_bucket(bucket_id, exist_ok=True) - - run_id = _get_or_create_run_id() - prefix = f"{subpath}/{run_id}" if subpath else run_id - add = [(local, f"{prefix}/{Path(local).name}") for local in saved_paths] - api.batch_bucket_files(bucket_id, add=add) - - uploaded = [f"hf://buckets/{bucket_id}/{dest}" for _, dest in add] - return {"bucket_id": bucket_id, "uploaded": uploaded} - - -# --------------------------------------------------------------------------- -# Remote execution (HF Sandbox) -# --------------------------------------------------------------------------- - - -def _build_task_kwargs(args: Namespace) -> dict[str, Any]: - """Pick out the kwargs the sandbox CLI should invoke the task with.""" - out: dict[str, Any] = {} - for key, value in vars(args).items(): - if key in REMOTE_KEYS or value is None or value is False: - continue - out[key] = value - return out - - -def _kwargs_to_argv(task: str, task_kwargs: dict[str, Any]) -> list[str]: - """Render `task_kwargs` as the argv list the sandbox CLI's argparse will see.""" - argv: list[str] = [task] - for key, value in task_kwargs.items(): - flag = "--" + key.replace("_", "-") - if value is True: - argv.append(flag) - elif isinstance(value, list): - for item in value: - argv.extend([flag, str(item)]) - else: - argv.extend([flag, str(value)]) - return argv - - -def _duration_to_seconds(value: str) -> float: - """Parse a duration like `30s`, `10m`, `2h` (or a bare number of seconds) into seconds.""" - value = value.strip() - units = {"s": 1, "m": 60, "h": 3600} - if value and value[-1] in units: - return float(value[:-1]) * units[value[-1]] - return float(value) - - -def _upload_inputs_to_sandbox(args: Namespace, sbx: Any, run_id: str) -> None: - """Upload local media paths in `--pipeline-kwargs` into the sandbox and rewrite the JSON in place. - - Walks known image/video/audio-input keys; any string value that resolves to a local file is uploaded to - `<_SANDBOX_INPUTS_DIR>//_` and the JSON path is rewritten to that in-sandbox path. URLs, - `hf://` URIs, and non-existent paths pass through untouched. - """ - if not args.pipeline_kwargs: - return - try: - parsed = json.loads(args.pipeline_kwargs) - except json.JSONDecodeError: - return # the sandbox CLI will fail loudly with a parse error later - if not isinstance(parsed, dict): - return - - def _upload_one(key: str, index: int | None, local_str: str) -> str: - # `index` is None for scalar entries, an int for list entries (used to disambiguate names). - local = Path(local_str) - suffix = f"_{index}" if index is not None else "" - remote_path = f"{_SANDBOX_INPUTS_DIR}/{run_id}/{key}{suffix}_{local.name}" - sbx.files.upload(str(local), remote_path) - return remote_path - - uploaded = 0 - for key in (*_IMAGE_INPUT_KEYS, *_VIDEO_INPUT_KEYS, *_AUDIO_INPUT_KEYS): - value = parsed.get(key) - if isinstance(value, str) and Path(value).is_file(): - parsed[key] = _upload_one(key, None, value) - uploaded += 1 - elif isinstance(value, list): - # Batched inputs: upload each local path, leave URLs/hf:// URIs alone. - new_list = list(value) - for i, entry in enumerate(value): - if isinstance(entry, str) and Path(entry).is_file(): - new_list[i] = _upload_one(key, i, entry) - uploaded += 1 - parsed[key] = new_list - - if uploaded: - logger.info(f"uploaded {uploaded} local input file(s) to the sandbox") - args.pipeline_kwargs = json.dumps(parsed) - - -def _download_outputs_from_sandbox(sbx: Any, sandbox_dir: str, local_dir: Path) -> list[str]: - """Download every file the sandbox CLI wrote under `sandbox_dir` into `local_dir`.""" - local_dir.mkdir(parents=True, exist_ok=True) - saved: list[str] = [] - for entry in sbx.files.list(sandbox_dir): - if entry.type != "file": - continue - target = local_dir / Path(entry.path).name - sbx.files.download(entry.path, str(target)) - saved.append(str(target)) - return saved - - -def _maybe_submit_remote(args: Namespace, task: str) -> bool: - """If `--remote` was set, run this invocation inside an HF Sandbox and return True.""" - if not args.remote: - return False - - import shlex - import time - - from huggingface_hub import get_token - from huggingface_hub.utils import send_telemetry - - import diffusers - - try: - from huggingface_hub import Sandbox - except ImportError: - raise SystemExit( - "--remote requires huggingface_hub>=1.23 for HF Sandbox support. " - "Upgrade with `pip install -U huggingface_hub`." - ) - - if Path(args.model).exists(): - raise SystemExit( - f"--model {args.model!r} is a local path; the sandbox can't see it. " - "Pass a Hub repo id so the sandbox can download it." - ) - - hf_token = args.token or get_token() - run_id = _get_or_create_run_id() - - # An explicit --push-to means the bucket is the user's destination, so skip the local - # download unless they also asked for a local path via --output. - user_bucket = bool(args.push_to) - download_locally = (not user_bucket) or (args.output is not None) - local_dir = Path(args.output) if args.output else Path(DEFAULT_OUTPUT_DIR) / run_id - - use_existing_sandbox = bool(args.sandbox_id) - keep_alive = args.keep_alive or use_existing_sandbox - if use_existing_sandbox and args.volume: - logger.warning( - "--volume is ignored when reconnecting to an existing sandbox (mounts are set at creation time)." - ) - if use_existing_sandbox: - logger.info(f"reconnecting to sandbox {args.sandbox_id!r}...") - sbx = Sandbox.connect(args.sandbox_id, token=hf_token) - else: - logger.info(f"creating sandbox on flavor={args.flavor!r}...") - create_kwargs: dict[str, Any] = { - "image": args.image or _DEFAULT_REMOTE_IMAGE, - "flavor": args.flavor, - "forward_hf_token": True, - "token": hf_token, - "env": { - "HF_ENABLE_PARALLEL_LOADING": "1", - "DIFFUSERS_VERBOSITY": os.environ.get("DIFFUSERS_VERBOSITY", "info"), - }, - "idle_timeout": args.idle_timeout, - } - if args.volume: - from huggingface_hub import Volume - - volumes = [] - for spec in args.volume: - bucket_id, sep, mount_path = spec.partition(":") - if not sep: - mount_path = f"/mnt/buckets/{bucket_id}" - if bucket_id.count("/") != 1: - raise SystemExit(f"--volume: bucket id must be /, got {bucket_id!r}") - if not mount_path.startswith("/"): - raise SystemExit(f"--volume: mount path must be absolute, got {mount_path!r}") - volumes.append(Volume(type="bucket", source=bucket_id, mount_path=mount_path)) - create_kwargs["volumes"] = volumes - if args.namespace is not None: - create_kwargs["namespace"] = args.namespace - sbx = Sandbox.create(**create_kwargs) - - def _stream(chunk: str) -> None: - sys.stderr.write(chunk) - sys.stderr.flush() - - exit_code = 0 - saved: list[str] = [] - run_seconds = 0.0 - try: - _upload_inputs_to_sandbox(args, sbx, run_id) - - dependencies = list(_DEFAULT_REMOTE_DEPS) - if args.dependencies: - dependencies.extend(args.dependencies) - # --break-system-packages bypasses PEP 668; harmless in a throwaway sandbox. uv is a - # near no-op when the deps are already satisfied, so this stays cheap on a reused sandbox. - install_cmd = shlex.join(["uv", "pip", "install", "--system", "--break-system-packages", *dependencies]) - logger.info("installing dependencies in the sandbox...") - sbx.run(install_cmd, on_stdout=_stream, on_stderr=_stream) - - # Per-run outputs subdirectory so a reused sandbox doesn't leak files from prior runs - # into this run's download set. - sandbox_output_dir = f"{_SANDBOX_OUTPUTS_DIR}/{run_id}" - task_kwargs = _build_task_kwargs(args) - task_kwargs["output"] = sandbox_output_dir + "/" - cli_argv = _kwargs_to_argv(task, task_kwargs) - # Suppress the container CLI's own `out.result(...)` payload — the outer wrapper owns the - # final structured output for --remote runs. - format_argv = ["--format", "quiet"] - # torchrun wraps the CLI for --context-parallel so torch.distributed initializes across - # every visible GPU before the run command starts. - if args.context_parallel: - cli_argv = [ - "torchrun", - "--nproc-per-node=gpu", - "-m", - "diffusers.commands.diffusers_cli", - *format_argv, - *cli_argv, - ] - else: - cli_argv = [_CONTAINER_CLI_BINARY, *format_argv, *cli_argv] - - started = time.perf_counter() - # Per-invocation env: RUN_ID_ENV must be fresh each run. Sandbox.create-time env is - # baked in and would go stale on reused sandboxes, silently reusing the initial run's - # bucket prefix in `_push_outputs`. - result = sbx.run( - cli_argv, - env={RUN_ID_ENV: run_id}, - on_stdout=_stream, - on_stderr=_stream, - timeout=_duration_to_seconds(args.timeout), - check=False, - ) - run_seconds = time.perf_counter() - started - exit_code = result.exit_code - - if exit_code == 0 and download_locally: - saved = _download_outputs_from_sandbox(sbx, sandbox_output_dir, local_dir) - finally: - if keep_alive: - logger.info( - f"sandbox {sbx.id} kept alive — reconnect with " - f"`--remote --sandbox-id {sbx.id}`, stop with `hf sandbox kill {sbx.id}`." - ) - else: - sbx.kill() - - send_telemetry( - topic="diffusers/cli/run/remote", - library_name="diffusers", - library_version=diffusers.__version__, - ) - - payload: dict[str, Any] = { - "exit_code": exit_code, - "run_seconds": round(run_seconds, 1), - } - if keep_alive: - payload["sandbox_id"] = sbx.id - if download_locally: - payload["outputs"] = saved - if args.push_to: - bucket_id, subpath = _parse_push_to(args.push_to) - prefix = f"{subpath}/{run_id}" if subpath else run_id - payload["pushed-to"] = f"hf://buckets/{bucket_id}/{prefix}/" - out.result("remote-run", **payload) - - if exit_code != 0: - raise SystemExit(f"remote run failed with exit code {exit_code}") - return True - - -# --------------------------------------------------------------------------- -# Subcommand -# --------------------------------------------------------------------------- - - -class RunCommand(BaseDiffusersCLICommand): - task = "run" - - @staticmethod - def register_subcommand(subparsers: _SubParsersAction) -> None: - from argparse import RawDescriptionHelpFormatter - - epilog = ( - "Examples\n" - " $ diffusers-cli run -m black-forest-labs/FLUX.1-dev --dtype bf16 \\\n" - ' --pipeline-kwargs \'{"prompt": "a cat on the moon"}\'\n' - " $ diffusers-cli run -m black-forest-labs/FLUX.1-dev --dtype bf16 \\\n" - ' --pipeline-kwargs \'{"prompt": "make the fur grey", "image": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/cat.png"}\'\n' - " $ diffusers-cli run -m black-forest-labs/FLUX.1-dev --dtype bf16 \\\n" - ' --pipeline-kwargs \'{"prompt": "a tiny cat"}\' \\\n' - ' --lora \'{"lora_id": "alvdansen/littletinies", "lora_scale": 0.8}\'\n' - " $ diffusers-cli run -m black-forest-labs/FLUX.1-dev --dtype bf16 \\\n" - ' --pipeline-kwargs \'{"prompt": "a cat"}\' --remote --flavor a100-large\n' - " $ diffusers-cli run -m black-forest-labs/FLUX.1-dev --dtype bf16 --context-parallel \\\n" - ' --pipeline-kwargs \'{"prompt": "a cat"}\' --remote --flavor 4xa100-large\n' - "\n" - "Learn more\n" - " Use `diffusers-cli --help` for more information about a command.\n" - " Read the documentation at https://huggingface.co/docs/diffusers\n" - ) - - parser: ArgumentParser = subparsers.add_parser( - "run", - help="Run any diffusers pipeline locally or remotely in an HF Sandbox.", - usage="\n diffusers-cli run [options]", - epilog=epilog, - formatter_class=RawDescriptionHelpFormatter, - ) - parser._optionals.title = "Options" - _add_loading_arguments(parser) - _add_optimization_arguments(parser) - parser.add_argument( - "--pipeline-kwargs", - default=None, - help=( - "JSON object of kwargs passed to the pipeline call. String values at known " - f"image-input keys ({', '.join(_IMAGE_INPUT_KEYS)}) are auto-loaded as PIL images; " - f"video-input keys ({', '.join(_VIDEO_INPUT_KEYS)}) are auto-loaded as frame lists; " - f"audio-input keys ({', '.join(_AUDIO_INPUT_KEYS)}) are auto-loaded via torchaudio." - ), - ) - parser.add_argument( - "--output-key", - default=None, - help="For modular pipelines: name of the intermediate to extract (passed as `output=` to the call).", - ) - parser.add_argument("--seed", type=int, default=None, help="Random seed for reproducibility.") - parser.add_argument( - "--fps", - type=int, - default=8, - help="FPS used when the output happens to be a frame sequence.", - ) - parser.add_argument( - "--sampling-rate", - type=int, - default=None, - help="Sample rate used when the output happens to be an audio array.", - ) - _add_remote_arguments(parser) - _add_output_arguments(parser) - parser.set_defaults(func=RunCommand) - - def __init__(self, args: Namespace): - self.args = args - - def run(self) -> None: - import diffusers - - _get_or_create_run_id() # populate RUN_ID_ENV so local output dir + remote bucket prefix agree - - call_kwargs = _parse_pipeline_kwargs(self.args.pipeline_kwargs) - - if _maybe_submit_remote(self.args, self.task): - return - - # Resolve media before loading pipeline weights so dead URLs / missing files fail - # fast — cheap to fetch, expensive to load a 20GB model just to hit a 404. - _resolve_media_inputs(call_kwargs) - pipeline = _load_pipeline(self.args) - is_modular = isinstance(pipeline, diffusers.ModularPipeline) - - if self.args.output_key is not None: - call_kwargs["output"] = self.args.output_key - - device = pipeline.device.type if hasattr(pipeline, "device") else "cpu" - generator = _get_generator(self.args.seed, device) - if generator is not None: - call_kwargs["generator"] = generator - - try: - result = pipeline(**call_kwargs) - - # Under torchrun, ranks > 0 produce identical output to rank 0 (CP shards the - # transformer compute but ranks reduce to the same final tensors). Save/push/print - # from rank 0 only to avoid clobbering bucket files 4x and printing 4x. - if os.environ.get("RANK", "0") == "0": - savable = result if is_modular else _unwrap_pipeline_output(result) - saved = _save_output(savable, self.args, self.task) - pushed = _push_outputs(self.args, saved, self.task) - - out.result( - self.task, - model=self.args.model, - device=device, - pipeline_class=type(pipeline).__name__, - modular=is_modular, - outputs=saved, - pushed=pushed, - seed=self.args.seed, - output_key=self.args.output_key, - ) - finally: - import torch - - if torch.distributed.is_available() and torch.distributed.is_initialized(): - torch.distributed.destroy_process_group() diff --git a/diffusers/commands/schema.py b/diffusers/commands/schema.py deleted file mode 100644 index dc5965e8adad189177208831627ff6213a4172e1..0000000000000000000000000000000000000000 --- a/diffusers/commands/schema.py +++ /dev/null @@ -1,287 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""`diffusers-cli schema` — print the input schema for any pipeline repo. - -Tries `DiffusionPipeline.config_name` first (so standard repos get their `__call__` signature introspected); falls back -to `ModularPipelineBlocks.from_pretrained` for modular repos. No weights are downloaded — only the small index file -(and any custom block code if `--trust-remote-code` is set). -""" - -from __future__ import annotations - -import inspect -import re -from argparse import ArgumentParser, Namespace, _SubParsersAction -from typing import Any - -from huggingface_hub.cli._output import OutputFormat, out - -from ..utils import logging -from . import BaseDiffusersCLICommand - - -logger = logging.get_logger("diffusers-cli/schema") - - -def _schema(args: Namespace) -> None: - """Print the pipeline's input schema. - - Tries `DiffusionPipeline.config_name` (= `model_index.json`) first; if present, introspects the declared pipeline - class's `__call__` signature. Otherwise falls back to `ModularPipelineBlocks.from_pretrained` and reads the - block-declared `inputs`. No weights downloaded either way. - """ - import diffusers - - try: - index = diffusers.DiffusionPipeline.load_config(args.model, token=args.token, revision=args.revision) - except OSError: - index = None - - if index is not None: - class_name = index.get("_class_name") - if class_name is None: - raise SystemExit( - f"{diffusers.DiffusionPipeline.config_name} for {args.model!r} has no `_class_name` field." - ) - pipeline_cls = getattr(diffusers, class_name, None) - if pipeline_cls is None: - raise SystemExit( - f"Pipeline class {class_name!r} declared in {diffusers.DiffusionPipeline.config_name} " - "is not exported by the installed diffusers." - ) - - sig = inspect.signature(pipeline_cls.__call__) - descriptions = _parse_docstring_args(pipeline_cls.__call__.__doc__) if args.verbose else {} - schema: list[dict[str, Any]] = [] - for name, param in sig.parameters.items(): - if name == "self": - continue - if param.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD): - continue - has_default = param.default is not inspect.Parameter.empty - schema.append( - { - "name": name, - "type_hint": str(param.annotation) if param.annotation is not inspect.Parameter.empty else None, - "default": param.default if has_default else None, - "required": not has_default, - "description": descriptions.get(name, ""), - } - ) - else: - kwargs: dict[str, Any] = {"trust_remote_code": args.trust_remote_code} - if args.revision: - kwargs["revision"] = args.revision - if args.token: - kwargs["token"] = args.token - - # If the repo declares custom code + external dependencies, surface them upfront so - # the user knows what to install before we hit an ImportError inside from_pretrained. - _warn_custom_block_requirements(args) - - try: - blocks = diffusers.ModularPipelineBlocks.from_pretrained(args.model, **kwargs) - except Exception as e: - hint = "\nPass --trust-remote-code if it ships custom block code." if not args.trust_remote_code else "" - raise SystemExit( - f"Could not read schema for {args.model!r}: no {diffusers.DiffusionPipeline.config_name} and " - f"loading as a modular pipeline failed with:\n {type(e).__name__}: {e}{hint}" - ) from e - - class_name = type(blocks).__name__ - schema = [ - { - "name": p.name, - "type_hint": str(p.type_hint) if p.type_hint is not None else None, - "default": p.default, - "required": p.required, - "description": p.description, - } - for p in blocks.inputs - ] - - if out.mode == OutputFormat.json: - out.dict({"task": "schema", "model": args.model, "pipeline_class": class_name, "inputs": schema}) - elif out.mode == OutputFormat.agent: - out.table(schema, headers=["name", "required", "type_hint", "default", "description"]) - else: - out.text(f"{class_name} ({args.model}) inputs:") - for entry in schema: - tag = "required" if entry["required"] else f"optional, default={entry['default']!r}" - out.text(f" {entry['name']} ({tag})") - if entry["type_hint"]: - out.text(f" type: {entry['type_hint']}") - if entry["description"]: - out.text(f" desc: {entry['description']}") - - -def _warn_custom_block_requirements(args: Namespace) -> None: - """Warn upfront when a modular block ships custom code with declared external dependencies. - - Reads `modular_config.json` if present; if it has an `auto_map` (custom code) and a non-empty `requirements` - list/dict, prints a heads-up. `from_pretrained` will otherwise fail with an `ImportError` deep in the loader stack - when a listed dep is missing. - """ - import diffusers - - try: - config = diffusers.ModularPipelineBlocks.load_config(args.model, token=args.token, revision=args.revision) - except Exception: - return # no modular_config.json or unreachable — nothing to warn about - if not isinstance(config, dict): - return - if not config.get("auto_map"): - return - requirements = config.get("requirements") - if not requirements: - return - - # `requirements` may be a dict {name: version} or (older repos) a list of [name, version] pairs. - if isinstance(requirements, dict): - pairs = list(requirements.items()) - elif isinstance(requirements, list): - pairs = [(item[0], item[1]) for item in requirements if isinstance(item, (list, tuple)) and len(item) >= 2] - else: - pairs = [] - if not pairs: - return - - formatted = ", ".join(f"{name}=={version}" for name, version in pairs) - logger.warning( - f"{args.model!r} ships custom block code with external dependencies: {formatted}. " - "You will need to install these in order to determine the pipeline schema." - ) - - -def _parse_docstring_args(docstring: str | None) -> dict[str, str]: - """Extract per-argument descriptions from a Google-style `Args:` block. - - Returns a `{name: description}` mapping. Best-effort — unrecognised formats just yield an empty dict rather than - raising. - """ - if not docstring: - return {} - - lines = docstring.expandtabs().splitlines() - start = None - section_indent = 0 - for i, line in enumerate(lines): - if line.strip() in ("Args:", "Arguments:", "Parameters:"): - start = i + 1 - section_indent = len(line) - len(line.lstrip()) - break - if start is None: - return {} - - descriptions: dict[str, str] = {} - current_name: str | None = None - current_lines: list[str] = [] - arg_indent: int | None = None - name_pattern = re.compile(r"^(\w+)\s*(?:\([^)]*\))?\s*:?\s*(.*)$") - - def _flush() -> None: - if current_name and current_lines: - descriptions[current_name] = " ".join(s.strip() for s in current_lines).strip() - - for line in lines[start:]: - if not line.strip(): - continue - indent = len(line) - len(line.lstrip()) - # A new top-level section ends the Args block. - if indent <= section_indent and line.strip().endswith(":"): - break - if arg_indent is None: - arg_indent = indent - if indent == arg_indent: - _flush() - current_lines = [] - match = name_pattern.match(line.strip()) - if match: - current_name = match.group(1) - tail = match.group(2).strip() - if tail: - current_lines.append(tail) - else: - current_name = None - elif current_name is not None and indent > arg_indent: - current_lines.append(line.strip()) - _flush() - return descriptions - - -class SchemaCommand(BaseDiffusersCLICommand): - task = "schema" - - @staticmethod - def register_subcommand(subparsers: _SubParsersAction) -> None: - from argparse import RawDescriptionHelpFormatter - - epilog = ( - "Examples\n" - " $ diffusers-cli schema -m stabilityai/stable-diffusion-xl-base-1.0\n" - " $ diffusers-cli schema -m black-forest-labs/FLUX.1-dev --verbose\n" - " $ diffusers-cli --format json schema -m stabilityai/stable-diffusion-xl-base-1.0\n" - "\n" - "Learn more\n" - " Use `diffusers-cli --help` for more information about a command.\n" - " Read the documentation at https://huggingface.co/docs/diffusers\n" - ) - - parser: ArgumentParser = subparsers.add_parser( - "schema", - help="Print the input schema for a diffusers pipeline repo. No weights downloaded.", - usage="\n diffusers-cli schema [options]", - epilog=epilog, - formatter_class=RawDescriptionHelpFormatter, - ) - parser._optionals.title = "Options" - parser.add_argument( - "--model", - "-m", - required=True, - help="Model id on the Hugging Face Hub or local path.", - ) - parser.add_argument( - "--revision", - default=None, - help="Model revision (branch, tag, or commit SHA).", - ) - parser.add_argument( - "--token", - default=None, - help="Hugging Face token for gated/private models.", - ) - parser.add_argument( - "--trust-remote-code", - action="store_true", - help="Allow custom code from the Hub (required for modular pipelines that ship block code).", - ) - parser.add_argument( - "--verbose", - "-v", - action="store_true", - help=( - "Also include per-argument descriptions from the pipeline's __call__ docstring. " - "Modular pipelines always include block-declared descriptions; --verbose populates " - "the equivalent field for standard pipelines by parsing the Google-style Args: block." - ), - ) - parser.set_defaults(func=SchemaCommand) - - def __init__(self, args: Namespace): - self.args = args - - def run(self) -> None: - _schema(self.args) diff --git a/diffusers/commands/skills.py b/diffusers/commands/skills.py deleted file mode 100644 index 60d2e40e883f7459165e1093a4199ed345c3b182..0000000000000000000000000000000000000000 --- a/diffusers/commands/skills.py +++ /dev/null @@ -1,344 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -"""`diffusers-cli skills` — install Agent Skills bundles. - -Skill bundles live under `.ai/skills//` in the diffusers repo and follow the Agent Skills standard: a directory -containing `SKILL.md` (plus optional resources). Installs to `.agents/skills//` which Claude, Codex, and Cursor -all discover. -""" - -from __future__ import annotations - -import os -import shutil -from argparse import ArgumentParser, Namespace, _SubParsersAction -from pathlib import Path - -import httpx -from huggingface_hub.cli._output import out - -from ..utils import logging -from ..utils.constants import DIFFUSERS_REQUEST_TIMEOUT -from . import BaseDiffusersCLICommand - - -logger = logging.get_logger("diffusers-cli/skills") - - -_REGISTRY_BASE = "https://api.github.com/repos/huggingface/diffusers/contents/.ai/skills" -_REGISTRY_REF = "main" - -# Native skill-discovery paths per agent. Claude Code reads only `.claude/skills/`; Codex and -# Cursor read `.agents/skills/` (Cursor also honors `.claude/skills/` via compat, but installing -# to `.agents/skills/` is the portable choice for both). -_CLAUDE_SKILLS_DIR = Path(".claude") / "skills" -_AGENTS_SKILLS_DIR = Path(".agents") / "skills" - -# Env vars set by each agent when it launches the CLI. Values are the install path to use. -_AGENT_ENV_TO_DIR: dict[str, Path] = { - "CLAUDECODE": _CLAUDE_SKILLS_DIR, - "CLAUDE_CODE": _CLAUDE_SKILLS_DIR, - "CODEX_SANDBOX": _AGENTS_SKILLS_DIR, - "CURSOR_AI": _AGENTS_SKILLS_DIR, -} -# When no agent env var is set, install to every native path so whichever agent the user -# later switches to picks the skill up. -_ALL_INSTALL_DIRS: tuple[Path, ...] = (_CLAUDE_SKILLS_DIR, _AGENTS_SKILLS_DIR) - -# Empty marker dropped inside each installed skill dir so `update` can distinguish our -# installs from user-placed skills at the same paths. -_MANAGED_MARKER_FILE = ".diffusers-skill-managed" - - -# --------------------------------------------------------------------------- -# Registry fetch -# --------------------------------------------------------------------------- - - -def _registry_url(name: str = "") -> str: - """API URL for the registry root, or for a single skill bundle when `name` is given.""" - path = f"/{name}" if name else "" - return f"{_REGISTRY_BASE}{path}?ref={_REGISTRY_REF}" - - -def _fetch_json(url: str) -> list[dict]: - try: - resp = httpx.get(url, timeout=DIFFUSERS_REQUEST_TIMEOUT) - resp.raise_for_status() - return resp.json() - except httpx.HTTPStatusError as e: - if e.response.status_code == 404: - raise SystemExit(f"Not found in registry: {url}") from e - raise SystemExit(f"Registry fetch failed: HTTP {e.response.status_code} {e.response.reason_phrase}") from e - except httpx.HTTPError as e: - raise SystemExit(f"Could not reach registry: {e}") from e - - -def _walk_skill_files(name: str) -> list[tuple[str, str]]: - files: list[tuple[str, str]] = [] - - def _walk(api_url: str, prefix: str) -> None: - for entry in _fetch_json(api_url): - if entry["type"] == "file": - files.append((f"{prefix}{entry['name']}", entry["download_url"])) - elif entry["type"] == "dir": - _walk(entry["url"], f"{prefix}{entry['name']}/") - - _walk(_registry_url(name), "") - return files - - -def _download_skill_bundle(name: str) -> dict[str, bytes]: - files = _walk_skill_files(name) - if not files: - raise SystemExit(f"Skill '{name}' has no files in the registry.") - bundle: dict[str, bytes] = {} - for rel_path, url in files: - resp = httpx.get(url, timeout=DIFFUSERS_REQUEST_TIMEOUT) - resp.raise_for_status() - bundle[rel_path] = resp.content - return bundle - - -# --------------------------------------------------------------------------- -# Install / discovery -# --------------------------------------------------------------------------- - - -def _detect_install_dirs() -> tuple[Path, ...]: - """Pick where to install based on the launching agent. - - If we detect a specific agent from its env var, install only there. If nothing is detected, install to every native - path so any agent picks the skill up later. - """ - for env_var, skills_dir in _AGENT_ENV_TO_DIR.items(): - if os.environ.get(env_var): - return (skills_dir,) - return _ALL_INSTALL_DIRS - - -def _install_skill(name: str, bundle: dict[str, bytes], root: Path, skills_dir: Path, force: bool) -> Path: - skill_dir = root / skills_dir / name - if skill_dir.exists(): - if not force: - raise SystemExit(f"Skill already installed at {skill_dir}. Use --force to reinstall.") - shutil.rmtree(skill_dir) - skill_dir.mkdir(parents=True, exist_ok=True) - for rel_path, data in bundle.items(): - target = skill_dir / rel_path - target.parent.mkdir(parents=True, exist_ok=True) - target.write_bytes(data) - (skill_dir / _MANAGED_MARKER_FILE).touch() - return skill_dir - - -def _has_local_changes(skill_dir: Path, bundle: dict[str, bytes]) -> bool: - """True if the installed skill has any file that differs from `bundle` or has extra files. - - The marker file is ignored. Compares raw bytes so a whitespace-only edit still counts as dirty. - """ - on_disk: dict[str, bytes] = {} - for path in skill_dir.rglob("*"): - if not path.is_file(): - continue - rel = str(path.relative_to(skill_dir)) - if rel == _MANAGED_MARKER_FILE: - continue - on_disk[rel] = path.read_bytes() - return on_disk != bundle - - -def _discover_installed(root: Path) -> list[tuple[Path, str]]: - """Return `(skills_dir, name)` pairs for every managed install under `root`.""" - found: list[tuple[Path, str]] = [] - for skills_dir in _ALL_INSTALL_DIRS: - skills_root = root / skills_dir - if not skills_root.exists(): - continue - for d in sorted(skills_root.iterdir()): - if d.is_dir() and (d / _MANAGED_MARKER_FILE).exists(): - found.append((skills_dir, d.name)) - return found - - -class SkillsCommand(BaseDiffusersCLICommand): - @staticmethod - def register_subcommand(subparsers: _SubParsersAction) -> None: - parser: ArgumentParser = subparsers.add_parser( - "skills", - help="Manage Agent Skills for AI assistants.", - usage="\n diffusers-cli skills [options]", - ) - parser._optionals.title = "Options" - actions = parser.add_subparsers(dest="skills_action", required=True, metavar="") - - add = actions.add_parser("add", help="Download and install a skill.") - add.add_argument( - "name", - nargs="?", - default=None, - help="Skill name (e.g. diffusers-cli, custom-blocks). Omit and pass --all to install every skill.", - ) - add.add_argument( - "--all", - dest="install_all", - action="store_true", - help="Install every skill in the registry. Mutually exclusive with a positional name.", - ) - add.add_argument( - "--global", - "-g", - dest="install_global", - action="store_true", - help="Install globally (user-level) instead of in the current project directory.", - ) - add.add_argument("--force", action="store_true", help="Overwrite existing skills in the destination.") - add.set_defaults(func=SkillsCommand) - - list_action = actions.add_parser("list", help="List available skills in the registry.") - list_action.set_defaults(func=SkillsCommand) - - update = actions.add_parser("update", help="Re-download and reinstall managed skills.") - update.add_argument( - "name", - nargs="?", - default=None, - help="Optional installed skill name to update. Omit to update every managed skill.", - ) - update.add_argument( - "--global", - "-g", - dest="install_global", - action="store_true", - help="Update skills installed globally (user-level) instead of the current project.", - ) - update.add_argument( - "--force", - action="store_true", - help="Overwrite skills even if they have local modifications since install.", - ) - update.set_defaults(func=SkillsCommand) - - preview = actions.add_parser("preview", help="Print a skill's SKILL.md from the registry.") - preview.add_argument("name", help="Skill name to preview.") - preview.set_defaults(func=SkillsCommand) - - def __init__(self, args: Namespace): - self.args = args - - def run(self) -> None: - if self.args.skills_action == "add": - self._add() - elif self.args.skills_action == "list": - self._list() - elif self.args.skills_action == "update": - self._update() - elif self.args.skills_action == "preview": - self._preview() - - def _add(self) -> None: - if self.args.install_all and self.args.name: - raise SystemExit("--all and a positional skill name are mutually exclusive.") - if not self.args.install_all and not self.args.name: - raise SystemExit("Pass a skill name (e.g. diffusers-cli) or --all to install every skill.") - - root = Path.home() if self.args.install_global else Path.cwd() - install_dirs = _detect_install_dirs() - names = self._resolve_names() - - installed: list[str] = [] - failed: list[str] = [] - for name in names: - try: - bundle = _download_skill_bundle(name) - for skills_dir in install_dirs: - _install_skill(name, bundle, root, skills_dir, self.args.force) - installed.append(name) - except (SystemExit, httpx.HTTPError) as e: - # Downgrade to a warning so one broken skill doesn't abort the batch. - logger.warning(f"Skipping skill {name!r}: {e}") - failed.append(name) - - if not installed: - raise SystemExit(f"No skills installed. Failed: {failed}") - out.result( - f"Installed {len(installed)} skill(s)", - installed=", ".join(installed), - failed=", ".join(failed) if failed else None, - paths=", ".join(str(root / d) for d in install_dirs), - ) - - def _update(self) -> None: - root = Path.home() if self.args.install_global else Path.cwd() - installed = _discover_installed(root) - if self.args.name is not None: - installed = [entry for entry in installed if entry[1] == self.args.name] - if not installed: - raise SystemExit(f"No installed skill named {self.args.name!r} found under {root}.") - if not installed: - raise SystemExit(f"No managed skills found under {root}.") - - # Group by skill name so we redownload each bundle once even if it's installed to - # multiple locations (e.g. both .claude/skills/ and .agents/skills/). - by_name: dict[str, list[Path]] = {} - for skills_dir, name in installed: - by_name.setdefault(name, []).append(skills_dir) - - updated: list[str] = [] - failed: list[str] = [] - skipped: list[str] = [] - for name, dirs in sorted(by_name.items()): - try: - bundle = _download_skill_bundle(name) - for skills_dir in dirs: - skill_dir = root / skills_dir / name - if not self.args.force and _has_local_changes(skill_dir, bundle): - logger.warning( - f"Skill {name!r} at {skill_dir} has local modifications; " - "skipping. Pass --force to overwrite them." - ) - skipped.append(name) - continue - _install_skill(name, bundle, root, skills_dir, force=True) - updated.append(name) - except (SystemExit, httpx.HTTPError) as e: - logger.warning(f"Skipping skill {name!r}: {e}") - failed.append(name) - - out.result( - f"Updated {len(updated)} skill(s)", - updated=", ".join(updated), - skipped=", ".join(skipped) if skipped else None, - failed=", ".join(failed) if failed else None, - ) - - def _preview(self) -> None: - bundle = _download_skill_bundle(self.args.name) - skill_md = bundle.get("SKILL.md") - if skill_md is None: - raise SystemExit(f"Skill {self.args.name!r} has no SKILL.md in the registry.") - print(skill_md.decode()) - - def _list(self) -> None: - entries = _fetch_json(_registry_url()) - skills = [{"name": e["name"]} for e in entries if e["type"] == "dir" and not e["name"].startswith(".")] - if not skills: - raise SystemExit("No skills found in registry.") - out.table(skills, headers=["name"]) - - def _resolve_names(self) -> list[str]: - if self.args.install_all: - entries = _fetch_json(_registry_url()) - return sorted(e["name"] for e in entries if e["type"] == "dir" and not e["name"].startswith(".")) - return [self.args.name] diff --git a/diffusers/configuration_utils.py b/diffusers/configuration_utils.py deleted file mode 100644 index f16871b9f56f3b340d66cabd35cac5ea04e1cfdd..0000000000000000000000000000000000000000 --- a/diffusers/configuration_utils.py +++ /dev/null @@ -1,752 +0,0 @@ -# coding=utf-8 -# Copyright 2025 The HuggingFace Inc. team. -# Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -"""ConfigMixin base class and utilities.""" - -import functools -import importlib -import inspect -import json -import os -import re -from collections import OrderedDict -from pathlib import Path -from typing import Any - -import numpy as np -from huggingface_hub import DDUFEntry, create_repo, hf_hub_download -from huggingface_hub.utils import ( - EntryNotFoundError, - HfHubHTTPError, - RepositoryNotFoundError, - RevisionNotFoundError, - validate_hf_hub_args, -) -from typing_extensions import Self - -from . import __version__ -from .utils import ( - HUGGINGFACE_CO_RESOLVE_ENDPOINT, - DummyObject, - deprecate, - extract_commit_hash, - http_user_agent, - logging, -) - - -logger = logging.get_logger(__name__) - -_re_configuration_file = re.compile(r"config\.(.*)\.json") - - -class FrozenDict(OrderedDict): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - - for key, value in self.items(): - setattr(self, key, value) - - self.__frozen = True - - def __delitem__(self, *args, **kwargs): - raise Exception(f"You cannot use ``__delitem__`` on a {self.__class__.__name__} instance.") - - def setdefault(self, *args, **kwargs): - raise Exception(f"You cannot use ``setdefault`` on a {self.__class__.__name__} instance.") - - def pop(self, *args, **kwargs): - raise Exception(f"You cannot use ``pop`` on a {self.__class__.__name__} instance.") - - def update(self, *args, **kwargs): - raise Exception(f"You cannot use ``update`` on a {self.__class__.__name__} instance.") - - def __setattr__(self, name, value): - if hasattr(self, "__frozen") and self.__frozen: - raise Exception(f"You cannot use ``__setattr__`` on a {self.__class__.__name__} instance.") - super().__setattr__(name, value) - - def __setitem__(self, name, value): - if hasattr(self, "__frozen") and self.__frozen: - raise Exception(f"You cannot use ``__setattr__`` on a {self.__class__.__name__} instance.") - super().__setitem__(name, value) - - -class ConfigMixin: - r""" - Base class for all configuration classes. All configuration parameters are stored under `self.config`. Also - provides the [`~ConfigMixin.from_config`] and [`~ConfigMixin.save_config`] methods for loading, downloading, and - saving classes that inherit from [`ConfigMixin`]. - - Class attributes: - - **config_name** (`str`) -- A filename under which the config should stored when calling - [`~ConfigMixin.save_config`] (should be overridden by parent class). - - **ignore_for_config** (`list[str]`) -- A list of attributes that should not be saved in the config (should be - overridden by subclass). - - **has_compatibles** (`bool`) -- Whether the class has compatible classes (should be overridden by subclass). - - **_deprecated_kwargs** (`list[str]`) -- Keyword arguments that are deprecated. Note that the `init` function - should only have a `kwargs` argument if at least one argument is deprecated (should be overridden by - subclass). - """ - - config_name = None - ignore_for_config = [] - has_compatibles = False - - _deprecated_kwargs = [] - _auto_class = None - - @classmethod - def register_for_auto_class(cls, auto_class="AutoModel"): - """ - Register this class with the given auto class so that it can be loaded with `AutoModel.from_pretrained(..., - trust_remote_code=True)`. - - When the config is saved, the resulting `config.json` will include an `auto_map` entry mapping the auto class - to this class's module and class name. - - Args: - auto_class (`str` or type, *optional*, defaults to `"AutoModel"`): - The auto class to register this class with. Can be a string (e.g. `"AutoModel"`) or the class itself. - Currently only `"AutoModel"` is supported. - - Example: - - ```python - from diffusers import ModelMixin, ConfigMixin - - - class MyCustomModel(ModelMixin, ConfigMixin): ... - - - MyCustomModel.register_for_auto_class("AutoModel") - ``` - """ - if auto_class != "AutoModel": - raise ValueError(f"Only 'AutoModel' is supported, got '{auto_class}'.") - - cls._auto_class = auto_class - - def register_to_config(self, **kwargs): - if self.config_name is None: - raise NotImplementedError(f"Make sure that {self.__class__} has defined a class name `config_name`") - # Special case for `kwargs` used in deprecation warning added to schedulers - # TODO: remove this when we remove the deprecation warning, and the `kwargs` argument, - # or solve in a more general way. - kwargs.pop("kwargs", None) - - if not hasattr(self, "_internal_dict"): - internal_dict = kwargs - else: - previous_dict = dict(self._internal_dict) - internal_dict = {**self._internal_dict, **kwargs} - logger.debug(f"Updating config from {previous_dict} to {internal_dict}") - - self._internal_dict = FrozenDict(internal_dict) - - def __getattr__(self, name: str) -> Any: - """The only reason we overwrite `getattr` here is to gracefully deprecate accessing - config attributes directly. See https://github.com/huggingface/diffusers/pull/3129 - - This function is mostly copied from PyTorch's __getattr__ overwrite: - https://pytorch.org/docs/stable/_modules/torch/nn/modules/module.html#Module - """ - - is_in_config = "_internal_dict" in self.__dict__ and hasattr(self.__dict__["_internal_dict"], name) - is_attribute = name in self.__dict__ - - if is_in_config and not is_attribute: - 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}'." - deprecate("direct config name access", "1.0.0", deprecation_message, standard_warn=False) - return self._internal_dict[name] - - raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") - - def save_config(self, save_directory: str | os.PathLike, push_to_hub: bool = False, **kwargs): - """ - Save a configuration object to the directory specified in `save_directory` so that it can be reloaded using the - [`~ConfigMixin.from_config`] class method. - - Args: - save_directory (`str` or `os.PathLike`): - Directory where the configuration JSON file is saved (will be created if it does not exist). - push_to_hub (`bool`, *optional*, defaults to `False`): - Whether or not to push your model to the Hugging Face Hub after saving it. You can specify the - repository you want to push to with `repo_id` (will default to the name of `save_directory` in your - namespace). - kwargs (`dict[str, Any]`, *optional*): - Additional keyword arguments passed along to the [`~utils.PushToHubMixin.push_to_hub`] method. - """ - if os.path.isfile(save_directory): - raise AssertionError(f"Provided path ({save_directory}) should be a directory, not a file") - - os.makedirs(save_directory, exist_ok=True) - - # If we save using the predefined names, we can load using `from_config` - output_config_file = os.path.join(save_directory, self.config_name) - - self.to_json_file(output_config_file) - logger.info(f"Configuration saved in {output_config_file}") - - if push_to_hub: - commit_message = kwargs.pop("commit_message", None) - private = kwargs.pop("private", None) - create_pr = kwargs.pop("create_pr", False) - token = kwargs.pop("token", None) - repo_id = kwargs.pop("repo_id", save_directory.split(os.path.sep)[-1]) - repo_id = create_repo(repo_id, exist_ok=True, private=private, token=token).repo_id - subfolder = kwargs.pop("subfolder", None) - - self._upload_folder( - save_directory, - repo_id, - token=token, - commit_message=commit_message, - create_pr=create_pr, - subfolder=subfolder, - ) - - @classmethod - def from_config( - cls, config: FrozenDict | dict[str, Any] = None, return_unused_kwargs=False, **kwargs - ) -> Self | tuple[Self, dict[str, Any]]: - r""" - Instantiate a Python class from a config dictionary. - - Parameters: - config (`dict[str, Any]`): - A config dictionary from which the Python class is instantiated. Make sure to only load configuration - files of compatible classes. - return_unused_kwargs (`bool`, *optional*, defaults to `False`): - Whether kwargs that are not consumed by the Python class should be returned or not. - kwargs (remaining dictionary of keyword arguments, *optional*): - Can be used to update the configuration object (after it is loaded) and initiate the Python class. - `**kwargs` are passed directly to the underlying scheduler/model's `__init__` method and eventually - overwrite the same named arguments in `config`. - - Returns: - [`ModelMixin`] or [`SchedulerMixin`]: - A model or scheduler object instantiated from a config dictionary. - - Examples: - - ```python - >>> from diffusers import DDPMScheduler, DDIMScheduler, PNDMScheduler - - >>> # Download scheduler from huggingface.co and cache. - >>> scheduler = DDPMScheduler.from_pretrained("google/ddpm-cifar10-32") - - >>> # Instantiate DDIM scheduler class with same config as DDPM - >>> scheduler = DDIMScheduler.from_config(scheduler.config) - - >>> # Instantiate PNDM scheduler class with same config as DDPM - >>> scheduler = PNDMScheduler.from_config(scheduler.config) - ``` - """ - # <===== TO BE REMOVED WITH DEPRECATION - # TODO(Patrick) - make sure to remove the following lines when config=="model_path" is deprecated - if "pretrained_model_name_or_path" in kwargs: - config = kwargs.pop("pretrained_model_name_or_path") - - if config is None: - raise ValueError("Please make sure to provide a config as the first positional argument.") - # ======> - - if not isinstance(config, dict): - deprecation_message = "It is deprecated to pass a pretrained model name or path to `from_config`." - if "Scheduler" in cls.__name__: - deprecation_message += ( - f"If you were trying to load a scheduler, please use {cls}.from_pretrained(...) instead." - " Otherwise, please make sure to pass a configuration dictionary instead. This functionality will" - " be removed in v1.0.0." - ) - elif "Model" in cls.__name__: - deprecation_message += ( - f"If you were trying to load a model, please use {cls}.load_config(...) followed by" - f" {cls}.from_config(...) instead. Otherwise, please make sure to pass a configuration dictionary" - " instead. This functionality will be removed in v1.0.0." - ) - deprecate("config-passed-as-path", "1.0.0", deprecation_message, standard_warn=False) - config, kwargs = cls.load_config(pretrained_model_name_or_path=config, return_unused_kwargs=True, **kwargs) - - init_dict, unused_kwargs, hidden_dict = cls.extract_init_dict(config, **kwargs) - - # Allow dtype to be specified on initialization - if "dtype" in unused_kwargs: - init_dict["dtype"] = unused_kwargs.pop("dtype") - - # add possible deprecated kwargs - for deprecated_kwarg in cls._deprecated_kwargs: - if deprecated_kwarg in unused_kwargs: - init_dict[deprecated_kwarg] = unused_kwargs.pop(deprecated_kwarg) - - # Return model and optionally state and/or unused_kwargs - model = cls(**init_dict) - - # make sure to also save config parameters that might be used for compatible classes - # update _class_name - if "_class_name" in hidden_dict: - hidden_dict["_class_name"] = cls.__name__ - - model.register_to_config(**hidden_dict) - - # add hidden kwargs of compatible classes to unused_kwargs - unused_kwargs = {**unused_kwargs, **hidden_dict} - - if return_unused_kwargs: - return (model, unused_kwargs) - else: - return model - - @classmethod - def get_config_dict(cls, *args, **kwargs): - deprecation_message = ( - f" The function get_config_dict is deprecated. Please use {cls}.load_config instead. This function will be" - " removed in version v1.0.0" - ) - deprecate("get_config_dict", "1.0.0", deprecation_message, standard_warn=False) - return cls.load_config(*args, **kwargs) - - @classmethod - @validate_hf_hub_args - def load_config( - cls, - pretrained_model_name_or_path: str | os.PathLike, - return_unused_kwargs=False, - return_commit_hash=False, - **kwargs, - ) -> tuple[dict[str, Any], dict[str, Any]]: - r""" - Load a model or scheduler configuration. - - Parameters: - pretrained_model_name_or_path (`str` or `os.PathLike`, *optional*): - Can be either: - - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing model weights saved with - [`~ConfigMixin.save_config`]. - - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - output_loading_info(`bool`, *optional*, defaults to `False`): - Whether or not to also return a dictionary containing missing keys, unexpected keys and error messages. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - subfolder (`str`, *optional*, defaults to `""`): - The subfolder location of a model file within a larger model repository on the Hub or locally. - return_unused_kwargs (`bool`, *optional*, defaults to `False): - Whether unused keyword arguments of the config are returned. - return_commit_hash (`bool`, *optional*, defaults to `False): - Whether the `commit_hash` of the loaded configuration are returned. - - Returns: - `dict`: - A dictionary of all the parameters stored in a JSON configuration file. - - """ - cache_dir = kwargs.pop("cache_dir", None) - local_dir = kwargs.pop("local_dir", None) - local_dir_use_symlinks = kwargs.pop("local_dir_use_symlinks", "auto") - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - token = kwargs.pop("token", None) - local_files_only = kwargs.pop("local_files_only", False) - revision = kwargs.pop("revision", None) - _ = kwargs.pop("mirror", None) - subfolder = kwargs.pop("subfolder", None) - user_agent = kwargs.pop("user_agent", {}) - dduf_entries: dict[str, DDUFEntry] | None = kwargs.pop("dduf_entries", None) - - user_agent = {**user_agent, "file_type": "config"} - user_agent = http_user_agent(user_agent) - - pretrained_model_name_or_path = str(pretrained_model_name_or_path) - - if cls.config_name is None: - raise ValueError( - "`self.config_name` is not defined. Note that one should not load a config from " - "`ConfigMixin`. Please make sure to define `config_name` in a class inheriting from `ConfigMixin`" - ) - # Custom path for now - if dduf_entries: - if subfolder is not None: - raise ValueError( - "DDUF file only allow for 1 level of directory (e.g transformer/model1/model.safetentors is not allowed). " - "Please check the DDUF structure" - ) - config_file = cls._get_config_file_from_dduf(pretrained_model_name_or_path, dduf_entries) - elif os.path.isfile(pretrained_model_name_or_path): - config_file = pretrained_model_name_or_path - elif os.path.isdir(pretrained_model_name_or_path): - if subfolder is not None and os.path.isfile( - os.path.join(pretrained_model_name_or_path, subfolder, cls.config_name) - ): - config_file = os.path.join(pretrained_model_name_or_path, subfolder, cls.config_name) - elif os.path.isfile(os.path.join(pretrained_model_name_or_path, cls.config_name)): - # Load from a PyTorch checkpoint - config_file = os.path.join(pretrained_model_name_or_path, cls.config_name) - else: - raise EnvironmentError( - f"Error no file named {cls.config_name} found in directory {pretrained_model_name_or_path}." - ) - else: - try: - # Load from URL or cache if already cached - config_file = hf_hub_download( - pretrained_model_name_or_path, - filename=cls.config_name, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - user_agent=user_agent, - subfolder=subfolder, - revision=revision, - local_dir=local_dir, - local_dir_use_symlinks=local_dir_use_symlinks, - ) - except RepositoryNotFoundError: - raise EnvironmentError( - f"{pretrained_model_name_or_path} is not a local folder and is not a valid model identifier" - " listed on 'https://huggingface.co/models'\nIf this is a private repository, make sure to pass a" - " token having permission to this repo with `token` or log in with `hf auth login`." - ) - except RevisionNotFoundError: - raise EnvironmentError( - f"{revision} is not a valid git identifier (branch name, tag name or commit id) that exists for" - " this model name. Check the model page at" - f" 'https://huggingface.co/{pretrained_model_name_or_path}' for available revisions." - ) - except EntryNotFoundError: - raise EnvironmentError( - f"{pretrained_model_name_or_path} does not appear to have a file named {cls.config_name}." - ) - except HfHubHTTPError as err: - raise EnvironmentError( - "There was a specific connection error when trying to load" - f" {pretrained_model_name_or_path}:\n{err}" - ) - except ValueError: - raise EnvironmentError( - f"We couldn't connect to '{HUGGINGFACE_CO_RESOLVE_ENDPOINT}' to load this model, couldn't find it" - f" in the cached files and it looks like {pretrained_model_name_or_path} is not the path to a" - f" directory containing a {cls.config_name} file.\nCheckout your internet connection or see how to" - " run the library in offline mode at" - " 'https://huggingface.co/docs/diffusers/installation#offline-mode'." - ) - except EnvironmentError: - raise EnvironmentError( - f"Can't load config for '{pretrained_model_name_or_path}'. If you were trying to load it from " - "'https://huggingface.co/models', make sure you don't have a local directory with the same name. " - f"Otherwise, make sure '{pretrained_model_name_or_path}' is the correct path to a directory " - f"containing a {cls.config_name} file" - ) - try: - config_dict = cls._dict_from_json_file(config_file, dduf_entries=dduf_entries) - - commit_hash = extract_commit_hash(config_file) - except (json.JSONDecodeError, UnicodeDecodeError): - raise EnvironmentError(f"It looks like the config file at '{config_file}' is not a valid JSON file.") - - if not (return_unused_kwargs or return_commit_hash): - return config_dict - - outputs = (config_dict,) - - if return_unused_kwargs: - outputs += (kwargs,) - - if return_commit_hash: - outputs += (commit_hash,) - - return outputs - - @staticmethod - def _get_init_keys(input_class): - return set(dict(inspect.signature(input_class.__init__).parameters).keys()) - - @classmethod - def extract_init_dict(cls, config_dict, **kwargs): - # Skip keys that were not present in the original config, so default __init__ values were used - used_defaults = config_dict.get("_use_default_values", []) - config_dict = {k: v for k, v in config_dict.items() if k not in used_defaults and k != "_use_default_values"} - - # 0. Copy origin config dict - original_dict = dict(config_dict.items()) - - # 1. Retrieve expected config attributes from __init__ signature - expected_keys = cls._get_init_keys(cls) - expected_keys.remove("self") - # remove general kwargs if present in dict - if "kwargs" in expected_keys: - expected_keys.remove("kwargs") - - # 2. Remove attributes that cannot be expected from expected config attributes - # remove keys to be ignored - if len(cls.ignore_for_config) > 0: - expected_keys = expected_keys - set(cls.ignore_for_config) - - # load diffusers library to import compatible and original scheduler - diffusers_library = importlib.import_module(__name__.split(".")[0]) - - if cls.has_compatibles: - compatible_classes = [c for c in cls._get_compatibles() if not isinstance(c, DummyObject)] - else: - compatible_classes = [] - - expected_keys_comp_cls = set() - for c in compatible_classes: - expected_keys_c = cls._get_init_keys(c) - expected_keys_comp_cls = expected_keys_comp_cls.union(expected_keys_c) - expected_keys_comp_cls = expected_keys_comp_cls - cls._get_init_keys(cls) - config_dict = {k: v for k, v in config_dict.items() if k not in expected_keys_comp_cls} - - # remove attributes from orig class that cannot be expected - orig_cls_name = config_dict.pop("_class_name", cls.__name__) - if ( - isinstance(orig_cls_name, str) - and orig_cls_name != cls.__name__ - and hasattr(diffusers_library, orig_cls_name) - ): - orig_cls = getattr(diffusers_library, orig_cls_name) - unexpected_keys_from_orig = cls._get_init_keys(orig_cls) - expected_keys - config_dict = {k: v for k, v in config_dict.items() if k not in unexpected_keys_from_orig} - elif not isinstance(orig_cls_name, str) and not isinstance(orig_cls_name, (list, tuple)): - raise ValueError( - "Make sure that the `_class_name` is of type string or list of string (for custom pipelines)." - ) - - # remove private attributes - config_dict = {k: v for k, v in config_dict.items() if not k.startswith("_")} - - # remove quantization_config - config_dict = {k: v for k, v in config_dict.items() if k != "quantization_config"} - - # 3. Create keyword arguments that will be passed to __init__ from expected keyword arguments - init_dict = {} - for key in expected_keys: - # if config param is passed to kwarg and is present in config dict - # it should overwrite existing config dict key - if key in kwargs and key in config_dict: - config_dict[key] = kwargs.pop(key) - - if key in kwargs: - # overwrite key - init_dict[key] = kwargs.pop(key) - elif key in config_dict: - # use value from config dict - init_dict[key] = config_dict.pop(key) - - # 4. Give nice warning if unexpected values have been passed - if len(config_dict) > 0: - logger.warning( - f"The config attributes {config_dict} were passed to {cls.__name__}, " - "but are not expected and will be ignored. Please verify your " - f"{cls.config_name} configuration file." - ) - - # 5. Give nice info if config attributes are initialized to default because they have not been passed - passed_keys = set(init_dict.keys()) - if len(expected_keys - passed_keys) > 0: - logger.info( - f"{expected_keys - passed_keys} was not found in config. Values will be initialized to default values." - ) - - # 6. Define unused keyword arguments - unused_kwargs = {**config_dict, **kwargs} - - # 7. Define "hidden" config parameters that were saved for compatible classes - hidden_config_dict = {k: v for k, v in original_dict.items() if k not in init_dict} - - return init_dict, unused_kwargs, hidden_config_dict - - @classmethod - def _dict_from_json_file(cls, json_file: str | os.PathLike, dduf_entries: dict[str, DDUFEntry] | None = None): - if dduf_entries: - text = dduf_entries[json_file].read_text() - else: - with open(json_file, "r", encoding="utf-8") as reader: - text = reader.read() - return json.loads(text) - - def __repr__(self): - return f"{self.__class__.__name__} {self.to_json_string()}" - - @property - def config(self) -> dict[str, Any]: - """ - Returns the config of the class as a frozen dictionary - - Returns: - `dict[str, Any]`: Config of the class. - """ - return self._internal_dict - - def to_json_string(self) -> str: - """ - Serializes the configuration instance to a JSON string. - - Returns: - `str`: - String containing all the attributes that make up the configuration instance in JSON format. - """ - config_dict = self._internal_dict if hasattr(self, "_internal_dict") else {} - config_dict["_class_name"] = self.__class__.__name__ - config_dict["_diffusers_version"] = __version__ - - def to_json_saveable(value): - if isinstance(value, np.ndarray): - value = value.tolist() - elif isinstance(value, Path): - value = value.as_posix() - elif hasattr(value, "to_dict") and callable(value.to_dict): - value = value.to_dict() - elif isinstance(value, list): - value = [to_json_saveable(v) for v in value] - return value - - if "quantization_config" in config_dict: - config_dict["quantization_config"] = ( - config_dict.quantization_config.to_dict() - if not isinstance(config_dict.quantization_config, dict) - else config_dict.quantization_config - ) - - config_dict = {k: to_json_saveable(v) for k, v in config_dict.items()} - # Don't save "_ignore_files" or "_use_default_values" - config_dict.pop("_ignore_files", None) - config_dict.pop("_use_default_values", None) - # pop the `_pre_quantization_dtype` as torch.dtypes are not serializable. - _ = config_dict.pop("_pre_quantization_dtype", None) - - if getattr(self, "_auto_class", None) is not None: - module = self.__class__.__module__.split(".")[-1] - auto_map = config_dict.get("auto_map", {}) - auto_map[self._auto_class] = f"{module}.{self.__class__.__name__}" - config_dict["auto_map"] = auto_map - - return json.dumps(config_dict, indent=2, sort_keys=True) + "\n" - - def to_json_file(self, json_file_path: str | os.PathLike): - """ - Save the configuration instance's parameters to a JSON file. - - Args: - json_file_path (`str` or `os.PathLike`): - Path to the JSON file to save a configuration instance's parameters. - """ - with open(json_file_path, "w", encoding="utf-8") as writer: - writer.write(self.to_json_string()) - - @classmethod - def _get_config_file_from_dduf(cls, pretrained_model_name_or_path: str, dduf_entries: dict[str, DDUFEntry]): - # paths inside a DDUF file must always be "/" - config_file = ( - cls.config_name - if pretrained_model_name_or_path == "" - else "/".join([pretrained_model_name_or_path, cls.config_name]) - ) - if config_file not in dduf_entries: - raise ValueError( - f"We did not manage to find the file {config_file} in the dduf file. We only have the following files {dduf_entries.keys()}" - ) - return config_file - - -def register_to_config(init): - r""" - Decorator to apply on the init of classes inheriting from [`ConfigMixin`] so that all the arguments are - automatically sent to `self.register_for_config`. To ignore a specific argument accepted by the init but that - shouldn't be registered in the config, use the `ignore_for_config` class variable - - Warning: Once decorated, all private arguments (beginning with an underscore) are trashed and not sent to the init! - """ - - @functools.wraps(init) - def inner_init(self, *args, **kwargs): - # Ignore private kwargs in the init. - init_kwargs = {k: v for k, v in kwargs.items() if not k.startswith("_")} - config_init_kwargs = {k: v for k, v in kwargs.items() if k.startswith("_")} - if not isinstance(self, ConfigMixin): - raise RuntimeError( - f"`@register_for_config` was applied to {self.__class__.__name__} init method, but this class does " - "not inherit from `ConfigMixin`." - ) - - ignore = getattr(self, "ignore_for_config", []) - # Get positional arguments aligned with kwargs - new_kwargs = {} - signature = inspect.signature(init) - parameters = { - name: p.default for i, (name, p) in enumerate(signature.parameters.items()) if i > 0 and name not in ignore - } - for arg, name in zip(args, parameters.keys()): - new_kwargs[name] = arg - - # Then add all kwargs - new_kwargs.update( - { - k: init_kwargs.get(k, default) - for k, default in parameters.items() - if k not in ignore and k not in new_kwargs - } - ) - - # Take note of the parameters that were not present in the loaded config - if len(set(new_kwargs.keys()) - set(init_kwargs)) > 0: - new_kwargs["_use_default_values"] = list(set(new_kwargs.keys()) - set(init_kwargs)) - - new_kwargs = {**config_init_kwargs, **new_kwargs} - getattr(self, "register_to_config")(**new_kwargs) - init(self, *args, **init_kwargs) - - return inner_init - - -class LegacyConfigMixin(ConfigMixin): - r""" - A subclass of `ConfigMixin` to resolve class mapping from legacy classes (like `Transformer2DModel`) to more - pipeline-specific classes (like `DiTTransformer2DModel`). - """ - - @classmethod - def from_config(cls, config: FrozenDict | dict[str, Any] = None, return_unused_kwargs=False, **kwargs): - # To prevent dependency import problem. - from .models.model_loading_utils import _fetch_remapped_cls_from_config - - # resolve remapping - remapped_class = _fetch_remapped_cls_from_config(config, cls) - - if remapped_class is cls: - return super(LegacyConfigMixin, remapped_class).from_config(config, return_unused_kwargs, **kwargs) - else: - return remapped_class.from_config(config, return_unused_kwargs, **kwargs) diff --git a/diffusers/dependency_versions_check.py b/diffusers/dependency_versions_check.py deleted file mode 100644 index 262b3941d87dc2a539b2ffbdb02cd332b42776d1..0000000000000000000000000000000000000000 --- a/diffusers/dependency_versions_check.py +++ /dev/null @@ -1,34 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from .dependency_versions_table import deps -from .utils.versions import require_version, require_version_core - - -# define which module versions we always want to check at run time -# (usually the ones defined in `install_requires` in setup.py) -# -# order specific notes: -# - tqdm must be checked before tokenizers - -pkgs_to_check_at_runtime = "python requests filelock numpy".split() -for pkg in pkgs_to_check_at_runtime: - if pkg in deps: - require_version_core(deps[pkg]) - else: - raise ValueError(f"can't find {pkg} in {deps.keys()}, check dependency_versions_table.py") - - -def dep_version_check(pkg, hint=None): - require_version(deps[pkg], hint) diff --git a/diffusers/dependency_versions_table.py b/diffusers/dependency_versions_table.py deleted file mode 100644 index 02e304d2ab02d3c90803dcc30b3675b1aaccbf57..0000000000000000000000000000000000000000 --- a/diffusers/dependency_versions_table.py +++ /dev/null @@ -1,57 +0,0 @@ -# THIS FILE HAS BEEN AUTOGENERATED. To update: -# 1. modify the `_deps` dict in setup.py -# 2. run `make deps_table_update` -deps = { - "Pillow": "Pillow", - "accelerate": "accelerate>=0.31.0", - "datasets": "datasets", - "filelock": "filelock", - "ftfy": "ftfy", - "hf-doc-builder": "hf-doc-builder>=0.3.0", - "httpx": "httpx<1.0.0", - "huggingface-hub": "huggingface-hub>=1.23.0,<2.0", - "requests-mock": "requests-mock==1.10.0", - "importlib_metadata": "importlib_metadata", - "invisible-watermark": "invisible-watermark>=0.2.0", - "isort": "isort>=5.5.4", - "Jinja2": "Jinja2", - "torchsde": "torchsde", - "note_seq": "note_seq", - "librosa": "librosa", - "llvmlite": "llvmlite>=0.40.0", - "numba": "numba>=0.57.0", - "numpy": "numpy", - "parameterized": "parameterized", - "peft": "peft>=0.17.0", - "protobuf": "protobuf>=3.20.3,<4", - "pytest": "pytest", - "pytest-timeout": "pytest-timeout", - "pytest-xdist": "pytest-xdist", - "python": "python>=3.10.0", - "ruff": "ruff==0.9.10", - "safetensors": "safetensors>=0.8.0", - "sentencepiece": "sentencepiece>=0.1.91,!=0.1.92", - "GitPython": "GitPython<3.1.19", - "scipy": "scipy", - "onnx": "onnx", - "optimum_quanto": "optimum_quanto>=0.2.6", - "gguf": "gguf>=0.10.0", - "auto-round": "auto-round>=0.13.0", - "torchao": "torchao>=0.7.0", - "bitsandbytes": "bitsandbytes>=0.43.3", - "nvidia_modelopt[hf]": "nvidia_modelopt[hf]>=0.33.1", - "sdnq": "sdnq>=0.2.2", - "regex": "regex!=2019.12.17", - "requests": "requests", - "tensorboard": "tensorboard", - "tiktoken": "tiktoken>=0.7.0", - "torch": "torch>=2.6", - "torchvision": "torchvision", - "transformers": "transformers>=4.41.2", - "urllib3": "urllib3<=2.0.0", - "black": "black", - "phonemizer": "phonemizer", - "opencv-python": "opencv-python", - "timm": "timm", - "flashpack": "flashpack", -} diff --git a/diffusers/experimental/README.md b/diffusers/experimental/README.md deleted file mode 100644 index 77594b14dbfc3131aa79f09fb1d64231c124ae7b..0000000000000000000000000000000000000000 --- a/diffusers/experimental/README.md +++ /dev/null @@ -1,5 +0,0 @@ -# 🧨 Diffusers Experimental - -We are adding experimental code to support novel applications and usages of the Diffusers library. -Currently, the following experiments are supported: -* Reinforcement learning via an implementation of the [Diffuser](https://huggingface.co/papers/2205.09991) model. \ No newline at end of file diff --git a/diffusers/experimental/__init__.py b/diffusers/experimental/__init__.py deleted file mode 100644 index ebc8155403016dfd8ad7fb78d246f9da9098ac50..0000000000000000000000000000000000000000 --- a/diffusers/experimental/__init__.py +++ /dev/null @@ -1 +0,0 @@ -from .rl import ValueGuidedRLPipeline diff --git a/diffusers/experimental/rl/__init__.py b/diffusers/experimental/rl/__init__.py deleted file mode 100644 index 7b338d3173e12d478b6b6d6fd0e50650a0ab5a4c..0000000000000000000000000000000000000000 --- a/diffusers/experimental/rl/__init__.py +++ /dev/null @@ -1 +0,0 @@ -from .value_guided_sampling import ValueGuidedRLPipeline diff --git a/diffusers/experimental/rl/value_guided_sampling.py b/diffusers/experimental/rl/value_guided_sampling.py deleted file mode 100644 index 273eeb84c50bfb1138af972af4bd461995a2ae62..0000000000000000000000000000000000000000 --- a/diffusers/experimental/rl/value_guided_sampling.py +++ /dev/null @@ -1,153 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import numpy as np -import torch -import tqdm - -from ...models.unets.unet_1d import UNet1DModel -from ...pipelines import DiffusionPipeline -from ...utils.dummy_pt_objects import DDPMScheduler -from ...utils.torch_utils import randn_tensor - - -class ValueGuidedRLPipeline(DiffusionPipeline): - r""" - Pipeline for value-guided sampling from a diffusion model trained to predict sequences of states. - - This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods - implemented for all pipelines (downloading, saving, running on a particular device, etc.). - - Parameters: - value_function ([`UNet1DModel`]): - A specialized UNet for fine-tuning trajectories base on reward. - unet ([`UNet1DModel`]): - UNet architecture to denoise the encoded trajectories. - scheduler ([`SchedulerMixin`]): - A scheduler to be used in combination with `unet` to denoise the encoded trajectories. Default for this - application is [`DDPMScheduler`]. - env (): - An environment following the OpenAI gym API to act in. For now only Hopper has pretrained models. - """ - - def __init__( - self, - value_function: UNet1DModel, - unet: UNet1DModel, - scheduler: DDPMScheduler, - env, - ): - super().__init__() - - self.register_modules(value_function=value_function, unet=unet, scheduler=scheduler, env=env) - - self.data = env.get_dataset() - self.means = {} - for key in self.data.keys(): - try: - self.means[key] = self.data[key].mean() - except: # noqa: E722 - pass - self.stds = {} - for key in self.data.keys(): - try: - self.stds[key] = self.data[key].std() - except: # noqa: E722 - pass - self.state_dim = env.observation_space.shape[0] - self.action_dim = env.action_space.shape[0] - - def normalize(self, x_in, key): - return (x_in - self.means[key]) / self.stds[key] - - def de_normalize(self, x_in, key): - return x_in * self.stds[key] + self.means[key] - - def to_torch(self, x_in): - if isinstance(x_in, dict): - return {k: self.to_torch(v) for k, v in x_in.items()} - elif torch.is_tensor(x_in): - return x_in.to(self.unet.device) - return torch.tensor(x_in, device=self.unet.device) - - def reset_x0(self, x_in, cond, act_dim): - for key, val in cond.items(): - x_in[:, key, act_dim:] = val.clone() - return x_in - - def run_diffusion(self, x, conditions, n_guide_steps, scale): - batch_size = x.shape[0] - y = None - for i in tqdm.tqdm(self.scheduler.timesteps): - # create batch of timesteps to pass into model - timesteps = torch.full((batch_size,), i, device=self.unet.device, dtype=torch.long) - for _ in range(n_guide_steps): - with torch.enable_grad(): - x.requires_grad_() - - # permute to match dimension for pre-trained models - y = self.value_function(x.permute(0, 2, 1), timesteps).sample - grad = torch.autograd.grad([y.sum()], [x])[0] - - posterior_variance = self.scheduler._get_variance(i) - model_std = torch.exp(0.5 * posterior_variance) - grad = model_std * grad - - grad[timesteps < 2] = 0 - x = x.detach() - x = x + scale * grad - x = self.reset_x0(x, conditions, self.action_dim) - - prev_x = self.unet(x.permute(0, 2, 1), timesteps).sample.permute(0, 2, 1) - - # TODO: verify deprecation of this kwarg - x = self.scheduler.step(prev_x, i, x)["prev_sample"] - - # apply conditions to the trajectory (set the initial state) - x = self.reset_x0(x, conditions, self.action_dim) - x = self.to_torch(x) - return x, y - - def __call__(self, obs, batch_size=64, planning_horizon=32, n_guide_steps=2, scale=0.1): - # normalize the observations and create batch dimension - obs = self.normalize(obs, "observations") - obs = obs[None].repeat(batch_size, axis=0) - - conditions = {0: self.to_torch(obs)} - shape = (batch_size, planning_horizon, self.state_dim + self.action_dim) - - # generate initial noise and apply our conditions (to make the trajectories start at current state) - x1 = randn_tensor(shape, device=self.unet.device) - x = self.reset_x0(x1, conditions, self.action_dim) - x = self.to_torch(x) - - # run the diffusion process - x, y = self.run_diffusion(x, conditions, n_guide_steps, scale) - - # sort output trajectories by value - sorted_idx = y.argsort(0, descending=True).squeeze() - sorted_values = x[sorted_idx] - actions = sorted_values[:, :, : self.action_dim] - actions = actions.detach().cpu().numpy() - denorm_actions = self.de_normalize(actions, key="actions") - - # select the action with the highest value - if y is not None: - selected_index = 0 - else: - # if we didn't run value guiding, select a random action - selected_index = np.random.randint(0, batch_size) - - denorm_actions = denorm_actions[selected_index, 0] - return denorm_actions diff --git a/diffusers/guiders/__init__.py b/diffusers/guiders/__init__.py deleted file mode 100644 index 88fae37f5d0096ffa03ed553557893725a180ecd..0000000000000000000000000000000000000000 --- a/diffusers/guiders/__init__.py +++ /dev/null @@ -1,31 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from ..utils import is_torch_available, logging - - -if is_torch_available(): - from .adaptive_projected_guidance import AdaptiveProjectedGuidance - from .adaptive_projected_guidance_mix import AdaptiveProjectedMixGuidance - from .auto_guidance import AutoGuidance - from .classifier_free_guidance import ClassifierFreeGuidance - from .classifier_free_zero_star_guidance import ClassifierFreeZeroStarGuidance - from .frequency_decoupled_guidance import FrequencyDecoupledGuidance - from .guider_utils import BaseGuidance - from .magnitude_aware_guidance import MagnitudeAwareGuidance - from .perturbed_attention_guidance import PerturbedAttentionGuidance - from .skip_layer_guidance import SkipLayerGuidance - from .smoothed_energy_guidance import SmoothedEnergyGuidance - from .tangential_classifier_free_guidance import TangentialClassifierFreeGuidance diff --git a/diffusers/guiders/adaptive_projected_guidance.py b/diffusers/guiders/adaptive_projected_guidance.py deleted file mode 100644 index dd6675fcb1901d1c11d8e7d7116c1ae09a5400d6..0000000000000000000000000000000000000000 --- a/diffusers/guiders/adaptive_projected_guidance.py +++ /dev/null @@ -1,253 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import math -from typing import TYPE_CHECKING - -import torch - -from ..configuration_utils import register_to_config -from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -class AdaptiveProjectedGuidance(BaseGuidance): - """ - Adaptive Projected Guidance (APG): https://huggingface.co/papers/2410.02416 - - Args: - guidance_scale (`float`, defaults to `7.5`): - The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text - prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and - deterioration of image quality. - adaptive_projected_guidance_momentum (`float`, defaults to `None`): - The momentum parameter for the adaptive projected guidance. Disabled if set to `None`. - adaptive_projected_guidance_rescale (`float`, defaults to `15.0`): - The rescale factor applied to the noise predictions. This is used to improve image quality and fix - adaptive_projected_guidance_norm_dim (`int` or `tuple[int]`, *optional*): - Dimension(s) over which to compute the APG norm and projection. If omitted, all non-batch dimensions are - used, preserving the original behavior. - guidance_rescale (`float`, defaults to `0.0`): - The rescale factor applied to the noise predictions. This is used to improve image quality and fix - overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are - Flawed](https://huggingface.co/papers/2305.08891). - use_original_formulation (`bool`, defaults to `False`): - Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default, - we use the diffusers-native implementation that has been in the codebase for a long time. See - [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details. - start (`float`, defaults to `0.0`): - The fraction of the total number of denoising steps after which guidance starts. - stop (`float`, defaults to `1.0`): - The fraction of the total number of denoising steps after which guidance stops. - """ - - _input_predictions = ["pred_cond", "pred_uncond"] - - @register_to_config - def __init__( - self, - guidance_scale: float = 7.5, - adaptive_projected_guidance_momentum: float | None = None, - adaptive_projected_guidance_rescale: float = 15.0, - adaptive_projected_guidance_norm_dim: int | tuple[int, ...] | None = None, - eta: float = 1.0, - guidance_rescale: float = 0.0, - use_original_formulation: bool = False, - start: float = 0.0, - stop: float = 1.0, - enabled: bool = True, - ): - super().__init__(start, stop, enabled) - - self.guidance_scale = guidance_scale - self.adaptive_projected_guidance_momentum = adaptive_projected_guidance_momentum - self.adaptive_projected_guidance_rescale = adaptive_projected_guidance_rescale - self.adaptive_projected_guidance_norm_dim = adaptive_projected_guidance_norm_dim - self.eta = eta - self.guidance_rescale = guidance_rescale - self.use_original_formulation = use_original_formulation - self.momentum_buffer = None - - def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: - if self._step == 0: - if self.adaptive_projected_guidance_momentum is not None: - self.momentum_buffer = MomentumBuffer(self.adaptive_projected_guidance_momentum) - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch(data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - if self._step == 0: - if self.adaptive_projected_guidance_momentum is not None: - self.momentum_buffer = MomentumBuffer(self.adaptive_projected_guidance_momentum) - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput: - pred = None - - if not self._is_apg_enabled(): - pred = pred_cond - else: - pred = normalized_guidance( - pred_cond, - pred_uncond, - self.guidance_scale, - self.momentum_buffer, - self.eta, - self.adaptive_projected_guidance_rescale, - self.use_original_formulation, - self.adaptive_projected_guidance_norm_dim, - ) - - if self.guidance_rescale > 0.0: - pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale) - - return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) - - @property - def is_conditional(self) -> bool: - return self._count_prepared == 1 - - @property - def num_conditions(self) -> int: - num_conditions = 1 - if self._is_apg_enabled(): - num_conditions += 1 - return num_conditions - - def _is_apg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self._start * self._num_inference_steps) - skip_stop_step = int(self._stop * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.guidance_scale, 0.0) - else: - is_close = math.isclose(self.guidance_scale, 1.0) - - return is_within_range and not is_close - - -class MomentumBuffer: - def __init__(self, momentum: float): - self.momentum = momentum - self.running_average = 0 - - def update(self, update_value: torch.Tensor): - new_average = self.momentum * self.running_average - self.running_average = update_value + new_average - - def __repr__(self) -> str: - """ - Returns a string representation showing momentum, shape, statistics, and a slice of the running_average. - """ - if isinstance(self.running_average, torch.Tensor): - shape = tuple(self.running_average.shape) - - # Calculate statistics - with torch.no_grad(): - stats = { - "mean": self.running_average.mean().item(), - "std": self.running_average.std().item(), - "min": self.running_average.min().item(), - "max": self.running_average.max().item(), - } - - # Get a slice (max 3 elements per dimension) - slice_indices = tuple(slice(None, min(3, dim)) for dim in shape) - sliced_data = self.running_average[slice_indices] - - # Format the slice for display (convert to float32 for numpy compatibility with bfloat16) - slice_str = str(sliced_data.detach().float().cpu().numpy()) - if len(slice_str) > 200: # Truncate if too long - slice_str = slice_str[:200] + "..." - - stats_str = ", ".join([f"{k}={v:.4f}" for k, v in stats.items()]) - - return ( - f"MomentumBuffer(\n" - f" momentum={self.momentum},\n" - f" shape={shape},\n" - f" stats=[{stats_str}],\n" - f" slice={slice_str}\n" - f")" - ) - else: - return f"MomentumBuffer(momentum={self.momentum}, running_average={self.running_average})" - - -def normalized_guidance( - pred_cond: torch.Tensor, - pred_uncond: torch.Tensor, - guidance_scale: float, - momentum_buffer: MomentumBuffer | None = None, - eta: float = 1.0, - norm_threshold: float = 0.0, - use_original_formulation: bool = False, - norm_dim: int | tuple[int, ...] | None = None, -): - diff = pred_cond - pred_uncond - if norm_dim is None: - dim = [-i for i in range(1, len(diff.shape))] - elif isinstance(norm_dim, int): - dim = [norm_dim] - else: - dim = list(norm_dim) - - if momentum_buffer is not None: - momentum_buffer.update(diff) - diff = momentum_buffer.running_average - - if norm_threshold > 0: - ones = torch.ones_like(diff) - diff_norm = diff.norm(p=2, dim=dim, keepdim=True) - scale_factor = torch.minimum(ones, norm_threshold / diff_norm) - diff = diff * scale_factor - - if diff.device.type in {"mps", "npu"}: - v0, v1 = diff.cpu().double(), pred_cond.cpu().double() - else: - v0, v1 = diff.double(), pred_cond.double() - v1 = torch.nn.functional.normalize(v1, dim=dim) - v0_parallel = (v0 * v1).sum(dim=dim, keepdim=True) * v1 - v0_orthogonal = v0 - v0_parallel - diff_parallel = v0_parallel.to(device=diff.device, dtype=diff.dtype) - diff_orthogonal = v0_orthogonal.to(device=diff.device, dtype=diff.dtype) - normalized_update = diff_orthogonal + eta * diff_parallel - - pred = pred_cond if use_original_formulation else pred_uncond - pred = pred + guidance_scale * normalized_update - - return pred diff --git a/diffusers/guiders/adaptive_projected_guidance_mix.py b/diffusers/guiders/adaptive_projected_guidance_mix.py deleted file mode 100644 index a44a49b61724d13402279fd2192f3fde11adae5c..0000000000000000000000000000000000000000 --- a/diffusers/guiders/adaptive_projected_guidance_mix.py +++ /dev/null @@ -1,297 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import TYPE_CHECKING - -import torch - -from ..configuration_utils import register_to_config -from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -class AdaptiveProjectedMixGuidance(BaseGuidance): - """ - Adaptive Projected Guidance (APG) https://huggingface.co/papers/2410.02416 combined with Classifier-Free Guidance - (CFG). This guider is used in HunyuanImage2.1 https://github.com/Tencent-Hunyuan/HunyuanImage-2.1 - - Args: - guidance_scale (`float`, defaults to `7.5`): - The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text - prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and - deterioration of image quality. - adaptive_projected_guidance_momentum (`float`, defaults to `None`): - The momentum parameter for the adaptive projected guidance. Disabled if set to `None`. - adaptive_projected_guidance_rescale (`float`, defaults to `15.0`): - The rescale factor applied to the noise predictions for adaptive projected guidance. This is used to - improve image quality and fix - guidance_rescale (`float`, defaults to `0.0`): - The rescale factor applied to the noise predictions for classifier-free guidance. This is used to improve - image quality and fix overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample - Steps are Flawed](https://huggingface.co/papers/2305.08891). - use_original_formulation (`bool`, defaults to `False`): - Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default, - we use the diffusers-native implementation that has been in the codebase for a long time. See - [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details. - start (`float`, defaults to `0.0`): - The fraction of the total number of denoising steps after which the classifier-free guidance starts. - stop (`float`, defaults to `1.0`): - The fraction of the total number of denoising steps after which the classifier-free guidance stops. - adaptive_projected_guidance_start_step (`int`, defaults to `5`): - The step at which the adaptive projected guidance starts (before this step, classifier-free guidance is - used, and momentum buffer is updated). - enabled (`bool`, defaults to `True`): - Whether this guidance is enabled. - """ - - _input_predictions = ["pred_cond", "pred_uncond"] - - @register_to_config - def __init__( - self, - guidance_scale: float = 3.5, - guidance_rescale: float = 0.0, - adaptive_projected_guidance_scale: float = 10.0, - adaptive_projected_guidance_momentum: float = -0.5, - adaptive_projected_guidance_rescale: float = 10.0, - eta: float = 0.0, - use_original_formulation: bool = False, - start: float = 0.0, - stop: float = 1.0, - adaptive_projected_guidance_start_step: int = 5, - enabled: bool = True, - ): - super().__init__(start, stop, enabled) - - self.guidance_scale = guidance_scale - self.guidance_rescale = guidance_rescale - self.adaptive_projected_guidance_scale = adaptive_projected_guidance_scale - self.adaptive_projected_guidance_momentum = adaptive_projected_guidance_momentum - self.adaptive_projected_guidance_rescale = adaptive_projected_guidance_rescale - self.eta = eta - self.adaptive_projected_guidance_start_step = adaptive_projected_guidance_start_step - self.use_original_formulation = use_original_formulation - self.momentum_buffer = None - - def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: - if self._step == 0: - if self.adaptive_projected_guidance_momentum is not None: - self.momentum_buffer = MomentumBuffer(self.adaptive_projected_guidance_momentum) - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch(data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - if self._step == 0: - if self.adaptive_projected_guidance_momentum is not None: - self.momentum_buffer = MomentumBuffer(self.adaptive_projected_guidance_momentum) - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput: - pred = None - - # no guidance - if not self._is_cfg_enabled(): - pred = pred_cond - - # CFG + update momentum buffer - elif not self._is_apg_enabled(): - if self.momentum_buffer is not None: - update_momentum_buffer(pred_cond, pred_uncond, self.momentum_buffer) - # CFG + update momentum buffer - shift = pred_cond - pred_uncond - pred = pred_cond if self.use_original_formulation else pred_uncond - pred = pred + self.guidance_scale * shift - - # APG - elif self._is_apg_enabled(): - pred = normalized_guidance( - pred_cond, - pred_uncond, - self.adaptive_projected_guidance_scale, - self.momentum_buffer, - self.eta, - self.adaptive_projected_guidance_rescale, - self.use_original_formulation, - ) - - if self.guidance_rescale > 0.0: - pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale) - - return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) - - @property - def is_conditional(self) -> bool: - return self._count_prepared == 1 - - @property - def num_conditions(self) -> int: - num_conditions = 1 - if self._is_apg_enabled() or self._is_cfg_enabled(): - num_conditions += 1 - return num_conditions - - # Copied from diffusers.guiders.classifier_free_guidance.ClassifierFreeGuidance._is_cfg_enabled - def _is_cfg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self._start * self._num_inference_steps) - skip_stop_step = int(self._stop * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.guidance_scale, 0.0) - else: - is_close = math.isclose(self.guidance_scale, 1.0) - - return is_within_range and not is_close - - def _is_apg_enabled(self) -> bool: - if not self._enabled: - return False - - if not self._is_cfg_enabled(): - return False - - is_within_range = False - if self._step is not None: - is_within_range = self._step > self.adaptive_projected_guidance_start_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.adaptive_projected_guidance_scale, 0.0) - else: - is_close = math.isclose(self.adaptive_projected_guidance_scale, 1.0) - - return is_within_range and not is_close - - def get_state(self): - state = super().get_state() - state["momentum_buffer"] = self.momentum_buffer - state["is_apg_enabled"] = self._is_apg_enabled() - state["is_cfg_enabled"] = self._is_cfg_enabled() - return state - - -# Copied from diffusers.guiders.adaptive_projected_guidance.MomentumBuffer -class MomentumBuffer: - def __init__(self, momentum: float): - self.momentum = momentum - self.running_average = 0 - - def update(self, update_value: torch.Tensor): - new_average = self.momentum * self.running_average - self.running_average = update_value + new_average - - def __repr__(self) -> str: - """ - Returns a string representation showing momentum, shape, statistics, and a slice of the running_average. - """ - if isinstance(self.running_average, torch.Tensor): - shape = tuple(self.running_average.shape) - - # Calculate statistics - with torch.no_grad(): - stats = { - "mean": self.running_average.mean().item(), - "std": self.running_average.std().item(), - "min": self.running_average.min().item(), - "max": self.running_average.max().item(), - } - - # Get a slice (max 3 elements per dimension) - slice_indices = tuple(slice(None, min(3, dim)) for dim in shape) - sliced_data = self.running_average[slice_indices] - - # Format the slice for display (convert to float32 for numpy compatibility with bfloat16) - slice_str = str(sliced_data.detach().float().cpu().numpy()) - if len(slice_str) > 200: # Truncate if too long - slice_str = slice_str[:200] + "..." - - stats_str = ", ".join([f"{k}={v:.4f}" for k, v in stats.items()]) - - return ( - f"MomentumBuffer(\n" - f" momentum={self.momentum},\n" - f" shape={shape},\n" - f" stats=[{stats_str}],\n" - f" slice={slice_str}\n" - f")" - ) - else: - return f"MomentumBuffer(momentum={self.momentum}, running_average={self.running_average})" - - -def update_momentum_buffer( - pred_cond: torch.Tensor, - pred_uncond: torch.Tensor, - momentum_buffer: MomentumBuffer | None = None, -): - diff = pred_cond - pred_uncond - if momentum_buffer is not None: - momentum_buffer.update(diff) - - -def normalized_guidance( - pred_cond: torch.Tensor, - pred_uncond: torch.Tensor, - guidance_scale: float, - momentum_buffer: MomentumBuffer | None = None, - eta: float = 1.0, - norm_threshold: float = 0.0, - use_original_formulation: bool = False, -): - if momentum_buffer is not None: - update_momentum_buffer(pred_cond, pred_uncond, momentum_buffer) - diff = momentum_buffer.running_average - else: - diff = pred_cond - pred_uncond - - dim = [-i for i in range(1, len(diff.shape))] - - if norm_threshold > 0: - ones = torch.ones_like(diff) - diff_norm = diff.norm(p=2, dim=dim, keepdim=True) - scale_factor = torch.minimum(ones, norm_threshold / diff_norm) - diff = diff * scale_factor - - v0, v1 = diff.double(), pred_cond.double() - v1 = torch.nn.functional.normalize(v1, dim=dim) - v0_parallel = (v0 * v1).sum(dim=dim, keepdim=True) * v1 - v0_orthogonal = v0 - v0_parallel - diff_parallel, diff_orthogonal = v0_parallel.type_as(diff), v0_orthogonal.type_as(diff) - normalized_update = diff_orthogonal + eta * diff_parallel - - pred = pred_cond if use_original_formulation else pred_uncond - pred = pred + guidance_scale * normalized_update - - return pred diff --git a/diffusers/guiders/auto_guidance.py b/diffusers/guiders/auto_guidance.py deleted file mode 100644 index aaea0784b46f3645b109631ab5aa0f1930c47971..0000000000000000000000000000000000000000 --- a/diffusers/guiders/auto_guidance.py +++ /dev/null @@ -1,198 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import math -from typing import TYPE_CHECKING, Any - -import torch - -from ..configuration_utils import register_to_config -from ..hooks import HookRegistry, LayerSkipConfig -from ..hooks.layer_skip import _apply_layer_skip_hook -from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -class AutoGuidance(BaseGuidance): - """ - AutoGuidance: https://huggingface.co/papers/2406.02507 - - Args: - guidance_scale (`float`, defaults to `7.5`): - The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text - prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and - deterioration of image quality. - auto_guidance_layers (`int` or `list[int]`, *optional*): - The layer indices to apply skip layer guidance to. Can be a single integer or a list of integers. If not - provided, `skip_layer_config` must be provided. - auto_guidance_config (`LayerSkipConfig` or `list[LayerSkipConfig]`, *optional*): - The configuration for the skip layer guidance. Can be a single `LayerSkipConfig` or a list of - `LayerSkipConfig`. If not provided, `skip_layer_guidance_layers` must be provided. - dropout (`float`, *optional*): - The dropout probability for autoguidance on the enabled skip layers (either with `auto_guidance_layers` or - `auto_guidance_config`). If not provided, the dropout probability will be set to 1.0. - guidance_rescale (`float`, defaults to `0.0`): - The rescale factor applied to the noise predictions. This is used to improve image quality and fix - overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are - Flawed](https://huggingface.co/papers/2305.08891). - use_original_formulation (`bool`, defaults to `False`): - Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default, - we use the diffusers-native implementation that has been in the codebase for a long time. See - [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details. - start (`float`, defaults to `0.0`): - The fraction of the total number of denoising steps after which guidance starts. - stop (`float`, defaults to `1.0`): - The fraction of the total number of denoising steps after which guidance stops. - """ - - _input_predictions = ["pred_cond", "pred_uncond"] - - @register_to_config - def __init__( - self, - guidance_scale: float = 7.5, - auto_guidance_layers: int | list[int] | None = None, - auto_guidance_config: LayerSkipConfig | list[LayerSkipConfig] | dict[str, Any] = None, - dropout: float | None = None, - guidance_rescale: float = 0.0, - use_original_formulation: bool = False, - start: float = 0.0, - stop: float = 1.0, - enabled: bool = True, - ): - super().__init__(start, stop, enabled) - - self.guidance_scale = guidance_scale - self.auto_guidance_layers = auto_guidance_layers - self.auto_guidance_config = auto_guidance_config - self.dropout = dropout - self.guidance_rescale = guidance_rescale - self.use_original_formulation = use_original_formulation - - is_layer_or_config_provided = auto_guidance_layers is not None or auto_guidance_config is not None - is_layer_and_config_provided = auto_guidance_layers is not None and auto_guidance_config is not None - if not is_layer_or_config_provided: - raise ValueError( - "Either `auto_guidance_layers` or `auto_guidance_config` must be provided to enable AutoGuidance." - ) - if is_layer_and_config_provided: - raise ValueError("Only one of `auto_guidance_layers` or `auto_guidance_config` can be provided.") - if auto_guidance_config is None and dropout is None: - raise ValueError("`dropout` must be provided if `auto_guidance_layers` is provided.") - - if auto_guidance_layers is not None: - if isinstance(auto_guidance_layers, int): - auto_guidance_layers = [auto_guidance_layers] - if not isinstance(auto_guidance_layers, list): - raise ValueError( - f"Expected `auto_guidance_layers` to be an int or a list of ints, but got {type(auto_guidance_layers)}." - ) - auto_guidance_config = [ - LayerSkipConfig(layer, fqn="auto", dropout=dropout) for layer in auto_guidance_layers - ] - - if isinstance(auto_guidance_config, dict): - auto_guidance_config = LayerSkipConfig.from_dict(auto_guidance_config) - - if isinstance(auto_guidance_config, LayerSkipConfig): - auto_guidance_config = [auto_guidance_config] - - if not isinstance(auto_guidance_config, list): - raise ValueError( - f"Expected `auto_guidance_config` to be a LayerSkipConfig or a list of LayerSkipConfig, but got {type(auto_guidance_config)}." - ) - elif isinstance(next(iter(auto_guidance_config), None), dict): - auto_guidance_config = [LayerSkipConfig.from_dict(config) for config in auto_guidance_config] - - self.auto_guidance_config = auto_guidance_config - self._auto_guidance_hook_names = [f"AutoGuidance_{i}" for i in range(len(self.auto_guidance_config))] - - def prepare_models(self, denoiser: torch.nn.Module) -> None: - self._count_prepared += 1 - if self._is_ag_enabled() and self.is_unconditional: - for name, config in zip(self._auto_guidance_hook_names, self.auto_guidance_config): - _apply_layer_skip_hook(denoiser, config, name=name) - - def cleanup_models(self, denoiser: torch.nn.Module) -> None: - if self._is_ag_enabled() and self.is_unconditional: - for name in self._auto_guidance_hook_names: - registry = HookRegistry.check_if_exists_or_initialize(denoiser) - registry.remove_hook(name, recurse=True) - - def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch(data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput: - pred = None - - if not self._is_ag_enabled(): - pred = pred_cond - else: - shift = pred_cond - pred_uncond - pred = pred_cond if self.use_original_formulation else pred_uncond - pred = pred + self.guidance_scale * shift - - if self.guidance_rescale > 0.0: - pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale) - - return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) - - @property - def is_conditional(self) -> bool: - return self._count_prepared == 1 - - @property - def num_conditions(self) -> int: - num_conditions = 1 - if self._is_ag_enabled(): - num_conditions += 1 - return num_conditions - - def _is_ag_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self._start * self._num_inference_steps) - skip_stop_step = int(self._stop * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.guidance_scale, 0.0) - else: - is_close = math.isclose(self.guidance_scale, 1.0) - - return is_within_range and not is_close diff --git a/diffusers/guiders/classifier_free_guidance.py b/diffusers/guiders/classifier_free_guidance.py deleted file mode 100644 index a669f61b465286ead814b2074968865a874ff909..0000000000000000000000000000000000000000 --- a/diffusers/guiders/classifier_free_guidance.py +++ /dev/null @@ -1,156 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import math -from typing import TYPE_CHECKING - -import torch - -from ..configuration_utils import register_to_config -from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -class ClassifierFreeGuidance(BaseGuidance): - """ - Implements Classifier-Free Guidance (CFG) for diffusion models. - - Reference: https://huggingface.co/papers/2207.12598 - - CFG improves generation quality and prompt adherence by jointly training models on both conditional and - unconditional data, then combining predictions during inference. This allows trading off between quality (high - guidance) and diversity (low guidance). - - **Two CFG Formulations:** - - 1. **Original formulation** (from paper): - ``` - x_pred = x_cond + guidance_scale * (x_cond - x_uncond) - ``` - Moves conditional predictions further from unconditional ones. - - 2. **Diffusers-native formulation** (default, from Imagen paper): - ``` - x_pred = x_uncond + guidance_scale * (x_cond - x_uncond) - ``` - Moves unconditional predictions toward conditional ones, effectively suppressing negative features (e.g., "bad - quality", "watermarks"). Equivalent in theory but more intuitive. - - Use `use_original_formulation=True` to switch to the original formulation. - - Args: - guidance_scale (`float`, defaults to `7.5`): - CFG scale applied by this guider during post-processing. Higher values = stronger prompt conditioning but - may reduce quality. Typical range: 1.0-20.0. - guidance_rescale (`float`, defaults to `0.0`): - Rescaling factor to prevent overexposure from high guidance scales. Based on [Common Diffusion Noise - Schedules and Sample Steps are Flawed](https://huggingface.co/papers/2305.08891). Range: 0.0 (no rescaling) - to 1.0 (full rescaling). - use_original_formulation (`bool`, defaults to `False`): - If `True`, uses the original CFG formulation from the paper. If `False` (default), uses the - diffusers-native formulation from the Imagen paper. - start (`float`, defaults to `0.0`): - Fraction of denoising steps (0.0-1.0) after which CFG starts. Use > 0.0 to disable CFG in early denoising - steps. - stop (`float`, defaults to `1.0`): - Fraction of denoising steps (0.0-1.0) after which CFG stops. Use < 1.0 to disable CFG in late denoising - steps. - enabled (`bool`, defaults to `True`): - Whether CFG is enabled. Set to `False` to disable CFG entirely (uses only conditional predictions). - """ - - _input_predictions = ["pred_cond", "pred_uncond"] - - @register_to_config - def __init__( - self, - guidance_scale: float = 7.5, - guidance_rescale: float = 0.0, - use_original_formulation: bool = False, - start: float = 0.0, - stop: float = 1.0, - enabled: bool = True, - ): - super().__init__(start, stop, enabled) - - self.guidance_scale = guidance_scale - self.guidance_rescale = guidance_rescale - self.use_original_formulation = use_original_formulation - - def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch(data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput: - pred = None - - if not self._is_cfg_enabled(): - pred = pred_cond - else: - shift = pred_cond - pred_uncond - pred = pred_cond if self.use_original_formulation else pred_uncond - pred = pred + self.guidance_scale * shift - - if self.guidance_rescale > 0.0: - pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale) - - return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) - - @property - def is_conditional(self) -> bool: - return self._count_prepared == 1 - - @property - def num_conditions(self) -> int: - num_conditions = 1 - if self._is_cfg_enabled(): - num_conditions += 1 - return num_conditions - - def _is_cfg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self._start * self._num_inference_steps) - skip_stop_step = int(self._stop * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.guidance_scale, 0.0) - else: - is_close = math.isclose(self.guidance_scale, 1.0) - - return is_within_range and not is_close diff --git a/diffusers/guiders/classifier_free_zero_star_guidance.py b/diffusers/guiders/classifier_free_zero_star_guidance.py deleted file mode 100644 index 83a31881ea07a66d5e98dc55cdfa4d89826f40b5..0000000000000000000000000000000000000000 --- a/diffusers/guiders/classifier_free_zero_star_guidance.py +++ /dev/null @@ -1,164 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import math -from typing import TYPE_CHECKING - -import torch - -from ..configuration_utils import register_to_config -from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -class ClassifierFreeZeroStarGuidance(BaseGuidance): - """ - Classifier-free Zero* (CFG-Zero*): https://huggingface.co/papers/2503.18886 - - This is an implementation of the Classifier-Free Zero* guidance technique, which is a variant of classifier-free - guidance. It proposes zero initialization of the noise predictions for the first few steps of the diffusion - process, and also introduces an optimal rescaling factor for the noise predictions, which can help in improving the - quality of generated images. - - The authors of the paper suggest setting zero initialization in the first 4% of the inference steps. - - Args: - guidance_scale (`float`, defaults to `7.5`): - The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text - prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and - deterioration of image quality. - zero_init_steps (`int`, defaults to `1`): - The number of inference steps for which the noise predictions are zeroed out (see Section 4.2). - guidance_rescale (`float`, defaults to `0.0`): - The rescale factor applied to the noise predictions. This is used to improve image quality and fix - overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are - Flawed](https://huggingface.co/papers/2305.08891). - use_original_formulation (`bool`, defaults to `False`): - Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default, - we use the diffusers-native implementation that has been in the codebase for a long time. See - [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details. - start (`float`, defaults to `0.01`): - The fraction of the total number of denoising steps after which guidance starts. - stop (`float`, defaults to `0.2`): - The fraction of the total number of denoising steps after which guidance stops. - """ - - _input_predictions = ["pred_cond", "pred_uncond"] - - @register_to_config - def __init__( - self, - guidance_scale: float = 7.5, - zero_init_steps: int = 1, - guidance_rescale: float = 0.0, - use_original_formulation: bool = False, - start: float = 0.0, - stop: float = 1.0, - enabled: bool = True, - ): - super().__init__(start, stop, enabled) - - self.guidance_scale = guidance_scale - self.zero_init_steps = zero_init_steps - self.guidance_rescale = guidance_rescale - self.use_original_formulation = use_original_formulation - - def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch(data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput: - pred = None - - # YiYi Notes: add default behavior for self._enabled == False - if not self._enabled: - pred = pred_cond - - elif self._step < self.zero_init_steps: - pred = torch.zeros_like(pred_cond) - elif not self._is_cfg_enabled(): - pred = pred_cond - else: - pred_cond_flat = pred_cond.flatten(1) - pred_uncond_flat = pred_uncond.flatten(1) - alpha = cfg_zero_star_scale(pred_cond_flat, pred_uncond_flat) - alpha = alpha.view(-1, *(1,) * (len(pred_cond.shape) - 1)) - pred_uncond = pred_uncond * alpha - shift = pred_cond - pred_uncond - pred = pred_cond if self.use_original_formulation else pred_uncond - pred = pred + self.guidance_scale * shift - - if self.guidance_rescale > 0.0: - pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale) - - return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) - - @property - def is_conditional(self) -> bool: - return self._count_prepared == 1 - - @property - def num_conditions(self) -> int: - num_conditions = 1 - if self._is_cfg_enabled(): - num_conditions += 1 - return num_conditions - - def _is_cfg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self._start * self._num_inference_steps) - skip_stop_step = int(self._stop * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.guidance_scale, 0.0) - else: - is_close = math.isclose(self.guidance_scale, 1.0) - - return is_within_range and not is_close - - -def cfg_zero_star_scale(cond: torch.Tensor, uncond: torch.Tensor, eps: float = 1e-8) -> torch.Tensor: - cond_dtype = cond.dtype - cond = cond.float() - uncond = uncond.float() - dot_product = torch.sum(cond * uncond, dim=1, keepdim=True) - squared_norm = torch.sum(uncond**2, dim=1, keepdim=True) + eps - # st_star = v_cond^T * v_uncond / ||v_uncond||^2 - scale = dot_product / squared_norm - return scale.to(dtype=cond_dtype) diff --git a/diffusers/guiders/frequency_decoupled_guidance.py b/diffusers/guiders/frequency_decoupled_guidance.py deleted file mode 100644 index f1786d0e603dfa399e7b8ffcced486f17c8bdcdf..0000000000000000000000000000000000000000 --- a/diffusers/guiders/frequency_decoupled_guidance.py +++ /dev/null @@ -1,335 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import math -from typing import TYPE_CHECKING - -import torch - -from ..configuration_utils import register_to_config -from ..utils import is_kornia_available -from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -_CAN_USE_KORNIA = is_kornia_available() - - -if _CAN_USE_KORNIA: - from kornia.geometry import pyrup as upsample_and_blur_func - from kornia.geometry.transform import build_laplacian_pyramid as build_laplacian_pyramid_func -else: - upsample_and_blur_func = None - build_laplacian_pyramid_func = None - - -def project(v0: torch.Tensor, v1: torch.Tensor, upcast_to_double: bool = True) -> tuple[torch.Tensor, torch.Tensor]: - """ - Project vector v0 onto vector v1, returning the parallel and orthogonal components of v0. Implementation from paper - (Algorithm 2). - """ - # v0 shape: [B, ...] - # v1 shape: [B, ...] - # Assume first dim is a batch dim and all other dims are channel or "spatial" dims - all_dims_but_first = list(range(1, len(v0.shape))) - if upcast_to_double: - dtype = v0.dtype - v0, v1 = v0.double(), v1.double() - v1 = torch.nn.functional.normalize(v1, dim=all_dims_but_first) - v0_parallel = (v0 * v1).sum(dim=all_dims_but_first, keepdim=True) * v1 - v0_orthogonal = v0 - v0_parallel - if upcast_to_double: - v0_parallel = v0_parallel.to(dtype) - v0_orthogonal = v0_orthogonal.to(dtype) - return v0_parallel, v0_orthogonal - - -def build_image_from_pyramid(pyramid: list[torch.Tensor]) -> torch.Tensor: - """ - Recovers the data space latents from the Laplacian pyramid frequency space. Implementation from the paper - (Algorithm 2). - """ - # pyramid shapes: [[B, C, H, W], [B, C, H/2, W/2], ...] - img = pyramid[-1] - for i in range(len(pyramid) - 2, -1, -1): - img = upsample_and_blur_func(img) + pyramid[i] - return img - - -class FrequencyDecoupledGuidance(BaseGuidance): - """ - Frequency-Decoupled Guidance (FDG): https://huggingface.co/papers/2506.19713 - - FDG is a technique similar to (and based on) classifier-free guidance (CFG) which is used to improve generation - quality and condition-following in diffusion models. Like CFG, during training we jointly train the model on both - conditional and unconditional data, and use a combination of the two during inference. (If you want more details on - how CFG works, you can check out the CFG guider.) - - FDG differs from CFG in that the normal CFG prediction is instead decoupled into low- and high-frequency components - using a frequency transform (such as a Laplacian pyramid). The CFG update is then performed in frequency space - separately for the low- and high-frequency components with different guidance scales. Finally, the inverse - frequency transform is used to map the CFG frequency predictions back to data space (e.g. pixel space for images) - to form the final FDG prediction. - - For images, the FDG authors found that using low guidance scales for the low-frequency components retains sample - diversity and realistic color composition, while using high guidance scales for high-frequency components enhances - sample quality (such as better visual details). Therefore, they recommend using low guidance scales (low w_low) for - the low-frequency components and high guidance scales (high w_high) for the high-frequency components. As an - example, they suggest w_low = 5.0 and w_high = 10.0 for Stable Diffusion XL (see Table 8 in the paper). - - As with CFG, Diffusers implements the scaling and shifting on the unconditional prediction based on the [Imagen - paper](https://huggingface.co/papers/2205.11487), which is equivalent to what the original CFG paper proposed in - theory. [x_pred = x_uncond + scale * (x_cond - x_uncond)] - - The `use_original_formulation` argument can be set to `True` to use the original CFG formulation mentioned in the - paper. By default, we use the diffusers-native implementation that has been in the codebase for a long time. - - Args: - guidance_scales (`list[float]`, defaults to `[10.0, 5.0]`): - The scale parameter for frequency-decoupled guidance for each frequency component, listed from highest - frequency level to lowest. Higher values result in stronger conditioning on the text prompt, while lower - values allow for more freedom in generation. Higher values may lead to saturation and deterioration of - image quality. The FDG authors recommend using higher guidance scales for higher frequency components and - lower guidance scales for lower frequency components (so `guidance_scales` should typically be sorted in - descending order). - guidance_rescale (`float` or `list[float]`, defaults to `0.0`): - The rescale factor applied to the noise predictions. This is used to improve image quality and fix - overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are - Flawed](https://huggingface.co/papers/2305.08891). If a list is supplied, it should be the same length as - `guidance_scales`. - parallel_weights (`float` or `list[float]`, *optional*): - Optional weights for the parallel component of each frequency component of the projected CFG shift. If not - set, the weights will default to `1.0` for all components, which corresponds to using the normal CFG shift - (that is, equal weights for the parallel and orthogonal components). If set, a value in `[0, 1]` is - recommended. If a list is supplied, it should be the same length as `guidance_scales`. - use_original_formulation (`bool`, defaults to `False`): - Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default, - we use the diffusers-native implementation that has been in the codebase for a long time. See - [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details. - start (`float` or `list[float]`, defaults to `0.0`): - The fraction of the total number of denoising steps after which guidance starts. If a list is supplied, it - should be the same length as `guidance_scales`. - stop (`float` or `list[float]`, defaults to `1.0`): - The fraction of the total number of denoising steps after which guidance stops. If a list is supplied, it - should be the same length as `guidance_scales`. - guidance_rescale_space (`str`, defaults to `"data"`): - Whether to performance guidance rescaling in `"data"` space (after the full FDG update in data space) or in - `"freq"` space (right after the CFG update, for each freq level). Note that frequency space rescaling is - speculative and may not produce expected results. If `"data"` is set, the first `guidance_rescale` value - will be used; otherwise, per-frequency-level guidance rescale values will be used if available. - upcast_to_double (`bool`, defaults to `True`): - Whether to upcast certain operations, such as the projection operation when using `parallel_weights`, to - float64 when performing guidance. This may result in better performance at the cost of increased runtime. - """ - - _input_predictions = ["pred_cond", "pred_uncond"] - - @register_to_config - def __init__( - self, - guidance_scales: list[float] | tuple[float] = [10.0, 5.0], - guidance_rescale: float | list[float] | tuple[float] = 0.0, - parallel_weights: float | list[float] | tuple[float] | None = None, - use_original_formulation: bool = False, - start: float | list[float] | tuple[float] = 0.0, - stop: float | list[float] | tuple[float] = 1.0, - guidance_rescale_space: str = "data", - upcast_to_double: bool = True, - enabled: bool = True, - ): - if not _CAN_USE_KORNIA: - raise ImportError( - "The `FrequencyDecoupledGuidance` guider cannot be instantiated because the `kornia` library on which " - "it depends is not available in the current environment. You can install `kornia` with `pip install " - "kornia`." - ) - - # Set start to earliest start for any freq component and stop to latest stop for any freq component - min_start = start if isinstance(start, float) else min(start) - max_stop = stop if isinstance(stop, float) else max(stop) - super().__init__(min_start, max_stop, enabled) - - self.guidance_scales = guidance_scales - self.levels = len(guidance_scales) - - if isinstance(guidance_rescale, float): - self.guidance_rescale = [guidance_rescale] * self.levels - elif len(guidance_rescale) == self.levels: - self.guidance_rescale = guidance_rescale - else: - raise ValueError( - f"`guidance_rescale` has length {len(guidance_rescale)} but should have the same length as " - f"`guidance_scales` ({len(self.guidance_scales)})" - ) - # Whether to perform guidance rescaling in frequency space (right after the CFG update) or data space (after - # transforming from frequency space back to data space) - if guidance_rescale_space not in ["data", "freq"]: - raise ValueError( - f"Guidance rescale space is {guidance_rescale_space} but must be one of `data` or `freq`." - ) - self.guidance_rescale_space = guidance_rescale_space - - if parallel_weights is None: - # Use normal CFG shift (equal weights for parallel and orthogonal components) - self.parallel_weights = [1.0] * self.levels - elif isinstance(parallel_weights, float): - self.parallel_weights = [parallel_weights] * self.levels - elif len(parallel_weights) == self.levels: - self.parallel_weights = parallel_weights - else: - raise ValueError( - f"`parallel_weights` has length {len(parallel_weights)} but should have the same length as " - f"`guidance_scales` ({len(self.guidance_scales)})" - ) - - self.use_original_formulation = use_original_formulation - self.upcast_to_double = upcast_to_double - - if isinstance(start, float): - self.guidance_start = [start] * self.levels - elif len(start) == self.levels: - self.guidance_start = start - else: - raise ValueError( - f"`start` has length {len(start)} but should have the same length as `guidance_scales` " - f"({len(self.guidance_scales)})" - ) - if isinstance(stop, float): - self.guidance_stop = [stop] * self.levels - elif len(stop) == self.levels: - self.guidance_stop = stop - else: - raise ValueError( - f"`stop` has length {len(stop)} but should have the same length as `guidance_scales` " - f"({len(self.guidance_scales)})" - ) - - def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch(data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput: - pred = None - - if not self._is_fdg_enabled(): - pred = pred_cond - else: - # Apply the frequency transform (e.g. Laplacian pyramid) to the conditional and unconditional predictions. - pred_cond_pyramid = build_laplacian_pyramid_func(pred_cond, self.levels) - pred_uncond_pyramid = build_laplacian_pyramid_func(pred_uncond, self.levels) - - # From high frequencies to low frequencies, following the paper implementation - pred_guided_pyramid = [] - parameters = zip(self.guidance_scales, self.parallel_weights, self.guidance_rescale) - for level, (guidance_scale, parallel_weight, guidance_rescale) in enumerate(parameters): - if self._is_fdg_enabled_for_level(level): - # Get the cond/uncond preds (in freq space) at the current frequency level - pred_cond_freq = pred_cond_pyramid[level] - pred_uncond_freq = pred_uncond_pyramid[level] - - shift = pred_cond_freq - pred_uncond_freq - - # Apply parallel weights, if used (1.0 corresponds to using the normal CFG shift) - if not math.isclose(parallel_weight, 1.0): - shift_parallel, shift_orthogonal = project(shift, pred_cond_freq, self.upcast_to_double) - shift = parallel_weight * shift_parallel + shift_orthogonal - - # Apply CFG update for the current frequency level - pred = pred_cond_freq if self.use_original_formulation else pred_uncond_freq - pred = pred + guidance_scale * shift - - if self.guidance_rescale_space == "freq" and guidance_rescale > 0.0: - pred = rescale_noise_cfg(pred, pred_cond_freq, guidance_rescale) - - # Add the current FDG guided level to the FDG prediction pyramid - pred_guided_pyramid.append(pred) - else: - # Add the current pred_cond_pyramid level as the "non-FDG" prediction - pred_guided_pyramid.append(pred_cond_freq) - - # Convert from frequency space back to data (e.g. pixel) space by applying inverse freq transform - pred = build_image_from_pyramid(pred_guided_pyramid) - - # If rescaling in data space, use the first elem of self.guidance_rescale as the "global" rescale value - # across all freq levels - if self.guidance_rescale_space == "data" and self.guidance_rescale[0] > 0.0: - pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale[0]) - - return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) - - @property - def is_conditional(self) -> bool: - return self._count_prepared == 1 - - @property - def num_conditions(self) -> int: - num_conditions = 1 - if self._is_fdg_enabled(): - num_conditions += 1 - return num_conditions - - def _is_fdg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self._start * self._num_inference_steps) - skip_stop_step = int(self._stop * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = all(math.isclose(guidance_scale, 0.0) for guidance_scale in self.guidance_scales) - else: - is_close = all(math.isclose(guidance_scale, 1.0) for guidance_scale in self.guidance_scales) - - return is_within_range and not is_close - - def _is_fdg_enabled_for_level(self, level: int) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self.guidance_start[level] * self._num_inference_steps) - skip_stop_step = int(self.guidance_stop[level] * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.guidance_scales[level], 0.0) - else: - is_close = math.isclose(self.guidance_scales[level], 1.0) - - return is_within_range and not is_close diff --git a/diffusers/guiders/guider_utils.py b/diffusers/guiders/guider_utils.py deleted file mode 100644 index 4af7abbe212ec265751680d9b885076b9c23f7b5..0000000000000000000000000000000000000000 --- a/diffusers/guiders/guider_utils.py +++ /dev/null @@ -1,396 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import os -from typing import TYPE_CHECKING, Any - -import torch -from huggingface_hub.utils import validate_hf_hub_args -from typing_extensions import Self - -from ..configuration_utils import ConfigMixin -from ..utils import BaseOutput, PushToHubMixin, get_logger - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -GUIDER_CONFIG_NAME = "guider_config.json" - - -logger = get_logger(__name__) # pylint: disable=invalid-name - - -class BaseGuidance(ConfigMixin, PushToHubMixin): - r"""Base class providing the skeleton for implementing guidance techniques.""" - - config_name = GUIDER_CONFIG_NAME - _input_predictions = None - _identifier_key = "__guidance_identifier__" - - def __init__(self, start: float = 0.0, stop: float = 1.0, enabled: bool = True): - logger.warning( - "Guiders are currently an experimental feature under active development. The API is subject to breaking changes in future releases." - ) - - self._start = start - self._stop = stop - self._step: int = None - self._num_inference_steps: int = None - self._timestep: torch.LongTensor = None - self._count_prepared = 0 - self._input_fields: dict[str, str | tuple[str, str]] = None - self._enabled = enabled - - if not (0.0 <= start < 1.0): - raise ValueError(f"Expected `start` to be between 0.0 and 1.0, but got {start}.") - if not (start <= stop <= 1.0): - raise ValueError(f"Expected `stop` to be between {start} and 1.0, but got {stop}.") - - if self._input_predictions is None or not isinstance(self._input_predictions, list): - raise ValueError( - "`_input_predictions` must be a list of required prediction names for the guidance technique." - ) - - def new(self, **kwargs): - """ - Creates a copy of this guider instance, optionally with modified configuration parameters. - - Args: - **kwargs: Configuration parameters to override in the new instance. If no kwargs are provided, - returns an exact copy with the same configuration. - - Returns: - A new guider instance with the same (or updated) configuration. - - Example: - ```python - # Create a CFG guider - guider = ClassifierFreeGuidance(guidance_scale=3.5) - - # Create an exact copy - same_guider = guider.new() - - # Create a copy with different start step, keeping other config the same - new_guider = guider.new(guidance_scale=5) - ``` - """ - return self.__class__.from_config(self.config, **kwargs) - - def disable(self): - self._enabled = False - - def enable(self): - self._enabled = True - - def set_state(self, step: int, num_inference_steps: int, timestep: torch.LongTensor) -> None: - self._step = step - self._num_inference_steps = num_inference_steps - self._timestep = timestep - self._count_prepared = 0 - - def get_state(self) -> dict[str, Any]: - """ - Returns the current state of the guidance technique as a dictionary. The state variables will be included in - the __repr__ method. Returns: - `dict[str, Any]`: A dictionary containing the current state variables including: - - step: Current inference step - - num_inference_steps: Total number of inference steps - - timestep: Current timestep tensor - - count_prepared: Number of times prepare_models has been called - - enabled: Whether the guidance is enabled - - num_conditions: Number of conditions - """ - state = { - "step": self._step, - "num_inference_steps": self._num_inference_steps, - "timestep": self._timestep, - "count_prepared": self._count_prepared, - "enabled": self._enabled, - "num_conditions": self.num_conditions, - } - return state - - def __repr__(self) -> str: - """ - Returns a string representation of the guidance object including both config and current state. - """ - # Get ConfigMixin's __repr__ - str_repr = super().__repr__() - - # Get current state - state = self.get_state() - - # Format each state variable on its own line with indentation - state_lines = [] - for k, v in state.items(): - # Convert value to string and handle multi-line values - v_str = str(v) - if "\n" in v_str: - # For multi-line values (like MomentumBuffer), indent subsequent lines - v_lines = v_str.split("\n") - v_str = v_lines[0] + "\n" + "\n".join([" " + line for line in v_lines[1:]]) - state_lines.append(f" {k}: {v_str}") - - state_str = "\n".join(state_lines) - - return f"{str_repr}\nState:\n{state_str}" - - def prepare_models(self, denoiser: torch.nn.Module) -> None: - """ - Prepares the models for the guidance technique on a given batch of data. This method should be overridden in - subclasses to implement specific model preparation logic. - """ - self._count_prepared += 1 - - def cleanup_models(self, denoiser: torch.nn.Module) -> None: - """ - Cleans up the models for the guidance technique after a given batch of data. This method should be overridden - in subclasses to implement specific model cleanup logic. It is useful for removing any hooks or other stateful - modifications made during `prepare_models`. - """ - pass - - def prepare_inputs(self, data: "BlockState") -> list["BlockState"]: - raise NotImplementedError("BaseGuidance::prepare_inputs must be implemented in subclasses.") - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - raise NotImplementedError("BaseGuidance::prepare_inputs_from_block_state must be implemented in subclasses.") - - def __call__(self, data: list["BlockState"]) -> Any: - if not all(hasattr(d, "noise_pred") for d in data): - raise ValueError("Expected all data to have `noise_pred` attribute.") - if len(data) != self.num_conditions: - raise ValueError( - f"Expected {self.num_conditions} data items, but got {len(data)}. Please check the input data." - ) - forward_inputs = {getattr(d, self._identifier_key): d.noise_pred for d in data} - return self.forward(**forward_inputs) - - def forward(self, *args, **kwargs) -> Any: - raise NotImplementedError("BaseGuidance::forward must be implemented in subclasses.") - - @property - def is_conditional(self) -> bool: - raise NotImplementedError("BaseGuidance::is_conditional must be implemented in subclasses.") - - @property - def is_unconditional(self) -> bool: - return not self.is_conditional - - @property - def num_conditions(self) -> int: - raise NotImplementedError("BaseGuidance::num_conditions must be implemented in subclasses.") - - @classmethod - def _prepare_batch( - cls, - data: dict[str, tuple[torch.Tensor, torch.Tensor]], - tuple_index: int, - identifier: str, - ) -> "BlockState": - """ - Prepares a batch of data for the guidance technique. This method is used in the `prepare_inputs` method of the - `BaseGuidance` class. It prepares the batch based on the provided tuple index. - - Args: - input_fields (`dict[str, str | tuple[str, str]]`): - A dictionary where the keys are the names of the fields that will be used to store the data once it is - prepared with `prepare_inputs`. The values can be either a string or a tuple of length 2, which is used - to look up the required data provided for preparation. If a string is provided, it will be used as the - conditional data (or unconditional if used with a guidance method that requires it). If a tuple of - length 2 is provided, the first element must be the conditional data identifier and the second element - must be the unconditional data identifier or None. - data (`BlockState`): - The input data to be prepared. - tuple_index (`int`): - The index to use when accessing input fields that are tuples. - - Returns: - `BlockState`: The prepared batch of data. - """ - from ..modular_pipelines.modular_pipeline import BlockState - - data_batch = {} - for key, value in data.items(): - try: - if isinstance(value, torch.Tensor): - data_batch[key] = value - elif isinstance(value, tuple): - data_batch[key] = value[tuple_index] - else: - raise ValueError(f"Invalid value type: {type(value)}") - except ValueError: - logger.debug(f"`data` does not have attribute(s) {value}, skipping.") - data_batch[cls._identifier_key] = identifier - return BlockState(**data_batch) - - @classmethod - def _prepare_batch_from_block_state( - cls, - input_fields: dict[str, str | tuple[str, str]], - data: "BlockState", - tuple_index: int, - identifier: str, - ) -> "BlockState": - """ - Prepares a batch of data for the guidance technique. This method is used in the `prepare_inputs` method of the - `BaseGuidance` class. It prepares the batch based on the provided tuple index. - - Args: - input_fields (`dict[str, str | tuple[str, str]]`): - A dictionary where the keys are the names of the fields that will be used to store the data once it is - prepared with `prepare_inputs`. The values can be either a string or a tuple of length 2, which is used - to look up the required data provided for preparation. If a string is provided, it will be used as the - conditional data (or unconditional if used with a guidance method that requires it). If a tuple of - length 2 is provided, the first element must be the conditional data identifier and the second element - must be the unconditional data identifier or None. - data (`BlockState`): - The input data to be prepared. - tuple_index (`int`): - The index to use when accessing input fields that are tuples. - - Returns: - `BlockState`: The prepared batch of data. - """ - from ..modular_pipelines.modular_pipeline import BlockState - - data_batch = {} - for key, value in input_fields.items(): - try: - if isinstance(value, str): - data_batch[key] = getattr(data, value) - elif isinstance(value, tuple): - data_batch[key] = getattr(data, value[tuple_index]) - else: - # We've already checked that value is a string or a tuple of strings with length 2 - pass - except AttributeError: - logger.debug(f"`data` does not have attribute(s) {value}, skipping.") - data_batch[cls._identifier_key] = identifier - return BlockState(**data_batch) - - @classmethod - @validate_hf_hub_args - def from_pretrained( - cls, - pretrained_model_name_or_path: str | os.PathLike | None = None, - subfolder: str | None = None, - return_unused_kwargs=False, - **kwargs, - ) -> Self: - r""" - Instantiate a guider from a pre-defined JSON configuration file in a local directory or Hub repository. - - Parameters: - pretrained_model_name_or_path (`str` or `os.PathLike`, *optional*): - Can be either: - - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing the guider configuration - saved with [`~BaseGuidance.save_pretrained`]. - subfolder (`str`, *optional*): - The subfolder location of a model file within a larger model repository on the Hub or locally. - return_unused_kwargs (`bool`, *optional*, defaults to `False`): - Whether kwargs that are not consumed by the Python class should be returned or not. - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - output_loading_info(`bool`, *optional*, defaults to `False`): - Whether or not to also return a dictionary containing missing keys, unexpected keys and error messages. - local_files_only(`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - - > [!TIP] > To use private or [gated models](https://huggingface.co/docs/hub/models-gated#gated-models), log-in - with `hf > auth login`. You can also activate the special > - ["offline-mode"](https://huggingface.co/diffusers/installation.html#offline-mode) to use this method in a > - firewalled environment. - - """ - config, kwargs, commit_hash = cls.load_config( - pretrained_model_name_or_path=pretrained_model_name_or_path, - subfolder=subfolder, - return_unused_kwargs=True, - return_commit_hash=True, - **kwargs, - ) - return cls.from_config(config, return_unused_kwargs=return_unused_kwargs, **kwargs) - - def save_pretrained(self, save_directory: str | os.PathLike, push_to_hub: bool = False, **kwargs): - """ - Save a guider configuration object to a directory so that it can be reloaded using the - [`~BaseGuidance.from_pretrained`] class method. - - Args: - save_directory (`str` or `os.PathLike`): - Directory where the configuration JSON file will be saved (will be created if it does not exist). - push_to_hub (`bool`, *optional*, defaults to `False`): - Whether or not to push your model to the Hugging Face Hub after saving it. You can specify the - repository you want to push to with `repo_id` (will default to the name of `save_directory` in your - namespace). - kwargs (`dict[str, Any]`, *optional*): - Additional keyword arguments passed along to the [`~utils.PushToHubMixin.push_to_hub`] method. - """ - self.save_config(save_directory=save_directory, push_to_hub=push_to_hub, **kwargs) - - -class GuiderOutput(BaseOutput): - pred: torch.Tensor - pred_cond: torch.Tensor | None - pred_uncond: torch.Tensor | None - - -def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0): - r""" - Rescales `noise_cfg` tensor based on `guidance_rescale` to improve image quality and fix overexposure. Based on - Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are - Flawed](https://huggingface.co/papers/2305.08891). - - Args: - noise_cfg (`torch.Tensor`): - The predicted noise tensor for the guided diffusion process. - noise_pred_text (`torch.Tensor`): - The predicted noise tensor for the text-guided diffusion process. - guidance_rescale (`float`, *optional*, defaults to 0.0): - A rescale factor applied to the noise predictions. - Returns: - noise_cfg (`torch.Tensor`): The rescaled noise prediction tensor. - """ - std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True) - std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True) - # rescale the results from guidance (fixes overexposure) - noise_pred_rescaled = noise_cfg * (std_text / std_cfg) - # mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images - noise_cfg = guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg - return noise_cfg diff --git a/diffusers/guiders/magnitude_aware_guidance.py b/diffusers/guiders/magnitude_aware_guidance.py deleted file mode 100644 index 5f3ee9bea95ade2725d1900951d2c862f057830a..0000000000000000000000000000000000000000 --- a/diffusers/guiders/magnitude_aware_guidance.py +++ /dev/null @@ -1,159 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import TYPE_CHECKING - -import torch - -from ..configuration_utils import register_to_config -from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -class MagnitudeAwareGuidance(BaseGuidance): - """ - Magnitude-Aware Mitigation for Boosted Guidance (MAMBO-G): https://huggingface.co/papers/2508.03442 - - Args: - guidance_scale (`float`, defaults to `10.0`): - The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text - prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and - deterioration of image quality. - alpha (`float`, defaults to `8.0`): - The alpha parameter for the magnitude-aware guidance. Higher values cause more aggressive supression of - guidance scale when the magnitude of the guidance update is large. - guidance_rescale (`float`, defaults to `0.0`): - The rescale factor applied to the noise predictions. This is used to improve image quality and fix - overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are - Flawed](https://huggingface.co/papers/2305.08891). - use_original_formulation (`bool`, defaults to `False`): - Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default, - we use the diffusers-native implementation that has been in the codebase for a long time. See - [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details. - start (`float`, defaults to `0.0`): - The fraction of the total number of denoising steps after which guidance starts. - stop (`float`, defaults to `1.0`): - The fraction of the total number of denoising steps after which guidance stops. - """ - - _input_predictions = ["pred_cond", "pred_uncond"] - - @register_to_config - def __init__( - self, - guidance_scale: float = 10.0, - alpha: float = 8.0, - guidance_rescale: float = 0.0, - use_original_formulation: bool = False, - start: float = 0.0, - stop: float = 1.0, - enabled: bool = True, - ): - super().__init__(start, stop, enabled) - - self.guidance_scale = guidance_scale - self.alpha = alpha - self.guidance_rescale = guidance_rescale - self.use_original_formulation = use_original_formulation - - def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch(data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput: - pred = None - - if not self._is_mambo_g_enabled(): - pred = pred_cond - else: - pred = mambo_guidance( - pred_cond, - pred_uncond, - self.guidance_scale, - self.alpha, - self.use_original_formulation, - ) - - if self.guidance_rescale > 0.0: - pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale) - - return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) - - @property - def is_conditional(self) -> bool: - return self._count_prepared == 1 - - @property - def num_conditions(self) -> int: - num_conditions = 1 - if self._is_mambo_g_enabled(): - num_conditions += 1 - return num_conditions - - def _is_mambo_g_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self._start * self._num_inference_steps) - skip_stop_step = int(self._stop * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.guidance_scale, 0.0) - else: - is_close = math.isclose(self.guidance_scale, 1.0) - - return is_within_range and not is_close - - -def mambo_guidance( - pred_cond: torch.Tensor, - pred_uncond: torch.Tensor, - guidance_scale: float, - alpha: float = 8.0, - use_original_formulation: bool = False, -): - dim = list(range(1, len(pred_cond.shape))) - diff = pred_cond - pred_uncond - ratio = torch.norm(diff, dim=dim, keepdim=True) / torch.norm(pred_uncond, dim=dim, keepdim=True) - guidance_scale_final = ( - guidance_scale * torch.exp(-alpha * ratio) - if use_original_formulation - else 1.0 + (guidance_scale - 1.0) * torch.exp(-alpha * ratio) - ) - pred = pred_cond if use_original_formulation else pred_uncond - pred = pred + guidance_scale_final * diff - - return pred diff --git a/diffusers/guiders/perturbed_attention_guidance.py b/diffusers/guiders/perturbed_attention_guidance.py deleted file mode 100644 index eff89299c4a0102563fcf9cf94b112693299b4f5..0000000000000000000000000000000000000000 --- a/diffusers/guiders/perturbed_attention_guidance.py +++ /dev/null @@ -1,289 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import math -from typing import TYPE_CHECKING, Any - -import torch - -from ..configuration_utils import register_to_config -from ..hooks import HookRegistry, LayerSkipConfig -from ..hooks.layer_skip import _apply_layer_skip_hook -from ..utils import get_logger -from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -logger = get_logger(__name__) # pylint: disable=invalid-name - - -class PerturbedAttentionGuidance(BaseGuidance): - """ - Perturbed Attention Guidance (PAG): https://huggingface.co/papers/2403.17377 - - The intution behind PAG can be thought of as moving the CFG predicted distribution estimates further away from - worse versions of the conditional distribution estimates. PAG was one of the first techniques to introduce the idea - of using a worse version of the trained model for better guiding itself in the denoising process. It perturbs the - attention scores of the latent stream by replacing the score matrix with an identity matrix for selectively chosen - layers. - - Additional reading: - - [Guiding a Diffusion Model with a Bad Version of Itself](https://huggingface.co/papers/2406.02507) - - PAG is implemented with similar implementation to SkipLayerGuidance due to overlap in the configuration parameters - and implementation details. - - Args: - guidance_scale (`float`, defaults to `7.5`): - The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text - prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and - deterioration of image quality. - perturbed_guidance_scale (`float`, defaults to `2.8`): - The scale parameter for perturbed attention guidance. - perturbed_guidance_start (`float`, defaults to `0.01`): - The fraction of the total number of denoising steps after which perturbed attention guidance starts. - perturbed_guidance_stop (`float`, defaults to `0.2`): - The fraction of the total number of denoising steps after which perturbed attention guidance stops. - perturbed_guidance_layers (`int` or `list[int]`, *optional*): - The layer indices to apply perturbed attention guidance to. Can be a single integer or a list of integers. - If not provided, `perturbed_guidance_config` must be provided. - perturbed_guidance_config (`LayerSkipConfig` or `list[LayerSkipConfig]`, *optional*): - The configuration for the perturbed attention guidance. Can be a single `LayerSkipConfig` or a list of - `LayerSkipConfig`. If not provided, `perturbed_guidance_layers` must be provided. - guidance_rescale (`float`, defaults to `0.0`): - The rescale factor applied to the noise predictions. This is used to improve image quality and fix - overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are - Flawed](https://huggingface.co/papers/2305.08891). - use_original_formulation (`bool`, defaults to `False`): - Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default, - we use the diffusers-native implementation that has been in the codebase for a long time. See - [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details. - start (`float`, defaults to `0.01`): - The fraction of the total number of denoising steps after which guidance starts. - stop (`float`, defaults to `0.2`): - The fraction of the total number of denoising steps after which guidance stops. - """ - - # NOTE: The current implementation does not account for joint latent conditioning (text + image/video tokens in - # the same latent stream). It assumes the entire latent is a single stream of visual tokens. It would be very - # complex to support joint latent conditioning in a model-agnostic manner without specializing the implementation - # for each model architecture. - - _input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"] - - @register_to_config - def __init__( - self, - guidance_scale: float = 7.5, - perturbed_guidance_scale: float = 2.8, - perturbed_guidance_start: float = 0.01, - perturbed_guidance_stop: float = 0.2, - perturbed_guidance_layers: int | list[int] | None = None, - perturbed_guidance_config: LayerSkipConfig | list[LayerSkipConfig] | dict[str, Any] = None, - guidance_rescale: float = 0.0, - use_original_formulation: bool = False, - start: float = 0.0, - stop: float = 1.0, - enabled: bool = True, - ): - super().__init__(start, stop, enabled) - - self.guidance_scale = guidance_scale - self.skip_layer_guidance_scale = perturbed_guidance_scale - self.skip_layer_guidance_start = perturbed_guidance_start - self.skip_layer_guidance_stop = perturbed_guidance_stop - self.guidance_rescale = guidance_rescale - self.use_original_formulation = use_original_formulation - - if perturbed_guidance_config is None: - if perturbed_guidance_layers is None: - raise ValueError( - "`perturbed_guidance_layers` must be provided if `perturbed_guidance_config` is not specified." - ) - perturbed_guidance_config = LayerSkipConfig( - indices=perturbed_guidance_layers, - fqn="auto", - skip_attention=False, - skip_attention_scores=True, - skip_ff=False, - ) - else: - if perturbed_guidance_layers is not None: - raise ValueError( - "`perturbed_guidance_layers` should not be provided if `perturbed_guidance_config` is specified." - ) - - if isinstance(perturbed_guidance_config, dict): - perturbed_guidance_config = LayerSkipConfig.from_dict(perturbed_guidance_config) - - if isinstance(perturbed_guidance_config, LayerSkipConfig): - perturbed_guidance_config = [perturbed_guidance_config] - - if not isinstance(perturbed_guidance_config, list): - raise ValueError( - "`perturbed_guidance_config` must be a `LayerSkipConfig`, a list of `LayerSkipConfig`, or a dict that can be converted to a `LayerSkipConfig`." - ) - elif isinstance(next(iter(perturbed_guidance_config), None), dict): - perturbed_guidance_config = [LayerSkipConfig.from_dict(config) for config in perturbed_guidance_config] - - for config in perturbed_guidance_config: - if config.skip_attention or not config.skip_attention_scores or config.skip_ff: - logger.warning( - "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. " - "Please check your configuration. Modifying the config to match the expected values." - ) - config.skip_attention = False - config.skip_attention_scores = True - config.skip_ff = False - - self.skip_layer_config = perturbed_guidance_config - self._skip_layer_hook_names = [f"SkipLayerGuidance_{i}" for i in range(len(self.skip_layer_config))] - - # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.prepare_models - def prepare_models(self, denoiser: torch.nn.Module) -> None: - self._count_prepared += 1 - if self._is_slg_enabled() and self.is_conditional and self._count_prepared > 1: - for name, config in zip(self._skip_layer_hook_names, self.skip_layer_config): - _apply_layer_skip_hook(denoiser, config, name=name) - - # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.cleanup_models - def cleanup_models(self, denoiser: torch.nn.Module) -> None: - if self._is_slg_enabled() and self.is_conditional and self._count_prepared > 1: - registry = HookRegistry.check_if_exists_or_initialize(denoiser) - # Remove the hooks after inference - for hook_name in self._skip_layer_hook_names: - registry.remove_hook(hook_name, recurse=True) - - # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.prepare_inputs - def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: - if self.num_conditions == 1: - tuple_indices = [0] - input_predictions = ["pred_cond"] - elif self.num_conditions == 2: - tuple_indices = [0, 1] - input_predictions = ( - ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_skip"] - ) - else: - tuple_indices = [0, 1, 0] - input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, input_predictions): - data_batch = self._prepare_batch(data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - if self.num_conditions == 1: - tuple_indices = [0] - input_predictions = ["pred_cond"] - elif self.num_conditions == 2: - tuple_indices = [0, 1] - input_predictions = ( - ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_skip"] - ) - else: - tuple_indices = [0, 1, 0] - input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, input_predictions): - data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.forward - def forward( - self, - pred_cond: torch.Tensor, - pred_uncond: torch.Tensor | None = None, - pred_cond_skip: torch.Tensor | None = None, - ) -> GuiderOutput: - pred = None - - if not self._is_cfg_enabled() and not self._is_slg_enabled(): - pred = pred_cond - elif not self._is_cfg_enabled(): - shift = pred_cond - pred_cond_skip - pred = pred_cond if self.use_original_formulation else pred_cond_skip - pred = pred + self.skip_layer_guidance_scale * shift - elif not self._is_slg_enabled(): - shift = pred_cond - pred_uncond - pred = pred_cond if self.use_original_formulation else pred_uncond - pred = pred + self.guidance_scale * shift - else: - shift = pred_cond - pred_uncond - shift_skip = pred_cond - pred_cond_skip - pred = pred_cond if self.use_original_formulation else pred_uncond - pred = pred + self.guidance_scale * shift + self.skip_layer_guidance_scale * shift_skip - - if self.guidance_rescale > 0.0: - pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale) - - return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) - - @property - # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.is_conditional - def is_conditional(self) -> bool: - return self._count_prepared == 1 or self._count_prepared == 3 - - @property - # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance.num_conditions - def num_conditions(self) -> int: - num_conditions = 1 - if self._is_cfg_enabled(): - num_conditions += 1 - if self._is_slg_enabled(): - num_conditions += 1 - return num_conditions - - # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance._is_cfg_enabled - def _is_cfg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self._start * self._num_inference_steps) - skip_stop_step = int(self._stop * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.guidance_scale, 0.0) - else: - is_close = math.isclose(self.guidance_scale, 1.0) - - return is_within_range and not is_close - - # Copied from diffusers.guiders.skip_layer_guidance.SkipLayerGuidance._is_slg_enabled - def _is_slg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self.skip_layer_guidance_start * self._num_inference_steps) - skip_stop_step = int(self.skip_layer_guidance_stop * self._num_inference_steps) - is_within_range = skip_start_step < self._step < skip_stop_step - - is_zero = math.isclose(self.skip_layer_guidance_scale, 0.0) - - return is_within_range and not is_zero diff --git a/diffusers/guiders/skip_layer_guidance.py b/diffusers/guiders/skip_layer_guidance.py deleted file mode 100644 index dd248135f74e74f893134124a69ccb7a44f2ac9d..0000000000000000000000000000000000000000 --- a/diffusers/guiders/skip_layer_guidance.py +++ /dev/null @@ -1,280 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import math -from typing import TYPE_CHECKING, Any - -import torch - -from ..configuration_utils import register_to_config -from ..hooks import HookRegistry, LayerSkipConfig -from ..hooks.layer_skip import _apply_layer_skip_hook -from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -class SkipLayerGuidance(BaseGuidance): - """ - Skip Layer Guidance (SLG): https://github.com/Stability-AI/sd3.5 - - Spatio-Temporal Guidance (STG): https://huggingface.co/papers/2411.18664 - - SLG was introduced by StabilityAI for improving structure and anotomy coherence in generated images. It works by - skipping the forward pass of specified transformer blocks during the denoising process on an additional conditional - batch of data, apart from the conditional and unconditional batches already used in CFG - ([~guiders.classifier_free_guidance.ClassifierFreeGuidance]), and then scaling and shifting the CFG predictions - based on the difference between conditional without skipping and conditional with skipping predictions. - - The intution behind SLG can be thought of as moving the CFG predicted distribution estimates further away from - worse versions of the conditional distribution estimates (because skipping layers is equivalent to using a worse - version of the model for the conditional prediction). - - STG is an improvement and follow-up work combining ideas from SLG, PAG and similar techniques for improving - generation quality in video diffusion models. - - Additional reading: - - [Guiding a Diffusion Model with a Bad Version of Itself](https://huggingface.co/papers/2406.02507) - - The values for `skip_layer_guidance_scale`, `skip_layer_guidance_start`, and `skip_layer_guidance_stop` are - defaulted to the recommendations by StabilityAI for Stable Diffusion 3.5 Medium. - - Args: - guidance_scale (`float`, defaults to `7.5`): - The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text - prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and - deterioration of image quality. - skip_layer_guidance_scale (`float`, defaults to `2.8`): - The scale parameter for skip layer guidance. Anatomy and structure coherence may improve with higher - values, but it may also lead to overexposure and saturation. - skip_layer_guidance_start (`float`, defaults to `0.01`): - The fraction of the total number of denoising steps after which skip layer guidance starts. - skip_layer_guidance_stop (`float`, defaults to `0.2`): - The fraction of the total number of denoising steps after which skip layer guidance stops. - skip_layer_guidance_layers (`int` or `list[int]`, *optional*): - The layer indices to apply skip layer guidance to. Can be a single integer or a list of integers. If not - provided, `skip_layer_config` must be provided. The recommended values are `[7, 8, 9]` for Stable Diffusion - 3.5 Medium. - skip_layer_config (`LayerSkipConfig` or `list[LayerSkipConfig]`, *optional*): - The configuration for the skip layer guidance. Can be a single `LayerSkipConfig` or a list of - `LayerSkipConfig`. If not provided, `skip_layer_guidance_layers` must be provided. - guidance_rescale (`float`, defaults to `0.0`): - The rescale factor applied to the noise predictions. This is used to improve image quality and fix - overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are - Flawed](https://huggingface.co/papers/2305.08891). - use_original_formulation (`bool`, defaults to `False`): - Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default, - we use the diffusers-native implementation that has been in the codebase for a long time. See - [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details. - start (`float`, defaults to `0.01`): - The fraction of the total number of denoising steps after which guidance starts. - stop (`float`, defaults to `0.2`): - The fraction of the total number of denoising steps after which guidance stops. - """ - - _input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"] - - @register_to_config - def __init__( - self, - guidance_scale: float = 7.5, - skip_layer_guidance_scale: float = 2.8, - skip_layer_guidance_start: float = 0.01, - skip_layer_guidance_stop: float = 0.2, - skip_layer_guidance_layers: int | list[int] | None = None, - skip_layer_config: LayerSkipConfig | list[LayerSkipConfig] | dict[str, Any] = None, - guidance_rescale: float = 0.0, - use_original_formulation: bool = False, - start: float = 0.0, - stop: float = 1.0, - enabled: bool = True, - ): - super().__init__(start, stop, enabled) - - self.guidance_scale = guidance_scale - self.skip_layer_guidance_scale = skip_layer_guidance_scale - self.skip_layer_guidance_start = skip_layer_guidance_start - self.skip_layer_guidance_stop = skip_layer_guidance_stop - self.guidance_rescale = guidance_rescale - self.use_original_formulation = use_original_formulation - - if not (0.0 <= skip_layer_guidance_start < 1.0): - raise ValueError( - f"Expected `skip_layer_guidance_start` to be between 0.0 and 1.0, but got {skip_layer_guidance_start}." - ) - if not (skip_layer_guidance_start <= skip_layer_guidance_stop <= 1.0): - raise ValueError( - f"Expected `skip_layer_guidance_stop` to be between 0.0 and 1.0, but got {skip_layer_guidance_stop}." - ) - - if skip_layer_guidance_layers is None and skip_layer_config is None: - raise ValueError( - "Either `skip_layer_guidance_layers` or `skip_layer_config` must be provided to enable Skip Layer Guidance." - ) - if skip_layer_guidance_layers is not None and skip_layer_config is not None: - raise ValueError("Only one of `skip_layer_guidance_layers` or `skip_layer_config` can be provided.") - - if skip_layer_guidance_layers is not None: - if isinstance(skip_layer_guidance_layers, int): - skip_layer_guidance_layers = [skip_layer_guidance_layers] - if not isinstance(skip_layer_guidance_layers, list): - raise ValueError( - f"Expected `skip_layer_guidance_layers` to be an int or a list of ints, but got {type(skip_layer_guidance_layers)}." - ) - skip_layer_config = [LayerSkipConfig(layer, fqn="auto") for layer in skip_layer_guidance_layers] - - if isinstance(skip_layer_config, dict): - skip_layer_config = LayerSkipConfig.from_dict(skip_layer_config) - - if isinstance(skip_layer_config, LayerSkipConfig): - skip_layer_config = [skip_layer_config] - - if not isinstance(skip_layer_config, list): - raise ValueError( - f"Expected `skip_layer_config` to be a LayerSkipConfig or a list of LayerSkipConfig, but got {type(skip_layer_config)}." - ) - elif isinstance(next(iter(skip_layer_config), None), dict): - skip_layer_config = [LayerSkipConfig.from_dict(config) for config in skip_layer_config] - - self.skip_layer_config = skip_layer_config - self._skip_layer_hook_names = [f"SkipLayerGuidance_{i}" for i in range(len(self.skip_layer_config))] - - def prepare_models(self, denoiser: torch.nn.Module) -> None: - self._count_prepared += 1 - if self._is_slg_enabled() and self.is_conditional and self._count_prepared > 1: - for name, config in zip(self._skip_layer_hook_names, self.skip_layer_config): - _apply_layer_skip_hook(denoiser, config, name=name) - - def cleanup_models(self, denoiser: torch.nn.Module) -> None: - if self._is_slg_enabled() and self.is_conditional and self._count_prepared > 1: - registry = HookRegistry.check_if_exists_or_initialize(denoiser) - # Remove the hooks after inference - for hook_name in self._skip_layer_hook_names: - registry.remove_hook(hook_name, recurse=True) - - def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: - if self.num_conditions == 1: - tuple_indices = [0] - input_predictions = ["pred_cond"] - elif self.num_conditions == 2: - tuple_indices = [0, 1] - input_predictions = ( - ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_skip"] - ) - else: - tuple_indices = [0, 1, 0] - input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, input_predictions): - data_batch = self._prepare_batch(data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - if self.num_conditions == 1: - tuple_indices = [0] - input_predictions = ["pred_cond"] - elif self.num_conditions == 2: - tuple_indices = [0, 1] - input_predictions = ( - ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_skip"] - ) - else: - tuple_indices = [0, 1, 0] - input_predictions = ["pred_cond", "pred_uncond", "pred_cond_skip"] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, input_predictions): - data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def forward( - self, - pred_cond: torch.Tensor, - pred_uncond: torch.Tensor | None = None, - pred_cond_skip: torch.Tensor | None = None, - ) -> GuiderOutput: - pred = None - - if not self._is_cfg_enabled() and not self._is_slg_enabled(): - pred = pred_cond - elif not self._is_cfg_enabled(): - shift = pred_cond - pred_cond_skip - pred = pred_cond if self.use_original_formulation else pred_cond_skip - pred = pred + self.skip_layer_guidance_scale * shift - elif not self._is_slg_enabled(): - shift = pred_cond - pred_uncond - pred = pred_cond if self.use_original_formulation else pred_uncond - pred = pred + self.guidance_scale * shift - else: - shift = pred_cond - pred_uncond - shift_skip = pred_cond - pred_cond_skip - pred = pred_cond if self.use_original_formulation else pred_uncond - pred = pred + self.guidance_scale * shift + self.skip_layer_guidance_scale * shift_skip - - if self.guidance_rescale > 0.0: - pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale) - - return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) - - @property - def is_conditional(self) -> bool: - return self._count_prepared == 1 or self._count_prepared == 3 - - @property - def num_conditions(self) -> int: - num_conditions = 1 - if self._is_cfg_enabled(): - num_conditions += 1 - if self._is_slg_enabled(): - num_conditions += 1 - return num_conditions - - def _is_cfg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self._start * self._num_inference_steps) - skip_stop_step = int(self._stop * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.guidance_scale, 0.0) - else: - is_close = math.isclose(self.guidance_scale, 1.0) - - return is_within_range and not is_close - - def _is_slg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self.skip_layer_guidance_start * self._num_inference_steps) - skip_stop_step = int(self.skip_layer_guidance_stop * self._num_inference_steps) - is_within_range = skip_start_step < self._step < skip_stop_step - - is_zero = math.isclose(self.skip_layer_guidance_scale, 0.0) - - return is_within_range and not is_zero diff --git a/diffusers/guiders/smoothed_energy_guidance.py b/diffusers/guiders/smoothed_energy_guidance.py deleted file mode 100644 index 86313ed1ac3ff408b0c79dec8b1913c79b7eb9af..0000000000000000000000000000000000000000 --- a/diffusers/guiders/smoothed_energy_guidance.py +++ /dev/null @@ -1,269 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import math -from typing import TYPE_CHECKING - -import torch - -from ..configuration_utils import register_to_config -from ..hooks import HookRegistry -from ..hooks.smoothed_energy_guidance_utils import SmoothedEnergyGuidanceConfig, _apply_smoothed_energy_guidance_hook -from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -class SmoothedEnergyGuidance(BaseGuidance): - """ - Smoothed Energy Guidance (SEG): https://huggingface.co/papers/2408.00760 - - SEG is only supported as an experimental prototype feature for now, so the implementation may be modified in the - future without warning or guarantee of reproducibility. This implementation assumes: - - Generated images are square (height == width) - - The model does not combine different modalities together (e.g., text and image latent streams are not combined - together such as Flux) - - Args: - guidance_scale (`float`, defaults to `7.5`): - The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text - prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and - deterioration of image quality. - seg_guidance_scale (`float`, defaults to `3.0`): - The scale parameter for smoothed energy guidance. Anatomy and structure coherence may improve with higher - values, but it may also lead to overexposure and saturation. - seg_blur_sigma (`float`, defaults to `9999999.0`): - The amount by which we blur the attention weights. Setting this value greater than 9999.0 results in - infinite blur, which means uniform queries. Controlling it exponentially is empirically effective. - seg_blur_threshold_inf (`float`, defaults to `9999.0`): - The threshold above which the blur is considered infinite. - seg_guidance_start (`float`, defaults to `0.0`): - The fraction of the total number of denoising steps after which smoothed energy guidance starts. - seg_guidance_stop (`float`, defaults to `1.0`): - The fraction of the total number of denoising steps after which smoothed energy guidance stops. - seg_guidance_layers (`int` or `list[int]`, *optional*): - The layer indices to apply smoothed energy guidance to. Can be a single integer or a list of integers. If - not provided, `seg_guidance_config` must be provided. The recommended values are `[7, 8, 9]` for Stable - Diffusion 3.5 Medium. - seg_guidance_config (`SmoothedEnergyGuidanceConfig` or `list[SmoothedEnergyGuidanceConfig]`, *optional*): - The configuration for the smoothed energy layer guidance. Can be a single `SmoothedEnergyGuidanceConfig` or - a list of `SmoothedEnergyGuidanceConfig`. If not provided, `seg_guidance_layers` must be provided. - guidance_rescale (`float`, defaults to `0.0`): - The rescale factor applied to the noise predictions. This is used to improve image quality and fix - overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are - Flawed](https://huggingface.co/papers/2305.08891). - use_original_formulation (`bool`, defaults to `False`): - Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default, - we use the diffusers-native implementation that has been in the codebase for a long time. See - [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details. - start (`float`, defaults to `0.01`): - The fraction of the total number of denoising steps after which guidance starts. - stop (`float`, defaults to `0.2`): - The fraction of the total number of denoising steps after which guidance stops. - """ - - _input_predictions = ["pred_cond", "pred_uncond", "pred_cond_seg"] - - @register_to_config - def __init__( - self, - guidance_scale: float = 7.5, - seg_guidance_scale: float = 2.8, - seg_blur_sigma: float = 9999999.0, - seg_blur_threshold_inf: float = 9999.0, - seg_guidance_start: float = 0.0, - seg_guidance_stop: float = 1.0, - seg_guidance_layers: int | list[int] | None = None, - seg_guidance_config: SmoothedEnergyGuidanceConfig | list[SmoothedEnergyGuidanceConfig] = None, - guidance_rescale: float = 0.0, - use_original_formulation: bool = False, - start: float = 0.0, - stop: float = 1.0, - enabled: bool = True, - ): - super().__init__(start, stop, enabled) - - self.guidance_scale = guidance_scale - self.seg_guidance_scale = seg_guidance_scale - self.seg_blur_sigma = seg_blur_sigma - self.seg_blur_threshold_inf = seg_blur_threshold_inf - self.seg_guidance_start = seg_guidance_start - self.seg_guidance_stop = seg_guidance_stop - self.guidance_rescale = guidance_rescale - self.use_original_formulation = use_original_formulation - - if not (0.0 <= seg_guidance_start < 1.0): - raise ValueError(f"Expected `seg_guidance_start` to be between 0.0 and 1.0, but got {seg_guidance_start}.") - if not (seg_guidance_start <= seg_guidance_stop <= 1.0): - raise ValueError(f"Expected `seg_guidance_stop` to be between 0.0 and 1.0, but got {seg_guidance_stop}.") - - if seg_guidance_layers is None and seg_guidance_config is None: - raise ValueError( - "Either `seg_guidance_layers` or `seg_guidance_config` must be provided to enable Smoothed Energy Guidance." - ) - if seg_guidance_layers is not None and seg_guidance_config is not None: - raise ValueError("Only one of `seg_guidance_layers` or `seg_guidance_config` can be provided.") - - if seg_guidance_layers is not None: - if isinstance(seg_guidance_layers, int): - seg_guidance_layers = [seg_guidance_layers] - if not isinstance(seg_guidance_layers, list): - raise ValueError( - f"Expected `seg_guidance_layers` to be an int or a list of ints, but got {type(seg_guidance_layers)}." - ) - seg_guidance_config = [SmoothedEnergyGuidanceConfig(layer, fqn="auto") for layer in seg_guidance_layers] - - if isinstance(seg_guidance_config, dict): - seg_guidance_config = SmoothedEnergyGuidanceConfig.from_dict(seg_guidance_config) - - if isinstance(seg_guidance_config, SmoothedEnergyGuidanceConfig): - seg_guidance_config = [seg_guidance_config] - - if not isinstance(seg_guidance_config, list): - raise ValueError( - f"Expected `seg_guidance_config` to be a SmoothedEnergyGuidanceConfig or a list of SmoothedEnergyGuidanceConfig, but got {type(seg_guidance_config)}." - ) - elif isinstance(next(iter(seg_guidance_config), None), dict): - seg_guidance_config = [SmoothedEnergyGuidanceConfig.from_dict(config) for config in seg_guidance_config] - - self.seg_guidance_config = seg_guidance_config - self._seg_layer_hook_names = [f"SmoothedEnergyGuidance_{i}" for i in range(len(self.seg_guidance_config))] - - def prepare_models(self, denoiser: torch.nn.Module) -> None: - if self._is_seg_enabled() and self.is_conditional and self._count_prepared > 1: - for name, config in zip(self._seg_layer_hook_names, self.seg_guidance_config): - _apply_smoothed_energy_guidance_hook(denoiser, config, self.seg_blur_sigma, name=name) - - def cleanup_models(self, denoiser: torch.nn.Module): - if self._is_seg_enabled() and self.is_conditional and self._count_prepared > 1: - registry = HookRegistry.check_if_exists_or_initialize(denoiser) - # Remove the hooks after inference - for hook_name in self._seg_layer_hook_names: - registry.remove_hook(hook_name, recurse=True) - - def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: - if self.num_conditions == 1: - tuple_indices = [0] - input_predictions = ["pred_cond"] - elif self.num_conditions == 2: - tuple_indices = [0, 1] - input_predictions = ( - ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_seg"] - ) - else: - tuple_indices = [0, 1, 0] - input_predictions = ["pred_cond", "pred_uncond", "pred_cond_seg"] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, input_predictions): - data_batch = self._prepare_batch(data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - if self.num_conditions == 1: - tuple_indices = [0] - input_predictions = ["pred_cond"] - elif self.num_conditions == 2: - tuple_indices = [0, 1] - input_predictions = ( - ["pred_cond", "pred_uncond"] if self._is_cfg_enabled() else ["pred_cond", "pred_cond_seg"] - ) - else: - tuple_indices = [0, 1, 0] - input_predictions = ["pred_cond", "pred_uncond", "pred_cond_seg"] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, input_predictions): - data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def forward( - self, - pred_cond: torch.Tensor, - pred_uncond: torch.Tensor | None = None, - pred_cond_seg: torch.Tensor | None = None, - ) -> GuiderOutput: - pred = None - - if not self._is_cfg_enabled() and not self._is_seg_enabled(): - pred = pred_cond - elif not self._is_cfg_enabled(): - shift = pred_cond - pred_cond_seg - pred = pred_cond if self.use_original_formulation else pred_cond_seg - pred = pred + self.seg_guidance_scale * shift - elif not self._is_seg_enabled(): - shift = pred_cond - pred_uncond - pred = pred_cond if self.use_original_formulation else pred_uncond - pred = pred + self.guidance_scale * shift - else: - shift = pred_cond - pred_uncond - shift_seg = pred_cond - pred_cond_seg - pred = pred_cond if self.use_original_formulation else pred_uncond - pred = pred + self.guidance_scale * shift + self.seg_guidance_scale * shift_seg - - if self.guidance_rescale > 0.0: - pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale) - - return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) - - @property - def is_conditional(self) -> bool: - return self._count_prepared == 1 or self._count_prepared == 3 - - @property - def num_conditions(self) -> int: - num_conditions = 1 - if self._is_cfg_enabled(): - num_conditions += 1 - if self._is_seg_enabled(): - num_conditions += 1 - return num_conditions - - def _is_cfg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self._start * self._num_inference_steps) - skip_stop_step = int(self._stop * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.guidance_scale, 0.0) - else: - is_close = math.isclose(self.guidance_scale, 1.0) - - return is_within_range and not is_close - - def _is_seg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self.seg_guidance_start * self._num_inference_steps) - skip_stop_step = int(self.seg_guidance_stop * self._num_inference_steps) - is_within_range = skip_start_step < self._step < skip_stop_step - - is_zero = math.isclose(self.seg_guidance_scale, 0.0) - - return is_within_range and not is_zero diff --git a/diffusers/guiders/tangential_classifier_free_guidance.py b/diffusers/guiders/tangential_classifier_free_guidance.py deleted file mode 100644 index 497cdc3c463d84c7075bb41aa9fbebd51ce6a926..0000000000000000000000000000000000000000 --- a/diffusers/guiders/tangential_classifier_free_guidance.py +++ /dev/null @@ -1,151 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import math -from typing import TYPE_CHECKING - -import torch - -from ..configuration_utils import register_to_config -from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg - - -if TYPE_CHECKING: - from ..modular_pipelines.modular_pipeline import BlockState - - -class TangentialClassifierFreeGuidance(BaseGuidance): - """ - Tangential Classifier Free Guidance (TCFG): https://huggingface.co/papers/2503.18137 - - Args: - guidance_scale (`float`, defaults to `7.5`): - The scale parameter for classifier-free guidance. Higher values result in stronger conditioning on the text - prompt, while lower values allow for more freedom in generation. Higher values may lead to saturation and - deterioration of image quality. - guidance_rescale (`float`, defaults to `0.0`): - The rescale factor applied to the noise predictions. This is used to improve image quality and fix - overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are - Flawed](https://huggingface.co/papers/2305.08891). - use_original_formulation (`bool`, defaults to `False`): - Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default, - we use the diffusers-native implementation that has been in the codebase for a long time. See - [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details. - start (`float`, defaults to `0.0`): - The fraction of the total number of denoising steps after which guidance starts. - stop (`float`, defaults to `1.0`): - The fraction of the total number of denoising steps after which guidance stops. - """ - - _input_predictions = ["pred_cond", "pred_uncond"] - - @register_to_config - def __init__( - self, - guidance_scale: float = 7.5, - guidance_rescale: float = 0.0, - use_original_formulation: bool = False, - start: float = 0.0, - stop: float = 1.0, - enabled: bool = True, - ): - super().__init__(start, stop, enabled) - - self.guidance_scale = guidance_scale - self.guidance_rescale = guidance_rescale - self.use_original_formulation = use_original_formulation - - def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch(data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def prepare_inputs_from_block_state( - self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] - ) -> list["BlockState"]: - tuple_indices = [0] if self.num_conditions == 1 else [0, 1] - data_batches = [] - for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): - data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) - data_batches.append(data_batch) - return data_batches - - def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput: - pred = None - - if not self._is_tcfg_enabled(): - pred = pred_cond - else: - pred = normalized_guidance(pred_cond, pred_uncond, self.guidance_scale, self.use_original_formulation) - - if self.guidance_rescale > 0.0: - pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale) - - return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) - - @property - def is_conditional(self) -> bool: - return self._num_outputs_prepared == 1 - - @property - def num_conditions(self) -> int: - num_conditions = 1 - if self._is_tcfg_enabled(): - num_conditions += 1 - return num_conditions - - def _is_tcfg_enabled(self) -> bool: - if not self._enabled: - return False - - is_within_range = True - if self._num_inference_steps is not None: - skip_start_step = int(self._start * self._num_inference_steps) - skip_stop_step = int(self._stop * self._num_inference_steps) - is_within_range = skip_start_step <= self._step < skip_stop_step - - is_close = False - if self.use_original_formulation: - is_close = math.isclose(self.guidance_scale, 0.0) - else: - is_close = math.isclose(self.guidance_scale, 1.0) - - return is_within_range and not is_close - - -def normalized_guidance( - pred_cond: torch.Tensor, pred_uncond: torch.Tensor, guidance_scale: float, use_original_formulation: bool = False -) -> torch.Tensor: - cond_dtype = pred_cond.dtype - preds = torch.stack([pred_cond, pred_uncond], dim=1).float() - preds = preds.flatten(2) - U, S, Vh = torch.linalg.svd(preds, full_matrices=False) - Vh_modified = Vh.clone() - Vh_modified[:, 1] = 0 - - uncond_flat = pred_uncond.reshape(pred_uncond.size(0), 1, -1).float() - x_Vh = torch.matmul(uncond_flat, Vh.transpose(-2, -1)) - x_Vh_V = torch.matmul(x_Vh, Vh_modified) - pred_uncond = x_Vh_V.reshape(pred_uncond.shape).to(cond_dtype) - - pred = pred_cond if use_original_formulation else pred_uncond - shift = pred_cond - pred_uncond - pred = pred + guidance_scale * shift - - return pred diff --git a/diffusers/hooks/__init__.py b/diffusers/hooks/__init__.py deleted file mode 100644 index 2a9aa81608e7225d183d391fcbd5acd1e5744c8b..0000000000000000000000000000000000000000 --- a/diffusers/hooks/__init__.py +++ /dev/null @@ -1,30 +0,0 @@ -# Copyright 2024 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from ..utils import is_torch_available - - -if is_torch_available(): - from .context_parallel import apply_context_parallel - from .faster_cache import FasterCacheConfig, apply_faster_cache - from .first_block_cache import FirstBlockCacheConfig, apply_first_block_cache - from .group_offloading import apply_group_offloading - from .hooks import HookRegistry, ModelHook - from .layer_skip import LayerSkipConfig, apply_layer_skip - from .layerwise_casting import apply_layerwise_casting, apply_layerwise_casting_hook - from .mag_cache import MagCacheConfig, apply_mag_cache - from .pyramid_attention_broadcast import PyramidAttentionBroadcastConfig, apply_pyramid_attention_broadcast - from .smoothed_energy_guidance_utils import SmoothedEnergyGuidanceConfig - from .taylorseer_cache import TaylorSeerCacheConfig, apply_taylorseer_cache - from .text_kv_cache import TextKVCacheConfig, apply_text_kv_cache diff --git a/diffusers/hooks/_common.py b/diffusers/hooks/_common.py deleted file mode 100644 index 26ae2b5d715f0e471207848e43fc0c24c8a7e830..0000000000000000000000000000000000000000 --- a/diffusers/hooks/_common.py +++ /dev/null @@ -1,61 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ..models.attention import AttentionModuleMixin, FeedForward, LuminaFeedForward -from ..models.attention_processor import Attention, MochiAttention - - -_ATTENTION_CLASSES = (Attention, MochiAttention, AttentionModuleMixin) -_FEEDFORWARD_CLASSES = (FeedForward, LuminaFeedForward) - -_SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS = ( - "blocks", - "transformer_blocks", - "single_transformer_blocks", - "layers", - "visual_transformer_blocks", -) -_TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS = ("temporal_transformer_blocks",) -_CROSS_TRANSFORMER_BLOCK_IDENTIFIERS = ("blocks", "transformer_blocks", "layers") - -_ALL_TRANSFORMER_BLOCK_IDENTIFIERS = tuple( - { - *_SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS, - *_TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS, - *_CROSS_TRANSFORMER_BLOCK_IDENTIFIERS, - } -) - -# Layers supported for group offloading and layerwise casting -_GO_LC_SUPPORTED_PYTORCH_LAYERS = ( - torch.nn.Conv1d, - torch.nn.Conv2d, - torch.nn.Conv3d, - torch.nn.ConvTranspose1d, - torch.nn.ConvTranspose2d, - torch.nn.ConvTranspose3d, - torch.nn.Linear, - torch.nn.Embedding, - # TODO(aryan): look into torch.nn.LayerNorm, torch.nn.GroupNorm later, seems to be causing some issues with CogVideoX - # because of double invocation of the same norm layer in CogVideoXLayerNorm -) - - -def _get_submodule_from_fqn(module: torch.nn.Module, fqn: str) -> torch.nn.Module | None: - for submodule_name, submodule in module.named_modules(): - if submodule_name == fqn: - return submodule - return None diff --git a/diffusers/hooks/_helpers.py b/diffusers/hooks/_helpers.py deleted file mode 100644 index 9cbe5bc8108f3cd0a839aecf9b75f7db7dc22dd2..0000000000000000000000000000000000000000 --- a/diffusers/hooks/_helpers.py +++ /dev/null @@ -1,401 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -from dataclasses import dataclass -from typing import Any, Callable, Type - - -@dataclass -class AttentionProcessorMetadata: - skip_processor_output_fn: Callable[[Any], Any] - - -@dataclass -class TransformerBlockMetadata: - return_hidden_states_index: int = None - return_encoder_hidden_states_index: int = None - hidden_states_argument_name: str = "hidden_states" - - _cls: Type = None - _cached_parameter_indices: dict[str, int] = None - - def _get_parameter_from_args_kwargs(self, identifier: str, args=(), kwargs=None): - kwargs = kwargs or {} - if identifier in kwargs: - return kwargs[identifier] - if self._cached_parameter_indices is not None: - return args[self._cached_parameter_indices[identifier]] - if self._cls is None: - raise ValueError("Model class is not set for metadata.") - parameters = list(inspect.signature(self._cls.forward).parameters.keys()) - parameters = parameters[1:] # skip `self` - self._cached_parameter_indices = {param: i for i, param in enumerate(parameters)} - if identifier not in self._cached_parameter_indices: - raise ValueError(f"Parameter '{identifier}' not found in function signature but was requested.") - index = self._cached_parameter_indices[identifier] - if index >= len(args): - raise ValueError(f"Expected {index} arguments but got {len(args)}.") - return args[index] - - -class AttentionProcessorRegistry: - _registry = {} - # TODO(aryan): this is only required for the time being because we need to do the registrations - # for classes. If we do it eagerly, i.e. call the functions in global scope, we will get circular - # import errors because of the models imported in this file. - _is_registered = False - - @classmethod - def register(cls, model_class: Type, metadata: AttentionProcessorMetadata): - cls._register() - cls._registry[model_class] = metadata - - @classmethod - def get(cls, model_class: Type) -> AttentionProcessorMetadata: - cls._register() - if model_class not in cls._registry: - raise ValueError(f"Model class {model_class} not registered.") - return cls._registry[model_class] - - @classmethod - def _register(cls): - if cls._is_registered: - return - cls._is_registered = True - _register_attention_processors_metadata() - - -class TransformerBlockRegistry: - _registry = {} - # TODO(aryan): this is only required for the time being because we need to do the registrations - # for classes. If we do it eagerly, i.e. call the functions in global scope, we will get circular - # import errors because of the models imported in this file. - _is_registered = False - - @classmethod - def register(cls, model_class: Type, metadata: TransformerBlockMetadata): - cls._register() - metadata._cls = model_class - cls._registry[model_class] = metadata - - @classmethod - def get(cls, model_class: Type) -> TransformerBlockMetadata: - cls._register() - if model_class not in cls._registry: - raise ValueError(f"Model class {model_class} not registered.") - return cls._registry[model_class] - - @classmethod - def _register(cls): - if cls._is_registered: - return - cls._is_registered = True - _register_transformer_blocks_metadata() - - -def _register_attention_processors_metadata(): - from ..models.attention_processor import AttnProcessor2_0 - from ..models.transformers.transformer_cogview4 import CogView4AttnProcessor - from ..models.transformers.transformer_flux import FluxAttnProcessor - from ..models.transformers.transformer_hunyuanimage import HunyuanImageAttnProcessor - from ..models.transformers.transformer_qwenimage import QwenDoubleStreamAttnProcessor2_0 - from ..models.transformers.transformer_wan import WanAttnProcessor2_0 - from ..models.transformers.transformer_z_image import ZSingleStreamAttnProcessor - - # AttnProcessor2_0 - AttentionProcessorRegistry.register( - model_class=AttnProcessor2_0, - metadata=AttentionProcessorMetadata( - skip_processor_output_fn=_skip_proc_output_fn_Attention_AttnProcessor2_0, - ), - ) - - # CogView4AttnProcessor - AttentionProcessorRegistry.register( - model_class=CogView4AttnProcessor, - metadata=AttentionProcessorMetadata( - skip_processor_output_fn=_skip_proc_output_fn_Attention_CogView4AttnProcessor, - ), - ) - - # WanAttnProcessor2_0 - AttentionProcessorRegistry.register( - model_class=WanAttnProcessor2_0, - metadata=AttentionProcessorMetadata( - skip_processor_output_fn=_skip_proc_output_fn_Attention_WanAttnProcessor2_0, - ), - ) - - # FluxAttnProcessor - AttentionProcessorRegistry.register( - model_class=FluxAttnProcessor, - metadata=AttentionProcessorMetadata(skip_processor_output_fn=_skip_proc_output_fn_Attention_FluxAttnProcessor), - ) - - # QwenDoubleStreamAttnProcessor2 - AttentionProcessorRegistry.register( - model_class=QwenDoubleStreamAttnProcessor2_0, - metadata=AttentionProcessorMetadata( - skip_processor_output_fn=_skip_proc_output_fn_Attention_QwenDoubleStreamAttnProcessor2_0 - ), - ) - - # HunyuanImageAttnProcessor - AttentionProcessorRegistry.register( - model_class=HunyuanImageAttnProcessor, - metadata=AttentionProcessorMetadata( - skip_processor_output_fn=_skip_proc_output_fn_Attention_HunyuanImageAttnProcessor, - ), - ) - - # ZSingleStreamAttnProcessor - AttentionProcessorRegistry.register( - model_class=ZSingleStreamAttnProcessor, - metadata=AttentionProcessorMetadata( - skip_processor_output_fn=_skip_proc_output_fn_Attention_ZSingleStreamAttnProcessor, - ), - ) - - -def _register_transformer_blocks_metadata(): - from ..models.attention import BasicTransformerBlock, JointTransformerBlock - from ..models.transformers.cogvideox_transformer_3d import CogVideoXBlock - from ..models.transformers.transformer_bria import BriaTransformerBlock - from ..models.transformers.transformer_cogview4 import CogView4TransformerBlock - from ..models.transformers.transformer_flux import FluxSingleTransformerBlock, FluxTransformerBlock - from ..models.transformers.transformer_hunyuan_video import ( - HunyuanVideoSingleTransformerBlock, - HunyuanVideoTokenReplaceSingleTransformerBlock, - HunyuanVideoTokenReplaceTransformerBlock, - HunyuanVideoTransformerBlock, - ) - from ..models.transformers.transformer_hunyuanimage import ( - HunyuanImageSingleTransformerBlock, - HunyuanImageTransformerBlock, - ) - from ..models.transformers.transformer_kandinsky import Kandinsky5TransformerDecoderBlock - from ..models.transformers.transformer_ltx import LTXVideoTransformerBlock - from ..models.transformers.transformer_mochi import MochiTransformerBlock - from ..models.transformers.transformer_motif_video import ( - MotifVideoSingleTransformerBlock, - MotifVideoTransformerBlock, - ) - from ..models.transformers.transformer_qwenimage import QwenImageTransformerBlock - from ..models.transformers.transformer_wan import WanTransformerBlock - from ..models.transformers.transformer_z_image import ZImageTransformerBlock - - # BasicTransformerBlock - TransformerBlockRegistry.register( - model_class=BasicTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=None, - ), - ) - TransformerBlockRegistry.register( - model_class=BriaTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=None, - ), - ) - - # CogVideoX - TransformerBlockRegistry.register( - model_class=CogVideoXBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=1, - ), - ) - - # CogView4 - TransformerBlockRegistry.register( - model_class=CogView4TransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=1, - ), - ) - - # Flux - TransformerBlockRegistry.register( - model_class=FluxTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=1, - return_encoder_hidden_states_index=0, - ), - ) - TransformerBlockRegistry.register( - model_class=FluxSingleTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=1, - return_encoder_hidden_states_index=0, - ), - ) - - # HunyuanVideo - TransformerBlockRegistry.register( - model_class=HunyuanVideoTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=1, - ), - ) - TransformerBlockRegistry.register( - model_class=HunyuanVideoSingleTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=1, - ), - ) - TransformerBlockRegistry.register( - model_class=HunyuanVideoTokenReplaceTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=1, - ), - ) - TransformerBlockRegistry.register( - model_class=HunyuanVideoTokenReplaceSingleTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=1, - ), - ) - - # LTXVideo - TransformerBlockRegistry.register( - model_class=LTXVideoTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=None, - ), - ) - - # Mochi - TransformerBlockRegistry.register( - model_class=MochiTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=1, - ), - ) - - # MotifVideo - TransformerBlockRegistry.register( - model_class=MotifVideoTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=1, - ), - ) - TransformerBlockRegistry.register( - model_class=MotifVideoSingleTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=1, - ), - ) - - # Wan - TransformerBlockRegistry.register( - model_class=WanTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=None, - ), - ) - - # QwenImage - TransformerBlockRegistry.register( - model_class=QwenImageTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=1, - return_encoder_hidden_states_index=0, - ), - ) - - # HunyuanImage2.1 - TransformerBlockRegistry.register( - model_class=HunyuanImageTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=1, - ), - ) - TransformerBlockRegistry.register( - model_class=HunyuanImageSingleTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=1, - ), - ) - - # ZImage - TransformerBlockRegistry.register( - model_class=ZImageTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=None, - ), - ) - - TransformerBlockRegistry.register( - model_class=JointTransformerBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=1, - return_encoder_hidden_states_index=0, - ), - ) - - # Kandinsky 5.0 (Kandinsky5TransformerDecoderBlock) - TransformerBlockRegistry.register( - model_class=Kandinsky5TransformerDecoderBlock, - metadata=TransformerBlockMetadata( - return_hidden_states_index=0, - return_encoder_hidden_states_index=None, - hidden_states_argument_name="visual_embed", - ), - ) - - -# fmt: off -def _skip_attention___ret___hidden_states(self, *args, **kwargs): - hidden_states = kwargs.get("hidden_states", None) - if hidden_states is None and len(args) > 0: - hidden_states = args[0] - return hidden_states - - -def _skip_attention___ret___hidden_states___encoder_hidden_states(self, *args, **kwargs): - hidden_states = kwargs.get("hidden_states", None) - encoder_hidden_states = kwargs.get("encoder_hidden_states", None) - if hidden_states is None and len(args) > 0: - hidden_states = args[0] - if encoder_hidden_states is None and len(args) > 1: - encoder_hidden_states = args[1] - return hidden_states, encoder_hidden_states - - -_skip_proc_output_fn_Attention_AttnProcessor2_0 = _skip_attention___ret___hidden_states -_skip_proc_output_fn_Attention_CogView4AttnProcessor = _skip_attention___ret___hidden_states___encoder_hidden_states -_skip_proc_output_fn_Attention_WanAttnProcessor2_0 = _skip_attention___ret___hidden_states -# not sure what this is yet. -_skip_proc_output_fn_Attention_FluxAttnProcessor = _skip_attention___ret___hidden_states -_skip_proc_output_fn_Attention_QwenDoubleStreamAttnProcessor2_0 = _skip_attention___ret___hidden_states -_skip_proc_output_fn_Attention_HunyuanImageAttnProcessor = _skip_attention___ret___hidden_states -_skip_proc_output_fn_Attention_ZSingleStreamAttnProcessor = _skip_attention___ret___hidden_states -# fmt: on diff --git a/diffusers/hooks/context_parallel.py b/diffusers/hooks/context_parallel.py deleted file mode 100644 index 1310b20c5c11febf61069a70b7e47982830304d4..0000000000000000000000000000000000000000 --- a/diffusers/hooks/context_parallel.py +++ /dev/null @@ -1,382 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import copy -import inspect -from dataclasses import dataclass -from typing import Type - -import torch -import torch.distributed as dist - - -if torch.distributed.is_available(): - import torch.distributed._functional_collectives as funcol - -from ..models._modeling_parallel import ( - ContextParallelConfig, - ContextParallelInput, - ContextParallelModelPlan, - ContextParallelOutput, - gather_size_by_comm, -) -from ..utils import get_logger -from ..utils.torch_utils import lru_cache_unless_export, maybe_allow_in_graph, unwrap_module -from .hooks import HookRegistry, ModelHook - - -logger = get_logger(__name__) # pylint: disable=invalid-name - -_CONTEXT_PARALLEL_INPUT_HOOK_TEMPLATE = "cp_input---{}" -_CONTEXT_PARALLEL_OUTPUT_HOOK_TEMPLATE = "cp_output---{}" - - -# TODO(aryan): consolidate with ._helpers.TransformerBlockMetadata -@dataclass -class ModuleForwardMetadata: - cached_parameter_indices: dict[str, int] = None - _cls: Type = None - - def _get_parameter_from_args_kwargs(self, identifier: str, args=(), kwargs=None): - kwargs = kwargs or {} - - if identifier in kwargs: - return kwargs[identifier], True, None - - if self.cached_parameter_indices is not None: - index = self.cached_parameter_indices.get(identifier, None) - if index is None: - raise ValueError(f"Parameter '{identifier}' not found in cached indices.") - return args[index], False, index - - if self._cls is None: - raise ValueError("Model class is not set for metadata.") - - parameters = list(inspect.signature(self._cls.forward).parameters.keys()) - parameters = parameters[1:] # skip `self` - self.cached_parameter_indices = {param: i for i, param in enumerate(parameters)} - - if identifier not in self.cached_parameter_indices: - raise ValueError(f"Parameter '{identifier}' not found in function signature but was requested.") - - index = self.cached_parameter_indices[identifier] - - if index >= len(args): - raise ValueError(f"Expected {index} arguments but got {len(args)}.") - - return args[index], False, index - - -def apply_context_parallel( - module: torch.nn.Module, - parallel_config: ContextParallelConfig, - plan: dict[str, ContextParallelModelPlan], -) -> None: - """Apply context parallel on a model.""" - logger.debug(f"Applying context parallel with CP mesh: {parallel_config._mesh} and plan: {plan}") - - for module_id, cp_model_plan in plan.items(): - submodule = _get_submodule_by_name(module, module_id) - if not isinstance(submodule, list): - submodule = [submodule] - - logger.debug(f"Applying ContextParallelHook to {module_id=} identifying a total of {len(submodule)} modules") - - for m in submodule: - if isinstance(cp_model_plan, dict): - hook = ContextParallelSplitHook(cp_model_plan, parallel_config) - hook_name = _CONTEXT_PARALLEL_INPUT_HOOK_TEMPLATE.format(module_id) - elif isinstance(cp_model_plan, (ContextParallelOutput, list, tuple)): - if isinstance(cp_model_plan, ContextParallelOutput): - cp_model_plan = [cp_model_plan] - if not all(isinstance(x, ContextParallelOutput) for x in cp_model_plan): - raise ValueError(f"Expected all elements of cp_model_plan to be CPOutput, but got {cp_model_plan}") - hook = ContextParallelGatherHook(cp_model_plan, parallel_config) - hook_name = _CONTEXT_PARALLEL_OUTPUT_HOOK_TEMPLATE.format(module_id) - else: - raise ValueError(f"Unsupported context parallel model plan type: {type(cp_model_plan)}") - registry = HookRegistry.check_if_exists_or_initialize(m) - registry.register_hook(hook, hook_name) - - -def remove_context_parallel(module: torch.nn.Module, plan: dict[str, ContextParallelModelPlan]) -> None: - for module_id, cp_model_plan in plan.items(): - submodule = _get_submodule_by_name(module, module_id) - if not isinstance(submodule, list): - submodule = [submodule] - - for m in submodule: - registry = HookRegistry.check_if_exists_or_initialize(m) - if isinstance(cp_model_plan, dict): - hook_name = _CONTEXT_PARALLEL_INPUT_HOOK_TEMPLATE.format(module_id) - elif isinstance(cp_model_plan, (ContextParallelOutput, list, tuple)): - hook_name = _CONTEXT_PARALLEL_OUTPUT_HOOK_TEMPLATE.format(module_id) - else: - raise ValueError(f"Unsupported context parallel model plan type: {type(cp_model_plan)}") - registry.remove_hook(hook_name) - - -class ContextParallelSplitHook(ModelHook): - def __init__(self, metadata: ContextParallelModelPlan, parallel_config: ContextParallelConfig) -> None: - super().__init__() - self.metadata = metadata - self.parallel_config = parallel_config - self.module_forward_metadata = None - - def initialize_hook(self, module): - cls = unwrap_module(module).__class__ - self.module_forward_metadata = ModuleForwardMetadata(_cls=cls) - return module - - def pre_forward(self, module, *args, **kwargs): - args_list = list(args) - - for name, cpm in self.metadata.items(): - if isinstance(cpm, ContextParallelInput) and cpm.split_output: - continue - - # Maybe the parameter was passed as a keyword argument - input_val, is_kwarg, index = self.module_forward_metadata._get_parameter_from_args_kwargs( - name, args_list, kwargs - ) - - if input_val is None: - continue - - # The input_val may be a tensor or list/tuple of tensors. In certain cases, user may specify to shard - # the output instead of input for a particular layer by setting split_output=True - if isinstance(input_val, torch.Tensor): - input_val = self._prepare_cp_input(input_val, cpm) - elif isinstance(input_val, (list, tuple)): - if len(input_val) != len(cpm): - raise ValueError( - f"Expected input model plan to have {len(input_val)} elements, but got {len(cpm)}." - ) - sharded_input_val = [] - for i, x in enumerate(input_val): - if torch.is_tensor(x) and not cpm[i].split_output: - x = self._prepare_cp_input(x, cpm[i]) - sharded_input_val.append(x) - input_val = sharded_input_val - else: - raise ValueError(f"Unsupported input type: {type(input_val)}") - - if is_kwarg: - kwargs[name] = input_val - elif index is not None and index < len(args_list): - args_list[index] = input_val - else: - raise ValueError( - f"An unexpected error occurred while processing the input '{name}'. Please open an " - f"issue at https://github.com/huggingface/diffusers/issues and provide a minimal reproducible " - f"example along with the full stack trace." - ) - - return tuple(args_list), kwargs - - def post_forward(self, module, output): - is_tensor = isinstance(output, torch.Tensor) - is_tensor_list = isinstance(output, (list, tuple)) and all(isinstance(x, torch.Tensor) for x in output) - - if not is_tensor and not is_tensor_list: - raise ValueError(f"Expected output to be a tensor or a list/tuple of tensors, but got {type(output)}.") - - output = [output] if is_tensor else list(output) - for index, cpm in self.metadata.items(): - if not isinstance(cpm, ContextParallelInput) or not cpm.split_output: - continue - if index >= len(output): - raise ValueError(f"Index {index} out of bounds for output of length {len(output)}.") - current_output = output[index] - current_output = self._prepare_cp_input(current_output, cpm) - output[index] = current_output - - return output[0] if is_tensor else tuple(output) - - def _prepare_cp_input(self, x: torch.Tensor, cp_input: ContextParallelInput) -> torch.Tensor: - if cp_input.expected_dims is not None and x.dim() != cp_input.expected_dims: - logger.warning_once( - f"Expected input tensor to have {cp_input.expected_dims} dimensions, but got {x.dim()} dimensions, split will not be applied." - ) - return x - else: - if self.parallel_config.ulysses_anything or self.parallel_config.ring_anything: - return PartitionAnythingSharder.shard_anything( - x, cp_input.split_dim, self.parallel_config._flattened_mesh - ) - return EquipartitionSharder.shard(x, cp_input.split_dim, self.parallel_config._flattened_mesh) - - -class ContextParallelGatherHook(ModelHook): - def __init__(self, metadata: ContextParallelModelPlan, parallel_config: ContextParallelConfig) -> None: - super().__init__() - self.metadata = metadata - self.parallel_config = parallel_config - - def post_forward(self, module, output): - is_tensor = isinstance(output, torch.Tensor) - - if is_tensor: - output = [output] - elif not (isinstance(output, (list, tuple)) and all(isinstance(x, torch.Tensor) for x in output)): - raise ValueError(f"Expected output to be a tensor or a list/tuple of tensors, but got {type(output)}.") - - output = list(output) - - if len(output) != len(self.metadata): - raise ValueError(f"Expected output to have {len(self.metadata)} elements, but got {len(output)}.") - - for i, cpm in enumerate(self.metadata): - if cpm is None: - continue - if self.parallel_config.ulysses_anything or self.parallel_config.ring_anything: - output[i] = PartitionAnythingSharder.unshard_anything( - output[i], cpm.gather_dim, self.parallel_config._flattened_mesh - ) - else: - output[i] = EquipartitionSharder.unshard( - output[i], cpm.gather_dim, self.parallel_config._flattened_mesh - ) - - return output[0] if is_tensor else tuple(output) - - -class AllGatherFunction(torch.autograd.Function): - @staticmethod - def forward(ctx, tensor, dim, group): - ctx.dim = dim - ctx.group = group - ctx.world_size = torch.distributed.get_world_size(group) - ctx.rank = torch.distributed.get_rank(group) - return funcol.all_gather_tensor(tensor, dim, group=group) - - @staticmethod - def backward(ctx, grad_output): - grad_chunks = torch.chunk(grad_output, ctx.world_size, dim=ctx.dim) - return grad_chunks[ctx.rank], None, None - - -class EquipartitionSharder: - @classmethod - def shard(cls, tensor: torch.Tensor, dim: int, mesh: torch.distributed.device_mesh.DeviceMesh) -> torch.Tensor: - # NOTE: the following assertion does not have to be true in general. We simply enforce it for now - # because the alternate case has not yet been tested/required for any model. - assert tensor.size()[dim] % mesh.size() == 0, ( - "Tensor size along dimension to be sharded must be divisible by mesh size" - ) - - # The following is not fullgraph compatible with Dynamo (fails in DeviceMesh.get_rank) - # return tensor.chunk(mesh.size(), dim=dim)[mesh.get_rank()] - - return tensor.chunk(mesh.size(), dim=dim)[torch.distributed.get_rank(mesh.get_group())] - - @classmethod - def unshard(cls, tensor: torch.Tensor, dim: int, mesh: torch.distributed.device_mesh.DeviceMesh) -> torch.Tensor: - tensor = tensor.contiguous() - tensor = AllGatherFunction.apply(tensor, dim, mesh.get_group()) - return tensor - - -class AllGatherAnythingFunction(torch.autograd.Function): - @staticmethod - def forward(ctx, tensor: torch.Tensor, dim: int, group: dist.device_mesh.DeviceMesh): - ctx.dim = dim - ctx.group = group - ctx.world_size = dist.get_world_size(group) - ctx.rank = dist.get_rank(group) - gathered_tensor = _all_gather_anything(tensor, dim, group) - return gathered_tensor - - @staticmethod - def backward(ctx, grad_output): - # NOTE: We use `tensor_split` instead of chunk, because the `chunk` - # function may return fewer than the specified number of chunks! - grad_splits = torch.tensor_split(grad_output, ctx.world_size, dim=ctx.dim) - return grad_splits[ctx.rank], None, None - - -class PartitionAnythingSharder: - @classmethod - def shard_anything( - cls, tensor: torch.Tensor, dim: int, mesh: torch.distributed.device_mesh.DeviceMesh - ) -> torch.Tensor: - assert tensor.size()[dim] >= mesh.size(), ( - f"Cannot shard tensor of size {tensor.size()} along dim {dim} across mesh of size {mesh.size()}." - ) - # NOTE: We use `tensor_split` instead of chunk, because the `chunk` - # function may return fewer than the specified number of chunks! - return tensor.tensor_split(mesh.size(), dim=dim)[dist.get_rank(mesh.get_group())] - - @classmethod - def unshard_anything( - cls, tensor: torch.Tensor, dim: int, mesh: torch.distributed.device_mesh.DeviceMesh - ) -> torch.Tensor: - tensor = tensor.contiguous() - tensor = AllGatherAnythingFunction.apply(tensor, dim, mesh.get_group()) - return tensor - - -@lru_cache_unless_export(maxsize=64) -def _fill_gather_shapes(shape: tuple[int], gather_dims: tuple[int], dim: int, world_size: int) -> list[list[int]]: - gather_shapes = [] - for i in range(world_size): - rank_shape = list(copy.deepcopy(shape)) - rank_shape[dim] = gather_dims[i] - gather_shapes.append(rank_shape) - return gather_shapes - - -@maybe_allow_in_graph -def _all_gather_anything(tensor: torch.Tensor, dim: int, group: dist.device_mesh.DeviceMesh) -> torch.Tensor: - world_size = dist.get_world_size(group=group) - - tensor = tensor.contiguous() - shape = tensor.shape - rank_dim = shape[dim] - gather_dims = gather_size_by_comm(rank_dim, group) - - gather_shapes = _fill_gather_shapes(tuple(shape), tuple(gather_dims), dim, world_size) - - gathered_tensors = [torch.empty(shape, device=tensor.device, dtype=tensor.dtype) for shape in gather_shapes] - - dist.all_gather(gathered_tensors, tensor, group=group) - gathered_tensor = torch.cat(gathered_tensors, dim=dim) - return gathered_tensor - - -def _get_submodule_by_name(model: torch.nn.Module, name: str) -> torch.nn.Module | list[torch.nn.Module]: - if name.count("*") > 1: - raise ValueError("Wildcard '*' can only be used once in the name") - return _find_submodule_by_name(model, name) - - -def _find_submodule_by_name(model: torch.nn.Module, name: str) -> torch.nn.Module | list[torch.nn.Module]: - if name == "": - return model - first_atom, remaining_name = name.split(".", 1) if "." in name else (name, "") - if first_atom == "*": - if not isinstance(model, torch.nn.ModuleList): - raise ValueError("Wildcard '*' can only be used with ModuleList") - submodules = [] - for submodule in model: - subsubmodules = _find_submodule_by_name(submodule, remaining_name) - if not isinstance(subsubmodules, list): - subsubmodules = [subsubmodules] - submodules.extend(subsubmodules) - return submodules - else: - if hasattr(model, first_atom): - submodule = getattr(model, first_atom) - return _find_submodule_by_name(submodule, remaining_name) - else: - raise ValueError(f"'{first_atom}' is not a submodule of '{model.__class__.__name__}'") diff --git a/diffusers/hooks/faster_cache.py b/diffusers/hooks/faster_cache.py deleted file mode 100644 index 01544aa4b43022225b31df7b879baa37ac8f0b72..0000000000000000000000000000000000000000 --- a/diffusers/hooks/faster_cache.py +++ /dev/null @@ -1,654 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import re -from dataclasses import dataclass -from typing import Any, Callable - -import torch - -from ..models.attention import AttentionModuleMixin -from ..models.modeling_outputs import Transformer2DModelOutput -from ..utils import logging -from ._common import _ATTENTION_CLASSES -from .hooks import HookRegistry, ModelHook - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -_FASTER_CACHE_DENOISER_HOOK = "faster_cache_denoiser" -_FASTER_CACHE_BLOCK_HOOK = "faster_cache_block" -_SPATIAL_ATTENTION_BLOCK_IDENTIFIERS = ( - "^blocks.*attn", - "^transformer_blocks.*attn", - "^single_transformer_blocks.*attn", -) -_TEMPORAL_ATTENTION_BLOCK_IDENTIFIERS = ("^temporal_transformer_blocks.*attn",) -_TRANSFORMER_BLOCK_IDENTIFIERS = _SPATIAL_ATTENTION_BLOCK_IDENTIFIERS + _TEMPORAL_ATTENTION_BLOCK_IDENTIFIERS -_UNCOND_COND_INPUT_KWARGS_IDENTIFIERS = ( - "hidden_states", - "encoder_hidden_states", - "timestep", - "attention_mask", - "encoder_attention_mask", -) - - -@dataclass -class FasterCacheConfig: - r""" - Configuration for [FasterCache](https://huggingface.co/papers/2410.19355). - - Attributes: - spatial_attention_block_skip_range (`int`, defaults to `2`): - Calculate the attention states every `N` iterations. If this is set to `N`, the attention computation will - be skipped `N - 1` times (i.e., cached attention states will be reused) before computing the new attention - states again. - temporal_attention_block_skip_range (`int`, *optional*, defaults to `None`): - Calculate the attention states every `N` iterations. If this is set to `N`, the attention computation will - be skipped `N - 1` times (i.e., cached attention states will be reused) before computing the new attention - states again. - spatial_attention_timestep_skip_range (`tuple[float, float]`, defaults to `(-1, 681)`): - The timestep range within which the spatial attention computation can be skipped without a significant loss - in quality. This is to be determined by the user based on the underlying model. The first value in the - tuple is the lower bound and the second value is the upper bound. Typically, diffusion timesteps for - denoising are in the reversed range of 0 to 1000 (i.e. denoising starts at timestep 1000 and ends at - timestep 0). For the default values, this would mean that the spatial attention computation skipping will - be applicable only after denoising timestep 681 is reached, and continue until the end of the denoising - process. - temporal_attention_timestep_skip_range (`tuple[float, float]`, *optional*, defaults to `None`): - The timestep range within which the temporal attention computation can be skipped without a significant - loss in quality. This is to be determined by the user based on the underlying model. The first value in the - tuple is the lower bound and the second value is the upper bound. Typically, diffusion timesteps for - denoising are in the reversed range of 0 to 1000 (i.e. denoising starts at timestep 1000 and ends at - timestep 0). - low_frequency_weight_update_timestep_range (`tuple[int, int]`, defaults to `(99, 901)`): - The timestep range within which the low frequency weight scaling update is applied. The first value in the - tuple is the lower bound and the second value is the upper bound of the timestep range. The callback - function for the update is called only within this range. - high_frequency_weight_update_timestep_range (`tuple[int, int]`, defaults to `(-1, 301)`): - The timestep range within which the high frequency weight scaling update is applied. The first value in the - tuple is the lower bound and the second value is the upper bound of the timestep range. The callback - function for the update is called only within this range. - alpha_low_frequency (`float`, defaults to `1.1`): - The weight to scale the low frequency updates by. This is used to approximate the unconditional branch from - the conditional branch outputs. - alpha_high_frequency (`float`, defaults to `1.1`): - The weight to scale the high frequency updates by. This is used to approximate the unconditional branch - from the conditional branch outputs. - unconditional_batch_skip_range (`int`, defaults to `5`): - Process the unconditional branch every `N` iterations. If this is set to `N`, the unconditional branch - computation will be skipped `N - 1` times (i.e., cached unconditional branch states will be reused) before - computing the new unconditional branch states again. - unconditional_batch_timestep_skip_range (`tuple[float, float]`, defaults to `(-1, 641)`): - The timestep range within which the unconditional branch computation can be skipped without a significant - loss in quality. This is to be determined by the user based on the underlying model. The first value in the - tuple is the lower bound and the second value is the upper bound. - spatial_attention_block_identifiers (`tuple[str, ...]`, defaults to `("blocks.*attn1", "transformer_blocks.*attn1", "single_transformer_blocks.*attn1")`): - The identifiers to match the spatial attention blocks in the model. If the name of the block contains any - of these identifiers, FasterCache will be applied to that block. This can either be the full layer names, - partial layer names, or regex patterns. Matching will always be done using a regex match. - temporal_attention_block_identifiers (`tuple[str, ...]`, defaults to `("temporal_transformer_blocks.*attn1",)`): - The identifiers to match the temporal attention blocks in the model. If the name of the block contains any - of these identifiers, FasterCache will be applied to that block. This can either be the full layer names, - partial layer names, or regex patterns. Matching will always be done using a regex match. - attention_weight_callback (`Callable[[torch.nn.Module], float]`, defaults to `None`): - The callback function to determine the weight to scale the attention outputs by. This function should take - the attention module as input and return a float value. This is used to approximate the unconditional - branch from the conditional branch outputs. If not provided, the default weight is 0.5 for all timesteps. - Typically, as described in the paper, this weight should gradually increase from 0 to 1 as the inference - progresses. Users are encouraged to experiment and provide custom weight schedules that take into account - the number of inference steps and underlying model behaviour as denoising progresses. - low_frequency_weight_callback (`Callable[[torch.nn.Module], float]`, defaults to `None`): - The callback function to determine the weight to scale the low frequency updates by. If not provided, the - default weight is 1.1 for timesteps within the range specified (as described in the paper). - high_frequency_weight_callback (`Callable[[torch.nn.Module], float]`, defaults to `None`): - The callback function to determine the weight to scale the high frequency updates by. If not provided, the - default weight is 1.1 for timesteps within the range specified (as described in the paper). - tensor_format (`str`, defaults to `"BCFHW"`): - The format of the input tensors. This should be one of `"BCFHW"`, `"BFCHW"`, or `"BCHW"`. The format is - used to split individual latent frames in order for low and high frequency components to be computed. - is_guidance_distilled (`bool`, defaults to `False`): - Whether the model is guidance distilled or not. If the model is guidance distilled, FasterCache will not be - applied at the denoiser-level to skip the unconditional branch computation (as there is none). - _unconditional_conditional_input_kwargs_identifiers (`list[str]`, defaults to `("hidden_states", "encoder_hidden_states", "timestep", "attention_mask", "encoder_attention_mask")`): - The identifiers to match the input kwargs that contain the batchwise-concatenated unconditional and - conditional inputs. If the name of the input kwargs contains any of these identifiers, FasterCache will - split the inputs into unconditional and conditional branches. This must be a list of exact input kwargs - names that contain the batchwise-concatenated unconditional and conditional inputs. - """ - - # In the paper and codebase, they hardcode these values to 2. However, it can be made configurable - # after some testing. We default to 2 if these parameters are not provided. - spatial_attention_block_skip_range: int = 2 - temporal_attention_block_skip_range: int | None = None - - spatial_attention_timestep_skip_range: tuple[int, int] = (-1, 681) - temporal_attention_timestep_skip_range: tuple[int, int] = (-1, 681) - - # Indicator functions for low/high frequency as mentioned in Equation 11 of the paper - low_frequency_weight_update_timestep_range: tuple[int, int] = (99, 901) - high_frequency_weight_update_timestep_range: tuple[int, int] = (-1, 301) - - # ⍺1 and ⍺2 as mentioned in Equation 11 of the paper - alpha_low_frequency: float = 1.1 - alpha_high_frequency: float = 1.1 - - # n as described in CFG-Cache explanation in the paper - dependent on the model - unconditional_batch_skip_range: int = 5 - unconditional_batch_timestep_skip_range: tuple[int, int] = (-1, 641) - - spatial_attention_block_identifiers: tuple[str, ...] = _SPATIAL_ATTENTION_BLOCK_IDENTIFIERS - temporal_attention_block_identifiers: tuple[str, ...] = _TEMPORAL_ATTENTION_BLOCK_IDENTIFIERS - - attention_weight_callback: Callable[[torch.nn.Module], float] = None - low_frequency_weight_callback: Callable[[torch.nn.Module], float] = None - high_frequency_weight_callback: Callable[[torch.nn.Module], float] = None - - tensor_format: str = "BCFHW" - is_guidance_distilled: bool = False - - current_timestep_callback: Callable[[], int] = None - - _unconditional_conditional_input_kwargs_identifiers: list[str] = _UNCOND_COND_INPUT_KWARGS_IDENTIFIERS - - def __repr__(self) -> str: - return ( - f"FasterCacheConfig(\n" - f" spatial_attention_block_skip_range={self.spatial_attention_block_skip_range},\n" - f" temporal_attention_block_skip_range={self.temporal_attention_block_skip_range},\n" - f" spatial_attention_timestep_skip_range={self.spatial_attention_timestep_skip_range},\n" - f" temporal_attention_timestep_skip_range={self.temporal_attention_timestep_skip_range},\n" - f" low_frequency_weight_update_timestep_range={self.low_frequency_weight_update_timestep_range},\n" - f" high_frequency_weight_update_timestep_range={self.high_frequency_weight_update_timestep_range},\n" - f" alpha_low_frequency={self.alpha_low_frequency},\n" - f" alpha_high_frequency={self.alpha_high_frequency},\n" - f" unconditional_batch_skip_range={self.unconditional_batch_skip_range},\n" - f" unconditional_batch_timestep_skip_range={self.unconditional_batch_timestep_skip_range},\n" - f" spatial_attention_block_identifiers={self.spatial_attention_block_identifiers},\n" - f" temporal_attention_block_identifiers={self.temporal_attention_block_identifiers},\n" - f" tensor_format={self.tensor_format},\n" - f")" - ) - - -class FasterCacheDenoiserState: - r""" - State for [FasterCache](https://huggingface.co/papers/2410.19355) top-level denoiser module. - """ - - def __init__(self) -> None: - self.iteration: int = 0 - self.low_frequency_delta: torch.Tensor = None - self.high_frequency_delta: torch.Tensor = None - - def reset(self): - self.iteration = 0 - self.low_frequency_delta = None - self.high_frequency_delta = None - - -class FasterCacheBlockState: - r""" - State for [FasterCache](https://huggingface.co/papers/2410.19355). Every underlying block that FasterCache is - applied to will have an instance of this state. - """ - - def __init__(self) -> None: - self.iteration: int = 0 - self.batch_size: int = None - self.cache: tuple[torch.Tensor, torch.Tensor] = None - - def reset(self): - self.iteration = 0 - self.batch_size = None - self.cache = None - - -class FasterCacheDenoiserHook(ModelHook): - _is_stateful = True - - def __init__( - self, - unconditional_batch_skip_range: int, - unconditional_batch_timestep_skip_range: tuple[int, int], - tensor_format: str, - is_guidance_distilled: bool, - uncond_cond_input_kwargs_identifiers: list[str], - current_timestep_callback: Callable[[], int], - low_frequency_weight_callback: Callable[[torch.nn.Module], torch.Tensor], - high_frequency_weight_callback: Callable[[torch.nn.Module], torch.Tensor], - ) -> None: - super().__init__() - - self.unconditional_batch_skip_range = unconditional_batch_skip_range - self.unconditional_batch_timestep_skip_range = unconditional_batch_timestep_skip_range - # We can't easily detect what args are to be split in unconditional and conditional branches. We - # can only do it for kwargs, hence they are the only ones we split. The args are passed as-is. - # If a model is to be made compatible with FasterCache, the user must ensure that the inputs that - # contain batchwise-concatenated unconditional and conditional inputs are passed as kwargs. - self.uncond_cond_input_kwargs_identifiers = uncond_cond_input_kwargs_identifiers - self.tensor_format = tensor_format - self.is_guidance_distilled = is_guidance_distilled - - self.current_timestep_callback = current_timestep_callback - self.low_frequency_weight_callback = low_frequency_weight_callback - self.high_frequency_weight_callback = high_frequency_weight_callback - - def initialize_hook(self, module): - self.state = FasterCacheDenoiserState() - return module - - @staticmethod - def _get_cond_input(input: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - # Note: this method assumes that the input tensor is batchwise-concatenated with unconditional inputs - # followed by conditional inputs. - _, cond = input.chunk(2, dim=0) - return cond - - def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any: - # Split the unconditional and conditional inputs. We only want to infer the conditional branch if the - # requirements for skipping the unconditional branch are met as described in the paper. - # We skip the unconditional branch only if the following conditions are met: - # 1. We have completed at least one iteration of the denoiser - # 2. The current timestep is within the range specified by the user. This is the optimal timestep range - # where approximating the unconditional branch from the computation of the conditional branch is possible - # without a significant loss in quality. - # 3. The current iteration is not a multiple of the unconditional batch skip range. This is done so that - # we compute the unconditional branch at least once every few iterations to ensure minimal quality loss. - is_within_timestep_range = ( - self.unconditional_batch_timestep_skip_range[0] - < self.current_timestep_callback() - < self.unconditional_batch_timestep_skip_range[1] - ) - should_skip_uncond = ( - self.state.iteration > 0 - and is_within_timestep_range - and self.state.iteration % self.unconditional_batch_skip_range != 0 - and not self.is_guidance_distilled - ) - - if should_skip_uncond: - is_any_kwarg_uncond = any(k in self.uncond_cond_input_kwargs_identifiers for k in kwargs.keys()) - if is_any_kwarg_uncond: - logger.debug("FasterCache - Skipping unconditional branch computation") - args = tuple([self._get_cond_input(arg) if torch.is_tensor(arg) else arg for arg in args]) - kwargs = { - k: v if k not in self.uncond_cond_input_kwargs_identifiers else self._get_cond_input(v) - for k, v in kwargs.items() - } - - output = self.fn_ref.original_forward(*args, **kwargs) - - if self.is_guidance_distilled: - self.state.iteration += 1 - return output - - if torch.is_tensor(output): - hidden_states = output - elif isinstance(output, (tuple, Transformer2DModelOutput)): - hidden_states = output[0] - - batch_size = hidden_states.size(0) - - if should_skip_uncond: - self.state.low_frequency_delta = self.state.low_frequency_delta * self.low_frequency_weight_callback( - module - ) - self.state.high_frequency_delta = self.state.high_frequency_delta * self.high_frequency_weight_callback( - module - ) - - if self.tensor_format == "BCFHW": - hidden_states = hidden_states.permute(0, 2, 1, 3, 4) - if self.tensor_format == "BCFHW" or self.tensor_format == "BFCHW": - hidden_states = hidden_states.flatten(0, 1) - - low_freq_cond, high_freq_cond = _split_low_high_freq(hidden_states.float()) - - # Approximate/compute the unconditional branch outputs as described in Equation 9 and 10 of the paper - low_freq_uncond = self.state.low_frequency_delta + low_freq_cond - high_freq_uncond = self.state.high_frequency_delta + high_freq_cond - uncond_freq = low_freq_uncond + high_freq_uncond - - uncond_states = torch.fft.ifftshift(uncond_freq) - uncond_states = torch.fft.ifft2(uncond_states).real - - if self.tensor_format == "BCFHW" or self.tensor_format == "BFCHW": - uncond_states = uncond_states.unflatten(0, (batch_size, -1)) - hidden_states = hidden_states.unflatten(0, (batch_size, -1)) - if self.tensor_format == "BCFHW": - uncond_states = uncond_states.permute(0, 2, 1, 3, 4) - hidden_states = hidden_states.permute(0, 2, 1, 3, 4) - - # Concatenate the approximated unconditional and predicted conditional branches - uncond_states = uncond_states.to(hidden_states.dtype) - hidden_states = torch.cat([uncond_states, hidden_states], dim=0) - else: - uncond_states, cond_states = hidden_states.chunk(2, dim=0) - if self.tensor_format == "BCFHW": - uncond_states = uncond_states.permute(0, 2, 1, 3, 4) - cond_states = cond_states.permute(0, 2, 1, 3, 4) - if self.tensor_format == "BCFHW" or self.tensor_format == "BFCHW": - uncond_states = uncond_states.flatten(0, 1) - cond_states = cond_states.flatten(0, 1) - - low_freq_uncond, high_freq_uncond = _split_low_high_freq(uncond_states.float()) - low_freq_cond, high_freq_cond = _split_low_high_freq(cond_states.float()) - self.state.low_frequency_delta = low_freq_uncond - low_freq_cond - self.state.high_frequency_delta = high_freq_uncond - high_freq_cond - - self.state.iteration += 1 - if torch.is_tensor(output): - output = hidden_states - elif isinstance(output, tuple): - output = (hidden_states, *output[1:]) - else: - output.sample = hidden_states - - return output - - def reset_state(self, module: torch.nn.Module) -> torch.nn.Module: - self.state.reset() - return module - - -class FasterCacheBlockHook(ModelHook): - _is_stateful = True - - def __init__( - self, - block_skip_range: int, - timestep_skip_range: tuple[int, int], - is_guidance_distilled: bool, - weight_callback: Callable[[torch.nn.Module], float], - current_timestep_callback: Callable[[], int], - ) -> None: - super().__init__() - - self.block_skip_range = block_skip_range - self.timestep_skip_range = timestep_skip_range - self.is_guidance_distilled = is_guidance_distilled - - self.weight_callback = weight_callback - self.current_timestep_callback = current_timestep_callback - - def initialize_hook(self, module): - self.state = FasterCacheBlockState() - return module - - def _compute_approximated_attention_output( - self, t_2_output: torch.Tensor, t_output: torch.Tensor, weight: float, batch_size: int - ) -> torch.Tensor: - if t_2_output.size(0) != batch_size: - # The cache t_2_output contains both batchwise-concatenated unconditional-conditional branch outputs. Just - # take the conditional branch outputs. - assert t_2_output.size(0) == 2 * batch_size - t_2_output = t_2_output[batch_size:] - if t_output.size(0) != batch_size: - # The cache t_output contains both batchwise-concatenated unconditional-conditional branch outputs. Just - # take the conditional branch outputs. - assert t_output.size(0) == 2 * batch_size - t_output = t_output[batch_size:] - return t_output + (t_output - t_2_output) * weight - - def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any: - batch_size = [ - *[arg.size(0) for arg in args if torch.is_tensor(arg)], - *[v.size(0) for v in kwargs.values() if torch.is_tensor(v)], - ][0] - if self.state.batch_size is None: - # Will be updated on first forward pass through the denoiser - self.state.batch_size = batch_size - - # If we have to skip due to the skip conditions, then let's skip as expected. - # But, we can't skip if the denoiser wants to infer both unconditional and conditional branches. This - # is because the expected output shapes of attention layer will not match if we only return values from - # the cache (which only caches conditional branch outputs). So, if state.batch_size (which is the true - # unconditional-conditional batch size) is same as the current batch size, we don't perform the layer - # skip. Otherwise, we conditionally skip the layer based on what state.skip_callback returns. - is_within_timestep_range = ( - self.timestep_skip_range[0] < self.current_timestep_callback() < self.timestep_skip_range[1] - ) - if not is_within_timestep_range: - should_skip_attention = False - else: - should_compute_attention = self.state.iteration > 0 and self.state.iteration % self.block_skip_range == 0 - should_skip_attention = not should_compute_attention - if should_skip_attention: - should_skip_attention = self.is_guidance_distilled or self.state.batch_size != batch_size - - if should_skip_attention: - logger.debug("FasterCache - Skipping attention and using approximation") - if torch.is_tensor(self.state.cache[-1]): - t_2_output, t_output = self.state.cache - weight = self.weight_callback(module) - output = self._compute_approximated_attention_output(t_2_output, t_output, weight, batch_size) - else: - # The cache contains multiple tensors from past N iterations (N=2 for FasterCache). We need to handle all of them. - # Diffusers blocks can return multiple tensors - let's call them [A, B, C, ...] for simplicity. - # 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 - # a forward pass of the block. We need to compute the approximated output for each of these tensors. - # The zip(*state.cache) operation will give us [(A_1, A_2, ...), (B_1, B_2, ...), (C_1, C_2, ...), ...] which - # allows us to compute the approximated attention output for each tensor in the cache. - output = () - for t_2_output, t_output in zip(*self.state.cache): - result = self._compute_approximated_attention_output( - t_2_output, t_output, self.weight_callback(module), batch_size - ) - output += (result,) - else: - logger.debug("FasterCache - Computing attention") - output = self.fn_ref.original_forward(*args, **kwargs) - - # Note that the following condition for getting hidden_states should suffice since Diffusers blocks either return - # a single hidden_states tensor, or a tuple of (hidden_states, encoder_hidden_states) tensors. We need to handle - # both cases. - if torch.is_tensor(output): - cache_output = output - if not self.is_guidance_distilled and cache_output.size(0) == self.state.batch_size: - # The output here can be both unconditional-conditional branch outputs or just conditional branch outputs. - # This is determined at the higher-level denoiser module. We only want to cache the conditional branch outputs. - cache_output = cache_output.chunk(2, dim=0)[1] - else: - # Cache all return values and perform the same operation as above - cache_output = () - for out in output: - if not self.is_guidance_distilled and out.size(0) == self.state.batch_size: - out = out.chunk(2, dim=0)[1] - cache_output += (out,) - - if self.state.cache is None: - self.state.cache = [cache_output, cache_output] - else: - self.state.cache = [self.state.cache[-1], cache_output] - - self.state.iteration += 1 - return output - - def reset_state(self, module: torch.nn.Module) -> torch.nn.Module: - self.state.reset() - return module - - -def apply_faster_cache(module: torch.nn.Module, config: FasterCacheConfig) -> None: - r""" - Applies [FasterCache](https://huggingface.co/papers/2410.19355) to a given pipeline. - - Args: - module (`torch.nn.Module`): - The pytorch module to apply FasterCache to. Typically, this should be a transformer architecture supported - in Diffusers, such as `CogVideoXTransformer3DModel`, but external implementations may also work. - config (`FasterCacheConfig`): - The configuration to use for FasterCache. - - Example: - ```python - >>> import torch - >>> from diffusers import CogVideoXPipeline, FasterCacheConfig, apply_faster_cache - - >>> pipe = CogVideoXPipeline.from_pretrained("THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16) - >>> pipe.to("cuda") - - >>> config = FasterCacheConfig( - ... spatial_attention_block_skip_range=2, - ... spatial_attention_timestep_skip_range=(-1, 681), - ... low_frequency_weight_update_timestep_range=(99, 641), - ... high_frequency_weight_update_timestep_range=(-1, 301), - ... spatial_attention_block_identifiers=["transformer_blocks"], - ... attention_weight_callback=lambda _: 0.3, - ... tensor_format="BFCHW", - ... ) - >>> apply_faster_cache(pipe.transformer, config) - ``` - """ - - logger.warning( - "FasterCache is a purely experimental feature and may not work as expected. Not all models support FasterCache. " - "The API is subject to change in future releases, with no guarantee of backward compatibility. Please report any issues at " - "https://github.com/huggingface/diffusers/issues." - ) - - if config.attention_weight_callback is None: - # If the user has not provided a weight callback, we default to 0.5 for all timesteps. - # In the paper, they recommend using a gradually increasing weight from 0 to 1 as the inference progresses, but - # this depends from model-to-model. It is required by the user to provide a weight callback if they want to - # use a different weight function. Defaulting to 0.5 works well in practice for most cases. - logger.warning( - "No `attention_weight_callback` provided when enabling FasterCache. Defaulting to using a weight of 0.5 for all timesteps." - ) - config.attention_weight_callback = lambda _: 0.5 - - if config.low_frequency_weight_callback is None: - logger.debug( - "Low frequency weight callback not provided when enabling FasterCache. Defaulting to behaviour described in the paper." - ) - - def low_frequency_weight_callback(module: torch.nn.Module) -> float: - is_within_range = ( - config.low_frequency_weight_update_timestep_range[0] - < config.current_timestep_callback() - < config.low_frequency_weight_update_timestep_range[1] - ) - return config.alpha_low_frequency if is_within_range else 1.0 - - config.low_frequency_weight_callback = low_frequency_weight_callback - - if config.high_frequency_weight_callback is None: - logger.debug( - "High frequency weight callback not provided when enabling FasterCache. Defaulting to behaviour described in the paper." - ) - - def high_frequency_weight_callback(module: torch.nn.Module) -> float: - is_within_range = ( - config.high_frequency_weight_update_timestep_range[0] - < config.current_timestep_callback() - < config.high_frequency_weight_update_timestep_range[1] - ) - return config.alpha_high_frequency if is_within_range else 1.0 - - config.high_frequency_weight_callback = high_frequency_weight_callback - - supported_tensor_formats = ["BCFHW", "BFCHW", "BCHW"] # TODO(aryan): Support BSC for LTX Video - if config.tensor_format not in supported_tensor_formats: - raise ValueError(f"`tensor_format` must be one of {supported_tensor_formats}, but got {config.tensor_format}.") - - _apply_faster_cache_on_denoiser(module, config) - - for name, submodule in module.named_modules(): - if not isinstance(submodule, _ATTENTION_CLASSES): - continue - if any(re.search(identifier, name) is not None for identifier in _TRANSFORMER_BLOCK_IDENTIFIERS): - _apply_faster_cache_on_attention_class(name, submodule, config) - - -def _apply_faster_cache_on_denoiser(module: torch.nn.Module, config: FasterCacheConfig) -> None: - hook = FasterCacheDenoiserHook( - config.unconditional_batch_skip_range, - config.unconditional_batch_timestep_skip_range, - config.tensor_format, - config.is_guidance_distilled, - config._unconditional_conditional_input_kwargs_identifiers, - config.current_timestep_callback, - config.low_frequency_weight_callback, - config.high_frequency_weight_callback, - ) - registry = HookRegistry.check_if_exists_or_initialize(module) - registry.register_hook(hook, _FASTER_CACHE_DENOISER_HOOK) - - -def _apply_faster_cache_on_attention_class(name: str, module: AttentionModuleMixin, config: FasterCacheConfig) -> None: - is_spatial_self_attention = ( - any(re.search(identifier, name) is not None for identifier in config.spatial_attention_block_identifiers) - and config.spatial_attention_block_skip_range is not None - and not getattr(module, "is_cross_attention", False) - ) - is_temporal_self_attention = ( - any(re.search(identifier, name) is not None for identifier in config.temporal_attention_block_identifiers) - and config.temporal_attention_block_skip_range is not None - and not module.is_cross_attention - ) - - block_skip_range, timestep_skip_range, block_type = None, None, None - if is_spatial_self_attention: - block_skip_range = config.spatial_attention_block_skip_range - timestep_skip_range = config.spatial_attention_timestep_skip_range - block_type = "spatial" - elif is_temporal_self_attention: - block_skip_range = config.temporal_attention_block_skip_range - timestep_skip_range = config.temporal_attention_timestep_skip_range - block_type = "temporal" - - if block_skip_range is None or timestep_skip_range is None: - logger.debug( - f'Unable to apply FasterCache to the selected layer: "{name}" because it does ' - f"not match any of the required criteria for spatial or temporal attention layers. Note, " - f"however, that this layer may still be valid for applying PAB. Please specify the correct " - f"block identifiers in the configuration or use the specialized `apply_faster_cache_on_module` " - f"function to apply FasterCache to this layer." - ) - return - - logger.debug(f"Enabling FasterCache ({block_type}) for layer: {name}") - hook = FasterCacheBlockHook( - block_skip_range, - timestep_skip_range, - config.is_guidance_distilled, - config.attention_weight_callback, - config.current_timestep_callback, - ) - registry = HookRegistry.check_if_exists_or_initialize(module) - registry.register_hook(hook, _FASTER_CACHE_BLOCK_HOOK) - - -# Reference: https://github.com/Vchitect/FasterCache/blob/fab32c15014636dc854948319c0a9a8d92c7acb4/scripts/latte/faster_cache_sample_latte.py#L127C1-L143C39 -@torch.no_grad() -def _split_low_high_freq(x): - fft = torch.fft.fft2(x) - fft_shifted = torch.fft.fftshift(fft) - height, width = x.shape[-2:] - radius = min(height, width) // 5 - - y_grid, x_grid = torch.meshgrid(torch.arange(height), torch.arange(width)) - center_x, center_y = width // 2, height // 2 - mask = (x_grid - center_x) ** 2 + (y_grid - center_y) ** 2 <= radius**2 - - low_freq_mask = mask.unsqueeze(0).unsqueeze(0).to(x.device) - high_freq_mask = ~low_freq_mask - - low_freq_fft = fft_shifted * low_freq_mask - high_freq_fft = fft_shifted * high_freq_mask - - return low_freq_fft, high_freq_fft diff --git a/diffusers/hooks/first_block_cache.py b/diffusers/hooks/first_block_cache.py deleted file mode 100644 index 685ccd3836742d140cfdacd69c604d24fb2284b3..0000000000000000000000000000000000000000 --- a/diffusers/hooks/first_block_cache.py +++ /dev/null @@ -1,258 +0,0 @@ -# Copyright 2024 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass - -import torch - -from ..utils import get_logger -from ..utils.torch_utils import unwrap_module -from ._common import _ALL_TRANSFORMER_BLOCK_IDENTIFIERS -from ._helpers import TransformerBlockRegistry -from .hooks import BaseState, HookRegistry, ModelHook, StateManager - - -logger = get_logger(__name__) # pylint: disable=invalid-name - -_FBC_LEADER_BLOCK_HOOK = "fbc_leader_block_hook" -_FBC_BLOCK_HOOK = "fbc_block_hook" - - -@dataclass -class FirstBlockCacheConfig: - r""" - Configuration for [First Block - Cache](https://github.com/chengzeyi/ParaAttention/blob/7a266123671b55e7e5a2fe9af3121f07a36afc78/README.md#first-block-cache-our-dynamic-caching). - - Args: - threshold (`float`, defaults to `0.05`): - The threshold to determine whether or not a forward pass through all layers of the model is required. A - higher threshold usually results in a forward pass through a lower number of layers and faster inference, - but might lead to poorer generation quality. A lower threshold may not result in significant generation - speedup. The threshold is compared against the absmean difference of the residuals between the current and - cached outputs from the first transformer block. If the difference is below the threshold, the forward pass - is skipped. - """ - - threshold: float = 0.05 - - -class FBCSharedBlockState(BaseState): - def __init__(self) -> None: - super().__init__() - - self.head_block_output: torch.Tensor | tuple[torch.Tensor, ...] = None - self.head_block_residual: torch.Tensor = None - self.tail_block_residuals: torch.Tensor | tuple[torch.Tensor, ...] = None - self.should_compute: bool = True - - def reset(self): - self.tail_block_residuals = None - self.should_compute = True - - -class FBCHeadBlockHook(ModelHook): - _is_stateful = True - - def __init__(self, state_manager: StateManager, threshold: float): - self.state_manager = state_manager - self.threshold = threshold - self._metadata = None - - def initialize_hook(self, module): - unwrapped_module = unwrap_module(module) - self._metadata = TransformerBlockRegistry.get(unwrapped_module.__class__) - return module - - def new_forward(self, module: torch.nn.Module, *args, **kwargs): - original_hidden_states = self._metadata._get_parameter_from_args_kwargs("hidden_states", args, kwargs) - - output = self.fn_ref.original_forward(*args, **kwargs) - is_output_tuple = isinstance(output, tuple) - - if is_output_tuple: - hidden_states_residual = output[self._metadata.return_hidden_states_index] - original_hidden_states - else: - hidden_states_residual = output - original_hidden_states - - shared_state: FBCSharedBlockState = self.state_manager.get_state() - hidden_states = encoder_hidden_states = None - should_compute = self._should_compute_remaining_blocks(hidden_states_residual) - shared_state.should_compute = should_compute - - if not should_compute: - # Apply caching - if is_output_tuple: - hidden_states = ( - shared_state.tail_block_residuals[0] + output[self._metadata.return_hidden_states_index] - ) - else: - hidden_states = shared_state.tail_block_residuals[0] + output - - if self._metadata.return_encoder_hidden_states_index is not None: - assert is_output_tuple - encoder_hidden_states = ( - shared_state.tail_block_residuals[1] + output[self._metadata.return_encoder_hidden_states_index] - ) - - if is_output_tuple: - return_output = [None] * len(output) - return_output[self._metadata.return_hidden_states_index] = hidden_states - return_output[self._metadata.return_encoder_hidden_states_index] = encoder_hidden_states - return_output = tuple(return_output) - else: - return_output = hidden_states - output = return_output - else: - if is_output_tuple: - head_block_output = [None] * len(output) - head_block_output[0] = output[self._metadata.return_hidden_states_index] - head_block_output[1] = output[self._metadata.return_encoder_hidden_states_index] - else: - head_block_output = output - shared_state.head_block_output = head_block_output - shared_state.head_block_residual = hidden_states_residual - - return output - - def reset_state(self, module): - self.state_manager.reset() - return module - - @torch.compiler.disable - def _should_compute_remaining_blocks(self, hidden_states_residual: torch.Tensor) -> bool: - shared_state = self.state_manager.get_state() - if shared_state.head_block_residual is None: - return True - prev_hidden_states_residual = shared_state.head_block_residual - absmean = (hidden_states_residual - prev_hidden_states_residual).abs().mean() - prev_hidden_states_absmean = prev_hidden_states_residual.abs().mean() - diff = (absmean / prev_hidden_states_absmean).item() - return diff > self.threshold - - -class FBCBlockHook(ModelHook): - def __init__(self, state_manager: StateManager, is_tail: bool = False): - super().__init__() - self.state_manager = state_manager - self.is_tail = is_tail - self._metadata = None - - def initialize_hook(self, module): - unwrapped_module = unwrap_module(module) - self._metadata = TransformerBlockRegistry.get(unwrapped_module.__class__) - return module - - def new_forward(self, module: torch.nn.Module, *args, **kwargs): - original_hidden_states = self._metadata._get_parameter_from_args_kwargs("hidden_states", args, kwargs) - original_encoder_hidden_states = None - if self._metadata.return_encoder_hidden_states_index is not None: - original_encoder_hidden_states = self._metadata._get_parameter_from_args_kwargs( - "encoder_hidden_states", args, kwargs - ) - - shared_state = self.state_manager.get_state() - - if shared_state.should_compute: - output = self.fn_ref.original_forward(*args, **kwargs) - if self.is_tail: - hidden_states_residual = encoder_hidden_states_residual = None - if isinstance(output, tuple): - hidden_states_residual = ( - output[self._metadata.return_hidden_states_index] - shared_state.head_block_output[0] - ) - encoder_hidden_states_residual = ( - output[self._metadata.return_encoder_hidden_states_index] - shared_state.head_block_output[1] - ) - else: - hidden_states_residual = output - shared_state.head_block_output - shared_state.tail_block_residuals = (hidden_states_residual, encoder_hidden_states_residual) - return output - - if original_encoder_hidden_states is None: - return_output = original_hidden_states - else: - return_output = [None, None] - return_output[self._metadata.return_hidden_states_index] = original_hidden_states - return_output[self._metadata.return_encoder_hidden_states_index] = original_encoder_hidden_states - return_output = tuple(return_output) - return return_output - - -def apply_first_block_cache(module: torch.nn.Module, config: FirstBlockCacheConfig) -> None: - """ - Applies [First Block - Cache](https://github.com/chengzeyi/ParaAttention/blob/4de137c5b96416489f06e43e19f2c14a772e28fd/README.md#first-block-cache-our-dynamic-caching) - to a given module. - - First Block Cache builds on the ideas of [TeaCache](https://huggingface.co/papers/2411.19108). It is much simpler - to implement generically for a wide range of models and has been integrated first for experimental purposes. - - Args: - module (`torch.nn.Module`): - The pytorch module to apply FBCache to. Typically, this should be a transformer architecture supported in - Diffusers, such as `CogVideoXTransformer3DModel`, but external implementations may also work. - config (`FirstBlockCacheConfig`): - The configuration to use for applying the FBCache method. - - Example: - ```python - >>> import torch - >>> from diffusers import CogView4Pipeline - >>> from diffusers.hooks import apply_first_block_cache, FirstBlockCacheConfig - - >>> pipe = CogView4Pipeline.from_pretrained("THUDM/CogView4-6B", torch_dtype=torch.bfloat16) - >>> pipe.to("cuda") - - >>> apply_first_block_cache(pipe.transformer, FirstBlockCacheConfig(threshold=0.2)) - - >>> prompt = "A photo of an astronaut riding a horse on mars" - >>> image = pipe(prompt, generator=torch.Generator().manual_seed(42)).images[0] - >>> image.save("output.png") - ``` - """ - - state_manager = StateManager(FBCSharedBlockState, (), {}) - remaining_blocks = [] - - for name, submodule in module.named_children(): - if name not in _ALL_TRANSFORMER_BLOCK_IDENTIFIERS or not isinstance(submodule, torch.nn.ModuleList): - continue - for index, block in enumerate(submodule): - remaining_blocks.append((f"{name}.{index}", block)) - - head_block_name, head_block = remaining_blocks.pop(0) - tail_block_name, tail_block = remaining_blocks.pop(-1) - - logger.debug(f"Applying FBCHeadBlockHook to '{head_block_name}'") - _apply_fbc_head_block_hook(head_block, state_manager, config.threshold) - - for name, block in remaining_blocks: - logger.debug(f"Applying FBCBlockHook to '{name}'") - _apply_fbc_block_hook(block, state_manager) - - logger.debug(f"Applying FBCBlockHook to tail block '{tail_block_name}'") - _apply_fbc_block_hook(tail_block, state_manager, is_tail=True) - - -def _apply_fbc_head_block_hook(block: torch.nn.Module, state_manager: StateManager, threshold: float) -> None: - registry = HookRegistry.check_if_exists_or_initialize(block) - hook = FBCHeadBlockHook(state_manager, threshold) - registry.register_hook(hook, _FBC_LEADER_BLOCK_HOOK) - - -def _apply_fbc_block_hook(block: torch.nn.Module, state_manager: StateManager, is_tail: bool = False) -> None: - registry = HookRegistry.check_if_exists_or_initialize(block) - hook = FBCBlockHook(state_manager, is_tail) - registry.register_hook(hook, _FBC_BLOCK_HOOK) diff --git a/diffusers/hooks/group_offloading.py b/diffusers/hooks/group_offloading.py deleted file mode 100644 index 10d3f0c245a1a4ace45f8f0e710b049ae2df35e6..0000000000000000000000000000000000000000 --- a/diffusers/hooks/group_offloading.py +++ /dev/null @@ -1,1056 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import hashlib -import os -from contextlib import contextmanager, nullcontext -from dataclasses import dataclass, replace -from enum import Enum -from typing import Set - -import safetensors.torch -import torch - -from ..utils import get_logger, is_accelerate_available, is_torchao_available -from ._common import _GO_LC_SUPPORTED_PYTORCH_LAYERS -from .hooks import HookRegistry, ModelHook - - -if is_accelerate_available(): - from accelerate.hooks import AlignDevicesHook, CpuOffload - from accelerate.utils import send_to_device - - -logger = get_logger(__name__) # pylint: disable=invalid-name - - -def _is_torchao_tensor(tensor: torch.Tensor) -> bool: - if not is_torchao_available(): - return False - from torchao.utils import TorchAOBaseTensor - - return isinstance(tensor, TorchAOBaseTensor) - - -def _get_torchao_inner_tensor_names(tensor: torch.Tensor) -> list[str]: - """Get names of all internal tensor data attributes from a TorchAO tensor.""" - cls = type(tensor) - names = list(getattr(cls, "tensor_data_names", [])) - for attr_name in getattr(cls, "optional_tensor_data_names", []): - if getattr(tensor, attr_name, None) is not None: - names.append(attr_name) - return names - - -def _swap_torchao_tensor(param: torch.Tensor, source: torch.Tensor) -> None: - """Move a TorchAO parameter to the device of `source` via `swap_tensors`. - - `param.data = source` does not work for `_make_wrapper_subclass` tensors because the `.data` setter only replaces - the outer wrapper storage while leaving the subclass's internal attributes (e.g. `.qdata`, `.scale`) on the - original device. `swap_tensors` swaps the full tensor contents in-place, preserving the parameter's identity so - that any dict keyed by `id(param)` remains valid. - - Refer to https://github.com/huggingface/diffusers/pull/13276#discussion_r2944471548 for the full discussion. - """ - torch.utils.swap_tensors(param, source) - - -def _restore_torchao_tensor(param: torch.Tensor, source: torch.Tensor) -> None: - """Restore internal tensor data of a TorchAO parameter from `source` without mutating `source`. - - Unlike `_swap_torchao_tensor` this copies attribute references one-by-one via `setattr` so that `source` is **not** - modified. Use this when `source` is a cached tensor that must remain unchanged (e.g. a pinned CPU copy in - `cpu_param_dict`). - """ - for attr_name in _get_torchao_inner_tensor_names(source): - setattr(param, attr_name, getattr(source, attr_name)) - - -def _record_stream_torchao_tensor(param: torch.Tensor, stream) -> None: - """Record stream for all internal tensors of a TorchAO parameter.""" - for attr_name in _get_torchao_inner_tensor_names(param): - getattr(param, attr_name).record_stream(stream) - - -# fmt: off -_GROUP_OFFLOADING = "group_offloading" -_LAYER_EXECUTION_TRACKER = "layer_execution_tracker" -_LAZY_PREFETCH_GROUP_OFFLOADING = "lazy_prefetch_group_offloading" -_GROUP_ID_LAZY_LEAF = "lazy_leafs" -# fmt: on - - -class GroupOffloadingType(str, Enum): - BLOCK_LEVEL = "block_level" - LEAF_LEVEL = "leaf_level" - - -@dataclass -class GroupOffloadingConfig: - onload_device: torch.device - offload_device: torch.device - offload_type: GroupOffloadingType - non_blocking: bool - record_stream: bool - low_cpu_mem_usage: bool - num_blocks_per_group: int | None = None - offload_to_disk_path: str | None = None - stream: torch.cuda.Stream | torch.Stream | None = None - block_modules: list[str] | None = None - exclude_kwargs: list[str] | None = None - module_prefix: str = "" - - -class ModuleGroup: - def __init__( - self, - modules: list[torch.nn.Module], - offload_device: torch.device, - onload_device: torch.device, - offload_leader: torch.nn.Module, - onload_leader: torch.nn.Module | None = None, - parameters: list[torch.nn.Parameter] | None = None, - buffers: list[torch.Tensor] | None = None, - non_blocking: bool = False, - stream: torch.cuda.Stream | torch.Stream | None = None, - record_stream: bool | None = False, - low_cpu_mem_usage: bool = False, - onload_self: bool = True, - offload_to_disk_path: str | None = None, - group_id: int | str | None = None, - ) -> None: - self.modules = modules - self.offload_device = offload_device - self.onload_device = onload_device - self.offload_leader = offload_leader - self.onload_leader = onload_leader - self.parameters = parameters or [] - self.buffers = buffers or [] - self.non_blocking = non_blocking or stream is not None - self.stream = stream - self.record_stream = record_stream - self.onload_self = onload_self - self.low_cpu_mem_usage = low_cpu_mem_usage - - self.offload_to_disk_path = offload_to_disk_path - self._is_offloaded_to_disk = False - - if self.offload_to_disk_path is not None: - # Instead of `group_id or str(id(self))` we do this because `group_id` can be "" as well. - self.group_id = group_id if group_id is not None else str(id(self)) - short_hash = _compute_group_hash(self.group_id) - self.safetensors_file_path = os.path.join(self.offload_to_disk_path, f"group_{short_hash}.safetensors") - - all_tensors = [] - for module in self.modules: - all_tensors.extend(list(module.parameters())) - all_tensors.extend(list(module.buffers())) - all_tensors.extend(self.parameters) - all_tensors.extend(self.buffers) - all_tensors = list(dict.fromkeys(all_tensors)) # Remove duplicates - - self.tensor_to_key = {tensor: f"tensor_{i}" for i, tensor in enumerate(all_tensors)} - self.key_to_tensor = {v: k for k, v in self.tensor_to_key.items()} - self.cpu_param_dict = {} - else: - self.cpu_param_dict = self._init_cpu_param_dict() - - self._torch_accelerator_module = ( - getattr(torch, torch.accelerator.current_accelerator().type) - if hasattr(torch, "accelerator") - else torch.cuda - ) - - @staticmethod - def _to_cpu(tensor, low_cpu_mem_usage): - # For TorchAO tensors, `.data` returns an incomplete wrapper without internal attributes - # (e.g. `.qdata`, `.scale`), so we must call `.cpu()` on the tensor directly. - t = tensor.cpu() if _is_torchao_tensor(tensor) else tensor.data.cpu() - return t if low_cpu_mem_usage else t.pin_memory() - - def _init_cpu_param_dict(self): - cpu_param_dict = {} - if self.stream is None: - return cpu_param_dict - - for module in self.modules: - for param in module.parameters(): - cpu_param_dict[param] = self._to_cpu(param, self.low_cpu_mem_usage) - for buffer in module.buffers(): - cpu_param_dict[buffer] = self._to_cpu(buffer, self.low_cpu_mem_usage) - - for param in self.parameters: - cpu_param_dict[param] = self._to_cpu(param, self.low_cpu_mem_usage) - - for buffer in self.buffers: - cpu_param_dict[buffer] = self._to_cpu(buffer, self.low_cpu_mem_usage) - - return cpu_param_dict - - @contextmanager - def _pinned_memory_tensors(self): - try: - pinned_dict = { - param: tensor.pin_memory() if not tensor.is_pinned() else tensor - for param, tensor in self.cpu_param_dict.items() - } - yield pinned_dict - finally: - pinned_dict = None - - def _transfer_tensor_to_device(self, tensor, source_tensor, default_stream): - moved = source_tensor.to(self.onload_device, non_blocking=self.non_blocking) - if _is_torchao_tensor(tensor): - _swap_torchao_tensor(tensor, moved) - else: - tensor.data = moved - if self.record_stream: - if _is_torchao_tensor(tensor): - _record_stream_torchao_tensor(tensor, default_stream) - else: - tensor.data.record_stream(default_stream) - - def _process_tensors_from_modules(self, pinned_memory=None, default_stream=None): - for group_module in self.modules: - for param in group_module.parameters(): - source = pinned_memory[param] if pinned_memory else param.data - self._transfer_tensor_to_device(param, source, default_stream) - for buffer in group_module.buffers(): - source = pinned_memory[buffer] if pinned_memory else buffer.data - self._transfer_tensor_to_device(buffer, source, default_stream) - - for param in self.parameters: - source = pinned_memory[param] if pinned_memory else param.data - self._transfer_tensor_to_device(param, source, default_stream) - - for buffer in self.buffers: - source = pinned_memory[buffer] if pinned_memory else buffer.data - self._transfer_tensor_to_device(buffer, source, default_stream) - - def _check_disk_offload_torchao(self): - all_tensors = list(self.tensor_to_key.keys()) - has_torchao = any(_is_torchao_tensor(t) for t in all_tensors) - if has_torchao: - raise ValueError( - "Disk offloading is not supported for TorchAO quantized tensors because safetensors " - "cannot serialize TorchAO subclass tensors. Use memory offloading instead by not " - "setting `offload_to_disk_path`." - ) - - def _onload_from_disk(self): - self._check_disk_offload_torchao() - - if self.stream is not None: - # Wait for previous Host->Device transfer to complete - self.stream.synchronize() - - context = nullcontext() if self.stream is None else self._torch_accelerator_module.stream(self.stream) - current_stream = self._torch_accelerator_module.current_stream() if self.record_stream else None - - with context: - if self.stream is not None: - # Load to CPU first, pin memory, then async copy to the target device - loaded_tensors = safetensors.torch.load_file(self.safetensors_file_path, device="cpu") - for key, tensor_obj in self.key_to_tensor.items(): - pinned_tensor = loaded_tensors[key].pin_memory() - tensor_obj.data = pinned_tensor.to(self.onload_device, non_blocking=self.non_blocking) - if self.record_stream: - tensor_obj.data.record_stream(current_stream) - else: - # Load directly to the target device - onload_device = ( - self.onload_device.type if isinstance(self.onload_device, torch.device) else self.onload_device - ) - loaded_tensors = safetensors.torch.load_file(self.safetensors_file_path, device=onload_device) - for key, tensor_obj in self.key_to_tensor.items(): - tensor_obj.data = loaded_tensors[key] - - def _onload_from_memory(self): - if self.stream is not None: - # Wait for previous Host->Device transfer to complete - self.stream.synchronize() - - context = nullcontext() if self.stream is None else self._torch_accelerator_module.stream(self.stream) - default_stream = self._torch_accelerator_module.current_stream() if self.stream is not None else None - - with context: - if self.stream is not None: - with self._pinned_memory_tensors() as pinned_memory: - self._process_tensors_from_modules(pinned_memory, default_stream=default_stream) - else: - self._process_tensors_from_modules(None) - - def _offload_to_disk(self): - self._check_disk_offload_torchao() - - # TODO: we can potentially optimize this code path by checking if the _all_ the desired - # safetensor files exist on the disk and if so, skip this step entirely, reducing IO - # overhead. Currently, we just check if the given `safetensors_file_path` exists and if not - # we perform a write. - # Check if the file has been saved in this session or if it already exists on disk. - if not self._is_offloaded_to_disk and not os.path.exists(self.safetensors_file_path): - os.makedirs(os.path.dirname(self.safetensors_file_path), exist_ok=True) - tensors_to_save = {key: tensor.data.to(self.offload_device) for tensor, key in self.tensor_to_key.items()} - safetensors.torch.save_file(tensors_to_save, self.safetensors_file_path) - - # The group is now considered offloaded to disk for the rest of the session. - self._is_offloaded_to_disk = True - - # We do this to free up the RAM which is still holding the up tensor data. - for tensor_obj in self.tensor_to_key.keys(): - tensor_obj.data = torch.empty_like(tensor_obj.data, device=self.offload_device) - - def _offload_to_memory(self): - if self.stream is not None: - if not self.record_stream: - self._torch_accelerator_module.current_stream().synchronize() - - for group_module in self.modules: - for param in group_module.parameters(): - if _is_torchao_tensor(param): - _restore_torchao_tensor(param, self.cpu_param_dict[param]) - else: - param.data = self.cpu_param_dict[param] - for param in self.parameters: - if _is_torchao_tensor(param): - _restore_torchao_tensor(param, self.cpu_param_dict[param]) - else: - param.data = self.cpu_param_dict[param] - for buffer in self.buffers: - if _is_torchao_tensor(buffer): - _restore_torchao_tensor(buffer, self.cpu_param_dict[buffer]) - else: - buffer.data = self.cpu_param_dict[buffer] - else: - for group_module in self.modules: - group_module.to(self.offload_device, non_blocking=False) - for param in self.parameters: - if _is_torchao_tensor(param): - moved = param.to(self.offload_device, non_blocking=False) - _swap_torchao_tensor(param, moved) - else: - param.data = param.data.to(self.offload_device, non_blocking=False) - for buffer in self.buffers: - if _is_torchao_tensor(buffer): - moved = buffer.to(self.offload_device, non_blocking=False) - _swap_torchao_tensor(buffer, moved) - else: - buffer.data = buffer.data.to(self.offload_device, non_blocking=False) - - @torch.compiler.disable() - def onload_(self): - r"""Onloads the group of parameters to the onload_device.""" - if self.offload_to_disk_path is not None: - self._onload_from_disk() - else: - self._onload_from_memory() - - @torch.compiler.disable() - def offload_(self): - r"""Offloads the group of parameters to the offload_device.""" - if self.offload_to_disk_path: - self._offload_to_disk() - else: - self._offload_to_memory() - - -class GroupOffloadingHook(ModelHook): - r""" - A hook that offloads groups of torch.nn.Module to the CPU for storage and onloads to accelerator device for - computation. Each group has one "onload leader" module that is responsible for onloading, and an "offload leader" - module that is responsible for offloading. If prefetching is enabled, the onload leader of the previous module - group is responsible for onloading the current module group. - """ - - _is_stateful = False - - def __init__(self, group: ModuleGroup, *, config: GroupOffloadingConfig) -> None: - self.group = group - self.next_group: ModuleGroup | None = None - self.config = config - - def initialize_hook(self, module: torch.nn.Module) -> torch.nn.Module: - if self.group.offload_leader == module: - self.group.offload_() - return module - - def pre_forward(self, module: torch.nn.Module, *args, **kwargs): - # If there wasn't an onload_leader assigned, we assume that the submodule that first called its forward - # method is the onload_leader of the group. - if self.group.onload_leader is None: - self.group.onload_leader = module - - # If the current module is the onload_leader of the group, we onload the group if it is supposed - # to onload itself. In the case of using prefetching with streams, we onload the next group if - # it is not supposed to onload itself. - if self.group.onload_leader == module: - if self.group.onload_self: - self.group.onload_() - else: - # onload_self=False means this group relies on prefetching from a previous group. - # However, for conditionally-executed modules (e.g. patch_short/patch_mid/patch_long in Helios), - # the prefetch chain may not cover them if they were absent during the first forward pass - # when the execution order was traced. In that case, their weights remain on offload_device, - # so we fall back to a synchronous onload here. - params = [p for m in self.group.modules for p in m.parameters()] + list(self.group.parameters) - if params and params[0].device == self.group.offload_device: - self.group.onload_() - if self.group.stream is not None: - self.group.stream.synchronize() - - should_onload_next_group = self.next_group is not None and not self.next_group.onload_self - if should_onload_next_group: - self.next_group.onload_() - - should_synchronize = ( - not self.group.onload_self and self.group.stream is not None and not should_onload_next_group - ) - if should_synchronize: - # If this group didn't onload itself, it means it was asynchronously onloaded by the - # previous group. We need to synchronize the side stream to ensure parameters - # are completely loaded to proceed with forward pass. Without this, uninitialized - # weights will be used in the computation, leading to incorrect results - # Also, we should only do this synchronization if we don't already do it from the sync call in - # self.next_group.onload_, hence the `not should_onload_next_group` check. - self.group.stream.synchronize() - - args = send_to_device(args, self.group.onload_device, non_blocking=self.group.non_blocking) - - # Some Autoencoder models use a feature cache that is passed through submodules - # and modified in place. The `send_to_device` call returns a copy of this feature cache object - # which breaks the inplace updates. Use `exclude_kwargs` to mark these cache features - exclude_kwargs = self.config.exclude_kwargs or [] - if exclude_kwargs: - moved_kwargs = send_to_device( - {k: v for k, v in kwargs.items() if k not in exclude_kwargs}, - self.group.onload_device, - non_blocking=self.group.non_blocking, - ) - kwargs.update(moved_kwargs) - else: - kwargs = send_to_device(kwargs, self.group.onload_device, non_blocking=self.group.non_blocking) - - return args, kwargs - - def post_forward(self, module: torch.nn.Module, output): - if self.group.offload_leader == module: - self.group.offload_() - return output - - -class LazyPrefetchGroupOffloadingHook(ModelHook): - r""" - A hook, used in conjunction with GroupOffloadingHook, that applies lazy prefetching to groups of torch.nn.Module. - This hook is used to determine the order in which the layers are executed during the forward pass. Once the layer - invocation order is known, assignments of the next_group attribute for prefetching can be made, which allows - prefetching groups in the correct order. - """ - - _is_stateful = False - - def __init__(self): - self.execution_order: list[tuple[str, torch.nn.Module]] = [] - self._layer_execution_tracker_module_names = set() - - def initialize_hook(self, module): - def make_execution_order_update_callback(current_name, current_submodule): - def callback(): - if not torch.compiler.is_compiling(): - logger.debug(f"Adding {current_name} to the execution order") - self.execution_order.append((current_name, current_submodule)) - - return callback - - # To every submodule that contains a group offloading hook (at this point, no prefetching is enabled for any - # of the groups), we add a layer execution tracker hook that will be used to determine the order in which the - # layers are executed during the forward pass. - for name, submodule in module.named_modules(): - if name == "" or not hasattr(submodule, "_diffusers_hook"): - continue - - registry = HookRegistry.check_if_exists_or_initialize(submodule) - group_offloading_hook = registry.get_hook(_GROUP_OFFLOADING) - - if group_offloading_hook is not None: - # For the first forward pass, we have to load in a blocking manner - group_offloading_hook.group.non_blocking = False - layer_tracker_hook = LayerExecutionTrackerHook(make_execution_order_update_callback(name, submodule)) - registry.register_hook(layer_tracker_hook, _LAYER_EXECUTION_TRACKER) - self._layer_execution_tracker_module_names.add(name) - - return module - - def post_forward(self, module, output): - # At this point, for the current modules' submodules, we know the execution order of the layers. We can now - # remove the layer execution tracker hooks and apply prefetching by setting the next_group attribute for each - # group offloading hook. - num_executed = len(self.execution_order) - execution_order_module_names = {name for name, _ in self.execution_order} - - # It may be possible that some layers were not executed during the forward pass. This can happen if the layer - # is not used in the forward pass, or if the layer is not executed due to some other reason. In such cases, we - # may not be able to apply prefetching in the correct order, which can lead to device-mismatch related errors - # if the missing layers end up being executed in the future. - if execution_order_module_names != self._layer_execution_tracker_module_names: - unexecuted_layers = list(self._layer_execution_tracker_module_names - execution_order_module_names) - if not torch.compiler.is_compiling(): - logger.warning( - "It seems like some layers were not executed during the forward pass. This may lead to problems when " - "applying lazy prefetching with automatic tracing and lead to device-mismatch related errors. Please " - "make sure that all layers are executed during the forward pass. The following layers were not executed:\n" - f"{unexecuted_layers=}" - ) - - # Remove the layer execution tracker hooks from the submodules - base_module_registry = module._diffusers_hook - registries = [submodule._diffusers_hook for _, submodule in self.execution_order] - group_offloading_hooks = [registry.get_hook(_GROUP_OFFLOADING) for registry in registries] - - for i in range(num_executed): - registries[i].remove_hook(_LAYER_EXECUTION_TRACKER, recurse=False) - - # Remove the current lazy prefetch group offloading hook so that it doesn't interfere with the next forward pass - base_module_registry.remove_hook(_LAZY_PREFETCH_GROUP_OFFLOADING, recurse=False) - - # LazyPrefetchGroupOffloadingHook is only used with streams, so we know that non_blocking should be True. - # We disable non_blocking for the first forward pass, but need to enable it for the subsequent passes to - # see the benefits of prefetching. - for hook in group_offloading_hooks: - hook.group.non_blocking = True - - # Set required attributes for prefetching - if num_executed > 0: - base_module_group_offloading_hook = base_module_registry.get_hook(_GROUP_OFFLOADING) - base_module_group_offloading_hook.next_group = group_offloading_hooks[0].group - base_module_group_offloading_hook.next_group.onload_self = False - - for i in range(num_executed - 1): - name1, _ = self.execution_order[i] - name2, _ = self.execution_order[i + 1] - if not torch.compiler.is_compiling(): - logger.debug(f"Applying lazy prefetch group offloading from {name1} to {name2}") - group_offloading_hooks[i].next_group = group_offloading_hooks[i + 1].group - group_offloading_hooks[i].next_group.onload_self = False - - return output - - -class LayerExecutionTrackerHook(ModelHook): - r""" - A hook that tracks the order in which the layers are executed during the forward pass by calling back to the - LazyPrefetchGroupOffloadingHook to update the execution order. - """ - - _is_stateful = False - - def __init__(self, execution_order_update_callback): - self.execution_order_update_callback = execution_order_update_callback - - def pre_forward(self, module, *args, **kwargs): - self.execution_order_update_callback() - return args, kwargs - - -def apply_group_offloading( - module: torch.nn.Module, - onload_device: str | torch.device, - offload_device: str | torch.device = torch.device("cpu"), - offload_type: str | GroupOffloadingType = "block_level", - num_blocks_per_group: int | None = None, - non_blocking: bool = False, - use_stream: bool = False, - record_stream: bool = False, - low_cpu_mem_usage: bool = False, - offload_to_disk_path: str | None = None, - block_modules: list[str] | None = None, - exclude_kwargs: list[str] | None = None, -) -> None: - r""" - Applies group offloading to the internal layers of a torch.nn.Module. To understand what group offloading is, and - where it is beneficial, we need to first provide some context on how other supported offloading methods work. - - Typically, offloading is done at two levels: - - Module-level: In Diffusers, this can be enabled using the `ModelMixin::enable_model_cpu_offload()` method. It - works by offloading each component of a pipeline to the CPU for storage, and onloading to the accelerator device - when needed for computation. This method is more memory-efficient than keeping all components on the accelerator, - but the memory requirements are still quite high. For this method to work, one needs memory equivalent to size of - the model in runtime dtype + size of largest intermediate activation tensors to be able to complete the forward - pass. - - Leaf-level: In Diffusers, this can be enabled using the `ModelMixin::enable_sequential_cpu_offload()` method. It - works by offloading the lowest leaf-level parameters of the computation graph to the CPU for storage, and - onloading only the leafs to the accelerator device for computation. This uses the lowest amount of accelerator - memory, but can be slower due to the excessive number of device synchronizations. - - Group offloading is a middle ground between the two methods. It works by offloading groups of internal layers, - (either `torch.nn.ModuleList` or `torch.nn.Sequential`). This method uses lower memory than module-level - offloading. It is also faster than leaf-level/sequential offloading, as the number of device synchronizations is - reduced. - - Another supported feature (for CUDA devices with support for asynchronous data transfer streams) is the ability to - overlap data transfer and computation to reduce the overall execution time compared to sequential offloading. This - is enabled using layer prefetching with streams, i.e., the layer that is to be executed next starts onloading to - the accelerator device while the current layer is being executed - this increases the memory requirements slightly. - Note that this implementation also supports leaf-level offloading but can be made much faster when using streams. - - Args: - module (`torch.nn.Module`): - The module to which group offloading is applied. - onload_device (`torch.device`): - The device to which the group of modules are onloaded. - offload_device (`torch.device`, defaults to `torch.device("cpu")`): - The device to which the group of modules are offloaded. This should typically be the CPU. Default is CPU. - offload_type (`str` or `GroupOffloadingType`, defaults to "block_level"): - The type of offloading to be applied. Can be one of "block_level" or "leaf_level". Default is - "block_level". - offload_to_disk_path (`str`, *optional*, defaults to `None`): - The path to the directory where parameters will be offloaded. Setting this option can be useful in limited - RAM environment settings where a reasonable speed-memory trade-off is desired. - num_blocks_per_group (`int`, *optional*): - The number of blocks per group when using offload_type="block_level". This is required when using - offload_type="block_level". - non_blocking (`bool`, defaults to `False`): - If True, offloading and onloading is done with non-blocking data transfer. - use_stream (`bool`, defaults to `False`): - If True, offloading and onloading is done asynchronously using a CUDA stream. This can be useful for - overlapping computation and data transfer. - record_stream (`bool`, defaults to `False`): When enabled with `use_stream`, it marks the current tensor - as having been used by this stream. It is faster at the expense of slightly more memory usage. Refer to the - [PyTorch official docs](https://pytorch.org/docs/stable/generated/torch.Tensor.record_stream.html) more - details. - low_cpu_mem_usage (`bool`, defaults to `False`): - If True, the CPU memory usage is minimized by pinning tensors on-the-fly instead of pre-pinning them. This - option only matters when using streamed CPU offloading (i.e. `use_stream=True`). This can be useful when - the CPU memory is a bottleneck but may counteract the benefits of using streams. - block_modules (`list[str]`, *optional*): - List of module names that should be treated as blocks for offloading. If provided, only these modules will - be considered for block-level offloading. If not provided, the default block detection logic will be used. - exclude_kwargs (`list[str]`, *optional*): - List of kwarg keys that should not be processed by send_to_device. This is useful for mutable state like - caching lists that need to maintain their object identity across forward passes. If not provided, will be - inferred from the module's `_skip_keys` attribute if it exists. - - Example: - ```python - >>> from diffusers import CogVideoXTransformer3DModel - >>> from diffusers.hooks import apply_group_offloading - - >>> transformer = CogVideoXTransformer3DModel.from_pretrained( - ... "THUDM/CogVideoX-5b", subfolder="transformer", torch_dtype=torch.bfloat16 - ... ) - - >>> apply_group_offloading( - ... transformer, - ... onload_device=torch.device("cuda"), - ... offload_device=torch.device("cpu"), - ... offload_type="block_level", - ... num_blocks_per_group=2, - ... use_stream=True, - ... ) - ``` - """ - - onload_device = torch.device(onload_device) if isinstance(onload_device, str) else onload_device - offload_device = torch.device(offload_device) if isinstance(offload_device, str) else offload_device - offload_type = GroupOffloadingType(offload_type) - - stream = None - if use_stream: - if torch.cuda.is_available(): - stream = torch.cuda.Stream() - elif hasattr(torch, "xpu") and torch.xpu.is_available(): - stream = torch.Stream() - else: - raise ValueError("Using streams for data transfer requires a CUDA device, or an Intel XPU device.") - - if not use_stream and record_stream: - raise ValueError("`record_stream` cannot be True when `use_stream=False`.") - if offload_type == GroupOffloadingType.BLOCK_LEVEL and num_blocks_per_group is None: - raise ValueError("`num_blocks_per_group` must be provided when using `offload_type='block_level'.") - - _raise_error_if_accelerate_model_or_sequential_hook_present(module) - - if block_modules is None: - block_modules = getattr(module, "_group_offload_block_modules", None) - - if exclude_kwargs is None: - exclude_kwargs = getattr(module, "_skip_keys", None) - - config = GroupOffloadingConfig( - onload_device=onload_device, - offload_device=offload_device, - offload_type=offload_type, - num_blocks_per_group=num_blocks_per_group, - non_blocking=non_blocking, - stream=stream, - record_stream=record_stream, - low_cpu_mem_usage=low_cpu_mem_usage, - offload_to_disk_path=offload_to_disk_path, - block_modules=block_modules, - exclude_kwargs=exclude_kwargs, - ) - _apply_group_offloading(module, config) - - -def _apply_group_offloading(module: torch.nn.Module, config: GroupOffloadingConfig) -> None: - if config.offload_type == GroupOffloadingType.BLOCK_LEVEL: - _apply_group_offloading_block_level(module, config) - elif config.offload_type == GroupOffloadingType.LEAF_LEVEL: - _apply_group_offloading_leaf_level(module, config) - else: - assert False - - -def _apply_group_offloading_block_level(module: torch.nn.Module, config: GroupOffloadingConfig) -> None: - r""" - This function applies offloading to groups of torch.nn.ModuleList or torch.nn.Sequential blocks, and explicitly - defined block modules. In comparison to the "leaf_level" offloading, which is more fine-grained, this offloading is - done at the top-level blocks and modules specified in block_modules. - - When block_modules is provided, only those modules will be treated as blocks for offloading. For each specified - module, recursively apply block offloading to it. - """ - if config.stream is not None and config.num_blocks_per_group != 1: - logger.warning( - f"Using streams is only supported for num_blocks_per_group=1. Got {config.num_blocks_per_group=}. Setting it to 1." - ) - config.num_blocks_per_group = 1 - - block_modules = set(config.block_modules) if config.block_modules is not None else set() - - # Create module groups for ModuleList and Sequential blocks, and explicitly defined block modules - modules_with_group_offloading = set() - unmatched_modules = [] - matched_module_groups = [] - - for name, submodule in module.named_children(): - # Check if this is an explicitly defined block module - if name in block_modules: - # Track submodule using a prefix to avoid filename collisions during disk offload. - # Without this, submodules sharing the same model class would be assigned identical - # filenames (derived from the class name). - prefix = f"{config.module_prefix}{name}." if config.module_prefix else f"{name}." - submodule_config = replace(config, module_prefix=prefix) - - _apply_group_offloading_block_level(submodule, submodule_config) - modules_with_group_offloading.add(name) - - elif isinstance(submodule, (torch.nn.ModuleList, torch.nn.Sequential)): - # Handle ModuleList and Sequential blocks as before - for i in range(0, len(submodule), config.num_blocks_per_group): - current_modules = list(submodule[i : i + config.num_blocks_per_group]) - if len(current_modules) == 0: - continue - - group_id = f"{config.module_prefix}{name}_{i}_{i + len(current_modules) - 1}" - group = ModuleGroup( - modules=current_modules, - offload_device=config.offload_device, - onload_device=config.onload_device, - offload_to_disk_path=config.offload_to_disk_path, - offload_leader=current_modules[-1], - onload_leader=current_modules[0], - non_blocking=config.non_blocking, - stream=config.stream, - record_stream=config.record_stream, - low_cpu_mem_usage=config.low_cpu_mem_usage, - onload_self=True, - group_id=group_id, - ) - matched_module_groups.append(group) - for j in range(i, i + len(current_modules)): - modules_with_group_offloading.add(f"{name}.{j}") - else: - # This is an unmatched module - unmatched_modules.append((name, submodule)) - - # Apply group offloading hooks to the module groups - for i, group in enumerate(matched_module_groups): - for group_module in group.modules: - _apply_group_offloading_hook(group_module, group, config=config) - - # Parameters and Buffers of the top-level module need to be offloaded/onloaded separately - # when the forward pass of this module is called. This is because the top-level module is not - # part of any group (as doing so would lead to no VRAM savings). - parameters = _gather_parameters_with_no_group_offloading_parent(module, modules_with_group_offloading) - buffers = _gather_buffers_with_no_group_offloading_parent(module, modules_with_group_offloading) - parameters = [param for _, param in parameters] - buffers = [buffer for _, buffer in buffers] - - # Create a group for the remaining unmatched submodules of the top-level - # module so that they are on the correct device when the forward pass is called. - unmatched_modules = [unmatched_module for _, unmatched_module in unmatched_modules] - if len(unmatched_modules) > 0 or len(parameters) > 0 or len(buffers) > 0: - unmatched_group = ModuleGroup( - modules=unmatched_modules, - offload_device=config.offload_device, - onload_device=config.onload_device, - offload_to_disk_path=config.offload_to_disk_path, - offload_leader=module, - onload_leader=module, - parameters=parameters, - buffers=buffers, - non_blocking=False, - stream=None, - record_stream=False, - onload_self=True, - group_id=f"{config.module_prefix}{module.__class__.__name__}_unmatched_group", - ) - if config.stream is None: - _apply_group_offloading_hook(module, unmatched_group, config=config) - else: - _apply_lazy_group_offloading_hook(module, unmatched_group, config=config) - - -def _apply_group_offloading_leaf_level(module: torch.nn.Module, config: GroupOffloadingConfig) -> None: - r""" - This function applies offloading to groups of leaf modules in a torch.nn.Module. This method has minimal memory - requirements. However, it can be slower compared to other offloading methods due to the excessive number of device - synchronizations. When using devices that support streams to overlap data transfer and computation, this method can - reduce memory usage without any performance degradation. - """ - # Create module groups for leaf modules and apply group offloading hooks - modules_with_group_offloading = set() - for name, submodule in module.named_modules(): - if not isinstance(submodule, _GO_LC_SUPPORTED_PYTORCH_LAYERS): - continue - group = ModuleGroup( - modules=[submodule], - offload_device=config.offload_device, - onload_device=config.onload_device, - offload_to_disk_path=config.offload_to_disk_path, - offload_leader=submodule, - onload_leader=submodule, - non_blocking=config.non_blocking, - stream=config.stream, - record_stream=config.record_stream, - low_cpu_mem_usage=config.low_cpu_mem_usage, - onload_self=True, - group_id=name, - ) - _apply_group_offloading_hook(submodule, group, config=config) - modules_with_group_offloading.add(name) - - # Parameters and Buffers at all non-leaf levels need to be offloaded/onloaded separately when the forward pass - # of the module is called - module_dict = dict(module.named_modules()) - parameters = _gather_parameters_with_no_group_offloading_parent(module, modules_with_group_offloading) - buffers = _gather_buffers_with_no_group_offloading_parent(module, modules_with_group_offloading) - - # Find closest module parent for each parameter and buffer, and attach group hooks - parent_to_parameters = {} - for name, param in parameters: - parent_name = _find_parent_module_in_module_dict(name, module_dict) - if parent_name in parent_to_parameters: - parent_to_parameters[parent_name].append(param) - else: - parent_to_parameters[parent_name] = [param] - - parent_to_buffers = {} - for name, buffer in buffers: - parent_name = _find_parent_module_in_module_dict(name, module_dict) - if parent_name in parent_to_buffers: - parent_to_buffers[parent_name].append(buffer) - else: - parent_to_buffers[parent_name] = [buffer] - - parent_names = set(parent_to_parameters.keys()) | set(parent_to_buffers.keys()) - for name in parent_names: - parameters = parent_to_parameters.get(name, []) - buffers = parent_to_buffers.get(name, []) - parent_module = module_dict[name] - group = ModuleGroup( - modules=[], - offload_device=config.offload_device, - onload_device=config.onload_device, - offload_leader=parent_module, - onload_leader=parent_module, - offload_to_disk_path=config.offload_to_disk_path, - parameters=parameters, - buffers=buffers, - non_blocking=config.non_blocking, - stream=config.stream, - record_stream=config.record_stream, - low_cpu_mem_usage=config.low_cpu_mem_usage, - onload_self=True, - group_id=name, - ) - _apply_group_offloading_hook(parent_module, group, config=config) - - if config.stream is not None: - # When using streams, we need to know the layer execution order for applying prefetching (to overlap data transfer - # and computation). Since we don't know the order beforehand, we apply a lazy prefetching hook that will find the - # execution order and apply prefetching in the correct order. - unmatched_group = ModuleGroup( - modules=[], - offload_device=config.offload_device, - onload_device=config.onload_device, - offload_to_disk_path=config.offload_to_disk_path, - offload_leader=module, - onload_leader=module, - parameters=None, - buffers=None, - non_blocking=False, - stream=None, - record_stream=False, - low_cpu_mem_usage=config.low_cpu_mem_usage, - onload_self=True, - group_id=_GROUP_ID_LAZY_LEAF, - ) - _apply_lazy_group_offloading_hook(module, unmatched_group, config=config) - - -def _apply_group_offloading_hook( - module: torch.nn.Module, - group: ModuleGroup, - *, - config: GroupOffloadingConfig, -) -> None: - registry = HookRegistry.check_if_exists_or_initialize(module) - - # We may have already registered a group offloading hook if the module had a torch.nn.Parameter whose parent - # is the current module. In such cases, we don't want to overwrite the existing group offloading hook. - if registry.get_hook(_GROUP_OFFLOADING) is None: - hook = GroupOffloadingHook(group, config=config) - registry.register_hook(hook, _GROUP_OFFLOADING) - - -def _apply_lazy_group_offloading_hook( - module: torch.nn.Module, - group: ModuleGroup, - *, - config: GroupOffloadingConfig, -) -> None: - registry = HookRegistry.check_if_exists_or_initialize(module) - - # We may have already registered a group offloading hook if the module had a torch.nn.Parameter whose parent - # is the current module. In such cases, we don't want to overwrite the existing group offloading hook. - if registry.get_hook(_GROUP_OFFLOADING) is None: - hook = GroupOffloadingHook(group, config=config) - registry.register_hook(hook, _GROUP_OFFLOADING) - - lazy_prefetch_hook = LazyPrefetchGroupOffloadingHook() - registry.register_hook(lazy_prefetch_hook, _LAZY_PREFETCH_GROUP_OFFLOADING) - - -def _gather_parameters_with_no_group_offloading_parent( - module: torch.nn.Module, modules_with_group_offloading: Set[str] -) -> list[torch.nn.Parameter]: - parameters = [] - for name, parameter in module.named_parameters(): - has_parent_with_group_offloading = False - atoms = name.split(".") - while len(atoms) > 0: - parent_name = ".".join(atoms) - if parent_name in modules_with_group_offloading: - has_parent_with_group_offloading = True - break - atoms.pop() - if not has_parent_with_group_offloading: - parameters.append((name, parameter)) - return parameters - - -def _gather_buffers_with_no_group_offloading_parent( - module: torch.nn.Module, modules_with_group_offloading: Set[str] -) -> list[torch.Tensor]: - buffers = [] - for name, buffer in module.named_buffers(): - has_parent_with_group_offloading = False - atoms = name.split(".") - while len(atoms) > 0: - parent_name = ".".join(atoms) - if parent_name in modules_with_group_offloading: - has_parent_with_group_offloading = True - break - atoms.pop() - if not has_parent_with_group_offloading: - buffers.append((name, buffer)) - return buffers - - -def _find_parent_module_in_module_dict(name: str, module_dict: dict[str, torch.nn.Module]) -> str: - atoms = name.split(".") - while len(atoms) > 0: - parent_name = ".".join(atoms) - if parent_name in module_dict: - return parent_name - atoms.pop() - return "" - - -def _raise_error_if_accelerate_model_or_sequential_hook_present(module: torch.nn.Module) -> None: - if not is_accelerate_available(): - return - for name, submodule in module.named_modules(): - if not hasattr(submodule, "_hf_hook"): - continue - if isinstance(submodule._hf_hook, (AlignDevicesHook, CpuOffload)): - raise ValueError( - f"Cannot apply group offloading to a module that is already applying an alternative " - f"offloading strategy from Accelerate. If you want to apply group offloading, please " - f"disable the existing offloading strategy first. Offending module: {name} ({type(submodule)})" - ) - - -def _get_top_level_group_offload_hook(module: torch.nn.Module) -> GroupOffloadingHook | None: - for submodule in module.modules(): - if hasattr(submodule, "_diffusers_hook"): - group_offloading_hook = submodule._diffusers_hook.get_hook(_GROUP_OFFLOADING) - if group_offloading_hook is not None: - return group_offloading_hook - return None - - -def _is_group_offload_enabled(module: torch.nn.Module) -> bool: - top_level_group_offload_hook = _get_top_level_group_offload_hook(module) - return top_level_group_offload_hook is not None - - -def _get_group_onload_device(module: torch.nn.Module) -> torch.device: - top_level_group_offload_hook = _get_top_level_group_offload_hook(module) - if top_level_group_offload_hook is not None: - return top_level_group_offload_hook.config.onload_device - raise ValueError("Group offloading is not enabled for the provided module.") - - -def _compute_group_hash(group_id): - hashed_id = hashlib.sha256(group_id.encode("utf-8")).hexdigest() - # first 16 characters for a reasonably short but unique name - return hashed_id[:16] - - -def _maybe_remove_and_reapply_group_offloading(module: torch.nn.Module) -> None: - r""" - Removes the group offloading hook from the module and re-applies it. This is useful when the module has been - modified in-place and the group offloading hook references-to-tensors needs to be updated. The in-place - modification can happen in a number of ways, for example, fusing QKV or unloading/loading LoRAs on-the-fly. - - In this implementation, we make an assumption that group offloading has only been applied at the top-level module, - and therefore all submodules have the same onload and offload devices. If this assumption is not true, say in the - case where user has applied group offloading at multiple levels, this function will not work as expected. - - There is some performance penalty associated with doing this when non-default streams are used, because we need to - retrace the execution order of the layers with `LazyPrefetchGroupOffloadingHook`. - """ - top_level_group_offload_hook = _get_top_level_group_offload_hook(module) - - if top_level_group_offload_hook is None: - return - - registry = HookRegistry.check_if_exists_or_initialize(module) - registry.remove_hook(_GROUP_OFFLOADING, recurse=True) - registry.remove_hook(_LAYER_EXECUTION_TRACKER, recurse=True) - registry.remove_hook(_LAZY_PREFETCH_GROUP_OFFLOADING, recurse=True) - - _apply_group_offloading(module, top_level_group_offload_hook.config) diff --git a/diffusers/hooks/hooks.py b/diffusers/hooks/hooks.py deleted file mode 100644 index f278b40c15d6343a3da0e7d8649dc547740edee0..0000000000000000000000000000000000000000 --- a/diffusers/hooks/hooks.py +++ /dev/null @@ -1,312 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import functools -from typing import Any - -import torch - -from ..utils.logging import get_logger -from ..utils.torch_utils import unwrap_module - - -logger = get_logger(__name__) # pylint: disable=invalid-name - - -class BaseState: - def reset(self, *args, **kwargs) -> None: - raise NotImplementedError( - "BaseState::reset is not implemented. Please implement this method in the derived class." - ) - - -class StateManager: - def __init__(self, state_cls: BaseState, init_args=None, init_kwargs=None): - self._state_cls = state_cls - self._init_args = init_args if init_args is not None else () - self._init_kwargs = init_kwargs if init_kwargs is not None else {} - self._state_cache = {} - self._current_context = None - - def get_state(self): - if self._current_context is None: - raise ValueError("No context is set. Please set a context before retrieving the state.") - if self._current_context not in self._state_cache.keys(): - self._state_cache[self._current_context] = self._state_cls(*self._init_args, **self._init_kwargs) - return self._state_cache[self._current_context] - - def set_context(self, name: str) -> None: - self._current_context = name - - def reset(self, *args, **kwargs) -> None: - for name, state in list(self._state_cache.items()): - state.reset(*args, **kwargs) - self._state_cache.pop(name) - self._current_context = None - - -class ModelHook: - r""" - A hook that contains callbacks to be executed just before and after the forward method of a model. - """ - - _is_stateful = False - - def __init__(self): - self.fn_ref: "HookFunctionReference" = None - - def initialize_hook(self, module: torch.nn.Module) -> torch.nn.Module: - r""" - Hook that is executed when a model is initialized. - - Args: - module (`torch.nn.Module`): - The module attached to this hook. - """ - return module - - def deinitalize_hook(self, module: torch.nn.Module) -> torch.nn.Module: - r""" - Hook that is executed when a model is deinitialized. - - Args: - module (`torch.nn.Module`): - The module attached to this hook. - """ - return module - - def pre_forward(self, module: torch.nn.Module, *args, **kwargs) -> tuple[tuple[Any], dict[str, Any]]: - r""" - Hook that is executed just before the forward method of the model. - - Args: - module (`torch.nn.Module`): - The module whose forward pass will be executed just after this event. - args (`tuple[Any]`): - The positional arguments passed to the module. - kwargs (`dict[Str, Any]`): - The keyword arguments passed to the module. - Returns: - `tuple[tuple[Any], dict[Str, Any]]`: - A tuple with the treated `args` and `kwargs`. - """ - return args, kwargs - - def post_forward(self, module: torch.nn.Module, output: Any) -> Any: - r""" - Hook that is executed just after the forward method of the model. - - Args: - module (`torch.nn.Module`): - The module whose forward pass been executed just before this event. - output (`Any`): - The output of the module. - Returns: - `Any`: The processed `output`. - """ - return output - - def detach_hook(self, module: torch.nn.Module) -> torch.nn.Module: - r""" - Hook that is executed when the hook is detached from a module. - - Args: - module (`torch.nn.Module`): - The module detached from this hook. - """ - return module - - def reset_state(self, module: torch.nn.Module): - if self._is_stateful: - raise NotImplementedError("This hook is stateful and needs to implement the `reset_state` method.") - return module - - def _set_context(self, module: torch.nn.Module, name: str) -> None: - # Iterate over all attributes of the hook to see if any of them have the type `StateManager`. If so, call `set_context` on them. - for attr_name in dir(self): - attr = getattr(self, attr_name) - if isinstance(attr, StateManager): - attr.set_context(name) - return module - - -class HookFunctionReference: - def __init__(self) -> None: - """A container class that maintains mutable references to forward pass functions in a hook chain. - - Its mutable nature allows the hook system to modify the execution chain dynamically without rebuilding the - entire forward pass structure. - - Attributes: - pre_forward: A callable that processes inputs before the main forward pass. - post_forward: A callable that processes outputs after the main forward pass. - forward: The current forward function in the hook chain. - original_forward: The original forward function, stored when a hook provides a custom new_forward. - - The class enables hook removal by allowing updates to the forward chain through reference modification rather - than requiring reconstruction of the entire chain. When a hook is removed, only the relevant references need to - be updated, preserving the execution order of the remaining hooks. - """ - self.pre_forward = None - self.post_forward = None - self.forward = None - self.original_forward = None - - -class HookRegistry: - def __init__(self, module_ref: torch.nn.Module) -> None: - super().__init__() - - self.hooks: dict[str, ModelHook] = {} - - self._module_ref = module_ref - self._hook_order = [] - self._fn_refs = [] - - def register_hook(self, hook: ModelHook, name: str) -> None: - if name in self.hooks.keys(): - raise ValueError( - f"Hook with name {name} already exists in the registry. Please use a different name or " - f"first remove the existing hook and then add a new one." - ) - - self._module_ref = hook.initialize_hook(self._module_ref) - - def create_new_forward(function_reference: HookFunctionReference): - def new_forward(module, *args, **kwargs): - args, kwargs = function_reference.pre_forward(module, *args, **kwargs) - output = function_reference.forward(*args, **kwargs) - return function_reference.post_forward(module, output) - - return new_forward - - forward = self._module_ref.forward - - fn_ref = HookFunctionReference() - fn_ref.pre_forward = hook.pre_forward - fn_ref.post_forward = hook.post_forward - fn_ref.forward = forward - - if hasattr(hook, "new_forward"): - fn_ref.original_forward = forward - fn_ref.forward = functools.update_wrapper( - functools.partial(hook.new_forward, self._module_ref), hook.new_forward - ) - - rewritten_forward = create_new_forward(fn_ref) - # Wrap from the original `forward` so `inspect.signature` follows `__wrapped__` to the real - # signature instead of the generic `(module, *args, **kwargs)`, which breaks `torch.export`. - self._module_ref.forward = functools.update_wrapper( - functools.partial(rewritten_forward, self._module_ref), forward - ) - - hook.fn_ref = fn_ref - self.hooks[name] = hook - self._hook_order.append(name) - self._fn_refs.append(fn_ref) - - def get_hook(self, name: str) -> ModelHook | None: - return self.hooks.get(name, None) - - def remove_hook(self, name: str, recurse: bool = True) -> None: - if name in self.hooks.keys(): - num_hooks = len(self._hook_order) - hook = self.hooks[name] - index = self._hook_order.index(name) - fn_ref = self._fn_refs[index] - - old_forward = fn_ref.forward - if fn_ref.original_forward is not None: - old_forward = fn_ref.original_forward - - if index == num_hooks - 1: - self._module_ref.forward = old_forward - else: - self._fn_refs[index + 1].forward = old_forward - - self._module_ref = hook.deinitalize_hook(self._module_ref) - del self.hooks[name] - self._hook_order.pop(index) - self._fn_refs.pop(index) - - if recurse: - for module_name, module in self._module_ref.named_modules(): - if module_name == "": - continue - if hasattr(module, "_diffusers_hook"): - module._diffusers_hook.remove_hook(name, recurse=False) - - def reset_stateful_hooks(self, recurse: bool = True) -> None: - for hook_name in reversed(self._hook_order): - hook = self.hooks[hook_name] - if hook._is_stateful: - hook.reset_state(self._module_ref) - - if recurse: - for module_name, module in unwrap_module(self._module_ref).named_modules(): - if module_name == "": - continue - module = unwrap_module(module) - if hasattr(module, "_diffusers_hook"): - module._diffusers_hook.reset_stateful_hooks(recurse=False) - - @classmethod - def check_if_exists_or_initialize(cls, module: torch.nn.Module) -> "HookRegistry": - if not hasattr(module, "_diffusers_hook"): - module._diffusers_hook = cls(module) - return module._diffusers_hook - - def _set_context(self, name: str | None = None) -> None: - for hook_name in reversed(self._hook_order): - hook = self.hooks[hook_name] - if hook._is_stateful: - hook._set_context(self._module_ref, name) - - for registry in self._get_child_registries(): - registry._set_context(name) - - def _get_child_registries(self) -> list["HookRegistry"]: - """Return registries of child modules, using a cached list when available. - - The cache is built on first call and reused for subsequent calls. This avoids the cost of walking the full - module tree via named_modules() on every _set_context call, which is significant for large models (e.g. ~2.7ms - per call on Flux2). - """ - if not hasattr(self, "_child_registries_cache"): - self._child_registries_cache = None - - if self._child_registries_cache is not None: - return self._child_registries_cache - - registries = [] - for module_name, module in unwrap_module(self._module_ref).named_modules(): - if module_name == "": - continue - module = unwrap_module(module) - if hasattr(module, "_diffusers_hook"): - registries.append(module._diffusers_hook) - self._child_registries_cache = registries - return registries - - def __repr__(self) -> str: - registry_repr = "" - for i, hook_name in enumerate(self._hook_order): - if self.hooks[hook_name].__class__.__repr__ is not object.__repr__: - hook_repr = self.hooks[hook_name].__repr__() - else: - hook_repr = self.hooks[hook_name].__class__.__name__ - registry_repr += f" ({i}) {hook_name} - {hook_repr}" - if i < len(self._hook_order) - 1: - registry_repr += "\n" - return f"HookRegistry(\n{registry_repr}\n)" diff --git a/diffusers/hooks/layer_skip.py b/diffusers/hooks/layer_skip.py deleted file mode 100644 index 8085a88d3371d21264a3537b963c64467f6616d6..0000000000000000000000000000000000000000 --- a/diffusers/hooks/layer_skip.py +++ /dev/null @@ -1,263 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from dataclasses import asdict, dataclass -from typing import Callable - -import torch - -from ..utils import get_logger -from ..utils.torch_utils import unwrap_module -from ._common import ( - _ALL_TRANSFORMER_BLOCK_IDENTIFIERS, - _ATTENTION_CLASSES, - _FEEDFORWARD_CLASSES, - _get_submodule_from_fqn, -) -from ._helpers import AttentionProcessorRegistry, TransformerBlockRegistry -from .hooks import HookRegistry, ModelHook - - -logger = get_logger(__name__) # pylint: disable=invalid-name - -_LAYER_SKIP_HOOK = "layer_skip_hook" - - -# Aryan/YiYi TODO: we need to make guider class a config mixin so I think this is not needed -# either remove or make it serializable -@dataclass -class LayerSkipConfig: - r""" - Configuration for skipping internal transformer blocks when executing a transformer model. - - Args: - indices (`list[int]`): - The indices of the layer to skip. This is typically the first layer in the transformer block. - fqn (`str`, defaults to `"auto"`): - The fully qualified name identifying the stack of transformer blocks. Typically, this is - `transformer_blocks`, `single_transformer_blocks`, `blocks`, `layers`, or `temporal_transformer_blocks`. - For automatic detection, set this to `"auto"`. "auto" only works on DiT models. For UNet models, you must - provide the correct fqn. - skip_attention (`bool`, defaults to `True`): - Whether to skip attention blocks. - skip_ff (`bool`, defaults to `True`): - Whether to skip feed-forward blocks. - skip_attention_scores (`bool`, defaults to `False`): - Whether to skip attention score computation in the attention blocks. This is equivalent to using `value` - projections as the output of scaled dot product attention. - dropout (`float`, defaults to `1.0`): - The dropout probability for dropping the outputs of the skipped layers. By default, this is set to `1.0`, - meaning that the outputs of the skipped layers are completely ignored. If set to `0.0`, the outputs of the - skipped layers are fully retained, which is equivalent to not skipping any layers. - """ - - indices: list[int] - fqn: str = "auto" - skip_attention: bool = True - skip_attention_scores: bool = False - skip_ff: bool = True - dropout: float = 1.0 - - def __post_init__(self): - if not (0 <= self.dropout <= 1): - raise ValueError(f"Expected `dropout` to be between 0.0 and 1.0, but got {self.dropout}.") - if not math.isclose(self.dropout, 1.0) and self.skip_attention_scores: - raise ValueError( - "Cannot set `skip_attention_scores` to True when `dropout` is not 1.0. Please set `dropout` to 1.0." - ) - - def to_dict(self): - return asdict(self) - - @staticmethod - def from_dict(data: dict) -> "LayerSkipConfig": - return LayerSkipConfig(**data) - - -class AttentionScoreSkipFunctionMode(torch.overrides.TorchFunctionMode): - def __torch_function__(self, func, types, args=(), kwargs=None): - if kwargs is None: - kwargs = {} - if func is torch.nn.functional.scaled_dot_product_attention: - query = kwargs.get("query", None) - key = kwargs.get("key", None) - value = kwargs.get("value", None) - query = query if query is not None else args[0] - key = key if key is not None else args[1] - value = value if value is not None else args[2] - # If the Q sequence length does not match KV sequence length, methods like - # Perturbed Attention Guidance cannot be used (because the caller expects - # the same sequence length as Q, but if we return V here, it will not match). - # When Q.shape[2] != V.shape[2], PAG will essentially not be applied and - # the overall effect would that be of normal CFG with a scale of (guidance_scale + perturbed_guidance_scale). - if query.shape[2] == value.shape[2]: - return value - return func(*args, **kwargs) - - -class AttentionProcessorSkipHook(ModelHook): - def __init__(self, skip_processor_output_fn: Callable, skip_attention_scores: bool = False, dropout: float = 1.0): - self.skip_processor_output_fn = skip_processor_output_fn - self.skip_attention_scores = skip_attention_scores - self.dropout = dropout - - def new_forward(self, module: torch.nn.Module, *args, **kwargs): - if self.skip_attention_scores: - if not math.isclose(self.dropout, 1.0): - raise ValueError( - "Cannot set `skip_attention_scores` to True when `dropout` is not 1.0. Please set `dropout` to 1.0." - ) - with AttentionScoreSkipFunctionMode(): - output = self.fn_ref.original_forward(*args, **kwargs) - else: - if math.isclose(self.dropout, 1.0): - output = self.skip_processor_output_fn(module, *args, **kwargs) - else: - output = self.fn_ref.original_forward(*args, **kwargs) - output = torch.nn.functional.dropout(output, p=self.dropout) - return output - - -class FeedForwardSkipHook(ModelHook): - def __init__(self, dropout: float): - super().__init__() - self.dropout = dropout - - def new_forward(self, module: torch.nn.Module, *args, **kwargs): - if math.isclose(self.dropout, 1.0): - output = kwargs.get("hidden_states", None) - if output is None: - output = kwargs.get("x", None) - if output is None and len(args) > 0: - output = args[0] - else: - output = self.fn_ref.original_forward(*args, **kwargs) - output = torch.nn.functional.dropout(output, p=self.dropout) - return output - - -class TransformerBlockSkipHook(ModelHook): - def __init__(self, dropout: float): - super().__init__() - self.dropout = dropout - - def initialize_hook(self, module): - self._metadata = TransformerBlockRegistry.get(unwrap_module(module).__class__) - return module - - def new_forward(self, module: torch.nn.Module, *args, **kwargs): - if math.isclose(self.dropout, 1.0): - original_hidden_states = self._metadata._get_parameter_from_args_kwargs("hidden_states", args, kwargs) - if self._metadata.return_encoder_hidden_states_index is None: - output = original_hidden_states - else: - original_encoder_hidden_states = self._metadata._get_parameter_from_args_kwargs( - "encoder_hidden_states", args, kwargs - ) - output = (original_hidden_states, original_encoder_hidden_states) - else: - output = self.fn_ref.original_forward(*args, **kwargs) - output = torch.nn.functional.dropout(output, p=self.dropout) - return output - - -def apply_layer_skip(module: torch.nn.Module, config: LayerSkipConfig) -> None: - r""" - Apply layer skipping to internal layers of a transformer. - - Args: - module (`torch.nn.Module`): - The transformer model to which the layer skip hook should be applied. - config (`LayerSkipConfig`): - The configuration for the layer skip hook. - - Example: - - ```python - >>> from diffusers import apply_layer_skip_hook, CogVideoXTransformer3DModel, LayerSkipConfig - - >>> transformer = CogVideoXTransformer3DModel.from_pretrained("THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16) - >>> config = LayerSkipConfig(layer_index=[10, 20], fqn="transformer_blocks") - >>> apply_layer_skip_hook(transformer, config) - ``` - """ - _apply_layer_skip_hook(module, config) - - -def _apply_layer_skip_hook(module: torch.nn.Module, config: LayerSkipConfig, name: str | None = None) -> None: - name = name or _LAYER_SKIP_HOOK - - if config.skip_attention and config.skip_attention_scores: - raise ValueError("Cannot set both `skip_attention` and `skip_attention_scores` to True. Please choose one.") - if not math.isclose(config.dropout, 1.0) and config.skip_attention_scores: - raise ValueError( - "Cannot set `skip_attention_scores` to True when `dropout` is not 1.0. Please set `dropout` to 1.0." - ) - - if config.fqn == "auto": - for identifier in _ALL_TRANSFORMER_BLOCK_IDENTIFIERS: - if hasattr(module, identifier): - config.fqn = identifier - break - else: - raise ValueError( - "Could not find a suitable identifier for the transformer blocks automatically. Please provide a valid " - "`fqn` (fully qualified name) that identifies a stack of transformer blocks." - ) - - transformer_blocks = _get_submodule_from_fqn(module, config.fqn) - if transformer_blocks is None or not isinstance(transformer_blocks, torch.nn.ModuleList): - raise ValueError( - f"Could not find {config.fqn} in the provided module, or configured `fqn` (fully qualified name) does not identify " - f"a `torch.nn.ModuleList`. Please provide a valid `fqn` that identifies a stack of transformer blocks." - ) - if len(config.indices) == 0: - raise ValueError("Layer index list is empty. Please provide a non-empty list of layer indices to skip.") - - blocks_found = False - for i, block in enumerate(transformer_blocks): - if i not in config.indices: - continue - - blocks_found = True - - if config.skip_attention and config.skip_ff: - logger.debug(f"Applying TransformerBlockSkipHook to '{config.fqn}.{i}'") - registry = HookRegistry.check_if_exists_or_initialize(block) - hook = TransformerBlockSkipHook(config.dropout) - registry.register_hook(hook, name) - - elif config.skip_attention or config.skip_attention_scores: - for submodule_name, submodule in block.named_modules(): - if isinstance(submodule, _ATTENTION_CLASSES) and not submodule.is_cross_attention: - logger.debug(f"Applying AttentionProcessorSkipHook to '{config.fqn}.{i}.{submodule_name}'") - output_fn = AttentionProcessorRegistry.get(submodule.processor.__class__).skip_processor_output_fn - registry = HookRegistry.check_if_exists_or_initialize(submodule) - hook = AttentionProcessorSkipHook(output_fn, config.skip_attention_scores, config.dropout) - registry.register_hook(hook, name) - - if config.skip_ff: - for submodule_name, submodule in block.named_modules(): - if isinstance(submodule, _FEEDFORWARD_CLASSES): - logger.debug(f"Applying FeedForwardSkipHook to '{config.fqn}.{i}.{submodule_name}'") - registry = HookRegistry.check_if_exists_or_initialize(submodule) - hook = FeedForwardSkipHook(config.dropout) - registry.register_hook(hook, name) - - if not blocks_found: - raise ValueError( - f"Could not find any transformer blocks matching the provided indices {config.indices} and " - f"fully qualified name '{config.fqn}'. Please check the indices and fqn for correctness." - ) diff --git a/diffusers/hooks/layerwise_casting.py b/diffusers/hooks/layerwise_casting.py deleted file mode 100644 index e6dbd73219e30533bcec524ba364b47381651eb6..0000000000000000000000000000000000000000 --- a/diffusers/hooks/layerwise_casting.py +++ /dev/null @@ -1,240 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import re -from typing import Type - -import torch - -from ..utils import get_logger, is_peft_available, is_peft_version -from ._common import _GO_LC_SUPPORTED_PYTORCH_LAYERS -from .hooks import HookRegistry, ModelHook - - -logger = get_logger(__name__) # pylint: disable=invalid-name - - -# fmt: off -_LAYERWISE_CASTING_HOOK = "layerwise_casting" -_PEFT_AUTOCAST_DISABLE_HOOK = "peft_autocast_disable" -DEFAULT_SKIP_MODULES_PATTERN = ("pos_embed", "patch_embed", "norm", "^proj_in$", "^proj_out$") -# fmt: on - -_SHOULD_DISABLE_PEFT_INPUT_AUTOCAST = is_peft_available() and is_peft_version(">", "0.14.0") -if _SHOULD_DISABLE_PEFT_INPUT_AUTOCAST: - from peft.helpers import disable_input_dtype_casting - from peft.tuners.tuners_utils import BaseTunerLayer - - -class LayerwiseCastingHook(ModelHook): - r""" - A hook that casts the weights of a module to a high precision dtype for computation, and to a low precision dtype - for storage. This process may lead to quality loss in the output, but can significantly reduce the memory - footprint. - """ - - _is_stateful = False - - def __init__(self, storage_dtype: torch.dtype, compute_dtype: torch.dtype, non_blocking: bool) -> None: - self.storage_dtype = storage_dtype - self.compute_dtype = compute_dtype - self.non_blocking = non_blocking - - def initialize_hook(self, module: torch.nn.Module): - module.to(dtype=self.storage_dtype, non_blocking=self.non_blocking) - return module - - def deinitalize_hook(self, module: torch.nn.Module): - raise NotImplementedError( - "LayerwiseCastingHook does not support deinitialization. A model once enabled with layerwise casting will " - "have casted its weights to a lower precision dtype for storage. Casting this back to the original dtype " - "will lead to precision loss, which might have an impact on the model's generation quality. The model should " - "be re-initialized and loaded in the original dtype." - ) - - def pre_forward(self, module: torch.nn.Module, *args, **kwargs): - module.to(dtype=self.compute_dtype, non_blocking=self.non_blocking) - return args, kwargs - - def post_forward(self, module: torch.nn.Module, output): - module.to(dtype=self.storage_dtype, non_blocking=self.non_blocking) - return output - - -class PeftInputAutocastDisableHook(ModelHook): - r""" - A hook that disables the casting of inputs to the module weight dtype during the forward pass. By default, PEFT - casts the inputs to the weight dtype of the module, which can lead to precision loss. - - The reasons for needing this are: - - If we don't add PEFT layers' weight names to `skip_modules_pattern` when applying layerwise casting, the - inputs will be casted to the, possibly lower precision, storage dtype. Reference: - https://github.com/huggingface/peft/blob/0facdebf6208139cbd8f3586875acb378813dd97/src/peft/tuners/lora/layer.py#L706 - - We can, on our end, use something like accelerate's `send_to_device` but for dtypes. This way, we can ensure - that the inputs are casted to the computation dtype correctly always. However, there are two goals we are - hoping to achieve: - 1. Making forward implementations independent of device/dtype casting operations as much as possible. - 2. Performing inference without losing information from casting to different precisions. With the current - PEFT implementation (as linked in the reference above), and assuming running layerwise casting inference - with storage_dtype=torch.float8_e4m3fn and compute_dtype=torch.bfloat16, inputs are cast to - torch.float8_e4m3fn in the lora layer. We will then upcast back to torch.bfloat16 when we continue the - forward pass in PEFT linear forward or Diffusers layer forward, with a `send_to_dtype` operation from - LayerwiseCastingHook. This will be a lossy operation and result in poorer generation quality. - """ - - def new_forward(self, module: torch.nn.Module, *args, **kwargs): - with disable_input_dtype_casting(module): - return self.fn_ref.original_forward(*args, **kwargs) - - -def apply_layerwise_casting( - module: torch.nn.Module, - storage_dtype: torch.dtype, - compute_dtype: torch.dtype, - skip_modules_pattern: str | tuple[str, ...] = "auto", - skip_modules_classes: tuple[Type[torch.nn.Module], ...] | None = None, - non_blocking: bool = False, -) -> None: - r""" - Applies layerwise casting to a given module. The module expected here is a Diffusers ModelMixin but it can be any - nn.Module using diffusers layers or pytorch primitives. - - Example: - - ```python - >>> import torch - >>> from diffusers import CogVideoXTransformer3DModel - - >>> transformer = CogVideoXTransformer3DModel.from_pretrained( - ... model_id, subfolder="transformer", torch_dtype=torch.bfloat16 - ... ) - - >>> apply_layerwise_casting( - ... transformer, - ... storage_dtype=torch.float8_e4m3fn, - ... compute_dtype=torch.bfloat16, - ... skip_modules_pattern=["patch_embed", "norm", "proj_out"], - ... non_blocking=True, - ... ) - ``` - - Args: - module (`torch.nn.Module`): - The module whose leaf modules will be cast to a high precision dtype for computation, and to a low - precision dtype for storage. - storage_dtype (`torch.dtype`): - The dtype to cast the module to before/after the forward pass for storage. - compute_dtype (`torch.dtype`): - The dtype to cast the module to during the forward pass for computation. - skip_modules_pattern (`tuple[str, ...]`, defaults to `"auto"`): - A list of patterns to match the names of the modules to skip during the layerwise casting process. If set - to `"auto"`, the default patterns are used. If set to `None`, no modules are skipped. If set to `None` - alongside `skip_modules_classes` being `None`, the layerwise casting is applied directly to the module - instead of its internal submodules. - skip_modules_classes (`tuple[Type[torch.nn.Module], ...]`, defaults to `None`): - A list of module classes to skip during the layerwise casting process. - non_blocking (`bool`, defaults to `False`): - If `True`, the weight casting operations are non-blocking. - """ - if skip_modules_pattern == "auto": - skip_modules_pattern = DEFAULT_SKIP_MODULES_PATTERN - - if skip_modules_classes is None and skip_modules_pattern is None: - apply_layerwise_casting_hook(module, storage_dtype, compute_dtype, non_blocking) - return - - _apply_layerwise_casting( - module, - storage_dtype, - compute_dtype, - skip_modules_pattern, - skip_modules_classes, - non_blocking, - ) - _disable_peft_input_autocast(module) - - -def _apply_layerwise_casting( - module: torch.nn.Module, - storage_dtype: torch.dtype, - compute_dtype: torch.dtype, - skip_modules_pattern: tuple[str, ...] | None = None, - skip_modules_classes: tuple[Type[torch.nn.Module], ...] | None = None, - non_blocking: bool = False, - _prefix: str = "", -) -> None: - should_skip = (skip_modules_classes is not None and isinstance(module, skip_modules_classes)) or ( - skip_modules_pattern is not None and any(re.search(pattern, _prefix) for pattern in skip_modules_pattern) - ) - if should_skip: - logger.debug(f'Skipping layerwise casting for layer "{_prefix}"') - return - - if isinstance(module, _GO_LC_SUPPORTED_PYTORCH_LAYERS): - logger.debug(f'Applying layerwise casting to layer "{_prefix}"') - apply_layerwise_casting_hook(module, storage_dtype, compute_dtype, non_blocking) - return - - for name, submodule in module.named_children(): - layer_name = f"{_prefix}.{name}" if _prefix else name - _apply_layerwise_casting( - submodule, - storage_dtype, - compute_dtype, - skip_modules_pattern, - skip_modules_classes, - non_blocking, - _prefix=layer_name, - ) - - -def apply_layerwise_casting_hook( - module: torch.nn.Module, storage_dtype: torch.dtype, compute_dtype: torch.dtype, non_blocking: bool -) -> None: - r""" - Applies a `LayerwiseCastingHook` to a given module. - - Args: - module (`torch.nn.Module`): - The module to attach the hook to. - storage_dtype (`torch.dtype`): - The dtype to cast the module to before the forward pass. - compute_dtype (`torch.dtype`): - The dtype to cast the module to during the forward pass. - non_blocking (`bool`): - If `True`, the weight casting operations are non-blocking. - """ - registry = HookRegistry.check_if_exists_or_initialize(module) - hook = LayerwiseCastingHook(storage_dtype, compute_dtype, non_blocking) - registry.register_hook(hook, _LAYERWISE_CASTING_HOOK) - - -def _is_layerwise_casting_active(module: torch.nn.Module) -> bool: - for submodule in module.modules(): - if ( - hasattr(submodule, "_diffusers_hook") - and submodule._diffusers_hook.get_hook(_LAYERWISE_CASTING_HOOK) is not None - ): - return True - return False - - -def _disable_peft_input_autocast(module: torch.nn.Module) -> None: - if not _SHOULD_DISABLE_PEFT_INPUT_AUTOCAST: - return - for submodule in module.modules(): - if isinstance(submodule, BaseTunerLayer) and _is_layerwise_casting_active(submodule): - registry = HookRegistry.check_if_exists_or_initialize(submodule) - hook = PeftInputAutocastDisableHook() - registry.register_hook(hook, _PEFT_AUTOCAST_DISABLE_HOOK) diff --git a/diffusers/hooks/mag_cache.py b/diffusers/hooks/mag_cache.py deleted file mode 100644 index e5f0aaebc01a25b96b5253245cca838db7622df4..0000000000000000000000000000000000000000 --- a/diffusers/hooks/mag_cache.py +++ /dev/null @@ -1,468 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass -from typing import List, Optional, Tuple, Union - -import torch - -from ..utils import get_logger -from ..utils.torch_utils import unwrap_module -from ._common import _ALL_TRANSFORMER_BLOCK_IDENTIFIERS -from ._helpers import TransformerBlockRegistry -from .hooks import BaseState, HookRegistry, ModelHook, StateManager - - -logger = get_logger(__name__) # pylint: disable=invalid-name - -_MAG_CACHE_LEADER_BLOCK_HOOK = "mag_cache_leader_block_hook" -_MAG_CACHE_BLOCK_HOOK = "mag_cache_block_hook" - -# Default Mag Ratios for Flux models (Dev/Schnell) are provided for convenience. -# Users must explicitly pass these to the config if using Flux. -# Reference: https://github.com/Zehong-Ma/MagCache -FLUX_MAG_RATIOS = torch.tensor( - [1.0] - + [ - 1.21094, - 1.11719, - 1.07812, - 1.0625, - 1.03906, - 1.03125, - 1.03906, - 1.02344, - 1.03125, - 1.02344, - 0.98047, - 1.01562, - 1.00781, - 1.0, - 1.00781, - 1.0, - 1.00781, - 1.0, - 1.0, - 0.99609, - 0.99609, - 0.98047, - 0.98828, - 0.96484, - 0.95703, - 0.93359, - 0.89062, - ] -) - - -def nearest_interp(src_array: torch.Tensor, target_length: int) -> torch.Tensor: - """ - Interpolate the source array to the target length using nearest neighbor interpolation. - """ - src_length = len(src_array) - if target_length == 1: - return src_array[-1:] - - scale = (src_length - 1) / (target_length - 1) - grid = torch.arange(target_length, device=src_array.device, dtype=torch.float32) - mapped_indices = torch.round(grid * scale).long() - return src_array[mapped_indices] - - -@dataclass -class MagCacheConfig: - r""" - Configuration for [MagCache](https://github.com/Zehong-Ma/MagCache). - - Args: - threshold (`float`, defaults to `0.06`): - The threshold for the accumulated error. If the accumulated error is below this threshold, the block - computation is skipped. A higher threshold allows for more aggressive skipping (faster) but may degrade - quality. - max_skip_steps (`int`, defaults to `3`): - The maximum number of consecutive steps that can be skipped (K in the paper). - retention_ratio (`float`, defaults to `0.2`): - The fraction of initial steps during which skipping is disabled to ensure stability. For example, if - `num_inference_steps` is 28 and `retention_ratio` is 0.2, the first 6 steps will never be skipped. - num_inference_steps (`int`, defaults to `28`): - The number of inference steps used in the pipeline. This is required to interpolate `mag_ratios` correctly. - mag_ratios (`torch.Tensor`, *optional*): - The pre-computed magnitude ratios for the model. These are checkpoint-dependent. If not provided, you must - set `calibrate=True` to calculate them for your specific model. For Flux models, you can use - `diffusers.hooks.mag_cache.FLUX_MAG_RATIOS`. - calibrate (`bool`, defaults to `False`): - If True, enables calibration mode. In this mode, no blocks are skipped. Instead, the hook calculates the - magnitude ratios for the current run and logs them at the end. Use this to obtain `mag_ratios` for new - models or schedulers. - """ - - threshold: float = 0.06 - max_skip_steps: int = 3 - retention_ratio: float = 0.2 - num_inference_steps: int = 28 - mag_ratios: Optional[Union[torch.Tensor, List[float]]] = None - calibrate: bool = False - - def __post_init__(self): - # User MUST provide ratios OR enable calibration. - if self.mag_ratios is None and not self.calibrate: - raise ValueError( - " `mag_ratios` must be provided for MagCache inference because these ratios are model-dependent.\n" - "To get them for your model:\n" - "1. Initialize `MagCacheConfig(calibrate=True, ...)`\n" - "2. Run inference on your model once.\n" - "3. Copy the printed ratios array and pass it to `mag_ratios` in the config.\n" - "For Flux models, you can import `FLUX_MAG_RATIOS` from `diffusers.hooks.mag_cache`." - ) - - if not self.calibrate and self.mag_ratios is not None: - if not torch.is_tensor(self.mag_ratios): - self.mag_ratios = torch.tensor(self.mag_ratios) - - if len(self.mag_ratios) != self.num_inference_steps: - logger.debug( - f"Interpolating mag_ratios from length {len(self.mag_ratios)} to {self.num_inference_steps}" - ) - self.mag_ratios = nearest_interp(self.mag_ratios, self.num_inference_steps) - - -class MagCacheState(BaseState): - def __init__(self) -> None: - super().__init__() - # Cache for the residual (output - input) from the *previous* timestep - self.previous_residual: torch.Tensor = None - - # State inputs/outputs for the current forward pass - self.head_block_input: Union[torch.Tensor, Tuple[torch.Tensor, ...]] = None - self.should_compute: bool = True - - # MagCache accumulators - self.accumulated_ratio: float = 1.0 - self.accumulated_err: float = 0.0 - self.accumulated_steps: int = 0 - - # Current step counter (timestep index) - self.step_index: int = 0 - - # Calibration storage - self.calibration_ratios: List[float] = [] - - def reset(self): - self.previous_residual = None - self.should_compute = True - self.accumulated_ratio = 1.0 - self.accumulated_err = 0.0 - self.accumulated_steps = 0 - self.step_index = 0 - self.calibration_ratios = [] - - -class MagCacheHeadHook(ModelHook): - _is_stateful = True - - def __init__(self, state_manager: StateManager, config: MagCacheConfig): - self.state_manager = state_manager - self.config = config - self._metadata = None - - def initialize_hook(self, module): - unwrapped_module = unwrap_module(module) - self._metadata = TransformerBlockRegistry.get(unwrapped_module.__class__) - return module - - @torch.compiler.disable - def new_forward(self, module: torch.nn.Module, *args, **kwargs): - if self.state_manager._current_context is None: - self.state_manager.set_context("inference") - - arg_name = self._metadata.hidden_states_argument_name - hidden_states = self._metadata._get_parameter_from_args_kwargs(arg_name, args, kwargs) - - state: MagCacheState = self.state_manager.get_state() - state.head_block_input = hidden_states - - should_compute = True - - if self.config.calibrate: - # Never skip during calibration - should_compute = True - else: - # MagCache Logic - current_step = state.step_index - if current_step >= len(self.config.mag_ratios): - current_scale = 1.0 - else: - current_scale = self.config.mag_ratios[current_step] - - retention_step = int(self.config.retention_ratio * self.config.num_inference_steps + 0.5) - - if current_step >= retention_step: - state.accumulated_ratio *= current_scale - state.accumulated_steps += 1 - state.accumulated_err += abs(1.0 - state.accumulated_ratio) - - if ( - state.previous_residual is not None - and state.accumulated_err <= self.config.threshold - and state.accumulated_steps <= self.config.max_skip_steps - ): - should_compute = False - else: - state.accumulated_ratio = 1.0 - state.accumulated_steps = 0 - state.accumulated_err = 0.0 - - state.should_compute = should_compute - - if not should_compute: - logger.debug(f"MagCache: Skipping step {state.step_index}") - # Apply MagCache: Output = Input + Previous Residual - - output = hidden_states - res = state.previous_residual - - if res.device != output.device: - res = res.to(output.device) - - # Attempt to apply residual handling shape mismatches (e.g., text+image vs image only) - if res.shape == output.shape: - output = output + res - elif ( - output.ndim == 3 - and res.ndim == 3 - and output.shape[0] == res.shape[0] - and output.shape[2] == res.shape[2] - ): - # Assuming concatenation where image part is at the end (standard in Flux/SD3) - diff = output.shape[1] - res.shape[1] - if diff > 0: - output = output.clone() - output[:, diff:, :] = output[:, diff:, :] + res - else: - logger.warning( - f"MagCache: Dimension mismatch. Input {output.shape}, Residual {res.shape}. " - "Cannot apply residual safely. Returning input without residual." - ) - else: - logger.warning( - f"MagCache: Dimension mismatch. Input {output.shape}, Residual {res.shape}. " - "Cannot apply residual safely. Returning input without residual." - ) - - if self._metadata.return_encoder_hidden_states_index is not None: - original_encoder_hidden_states = self._metadata._get_parameter_from_args_kwargs( - "encoder_hidden_states", args, kwargs - ) - max_idx = max( - self._metadata.return_hidden_states_index, self._metadata.return_encoder_hidden_states_index - ) - ret_list = [None] * (max_idx + 1) - ret_list[self._metadata.return_hidden_states_index] = output - ret_list[self._metadata.return_encoder_hidden_states_index] = original_encoder_hidden_states - return tuple(ret_list) - else: - return output - - else: - # Compute original forward - output = self.fn_ref.original_forward(*args, **kwargs) - return output - - def reset_state(self, module): - self.state_manager.reset() - return module - - -class MagCacheBlockHook(ModelHook): - def __init__(self, state_manager: StateManager, is_tail: bool = False, config: MagCacheConfig = None): - super().__init__() - self.state_manager = state_manager - self.is_tail = is_tail - self.config = config - self._metadata = None - - def initialize_hook(self, module): - unwrapped_module = unwrap_module(module) - self._metadata = TransformerBlockRegistry.get(unwrapped_module.__class__) - return module - - @torch.compiler.disable - def new_forward(self, module: torch.nn.Module, *args, **kwargs): - if self.state_manager._current_context is None: - self.state_manager.set_context("inference") - state: MagCacheState = self.state_manager.get_state() - - if not state.should_compute: - arg_name = self._metadata.hidden_states_argument_name - hidden_states = self._metadata._get_parameter_from_args_kwargs(arg_name, args, kwargs) - - if self.is_tail: - # Still need to advance step index even if we skip - self._advance_step(state) - - if self._metadata.return_encoder_hidden_states_index is not None: - encoder_hidden_states = self._metadata._get_parameter_from_args_kwargs( - "encoder_hidden_states", args, kwargs - ) - max_idx = max( - self._metadata.return_hidden_states_index, self._metadata.return_encoder_hidden_states_index - ) - ret_list = [None] * (max_idx + 1) - ret_list[self._metadata.return_hidden_states_index] = hidden_states - ret_list[self._metadata.return_encoder_hidden_states_index] = encoder_hidden_states - return tuple(ret_list) - - return hidden_states - - output = self.fn_ref.original_forward(*args, **kwargs) - - if self.is_tail: - # Calculate residual for next steps - if isinstance(output, tuple): - out_hidden = output[self._metadata.return_hidden_states_index] - else: - out_hidden = output - - in_hidden = state.head_block_input - - if in_hidden is None: - return output - - # Determine residual - if out_hidden.shape == in_hidden.shape: - residual = out_hidden - in_hidden - elif out_hidden.ndim == 3 and in_hidden.ndim == 3 and out_hidden.shape[2] == in_hidden.shape[2]: - diff = in_hidden.shape[1] - out_hidden.shape[1] - if diff == 0: - residual = out_hidden - in_hidden - else: - residual = out_hidden - in_hidden # Fallback to matching tail - else: - # Fallback for completely mismatched shapes - residual = out_hidden - - if self.config.calibrate: - self._perform_calibration_step(state, residual) - - state.previous_residual = residual - self._advance_step(state) - - return output - - def _perform_calibration_step(self, state: MagCacheState, current_residual: torch.Tensor): - if state.previous_residual is None: - # First step has no previous residual to compare against. - # log 1.0 as a neutral starting point. - ratio = 1.0 - else: - # MagCache Calibration Formula: mean(norm(curr) / norm(prev)) - # norm(dim=-1) gives magnitude of each token vector - curr_norm = torch.linalg.norm(current_residual.float(), dim=-1) - prev_norm = torch.linalg.norm(state.previous_residual.float(), dim=-1) - - # Avoid division by zero - ratio = (curr_norm / (prev_norm + 1e-8)).mean().item() - - state.calibration_ratios.append(ratio) - - def _advance_step(self, state: MagCacheState): - state.step_index += 1 - if state.step_index >= self.config.num_inference_steps: - # End of inference loop - if self.config.calibrate: - print("\n[MagCache] Calibration Complete. Copy these values to MagCacheConfig(mag_ratios=...):") - print(f"{state.calibration_ratios}\n") - logger.info(f"MagCache Calibration Results: {state.calibration_ratios}") - - # Reset state - state.step_index = 0 - state.accumulated_ratio = 1.0 - state.accumulated_steps = 0 - state.accumulated_err = 0.0 - state.previous_residual = None - state.calibration_ratios = [] - - -def apply_mag_cache(module: torch.nn.Module, config: MagCacheConfig) -> None: - """ - Applies MagCache to a given module (typically a Transformer). - - Args: - module (`torch.nn.Module`): - The module to apply MagCache to. - config (`MagCacheConfig`): - The configuration for MagCache. - """ - # Initialize registry on the root module so the Pipeline can set context. - HookRegistry.check_if_exists_or_initialize(module) - - state_manager = StateManager(MagCacheState, (), {}) - remaining_blocks = [] - - for name, submodule in module.named_children(): - if name not in _ALL_TRANSFORMER_BLOCK_IDENTIFIERS or not isinstance(submodule, torch.nn.ModuleList): - continue - for index, block in enumerate(submodule): - remaining_blocks.append((f"{name}.{index}", block)) - - if not remaining_blocks: - logger.warning("MagCache: No transformer blocks found to apply hooks.") - return - - # Handle single-block models - if len(remaining_blocks) == 1: - name, block = remaining_blocks[0] - logger.info(f"MagCache: Applying Head+Tail Hooks to single block '{name}'") - _apply_mag_cache_block_hook(block, state_manager, config, is_tail=True) - _apply_mag_cache_head_hook(block, state_manager, config) - return - - head_block_name, head_block = remaining_blocks.pop(0) - tail_block_name, tail_block = remaining_blocks.pop(-1) - - logger.info(f"MagCache: Applying Head Hook to {head_block_name}") - _apply_mag_cache_head_hook(head_block, state_manager, config) - - for name, block in remaining_blocks: - _apply_mag_cache_block_hook(block, state_manager, config) - - logger.info(f"MagCache: Applying Tail Hook to {tail_block_name}") - _apply_mag_cache_block_hook(tail_block, state_manager, config, is_tail=True) - - -def _apply_mag_cache_head_hook(block: torch.nn.Module, state_manager: StateManager, config: MagCacheConfig) -> None: - registry = HookRegistry.check_if_exists_or_initialize(block) - - # Automatically remove existing hook to allow re-application (e.g. switching modes) - if registry.get_hook(_MAG_CACHE_LEADER_BLOCK_HOOK) is not None: - registry.remove_hook(_MAG_CACHE_LEADER_BLOCK_HOOK) - - hook = MagCacheHeadHook(state_manager, config) - registry.register_hook(hook, _MAG_CACHE_LEADER_BLOCK_HOOK) - - -def _apply_mag_cache_block_hook( - block: torch.nn.Module, - state_manager: StateManager, - config: MagCacheConfig, - is_tail: bool = False, -) -> None: - registry = HookRegistry.check_if_exists_or_initialize(block) - - # Automatically remove existing hook to allow re-application - if registry.get_hook(_MAG_CACHE_BLOCK_HOOK) is not None: - registry.remove_hook(_MAG_CACHE_BLOCK_HOOK) - - hook = MagCacheBlockHook(state_manager, is_tail, config) - registry.register_hook(hook, _MAG_CACHE_BLOCK_HOOK) diff --git a/diffusers/hooks/pyramid_attention_broadcast.py b/diffusers/hooks/pyramid_attention_broadcast.py deleted file mode 100644 index e7ed26b28778f57928bbf05b30ea3ac3adeec41b..0000000000000000000000000000000000000000 --- a/diffusers/hooks/pyramid_attention_broadcast.py +++ /dev/null @@ -1,314 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import re -from dataclasses import dataclass -from typing import Any, Callable - -import torch - -from ..models.attention import AttentionModuleMixin -from ..models.attention_processor import Attention, MochiAttention -from ..utils import logging -from ._common import ( - _ATTENTION_CLASSES, - _CROSS_TRANSFORMER_BLOCK_IDENTIFIERS, - _SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS, - _TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS, -) -from .hooks import HookRegistry, ModelHook - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -_PYRAMID_ATTENTION_BROADCAST_HOOK = "pyramid_attention_broadcast" - - -@dataclass -class PyramidAttentionBroadcastConfig: - r""" - Configuration for Pyramid Attention Broadcast. - - Args: - spatial_attention_block_skip_range (`int`, *optional*, defaults to `None`): - The number of times a specific spatial attention broadcast is skipped before computing the attention states - to re-use. If this is set to the value `N`, the attention computation will be skipped `N - 1` times (i.e., - old attention states will be reused) before computing the new attention states again. - temporal_attention_block_skip_range (`int`, *optional*, defaults to `None`): - The number of times a specific temporal attention broadcast is skipped before computing the attention - states to re-use. If this is set to the value `N`, the attention computation will be skipped `N - 1` times - (i.e., old attention states will be reused) before computing the new attention states again. - cross_attention_block_skip_range (`int`, *optional*, defaults to `None`): - The number of times a specific cross-attention broadcast is skipped before computing the attention states - to re-use. If this is set to the value `N`, the attention computation will be skipped `N - 1` times (i.e., - old attention states will be reused) before computing the new attention states again. - spatial_attention_timestep_skip_range (`tuple[int, int]`, defaults to `(100, 800)`): - The range of timesteps to skip in the spatial attention layer. The attention computations will be - conditionally skipped if the current timestep is within the specified range. - temporal_attention_timestep_skip_range (`tuple[int, int]`, defaults to `(100, 800)`): - The range of timesteps to skip in the temporal attention layer. The attention computations will be - conditionally skipped if the current timestep is within the specified range. - cross_attention_timestep_skip_range (`tuple[int, int]`, defaults to `(100, 800)`): - The range of timesteps to skip in the cross-attention layer. The attention computations will be - conditionally skipped if the current timestep is within the specified range. - spatial_attention_block_identifiers (`tuple[str, ...]`): - The identifiers to match against the layer names to determine if the layer is a spatial attention layer. - temporal_attention_block_identifiers (`tuple[str, ...]`): - The identifiers to match against the layer names to determine if the layer is a temporal attention layer. - cross_attention_block_identifiers (`tuple[str, ...]`): - The identifiers to match against the layer names to determine if the layer is a cross-attention layer. - """ - - spatial_attention_block_skip_range: int | None = None - temporal_attention_block_skip_range: int | None = None - cross_attention_block_skip_range: int | None = None - - spatial_attention_timestep_skip_range: tuple[int, int] = (100, 800) - temporal_attention_timestep_skip_range: tuple[int, int] = (100, 800) - cross_attention_timestep_skip_range: tuple[int, int] = (100, 800) - - spatial_attention_block_identifiers: tuple[str, ...] = _SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS - temporal_attention_block_identifiers: tuple[str, ...] = _TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS - cross_attention_block_identifiers: tuple[str, ...] = _CROSS_TRANSFORMER_BLOCK_IDENTIFIERS - - current_timestep_callback: Callable[[], int] = None - - # TODO(aryan): add PAB for MLP layers (very limited speedup from testing with original codebase - # so not added for now) - - def __repr__(self) -> str: - return ( - f"PyramidAttentionBroadcastConfig(\n" - f" spatial_attention_block_skip_range={self.spatial_attention_block_skip_range},\n" - f" temporal_attention_block_skip_range={self.temporal_attention_block_skip_range},\n" - f" cross_attention_block_skip_range={self.cross_attention_block_skip_range},\n" - f" spatial_attention_timestep_skip_range={self.spatial_attention_timestep_skip_range},\n" - f" temporal_attention_timestep_skip_range={self.temporal_attention_timestep_skip_range},\n" - f" cross_attention_timestep_skip_range={self.cross_attention_timestep_skip_range},\n" - f" spatial_attention_block_identifiers={self.spatial_attention_block_identifiers},\n" - f" temporal_attention_block_identifiers={self.temporal_attention_block_identifiers},\n" - f" cross_attention_block_identifiers={self.cross_attention_block_identifiers},\n" - f" current_timestep_callback={self.current_timestep_callback}\n" - ")" - ) - - -class PyramidAttentionBroadcastState: - r""" - State for Pyramid Attention Broadcast. - - Attributes: - iteration (`int`): - The current iteration of the Pyramid Attention Broadcast. It is necessary to ensure that `reset_state` is - called before starting a new inference forward pass for PAB to work correctly. - cache (`Any`): - The cached output from the previous forward pass. This is used to re-use the attention states when the - attention computation is skipped. It is either a tensor or a tuple of tensors, depending on the module. - """ - - def __init__(self) -> None: - self.iteration = 0 - self.cache = None - - def reset(self): - self.iteration = 0 - self.cache = None - - def __repr__(self): - cache_repr = "" - if self.cache is None: - cache_repr = "None" - else: - cache_repr = f"Tensor(shape={self.cache.shape}, dtype={self.cache.dtype})" - return f"PyramidAttentionBroadcastState(iteration={self.iteration}, cache={cache_repr})" - - -class PyramidAttentionBroadcastHook(ModelHook): - r"""A hook that applies Pyramid Attention Broadcast to a given module.""" - - _is_stateful = True - - def __init__( - self, timestep_skip_range: tuple[int, int], block_skip_range: int, current_timestep_callback: Callable[[], int] - ) -> None: - super().__init__() - - self.timestep_skip_range = timestep_skip_range - self.block_skip_range = block_skip_range - self.current_timestep_callback = current_timestep_callback - - def initialize_hook(self, module): - self.state = PyramidAttentionBroadcastState() - return module - - def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any: - is_within_timestep_range = ( - self.timestep_skip_range[0] < self.current_timestep_callback() < self.timestep_skip_range[1] - ) - should_compute_attention = ( - self.state.cache is None - or self.state.iteration == 0 - or not is_within_timestep_range - or self.state.iteration % self.block_skip_range == 0 - ) - - if should_compute_attention: - output = self.fn_ref.original_forward(*args, **kwargs) - else: - output = self.state.cache - - self.state.cache = output - self.state.iteration += 1 - return output - - def reset_state(self, module: torch.nn.Module) -> None: - self.state.reset() - return module - - -def apply_pyramid_attention_broadcast(module: torch.nn.Module, config: PyramidAttentionBroadcastConfig): - r""" - Apply [Pyramid Attention Broadcast](https://huggingface.co/papers/2408.12588) to a given pipeline. - - PAB is an attention approximation method that leverages the similarity in attention states between timesteps to - reduce the computational cost of attention computation. The key takeaway from the paper is that the attention - similarity in the cross-attention layers between timesteps is high, followed by less similarity in the temporal and - spatial layers. This allows for the skipping of attention computation in the cross-attention layers more frequently - than in the temporal and spatial layers. Applying PAB will, therefore, speedup the inference process. - - Args: - module (`torch.nn.Module`): - The module to apply Pyramid Attention Broadcast to. - config (`PyramidAttentionBroadcastConfig | None`, `optional`, defaults to `None`): - The configuration to use for Pyramid Attention Broadcast. - - Example: - - ```python - >>> import torch - >>> from diffusers import CogVideoXPipeline, PyramidAttentionBroadcastConfig, apply_pyramid_attention_broadcast - >>> from diffusers.utils import export_to_video - - >>> pipe = CogVideoXPipeline.from_pretrained("THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16) - >>> pipe.to("cuda") - - >>> config = PyramidAttentionBroadcastConfig( - ... spatial_attention_block_skip_range=2, - ... spatial_attention_timestep_skip_range=(100, 800), - ... current_timestep_callback=lambda: pipe.current_timestep, - ... ) - >>> apply_pyramid_attention_broadcast(pipe.transformer, config) - ``` - """ - if config.current_timestep_callback is None: - raise ValueError( - "The `current_timestep_callback` function must be provided in the configuration to apply Pyramid Attention Broadcast." - ) - - if ( - config.spatial_attention_block_skip_range is None - and config.temporal_attention_block_skip_range is None - and config.cross_attention_block_skip_range is None - ): - logger.warning( - "Pyramid Attention Broadcast requires one or more of `spatial_attention_block_skip_range`, `temporal_attention_block_skip_range` " - "or `cross_attention_block_skip_range` parameters to be set to an integer, not `None`. Defaulting to using `spatial_attention_block_skip_range=2`. " - "To avoid this warning, please set one of the above parameters." - ) - config.spatial_attention_block_skip_range = 2 - - for name, submodule in module.named_modules(): - if not isinstance(submodule, (*_ATTENTION_CLASSES, AttentionModuleMixin)): - # PAB has been implemented specific to Diffusers' Attention classes. However, this does not mean that PAB - # cannot be applied to this layer. For custom layers, users can extend this functionality and implement - # their own PAB logic similar to `_apply_pyramid_attention_broadcast_on_attention_class`. - continue - _apply_pyramid_attention_broadcast_on_attention_class(name, submodule, config) - - -def _apply_pyramid_attention_broadcast_on_attention_class( - name: str, module: Attention, config: PyramidAttentionBroadcastConfig -) -> bool: - is_spatial_self_attention = ( - any(re.search(identifier, name) is not None for identifier in config.spatial_attention_block_identifiers) - and config.spatial_attention_block_skip_range is not None - and not getattr(module, "is_cross_attention", False) - ) - is_temporal_self_attention = ( - any(re.search(identifier, name) is not None for identifier in config.temporal_attention_block_identifiers) - and config.temporal_attention_block_skip_range is not None - and not getattr(module, "is_cross_attention", False) - ) - is_cross_attention = ( - any(re.search(identifier, name) is not None for identifier in config.cross_attention_block_identifiers) - and config.cross_attention_block_skip_range is not None - and getattr(module, "is_cross_attention", False) - ) - - block_skip_range, timestep_skip_range, block_type = None, None, None - if is_spatial_self_attention: - block_skip_range = config.spatial_attention_block_skip_range - timestep_skip_range = config.spatial_attention_timestep_skip_range - block_type = "spatial" - elif is_temporal_self_attention: - block_skip_range = config.temporal_attention_block_skip_range - timestep_skip_range = config.temporal_attention_timestep_skip_range - block_type = "temporal" - elif is_cross_attention: - block_skip_range = config.cross_attention_block_skip_range - timestep_skip_range = config.cross_attention_timestep_skip_range - block_type = "cross" - - if block_skip_range is None or timestep_skip_range is None: - logger.info( - f'Unable to apply Pyramid Attention Broadcast to the selected layer: "{name}" because it does ' - f"not match any of the required criteria for spatial, temporal or cross attention layers. Note, " - f"however, that this layer may still be valid for applying PAB. Please specify the correct " - f"block identifiers in the configuration." - ) - return False - - logger.debug(f"Enabling Pyramid Attention Broadcast ({block_type}) in layer: {name}") - _apply_pyramid_attention_broadcast_hook( - module, timestep_skip_range, block_skip_range, config.current_timestep_callback - ) - return True - - -def _apply_pyramid_attention_broadcast_hook( - module: Attention | MochiAttention, - timestep_skip_range: tuple[int, int], - block_skip_range: int, - current_timestep_callback: Callable[[], int], -): - r""" - Apply [Pyramid Attention Broadcast](https://huggingface.co/papers/2408.12588) to a given torch.nn.Module. - - Args: - module (`torch.nn.Module`): - The module to apply Pyramid Attention Broadcast to. - timestep_skip_range (`tuple[int, int]`): - The range of timesteps to skip in the attention layer. The attention computations will be conditionally - skipped if the current timestep is within the specified range. - block_skip_range (`int`): - The number of times a specific attention broadcast is skipped before computing the attention states to - re-use. If this is set to the value `N`, the attention computation will be skipped `N - 1` times (i.e., old - attention states will be reused) before computing the new attention states again. - current_timestep_callback (`Callable[[], int]`): - A callback function that returns the current inference timestep. - """ - registry = HookRegistry.check_if_exists_or_initialize(module) - hook = PyramidAttentionBroadcastHook(timestep_skip_range, block_skip_range, current_timestep_callback) - registry.register_hook(hook, _PYRAMID_ATTENTION_BROADCAST_HOOK) diff --git a/diffusers/hooks/smoothed_energy_guidance_utils.py b/diffusers/hooks/smoothed_energy_guidance_utils.py deleted file mode 100644 index 868f4d07c765a847c43b5b1a96bfcd5c19747551..0000000000000000000000000000000000000000 --- a/diffusers/hooks/smoothed_energy_guidance_utils.py +++ /dev/null @@ -1,166 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from dataclasses import asdict, dataclass - -import torch -import torch.nn.functional as F - -from ..utils import get_logger -from ._common import _ALL_TRANSFORMER_BLOCK_IDENTIFIERS, _ATTENTION_CLASSES, _get_submodule_from_fqn -from .hooks import HookRegistry, ModelHook - - -logger = get_logger(__name__) # pylint: disable=invalid-name - -_SMOOTHED_ENERGY_GUIDANCE_HOOK = "smoothed_energy_guidance_hook" - - -@dataclass -class SmoothedEnergyGuidanceConfig: - r""" - Configuration for skipping internal transformer blocks when executing a transformer model. - - Args: - indices (`list[int]`): - The indices of the layer to skip. This is typically the first layer in the transformer block. - fqn (`str`, defaults to `"auto"`): - The fully qualified name identifying the stack of transformer blocks. Typically, this is - `transformer_blocks`, `single_transformer_blocks`, `blocks`, `layers`, or `temporal_transformer_blocks`. - For automatic detection, set this to `"auto"`. "auto" only works on DiT models. For UNet models, you must - provide the correct fqn. - _query_proj_identifiers (`list[str]`, defaults to `None`): - The identifiers for the query projection layers. Typically, these are `to_q`, `query`, or `q_proj`. If - `None`, `to_q` is used by default. - """ - - indices: list[int] - fqn: str = "auto" - _query_proj_identifiers: list[str] = None - - def to_dict(self): - return asdict(self) - - @staticmethod - def from_dict(data: dict) -> "SmoothedEnergyGuidanceConfig": - return SmoothedEnergyGuidanceConfig(**data) - - -class SmoothedEnergyGuidanceHook(ModelHook): - def __init__(self, blur_sigma: float = 1.0, blur_threshold_inf: float = 9999.9) -> None: - super().__init__() - self.blur_sigma = blur_sigma - self.blur_threshold_inf = blur_threshold_inf - - def post_forward(self, module: torch.nn.Module, output: torch.Tensor) -> torch.Tensor: - # Copied from https://github.com/SusungHong/SEG-SDXL/blob/cf8256d640d5373541cfea3b3b6caf93272cf986/pipeline_seg.py#L172C31-L172C102 - kernel_size = math.ceil(6 * self.blur_sigma) + 1 - math.ceil(6 * self.blur_sigma) % 2 - smoothed_output = _gaussian_blur_2d(output, kernel_size, self.blur_sigma, self.blur_threshold_inf) - return smoothed_output - - -def _apply_smoothed_energy_guidance_hook( - module: torch.nn.Module, config: SmoothedEnergyGuidanceConfig, blur_sigma: float, name: str | None = None -) -> None: - name = name or _SMOOTHED_ENERGY_GUIDANCE_HOOK - - if config.fqn == "auto": - for identifier in _ALL_TRANSFORMER_BLOCK_IDENTIFIERS: - if hasattr(module, identifier): - config.fqn = identifier - break - else: - raise ValueError( - "Could not find a suitable identifier for the transformer blocks automatically. Please provide a valid " - "`fqn` (fully qualified name) that identifies a stack of transformer blocks." - ) - - if config._query_proj_identifiers is None: - config._query_proj_identifiers = ["to_q"] - - transformer_blocks = _get_submodule_from_fqn(module, config.fqn) - blocks_found = False - for i, block in enumerate(transformer_blocks): - if i not in config.indices: - continue - - blocks_found = True - - for submodule_name, submodule in block.named_modules(): - if not isinstance(submodule, _ATTENTION_CLASSES) or submodule.is_cross_attention: - continue - for identifier in config._query_proj_identifiers: - query_proj = getattr(submodule, identifier, None) - if query_proj is None or not isinstance(query_proj, torch.nn.Linear): - continue - logger.debug( - f"Registering smoothed energy guidance hook on {config.fqn}.{i}.{submodule_name}.{identifier}" - ) - registry = HookRegistry.check_if_exists_or_initialize(query_proj) - hook = SmoothedEnergyGuidanceHook(blur_sigma) - registry.register_hook(hook, name) - - if not blocks_found: - raise ValueError( - f"Could not find any transformer blocks matching the provided indices {config.indices} and " - f"fully qualified name '{config.fqn}'. Please check the indices and fqn for correctness." - ) - - -# Modified from https://github.com/SusungHong/SEG-SDXL/blob/cf8256d640d5373541cfea3b3b6caf93272cf986/pipeline_seg.py#L71 -def _gaussian_blur_2d(query: torch.Tensor, kernel_size: int, sigma: float, sigma_threshold_inf: float) -> torch.Tensor: - """ - This implementation assumes that the input query is for visual (image/videos) tokens to apply the 2D gaussian blur. - However, some models use joint text-visual token attention for which this may not be suitable. Additionally, this - implementation also assumes that the visual tokens come from a square image/video. In practice, despite these - assumptions, applying the 2D square gaussian blur on the query projections generates reasonable results for - Smoothed Energy Guidance. - - SEG is only supported as an experimental prototype feature for now, so the implementation may be modified in the - future without warning or guarantee of reproducibility. - """ - assert query.ndim == 3 - - is_inf = sigma > sigma_threshold_inf - batch_size, seq_len, embed_dim = query.shape - - seq_len_sqrt = int(math.sqrt(seq_len)) - num_square_tokens = seq_len_sqrt * seq_len_sqrt - query_slice = query[:, :num_square_tokens, :] - query_slice = query_slice.permute(0, 2, 1) - query_slice = query_slice.reshape(batch_size, embed_dim, seq_len_sqrt, seq_len_sqrt) - - if is_inf: - kernel_size = min(kernel_size, seq_len_sqrt - (seq_len_sqrt % 2 - 1)) - kernel_size_half = (kernel_size - 1) / 2 - - x = torch.linspace(-kernel_size_half, kernel_size_half, steps=kernel_size) - pdf = torch.exp(-0.5 * (x / sigma).pow(2)) - kernel1d = pdf / pdf.sum() - kernel1d = kernel1d.to(query) - kernel2d = torch.matmul(kernel1d[:, None], kernel1d[None, :]) - kernel2d = kernel2d.expand(embed_dim, 1, kernel2d.shape[0], kernel2d.shape[1]) - - padding = [kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2] - query_slice = F.pad(query_slice, padding, mode="reflect") - query_slice = F.conv2d(query_slice, kernel2d, groups=embed_dim) - else: - query_slice[:] = query_slice.mean(dim=(-2, -1), keepdim=True) - - query_slice = query_slice.reshape(batch_size, embed_dim, num_square_tokens) - query_slice = query_slice.permute(0, 2, 1) - query[:, :num_square_tokens, :] = query_slice.clone() - - return query diff --git a/diffusers/hooks/taylorseer_cache.py b/diffusers/hooks/taylorseer_cache.py deleted file mode 100644 index 303155105e71e397d6a2990fa2264aa871ef980f..0000000000000000000000000000000000000000 --- a/diffusers/hooks/taylorseer_cache.py +++ /dev/null @@ -1,345 +0,0 @@ -import math -import re -from dataclasses import dataclass - -import torch -import torch.nn as nn - -from ..utils import logging -from .hooks import HookRegistry, ModelHook, StateManager - - -logger = logging.get_logger(__name__) -_TAYLORSEER_CACHE_HOOK = "taylorseer_cache" -_SPATIAL_ATTENTION_BLOCK_IDENTIFIERS = ( - "^blocks.*attn", - "^transformer_blocks.*attn", - "^single_transformer_blocks.*attn", -) -_TEMPORAL_ATTENTION_BLOCK_IDENTIFIERS = ("^temporal_transformer_blocks.*attn",) -_TRANSFORMER_BLOCK_IDENTIFIERS = _SPATIAL_ATTENTION_BLOCK_IDENTIFIERS + _TEMPORAL_ATTENTION_BLOCK_IDENTIFIERS -_BLOCK_IDENTIFIERS = ("^[^.]*block[^.]*\\.[^.]+$",) -_PROJ_OUT_IDENTIFIERS = ("^proj_out$",) - - -@dataclass -class TaylorSeerCacheConfig: - """ - Configuration for TaylorSeer cache. See: https://huggingface.co/papers/2503.06923 - - Attributes: - cache_interval (`int`, defaults to `5`): - The interval between full computation steps. After a full computation, the cached (predicted) outputs are - reused for this many subsequent denoising steps before refreshing with a new full forward pass. - - disable_cache_before_step (`int`, defaults to `3`): - The denoising step index before which caching is disabled, meaning full computation is performed for the - initial steps (0 to disable_cache_before_step - 1) to gather data for Taylor series approximations. During - these steps, Taylor factors are updated, but caching/predictions are not applied. Caching begins at this - step. - - disable_cache_after_step (`int`, *optional*, defaults to `None`): - The denoising step index after which caching is disabled. If set, for steps >= this value, all modules run - full computations without predictions or state updates, ensuring accuracy in later stages if needed. - - max_order (`int`, defaults to `1`): - The highest order in the Taylor series expansion for approximating module outputs. Higher orders provide - better approximations but increase computation and memory usage. - - taylor_factors_dtype (`torch.dtype`, defaults to `torch.bfloat16`): - Data type used for storing and computing Taylor series factors. Lower precision reduces memory but may - affect stability; higher precision improves accuracy at the cost of more memory. - - skip_predict_identifiers (`list[str]`, *optional*, defaults to `None`): - Regex patterns (using `re.fullmatch`) for module names to place as "skip" in "cache" mode. In this mode, - the module computes fully during initial or refresh steps but returns a zero tensor (matching recorded - shape) during prediction steps to skip computation cheaply. - - cache_identifiers (`list[str]`, *optional*, defaults to `None`): - Regex patterns (using `re.fullmatch`) for module names to place in Taylor-series caching mode, where - outputs are approximated and cached for reuse. - - use_lite_mode (`bool`, *optional*, defaults to `False`): - Enables a lightweight TaylorSeer variant that minimizes memory usage by applying predefined patterns for - skipping and caching (e.g., skipping blocks and caching projections). This overrides any custom - `inactive_identifiers` or `active_identifiers`. - - Notes: - - Patterns are matched using `re.fullmatch` on the module name. - - If `skip_predict_identifiers` or `cache_identifiers` are provided, only matching modules are hooked. - - If neither is provided, all attention-like modules are hooked by default. - - Example of inactive and active usage: - - ```py - def forward(x): - x = self.module1(x) # inactive module: returns zeros tensor based on shape recorded during full compute - x = self.module2(x) # active module: caches output here, avoiding recomputation of prior steps - return x - ``` - """ - - cache_interval: int = 5 - disable_cache_before_step: int = 3 - disable_cache_after_step: int | None = None - max_order: int = 1 - taylor_factors_dtype: torch.dtype | None = torch.bfloat16 - skip_predict_identifiers: list[str] | None = None - cache_identifiers: list[str] | None = None - use_lite_mode: bool = False - - def __repr__(self) -> str: - return ( - "TaylorSeerCacheConfig(" - f"cache_interval={self.cache_interval}, " - f"disable_cache_before_step={self.disable_cache_before_step}, " - f"disable_cache_after_step={self.disable_cache_after_step}, " - f"max_order={self.max_order}, " - f"taylor_factors_dtype={self.taylor_factors_dtype}, " - f"skip_predict_identifiers={self.skip_predict_identifiers}, " - f"cache_identifiers={self.cache_identifiers}, " - f"use_lite_mode={self.use_lite_mode})" - ) - - -class TaylorSeerState: - def __init__( - self, - taylor_factors_dtype: torch.dtype | None = torch.bfloat16, - max_order: int = 1, - is_inactive: bool = False, - ): - self.taylor_factors_dtype = taylor_factors_dtype - self.max_order = max_order - self.is_inactive = is_inactive - - self.module_dtypes: tuple[torch.dtype, ...] = () - self.last_update_step: int | None = None - self.taylor_factors: dict[int, dict[int, torch.Tensor]] = {} - self.inactive_shapes: tuple[tuple[int, ...], ...] | None = None - self.device: torch.device | None = None - self.current_step: int = -1 - - def reset(self) -> None: - self.current_step = -1 - self.last_update_step = None - self.taylor_factors = {} - self.inactive_shapes = None - self.device = None - - def update( - self, - outputs: tuple[torch.Tensor, ...], - ) -> None: - self.module_dtypes = tuple(output.dtype for output in outputs) - self.device = outputs[0].device - - if self.is_inactive: - self.inactive_shapes = tuple(output.shape for output in outputs) - else: - for i, features in enumerate(outputs): - new_factors: dict[int, torch.Tensor] = {0: features} - is_first_update = self.last_update_step is None - if not is_first_update: - delta_step = self.current_step - self.last_update_step - if delta_step == 0: - raise ValueError("Delta step cannot be zero for TaylorSeer update.") - - # Recursive divided differences up to max_order - prev_factors = self.taylor_factors.get(i, {}) - for j in range(self.max_order): - prev = prev_factors.get(j) - if prev is None: - break - new_factors[j + 1] = (new_factors[j] - prev.to(features.dtype)) / delta_step - self.taylor_factors[i] = { - order: factor.to(self.taylor_factors_dtype) for order, factor in new_factors.items() - } - - self.last_update_step = self.current_step - - @torch.compiler.disable - def predict(self) -> list[torch.Tensor]: - if self.last_update_step is None: - raise ValueError("Cannot predict without prior initialization/update.") - - step_offset = self.current_step - self.last_update_step - - outputs = [] - if self.is_inactive: - if self.inactive_shapes is None: - raise ValueError("Inactive shapes not set during prediction.") - for i in range(len(self.module_dtypes)): - outputs.append( - torch.zeros( - self.inactive_shapes[i], - dtype=self.module_dtypes[i], - device=self.device, - ) - ) - else: - if not self.taylor_factors: - raise ValueError("Taylor factors empty during prediction.") - num_outputs = len(self.taylor_factors) - num_orders = len(self.taylor_factors[0]) - for i in range(num_outputs): - output_dtype = self.module_dtypes[i] - taylor_factors = self.taylor_factors[i] - output = torch.zeros_like(taylor_factors[0], dtype=output_dtype) - for order in range(num_orders): - coeff = (step_offset**order) / math.factorial(order) - factor = taylor_factors[order] - output = output + factor.to(output_dtype) * coeff - outputs.append(output) - return outputs - - -class TaylorSeerCacheHook(ModelHook): - _is_stateful = True - - def __init__( - self, - cache_interval: int, - disable_cache_before_step: int, - taylor_factors_dtype: torch.dtype, - state_manager: StateManager, - disable_cache_after_step: int | None = None, - ): - super().__init__() - self.cache_interval = cache_interval - self.disable_cache_before_step = disable_cache_before_step - self.disable_cache_after_step = disable_cache_after_step - self.taylor_factors_dtype = taylor_factors_dtype - self.state_manager = state_manager - - def initialize_hook(self, module: torch.nn.Module): - return module - - def reset_state(self, module: torch.nn.Module) -> None: - """ - Reset state between sampling runs. - """ - self.state_manager.reset() - - @torch.compiler.disable - def _measure_should_compute(self) -> bool: - state: TaylorSeerState = self.state_manager.get_state() - state.current_step += 1 - current_step = state.current_step - is_warmup_phase = current_step < self.disable_cache_before_step - is_compute_interval = (current_step - self.disable_cache_before_step - 1) % self.cache_interval == 0 - is_cooldown_phase = self.disable_cache_after_step is not None and current_step >= self.disable_cache_after_step - should_compute = is_warmup_phase or is_compute_interval or is_cooldown_phase - return should_compute, state - - def new_forward(self, module: torch.nn.Module, *args, **kwargs): - should_compute, state = self._measure_should_compute() - if should_compute: - outputs = self.fn_ref.original_forward(*args, **kwargs) - wrapped_outputs = (outputs,) if isinstance(outputs, torch.Tensor) else outputs - state.update(wrapped_outputs) - return outputs - - outputs_list = state.predict() - return outputs_list[0] if len(outputs_list) == 1 else tuple(outputs_list) - - -def _resolve_patterns(config: TaylorSeerCacheConfig) -> tuple[list[str], list[str]]: - """ - Resolve effective inactive and active pattern lists from config + templates. - """ - - inactive_patterns = config.skip_predict_identifiers if config.skip_predict_identifiers is not None else None - active_patterns = config.cache_identifiers if config.cache_identifiers is not None else None - - return inactive_patterns or [], active_patterns or [] - - -def apply_taylorseer_cache(module: torch.nn.Module, config: TaylorSeerCacheConfig): - """ - Applies the TaylorSeer cache to a given pipeline (typically the transformer / UNet). - - This function hooks selected modules in the model to enable caching or skipping based on the provided - configuration, reducing redundant computations in diffusion denoising loops. - - Args: - module (torch.nn.Module): The model subtree to apply the hooks to. - config (TaylorSeerCacheConfig): Configuration for the cache. - - Example: - ```python - >>> import torch - >>> from diffusers import FluxPipeline, TaylorSeerCacheConfig - - >>> pipe = FluxPipeline.from_pretrained( - ... "black-forest-labs/FLUX.1-dev", - ... torch_dtype=torch.bfloat16, - ... ) - >>> pipe.to("cuda") - - >>> config = TaylorSeerCacheConfig( - ... cache_interval=5, - ... max_order=1, - ... disable_cache_before_step=3, - ... taylor_factors_dtype=torch.float32, - ... ) - >>> pipe.transformer.enable_cache(config) - ``` - """ - inactive_patterns, active_patterns = _resolve_patterns(config) - - active_patterns = active_patterns or _TRANSFORMER_BLOCK_IDENTIFIERS - - if config.use_lite_mode: - logger.info("Using TaylorSeer Lite variant for cache.") - active_patterns = _PROJ_OUT_IDENTIFIERS - inactive_patterns = _BLOCK_IDENTIFIERS - if config.skip_predict_identifiers or config.cache_identifiers: - logger.warning("Lite mode overrides user patterns.") - - for name, submodule in module.named_modules(): - matches_inactive = any(re.fullmatch(pattern, name) for pattern in inactive_patterns) - matches_active = any(re.fullmatch(pattern, name) for pattern in active_patterns) - if not (matches_inactive or matches_active): - continue - _apply_taylorseer_cache_hook( - module=submodule, - config=config, - is_inactive=matches_inactive, - ) - - -def _apply_taylorseer_cache_hook( - module: nn.Module, - config: TaylorSeerCacheConfig, - is_inactive: bool, -): - """ - Registers the TaylorSeer hook on the specified nn.Module. - - Args: - name: Name of the module. - module: The nn.Module to be hooked. - config: Cache configuration. - is_inactive: Whether this module should operate in "inactive" mode. - """ - state_manager = StateManager( - TaylorSeerState, - init_kwargs={ - "taylor_factors_dtype": config.taylor_factors_dtype, - "max_order": config.max_order, - "is_inactive": is_inactive, - }, - ) - - registry = HookRegistry.check_if_exists_or_initialize(module) - - hook = TaylorSeerCacheHook( - cache_interval=config.cache_interval, - disable_cache_before_step=config.disable_cache_before_step, - taylor_factors_dtype=config.taylor_factors_dtype, - disable_cache_after_step=config.disable_cache_after_step, - state_manager=state_manager, - ) - - registry.register_hook(hook, _TAYLORSEER_CACHE_HOOK) diff --git a/diffusers/hooks/text_kv_cache.py b/diffusers/hooks/text_kv_cache.py deleted file mode 100644 index b2772eaa3db22915e93459ee1399117e86d347e7..0000000000000000000000000000000000000000 --- a/diffusers/hooks/text_kv_cache.py +++ /dev/null @@ -1,173 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass - -import torch - -from .hooks import BaseState, HookRegistry, ModelHook, StateManager - - -_TEXT_KV_CACHE_TRANSFORMER_HOOK = "text_kv_cache_transformer" -_TEXT_KV_CACHE_BLOCK_HOOK = "text_kv_cache_block" - - -@dataclass -class TextKVCacheConfig: - """Enable exact (lossless) text K/V caching for transformer models. - - Pre-computes per-block text key and value projections once before the denoising loop and reuses them across all - steps. Positive and negative prompts are distinguished via a stable cache key captured by a transformer-level hook - before any intermediate tensor allocations. - """ - - pass - - -class TextKVCacheState(BaseState): - """Shared state between the transformer-level and block-level hooks. - - The transformer hook writes the stable ``encoder_hidden_states`` ``data_ptr()`` (captured *before* ``txt_norm``) so - that block hooks can use it as a reliable cache key across denoising steps. - """ - - def __init__(self): - self.key: int | None = None - - def reset(self): - self.key = None - - -class TextKVCacheBlockState(BaseState): - """Per-block state holding cached text key/value projections.""" - - def __init__(self): - self.kv_cache: dict[int, tuple[torch.Tensor, torch.Tensor]] = {} - - def reset(self): - self.kv_cache.clear() - - -class TextKVCacheTransformerHook(ModelHook): - """Captures ``encoder_hidden_states.data_ptr()`` before ``txt_norm`` - and writes it to shared state for the block hooks to read.""" - - _is_stateful = True - - def __init__(self, state_manager: StateManager): - super().__init__() - self.state_manager = state_manager - - def new_forward(self, module: torch.nn.Module, *args, **kwargs): - if self.state_manager._current_context is None: - self.state_manager.set_context("inference") - - encoder_hidden_states = kwargs.get("encoder_hidden_states") - if encoder_hidden_states is not None: - state: TextKVCacheState = self.state_manager.get_state() - state.key = encoder_hidden_states.data_ptr() - return self.fn_ref.original_forward(*args, **kwargs) - - def reset_state(self, module: torch.nn.Module): - self.state_manager.reset() - return module - - -class TextKVCacheBlockHook(ModelHook): - """Caches ``(txt_key, txt_value)`` per block per unique prompt using - the stable cache key from the shared state.""" - - _is_stateful = True - - def __init__(self, state_manager: StateManager, block_state_manager: StateManager): - super().__init__() - self.state_manager = state_manager - self.block_state_manager = block_state_manager - - def new_forward(self, module: torch.nn.Module, *args, **kwargs): - from ..models.transformers.transformer_nucleusmoe_image import _apply_rotary_emb_nucleus - - if self.state_manager._current_context is None: - self.state_manager.set_context("inference") - - if self.block_state_manager._current_context is None: - self.block_state_manager.set_context("inference") - - if "encoder_hidden_states" in kwargs: - encoder_hidden_states = kwargs["encoder_hidden_states"] - else: - encoder_hidden_states = args[1] - - if "image_rotary_emb" in kwargs: - image_rotary_emb = kwargs["image_rotary_emb"] - elif len(args) > 3: - image_rotary_emb = args[3] - else: - image_rotary_emb = None - - state: TextKVCacheState = self.state_manager.get_state() - cache_key = state.key - - block_state: TextKVCacheBlockState = self.block_state_manager.get_state() - - if cache_key not in block_state.kv_cache: - context = module.encoder_proj(encoder_hidden_states) - - attn = module.attn - head_dim = attn.inner_dim // attn.heads - num_kv_heads = attn.inner_kv_dim // head_dim - - txt_key = attn.add_k_proj(context).unflatten(-1, (num_kv_heads, -1)) - txt_value = attn.add_v_proj(context).unflatten(-1, (num_kv_heads, -1)) - - if attn.norm_added_k is not None: - txt_key = attn.norm_added_k(txt_key) - - if image_rotary_emb is not None: - _, txt_freqs = image_rotary_emb - txt_key = _apply_rotary_emb_nucleus(txt_key, txt_freqs, use_real=False) - - block_state.kv_cache[cache_key] = (txt_key, txt_value) - - txt_key, txt_value = block_state.kv_cache[cache_key] - - attn_kwargs = kwargs.get("attention_kwargs") or {} - attn_kwargs["cached_txt_key"] = txt_key - attn_kwargs["cached_txt_value"] = txt_value - kwargs["attention_kwargs"] = attn_kwargs - - return self.fn_ref.original_forward(*args, **kwargs) - - def reset_state(self, module: torch.nn.Module): - self.block_state_manager.reset() - return module - - -def apply_text_kv_cache(module: torch.nn.Module, config: TextKVCacheConfig) -> None: - from ..models.transformers.transformer_nucleusmoe_image import NucleusMoEImageTransformerBlock - - HookRegistry.check_if_exists_or_initialize(module) - - state_manager = StateManager(TextKVCacheState) - - transformer_hook = TextKVCacheTransformerHook(state_manager) - registry = HookRegistry.check_if_exists_or_initialize(module) - registry.register_hook(transformer_hook, _TEXT_KV_CACHE_TRANSFORMER_HOOK) - - for _, submodule in module.named_modules(): - if isinstance(submodule, NucleusMoEImageTransformerBlock): - block_state_manager = StateManager(TextKVCacheBlockState) - hook = TextKVCacheBlockHook(state_manager, block_state_manager) - block_registry = HookRegistry.check_if_exists_or_initialize(submodule) - block_registry.register_hook(hook, _TEXT_KV_CACHE_BLOCK_HOOK) diff --git a/diffusers/hooks/utils.py b/diffusers/hooks/utils.py deleted file mode 100644 index d3fb97709e736c55946de0faba0eac32ba6b0930..0000000000000000000000000000000000000000 --- a/diffusers/hooks/utils.py +++ /dev/null @@ -1,43 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ._common import _ALL_TRANSFORMER_BLOCK_IDENTIFIERS, _ATTENTION_CLASSES, _FEEDFORWARD_CLASSES - - -def _get_identifiable_transformer_blocks_in_module(module: torch.nn.Module): - module_list_with_transformer_blocks = [] - for name, submodule in module.named_modules(): - name_endswith_identifier = any(name.endswith(identifier) for identifier in _ALL_TRANSFORMER_BLOCK_IDENTIFIERS) - is_ModuleList = isinstance(submodule, torch.nn.ModuleList) - if name_endswith_identifier and is_ModuleList: - module_list_with_transformer_blocks.append((name, submodule)) - return module_list_with_transformer_blocks - - -def _get_identifiable_attention_layers_in_module(module: torch.nn.Module): - attention_layers = [] - for name, submodule in module.named_modules(): - if isinstance(submodule, _ATTENTION_CLASSES): - attention_layers.append((name, submodule)) - return attention_layers - - -def _get_identifiable_feedforward_layers_in_module(module: torch.nn.Module): - feedforward_layers = [] - for name, submodule in module.named_modules(): - if isinstance(submodule, _FEEDFORWARD_CLASSES): - feedforward_layers.append((name, submodule)) - return feedforward_layers diff --git a/diffusers/image_processor.py b/diffusers/image_processor.py deleted file mode 100644 index 4f6f4bd52b9c2c6efd4a35fa50706a8642bc6c75..0000000000000000000000000000000000000000 --- a/diffusers/image_processor.py +++ /dev/null @@ -1,1468 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -import warnings - -import numpy as np -import PIL.Image -import torch -import torch.nn.functional as F -from PIL import Image, ImageFilter, ImageOps - -from .configuration_utils import ConfigMixin, register_to_config -from .utils import CONFIG_NAME, PIL_INTERPOLATION, deprecate - - -PipelineImageInput = ( - PIL.Image.Image | np.ndarray | torch.Tensor | list[PIL.Image.Image] | list[np.ndarray] | list[torch.Tensor] -) - -PipelineDepthInput = PipelineImageInput - - -def is_valid_image(image) -> bool: - r""" - Checks if the input is a valid image. - - A valid image can be: - - A `PIL.Image.Image`. - - A 2D or 3D `np.ndarray` or `torch.Tensor` (grayscale or color image). - - Args: - image (`PIL.Image.Image | np.ndarray | torch.Tensor`): - The image to validate. It can be a PIL image, a NumPy array, or a torch tensor. - - Returns: - `bool`: - `True` if the input is a valid image, `False` otherwise. - """ - return isinstance(image, PIL.Image.Image) or isinstance(image, (np.ndarray, torch.Tensor)) and image.ndim in (2, 3) - - -def is_valid_image_imagelist(images): - r""" - Checks if the input is a valid image or list of images. - - The input can be one of the following formats: - - A 4D tensor or numpy array (batch of images). - - A valid single image: `PIL.Image.Image`, 2D `np.ndarray` or `torch.Tensor` (grayscale image), 3D `np.ndarray` or - `torch.Tensor`. - - A list of valid images. - - Args: - images (`np.ndarray | torch.Tensor | PIL.Image.Image | list`): - The image(s) to check. Can be a batch of images (4D tensor/array), a single image, or a list of valid - images. - - Returns: - `bool`: - `True` if the input is valid, `False` otherwise. - """ - if isinstance(images, (np.ndarray, torch.Tensor)) and images.ndim == 4: - return True - elif is_valid_image(images): - return True - elif isinstance(images, list): - return all(is_valid_image(image) for image in images) - return False - - -class VaeImageProcessor(ConfigMixin): - """ - Image processor for VAE. - - Args: - do_resize (`bool`, *optional*, defaults to `True`): - Whether to downscale the image's (height, width) dimensions to multiples of `vae_scale_factor`. Can accept - `height` and `width` arguments from [`image_processor.VaeImageProcessor.preprocess`] method. - vae_scale_factor (`int`, *optional*, defaults to `8`): - VAE scale factor. If `do_resize` is `True`, the image is automatically resized to multiples of this factor. - resample (`str`, *optional*, defaults to `lanczos`): - Resampling filter to use when resizing the image. - do_normalize (`bool`, *optional*, defaults to `True`): - Whether to normalize the image to [-1,1]. - do_binarize (`bool`, *optional*, defaults to `False`): - Whether to binarize the image to 0/1. - do_convert_rgb (`bool`, *optional*, defaults to be `False`): - Whether to convert the images to RGB format. - do_convert_grayscale (`bool`, *optional*, defaults to be `False`): - Whether to convert the images to grayscale format. - """ - - config_name = CONFIG_NAME - - @register_to_config - def __init__( - self, - do_resize: bool = True, - vae_scale_factor: int = 8, - vae_latent_channels: int = 4, - resample: str = "lanczos", - reducing_gap: int | None = None, - do_normalize: bool = True, - do_binarize: bool = False, - do_convert_rgb: bool = False, - do_convert_grayscale: bool = False, - ): - super().__init__() - if do_convert_rgb and do_convert_grayscale: - raise ValueError( - "`do_convert_rgb` and `do_convert_grayscale` can not both be set to `True`," - " if you intended to convert the image into RGB format, please set `do_convert_grayscale = False`.", - " if you intended to convert the image into grayscale format, please set `do_convert_rgb = False`", - ) - - @staticmethod - def numpy_to_pil(images: np.ndarray) -> list[PIL.Image.Image]: - r""" - Convert a numpy image or a batch of images to a PIL image. - - Args: - images (`np.ndarray`): - The image array to convert to PIL format. - - Returns: - `list[PIL.Image.Image]`: - A list of PIL images. - """ - if images.ndim == 3: - images = images[None, ...] - images = (images * 255).round().astype("uint8") - if images.shape[-1] == 1: - # special case for grayscale (single channel) images - pil_images = [Image.fromarray(image.squeeze(), mode="L") for image in images] - else: - pil_images = [Image.fromarray(image) for image in images] - - return pil_images - - @staticmethod - def pil_to_numpy(images: list[PIL.Image.Image] | PIL.Image.Image) -> np.ndarray: - r""" - Convert a PIL image or a list of PIL images to NumPy arrays. - - Args: - images (`PIL.Image.Image` or `list[PIL.Image.Image]`): - The PIL image or list of images to convert to NumPy format. - - Returns: - `np.ndarray`: - A NumPy array representation of the images. - """ - if not isinstance(images, list): - images = [images] - images = [np.array(image).astype(np.float32) / 255.0 for image in images] - images = np.stack(images, axis=0) - - return images - - @staticmethod - def numpy_to_pt(images: np.ndarray) -> torch.Tensor: - r""" - Convert a NumPy image to a PyTorch tensor. - - Args: - images (`np.ndarray`): - The NumPy image array to convert to PyTorch format. - - Returns: - `torch.Tensor`: - A PyTorch tensor representation of the images. - """ - if images.ndim == 3: - images = images[..., None] - - images = torch.from_numpy(images.transpose(0, 3, 1, 2)) - return images - - @staticmethod - def pt_to_numpy(images: torch.Tensor) -> np.ndarray: - r""" - Convert a PyTorch tensor to a NumPy image. - - Args: - images (`torch.Tensor`): - The PyTorch tensor to convert to NumPy format. - - Returns: - `np.ndarray`: - A NumPy array representation of the images. - """ - images = images.cpu().permute(0, 2, 3, 1).float().numpy() - return images - - @staticmethod - def normalize(images: np.ndarray | torch.Tensor) -> np.ndarray | torch.Tensor: - r""" - Normalize an image array to [-1,1]. - - Args: - images (`np.ndarray` or `torch.Tensor`): - The image array to normalize. - - Returns: - `np.ndarray` or `torch.Tensor`: - The normalized image array. - """ - return 2.0 * images - 1.0 - - @staticmethod - def denormalize(images: np.ndarray | torch.Tensor) -> np.ndarray | torch.Tensor: - r""" - Denormalize an image array to [0,1]. - - Args: - images (`np.ndarray` or `torch.Tensor`): - The image array to denormalize. - - Returns: - `np.ndarray` or `torch.Tensor`: - The denormalized image array. - """ - return (images * 0.5 + 0.5).clamp(0, 1) - - @staticmethod - def convert_to_rgb(image: PIL.Image.Image) -> PIL.Image.Image: - r""" - Converts a PIL image to RGB format. - - Args: - image (`PIL.Image.Image`): - The PIL image to convert to RGB. - - Returns: - `PIL.Image.Image`: - The RGB-converted PIL image. - """ - image = image.convert("RGB") - - return image - - @staticmethod - def convert_to_grayscale(image: PIL.Image.Image) -> PIL.Image.Image: - r""" - Converts a given PIL image to grayscale. - - Args: - image (`PIL.Image.Image`): - The input image to convert. - - Returns: - `PIL.Image.Image`: - The image converted to grayscale. - """ - image = image.convert("L") - - return image - - @staticmethod - def blur(image: PIL.Image.Image, blur_factor: int = 4) -> PIL.Image.Image: - r""" - Applies Gaussian blur to an image. - - Args: - image (`PIL.Image.Image`): - The PIL image to convert to grayscale. - - Returns: - `PIL.Image.Image`: - The grayscale-converted PIL image. - """ - image = image.filter(ImageFilter.GaussianBlur(blur_factor)) - - return image - - @staticmethod - def get_crop_region(mask_image: PIL.Image.Image, width: int, height: int, pad=0): - r""" - Finds a rectangular region that contains all masked ares in an image, and expands region to match the aspect - ratio of the original image; for example, if user drew mask in a 128x32 region, and the dimensions for - processing are 512x512, the region will be expanded to 128x128. - - Args: - mask_image (PIL.Image.Image): Mask image. - width (int): Width of the image to be processed. - height (int): Height of the image to be processed. - pad (int, optional): Padding to be added to the crop region. Defaults to 0. - - Returns: - tuple: (x1, y1, x2, y2) represent a rectangular region that contains all masked ares in an image and - matches the original aspect ratio. - """ - - mask_image = mask_image.convert("L") - mask = np.array(mask_image) - - # 1. find a rectangular region that contains all masked ares in an image - h, w = mask.shape - crop_left = 0 - for i in range(w): - if not (mask[:, i] == 0).all(): - break - crop_left += 1 - - crop_right = 0 - for i in reversed(range(w)): - if not (mask[:, i] == 0).all(): - break - crop_right += 1 - - crop_top = 0 - for i in range(h): - if not (mask[i] == 0).all(): - break - crop_top += 1 - - crop_bottom = 0 - for i in reversed(range(h)): - if not (mask[i] == 0).all(): - break - crop_bottom += 1 - - # 2. add padding to the crop region - x1, y1, x2, y2 = ( - int(max(crop_left - pad, 0)), - int(max(crop_top - pad, 0)), - int(min(w - crop_right + pad, w)), - int(min(h - crop_bottom + pad, h)), - ) - - # 3. expands crop region to match the aspect ratio of the image to be processed - ratio_crop_region = (x2 - x1) / (y2 - y1) - ratio_processing = width / height - - if ratio_crop_region > ratio_processing: - desired_height = (x2 - x1) / ratio_processing - desired_height_diff = int(desired_height - (y2 - y1)) - y1 -= desired_height_diff // 2 - y2 += desired_height_diff - desired_height_diff // 2 - if y2 >= mask_image.height: - diff = y2 - mask_image.height - y2 -= diff - y1 -= diff - if y1 < 0: - y2 -= y1 - y1 -= y1 - if y2 >= mask_image.height: - y2 = mask_image.height - else: - desired_width = (y2 - y1) * ratio_processing - desired_width_diff = int(desired_width - (x2 - x1)) - x1 -= desired_width_diff // 2 - x2 += desired_width_diff - desired_width_diff // 2 - if x2 >= mask_image.width: - diff = x2 - mask_image.width - x2 -= diff - x1 -= diff - if x1 < 0: - x2 -= x1 - x1 -= x1 - if x2 >= mask_image.width: - x2 = mask_image.width - - return x1, y1, x2, y2 - - def _resize_and_fill( - self, - image: PIL.Image.Image, - width: int, - height: int, - ) -> PIL.Image.Image: - r""" - Resize the image to fit within the specified width and height, maintaining the aspect ratio, and then center - the image within the dimensions, filling empty with data from image. - - Args: - image (`PIL.Image.Image`): - The image to resize and fill. - width (`int`): - The width to resize the image to. - height (`int`): - The height to resize the image to. - - Returns: - `PIL.Image.Image`: - The resized and filled image. - """ - - ratio = width / height - src_ratio = image.width / image.height - - src_w = width if ratio < src_ratio else image.width * height // image.height - src_h = height if ratio >= src_ratio else image.height * width // image.width - - resized = image.resize((src_w, src_h), resample=PIL_INTERPOLATION[self.config.resample]) - res = Image.new("RGB", (width, height)) - res.paste(resized, box=(width // 2 - src_w // 2, height // 2 - src_h // 2)) - - if ratio < src_ratio: - fill_height = height // 2 - src_h // 2 - if fill_height > 0: - res.paste(resized.resize((width, fill_height), box=(0, 0, width, 0)), box=(0, 0)) - res.paste( - resized.resize((width, fill_height), box=(0, resized.height, width, resized.height)), - box=(0, fill_height + src_h), - ) - elif ratio > src_ratio: - fill_width = width // 2 - src_w // 2 - if fill_width > 0: - res.paste(resized.resize((fill_width, height), box=(0, 0, 0, height)), box=(0, 0)) - res.paste( - resized.resize((fill_width, height), box=(resized.width, 0, resized.width, height)), - box=(fill_width + src_w, 0), - ) - - return res - - def _resize_and_crop( - self, - image: PIL.Image.Image, - width: int, - height: int, - ) -> PIL.Image.Image: - r""" - Resize the image to fit within the specified width and height, maintaining the aspect ratio, and then center - the image within the dimensions, cropping the excess. - - Args: - image (`PIL.Image.Image`): - The image to resize and crop. - width (`int`): - The width to resize the image to. - height (`int`): - The height to resize the image to. - - Returns: - `PIL.Image.Image`: - The resized and cropped image. - """ - ratio = width / height - src_ratio = image.width / image.height - - src_w = width if ratio > src_ratio else image.width * height // image.height - src_h = height if ratio <= src_ratio else image.height * width // image.width - - resized = image.resize((src_w, src_h), resample=PIL_INTERPOLATION[self.config.resample]) - res = Image.new("RGB", (width, height)) - res.paste(resized, box=(width // 2 - src_w // 2, height // 2 - src_h // 2)) - return res - - def resize( - self, - image: PIL.Image.Image | np.ndarray | torch.Tensor, - height: int, - width: int, - resize_mode: str = "default", # "default", "fill", "crop" - ) -> PIL.Image.Image | np.ndarray | torch.Tensor: - """ - Resize image. - - Args: - image (`PIL.Image.Image`, `np.ndarray` or `torch.Tensor`): - The image input, can be a PIL image, numpy array or pytorch tensor. - height (`int`): - The height to resize to. - width (`int`): - The width to resize to. - resize_mode (`str`, *optional*, defaults to `default`): - The resize mode to use, can be one of `default` or `fill`. If `default`, will resize the image to fit - within the specified width and height, and it may not maintaining the original aspect ratio. If `fill`, - will resize the image to fit within the specified width and height, maintaining the aspect ratio, and - then center the image within the dimensions, filling empty with data from image. If `crop`, will resize - the image to fit within the specified width and height, maintaining the aspect ratio, and then center - the image within the dimensions, cropping the excess. Note that resize_mode `fill` and `crop` are only - supported for PIL image input. - - Returns: - `PIL.Image.Image`, `np.ndarray` or `torch.Tensor`: - The resized image. - """ - if resize_mode != "default" and not isinstance(image, PIL.Image.Image): - raise ValueError(f"Only PIL image input is supported for resize_mode {resize_mode}") - if isinstance(image, PIL.Image.Image): - if resize_mode == "default": - image = image.resize( - (width, height), - resample=PIL_INTERPOLATION[self.config.resample], - reducing_gap=self.config.reducing_gap, - ) - elif resize_mode == "fill": - image = self._resize_and_fill(image, width, height) - elif resize_mode == "crop": - image = self._resize_and_crop(image, width, height) - else: - raise ValueError(f"resize_mode {resize_mode} is not supported") - - elif isinstance(image, torch.Tensor): - image = torch.nn.functional.interpolate( - image, - size=(height, width), - ) - elif isinstance(image, np.ndarray): - image = self.numpy_to_pt(image) - image = torch.nn.functional.interpolate( - image, - size=(height, width), - ) - image = self.pt_to_numpy(image) - - return image - - def binarize(self, image: PIL.Image.Image) -> PIL.Image.Image: - """ - Create a mask. - - Args: - image (`PIL.Image.Image`): - The image input, should be a PIL image. - - Returns: - `PIL.Image.Image`: - The binarized image. Values less than 0.5 are set to 0, values greater than 0.5 are set to 1. - """ - image[image < 0.5] = 0 - image[image >= 0.5] = 1 - - return image - - def _denormalize_conditionally( - self, images: torch.Tensor, do_denormalize: list[bool] | None = None - ) -> torch.Tensor: - r""" - Denormalize a batch of images based on a condition list. - - Args: - images (`torch.Tensor`): - The input image tensor. - do_denormalize (`Optional[list[bool]`, *optional*, defaults to `None`): - A list of booleans indicating whether to denormalize each image in the batch. If `None`, will use the - value of `do_normalize` in the `VaeImageProcessor` config. - """ - if do_denormalize is None: - return self.denormalize(images) if self.config.do_normalize else images - - return torch.stack( - [self.denormalize(images[i]) if do_denormalize[i] else images[i] for i in range(images.shape[0])] - ) - - def get_default_height_width( - self, - image: PIL.Image.Image | np.ndarray | torch.Tensor, - height: int | None = None, - width: int | None = None, - ) -> tuple[int, int]: - r""" - Returns the height and width of the image, downscaled to the next integer multiple of `vae_scale_factor`. - - Args: - image (`PIL.Image.Image | np.ndarray | torch.Tensor`): - The image input, which can be a PIL image, NumPy array, or PyTorch tensor. If it is a NumPy array, it - should have shape `[batch, height, width]` or `[batch, height, width, channels]`. If it is a PyTorch - tensor, it should have shape `[batch, channels, height, width]`. - height (`int | None`, *optional*, defaults to `None`): - The height of the preprocessed image. If `None`, the height of the `image` input will be used. - width (`int | None`, *optional*, defaults to `None`): - The width of the preprocessed image. If `None`, the width of the `image` input will be used. - - Returns: - `tuple[int, int]`: - A tuple containing the height and width, both resized to the nearest integer multiple of - `vae_scale_factor`. - """ - - if height is None: - if isinstance(image, PIL.Image.Image): - height = image.height - elif isinstance(image, torch.Tensor): - height = image.shape[2] - else: - height = image.shape[1] - - if width is None: - if isinstance(image, PIL.Image.Image): - width = image.width - elif isinstance(image, torch.Tensor): - width = image.shape[3] - else: - width = image.shape[2] - - width, height = ( - x - x % self.config.vae_scale_factor for x in (width, height) - ) # resize to integer multiple of vae_scale_factor - - return height, width - - def preprocess( - self, - image: PipelineImageInput, - height: int | None = None, - width: int | None = None, - resize_mode: str = "default", # "default", "fill", "crop" - crops_coords: tuple[int, int, int, int] | None = None, - ) -> torch.Tensor: - """ - Preprocess the image input. - - Args: - image (`PipelineImageInput`): - The image input, accepted formats are PIL images, NumPy arrays, PyTorch tensors; Also accept list of - supported formats. - height (`int`, *optional*): - The height in preprocessed image. If `None`, will use the `get_default_height_width()` to get default - height. - width (`int`, *optional*): - The width in preprocessed. If `None`, will use get_default_height_width()` to get the default width. - resize_mode (`str`, *optional*, defaults to `default`): - The resize mode, can be one of `default` or `fill`. If `default`, will resize the image to fit within - the specified width and height, and it may not maintaining the original aspect ratio. If `fill`, will - resize the image to fit within the specified width and height, maintaining the aspect ratio, and then - center the image within the dimensions, filling empty with data from image. If `crop`, will resize the - image to fit within the specified width and height, maintaining the aspect ratio, and then center the - image within the dimensions, cropping the excess. Note that resize_mode `fill` and `crop` are only - supported for PIL image input. - crops_coords (`list[tuple[int, int, int, int]]`, *optional*, defaults to `None`): - The crop coordinates for each image in the batch. If `None`, will not crop the image. - - Returns: - `torch.Tensor`: - The preprocessed image. - """ - supported_formats = (PIL.Image.Image, np.ndarray, torch.Tensor) - - # Expand the missing dimension for 3-dimensional pytorch tensor or numpy array that represents grayscale image - if self.config.do_convert_grayscale and isinstance(image, (torch.Tensor, np.ndarray)) and image.ndim == 3: - if isinstance(image, torch.Tensor): - # if image is a pytorch tensor could have 2 possible shapes: - # 1. batch x height x width: we should insert the channel dimension at position 1 - # 2. channel x height x width: we should insert batch dimension at position 0, - # however, since both channel and batch dimension has same size 1, it is same to insert at position 1 - # for simplicity, we insert a dimension of size 1 at position 1 for both cases - image = image.unsqueeze(1) - else: - # if it is a numpy array, it could have 2 possible shapes: - # 1. batch x height x width: insert channel dimension on last position - # 2. height x width x channel: insert batch dimension on first position - if image.shape[-1] == 1: - image = np.expand_dims(image, axis=0) - else: - image = np.expand_dims(image, axis=-1) - - if isinstance(image, list) and isinstance(image[0], np.ndarray) and image[0].ndim == 4: - warnings.warn( - "Passing `image` as a list of 4d np.ndarray is deprecated." - "Please concatenate the list along the batch dimension and pass it as a single 4d np.ndarray", - FutureWarning, - ) - image = np.concatenate(image, axis=0) - if isinstance(image, list) and isinstance(image[0], torch.Tensor) and image[0].ndim == 4: - warnings.warn( - "Passing `image` as a list of 4d torch.Tensor is deprecated." - "Please concatenate the list along the batch dimension and pass it as a single 4d torch.Tensor", - FutureWarning, - ) - image = torch.cat(image, axis=0) - - if not is_valid_image_imagelist(image): - raise ValueError( - f"Input is in incorrect format. Currently, we only support {', '.join(str(x) for x in supported_formats)}" - ) - if not isinstance(image, list): - image = [image] - - if isinstance(image[0], PIL.Image.Image): - if crops_coords is not None: - image = [i.crop(crops_coords) for i in image] - if self.config.do_resize: - height, width = self.get_default_height_width(image[0], height, width) - image = [self.resize(i, height, width, resize_mode=resize_mode) for i in image] - if self.config.do_convert_rgb: - image = [self.convert_to_rgb(i) for i in image] - elif self.config.do_convert_grayscale: - image = [self.convert_to_grayscale(i) for i in image] - image = self.pil_to_numpy(image) # to np - image = self.numpy_to_pt(image) # to pt - - elif isinstance(image[0], np.ndarray): - image = np.concatenate(image, axis=0) if image[0].ndim == 4 else np.stack(image, axis=0) - - image = self.numpy_to_pt(image) - - height, width = self.get_default_height_width(image, height, width) - if self.config.do_resize: - image = self.resize(image, height, width) - - elif isinstance(image[0], torch.Tensor): - image = torch.cat(image, axis=0) if image[0].ndim == 4 else torch.stack(image, axis=0) - - if self.config.do_convert_grayscale and image.ndim == 3: - image = image.unsqueeze(1) - - channel = image.shape[1] - # don't need any preprocess if the image is latents - if channel == self.config.vae_latent_channels: - return image - - height, width = self.get_default_height_width(image, height, width) - if self.config.do_resize: - image = self.resize(image, height, width) - - # expected range [0,1], normalize to [-1,1] - do_normalize = self.config.do_normalize - if do_normalize and image.min() < 0: - warnings.warn( - "Passing `image` as torch tensor with value range in [-1,1] is deprecated. The expected value range for image tensor is [0,1] " - f"when passing as pytorch tensor or numpy Array. You passed `image` with value range [{image.min()},{image.max()}]", - FutureWarning, - ) - do_normalize = False - if do_normalize: - image = self.normalize(image) - - if self.config.do_binarize: - image = self.binarize(image) - - return image - - def postprocess( - self, - image: torch.Tensor, - output_type: str = "pil", - do_denormalize: list[bool] | None = None, - ) -> PIL.Image.Image | np.ndarray | torch.Tensor: - """ - Postprocess the image output from tensor to `output_type`. - - Args: - image (`torch.Tensor`): - The image input, should be a pytorch tensor with shape `B x C x H x W`. - output_type (`str`, *optional*, defaults to `pil`): - The output type of the image, can be one of `pil`, `np`, `pt`, `latent`. - do_denormalize (`list[bool]`, *optional*, defaults to `None`): - Whether to denormalize the image to [0,1]. If `None`, will use the value of `do_normalize` in the - `VaeImageProcessor` config. - - Returns: - `PIL.Image.Image`, `np.ndarray` or `torch.Tensor`: - The postprocessed image. - """ - if not isinstance(image, torch.Tensor): - raise ValueError( - f"Input for postprocessing is in incorrect format: {type(image)}. We only support pytorch tensor" - ) - if output_type not in ["latent", "pt", "np", "pil"]: - deprecation_message = ( - 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: " - "`pil`, `np`, `pt`, `latent`" - ) - deprecate("Unsupported output_type", "1.0.0", deprecation_message, standard_warn=False) - output_type = "np" - - if output_type == "latent": - return image - - image = self._denormalize_conditionally(image, do_denormalize) - - if output_type == "pt": - return image - - image = self.pt_to_numpy(image) - - if output_type == "np": - return image - - if output_type == "pil": - return self.numpy_to_pil(image) - - def apply_overlay( - self, - mask: PIL.Image.Image, - init_image: PIL.Image.Image, - image: PIL.Image.Image, - crop_coords: tuple[int, int, int, int] | None = None, - ) -> PIL.Image.Image: - r""" - Applies an overlay of the mask and the inpainted image on the original image. - - Args: - mask (`PIL.Image.Image`): - The mask image that highlights regions to overlay. - init_image (`PIL.Image.Image`): - The original image to which the overlay is applied. - image (`PIL.Image.Image`): - The image to overlay onto the original. - crop_coords (`tuple[int, int, int, int]`, *optional*): - Coordinates to crop the image. If provided, the image will be cropped accordingly. - - Returns: - `PIL.Image.Image`: - The final image with the overlay applied. - """ - - width, height = init_image.width, init_image.height - - init_image_masked = PIL.Image.new("RGBa", (width, height)) - init_image_masked.paste(init_image.convert("RGBA").convert("RGBa"), mask=ImageOps.invert(mask.convert("L"))) - - init_image_masked = init_image_masked.convert("RGBA") - - if crop_coords is not None: - x, y, x2, y2 = crop_coords - w = x2 - x - h = y2 - y - base_image = PIL.Image.new("RGBA", (width, height)) - image = self.resize(image, height=h, width=w, resize_mode="crop") - base_image.paste(image, (x, y)) - image = base_image.convert("RGB") - - image = image.convert("RGBA") - image.alpha_composite(init_image_masked) - image = image.convert("RGB") - - return image - - -class InpaintProcessor(ConfigMixin): - """ - Image processor for inpainting image and mask. - """ - - config_name = CONFIG_NAME - - @register_to_config - def __init__( - self, - do_resize: bool = True, - vae_scale_factor: int = 8, - vae_latent_channels: int = 4, - resample: str = "lanczos", - reducing_gap: int | None = None, - do_normalize: bool = True, - do_binarize: bool = False, - do_convert_grayscale: bool = False, - mask_do_normalize: bool = False, - mask_do_binarize: bool = True, - mask_do_convert_grayscale: bool = True, - ): - super().__init__() - - self._image_processor = VaeImageProcessor( - do_resize=do_resize, - vae_scale_factor=vae_scale_factor, - vae_latent_channels=vae_latent_channels, - resample=resample, - reducing_gap=reducing_gap, - do_normalize=do_normalize, - do_binarize=do_binarize, - do_convert_grayscale=do_convert_grayscale, - ) - self._mask_processor = VaeImageProcessor( - do_resize=do_resize, - vae_scale_factor=vae_scale_factor, - vae_latent_channels=vae_latent_channels, - resample=resample, - reducing_gap=reducing_gap, - do_normalize=mask_do_normalize, - do_binarize=mask_do_binarize, - do_convert_grayscale=mask_do_convert_grayscale, - ) - - def preprocess( - self, - image: PIL.Image.Image, - mask: PIL.Image.Image | None = None, - height: int | None = None, - width: int | None = None, - padding_mask_crop: int | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - """ - Preprocess the image and mask. - """ - if mask is None and padding_mask_crop is not None: - raise ValueError("mask must be provided if padding_mask_crop is provided") - - # if mask is None, same behavior as regular image processor - if mask is None: - return self._image_processor.preprocess(image, height=height, width=width) - - if padding_mask_crop is not None: - crops_coords = self._image_processor.get_crop_region(mask, width, height, pad=padding_mask_crop) - resize_mode = "fill" - else: - crops_coords = None - resize_mode = "default" - - processed_image = self._image_processor.preprocess( - image, - height=height, - width=width, - crops_coords=crops_coords, - resize_mode=resize_mode, - ) - - processed_mask = self._mask_processor.preprocess( - mask, - height=height, - width=width, - resize_mode=resize_mode, - crops_coords=crops_coords, - ) - - if crops_coords is not None: - postprocessing_kwargs = { - "crops_coords": crops_coords, - "original_image": image, - "original_mask": mask, - } - else: - postprocessing_kwargs = { - "crops_coords": None, - "original_image": None, - "original_mask": None, - } - - return processed_image, processed_mask, postprocessing_kwargs - - def postprocess( - self, - image: torch.Tensor, - output_type: str = "pil", - original_image: PIL.Image.Image | None = None, - original_mask: PIL.Image.Image | None = None, - crops_coords: tuple[int, int, int, int] | None = None, - ) -> tuple[PIL.Image.Image, PIL.Image.Image]: - """ - Postprocess the image, optionally apply mask overlay - """ - image = self._image_processor.postprocess( - image, - output_type=output_type, - ) - # optionally apply the mask overlay - if crops_coords is not None and (original_image is None or original_mask is None): - raise ValueError("original_image and original_mask must be provided if crops_coords is provided") - - elif crops_coords is not None and output_type != "pil": - raise ValueError("output_type must be 'pil' if crops_coords is provided") - - elif crops_coords is not None: - image = [ - self._image_processor.apply_overlay(original_mask, original_image, i, crops_coords) for i in image - ] - - return image - - -class VaeImageProcessorLDM3D(VaeImageProcessor): - """ - Image processor for VAE LDM3D. - - Args: - do_resize (`bool`, *optional*, defaults to `True`): - Whether to downscale the image's (height, width) dimensions to multiples of `vae_scale_factor`. - vae_scale_factor (`int`, *optional*, defaults to `8`): - VAE scale factor. If `do_resize` is `True`, the image is automatically resized to multiples of this factor. - resample (`str`, *optional*, defaults to `lanczos`): - Resampling filter to use when resizing the image. - do_normalize (`bool`, *optional*, defaults to `True`): - Whether to normalize the image to [-1,1]. - """ - - config_name = CONFIG_NAME - - @register_to_config - def __init__( - self, - do_resize: bool = True, - vae_scale_factor: int = 8, - resample: str = "lanczos", - do_normalize: bool = True, - ): - super().__init__() - - @staticmethod - def numpy_to_pil(images: np.ndarray) -> list[PIL.Image.Image]: - r""" - Convert a NumPy image or a batch of images to a list of PIL images. - - Args: - images (`np.ndarray`): - The input NumPy array of images, which can be a single image or a batch. - - Returns: - `list[PIL.Image.Image]`: - A list of PIL images converted from the input NumPy array. - """ - if images.ndim == 3: - images = images[None, ...] - images = (images * 255).round().astype("uint8") - if images.shape[-1] == 1: - # special case for grayscale (single channel) images - pil_images = [Image.fromarray(image.squeeze(), mode="L") for image in images] - else: - pil_images = [Image.fromarray(image[:, :, :3]) for image in images] - - return pil_images - - @staticmethod - def depth_pil_to_numpy(images: list[PIL.Image.Image] | PIL.Image.Image) -> np.ndarray: - r""" - Convert a PIL image or a list of PIL images to NumPy arrays. - - Args: - images (`list[PIL.Image.Image, PIL.Image.Image]`): - The input image or list of images to be converted. - - Returns: - `np.ndarray`: - A NumPy array of the converted images. - """ - if not isinstance(images, list): - images = [images] - - images = [np.array(image).astype(np.float32) / (2**16 - 1) for image in images] - images = np.stack(images, axis=0) - return images - - @staticmethod - def rgblike_to_depthmap(image: np.ndarray | torch.Tensor) -> np.ndarray | torch.Tensor: - r""" - Convert an RGB-like depth image to a depth map. - """ - # 1. Cast the tensor to a larger integer type (e.g., int32) - # to safely perform the multiplication by 256. - # 2. Perform the 16-bit combination: High-byte * 256 + Low-byte. - # 3. Cast the final result to the desired depth map type (uint16) if needed - # before returning, though leaving it as int32/int64 is often safer - # for return value from a library function. - - if isinstance(image, torch.Tensor): - # Cast to a safe dtype (e.g., int32 or int64) for the calculation - original_dtype = image.dtype - image_safe = image.to(torch.int32) - - # Calculate the depth map - depth_map = image_safe[:, :, 1] * 256 + image_safe[:, :, 2] - - # You may want to cast the final result to uint16, but casting to a - # larger int type (like int32) is sufficient to fix the overflow. - # depth_map = depth_map.to(torch.uint16) # Uncomment if uint16 is strictly required - return depth_map.to(original_dtype) - - elif isinstance(image, np.ndarray): - # NumPy equivalent: Cast to a safe dtype (e.g., np.int32) - original_dtype = image.dtype - image_safe = image.astype(np.int32) - - # Calculate the depth map - depth_map = image_safe[:, :, 1] * 256 + image_safe[:, :, 2] - - # depth_map = depth_map.astype(np.uint16) # Uncomment if uint16 is strictly required - return depth_map.astype(original_dtype) - else: - raise TypeError("Input image must be a torch.Tensor or np.ndarray") - - def numpy_to_depth(self, images: np.ndarray) -> list[PIL.Image.Image]: - r""" - Convert a NumPy depth image or a batch of images to a list of PIL images. - - Args: - images (`np.ndarray`): - The input NumPy array of depth images, which can be a single image or a batch. - - Returns: - `list[PIL.Image.Image]`: - A list of PIL images converted from the input NumPy depth images. - """ - if images.ndim == 3: - images = images[None, ...] - images_depth = images[:, :, :, 3:] - if images.shape[-1] == 6: - images_depth = (images_depth * 255).round().astype("uint8") - pil_images = [ - Image.fromarray(self.rgblike_to_depthmap(image_depth), mode="I;16") for image_depth in images_depth - ] - elif images.shape[-1] == 4: - images_depth = (images_depth * 65535.0).astype(np.uint16) - pil_images = [Image.fromarray(image_depth, mode="I;16") for image_depth in images_depth] - else: - raise Exception("Not supported") - - return pil_images - - def postprocess( - self, - image: torch.Tensor, - output_type: str = "pil", - do_denormalize: list[bool] | None = None, - ) -> PIL.Image.Image | np.ndarray | torch.Tensor: - """ - Postprocess the image output from tensor to `output_type`. - - Args: - image (`torch.Tensor`): - The image input, should be a pytorch tensor with shape `B x C x H x W`. - output_type (`str`, *optional*, defaults to `pil`): - The output type of the image, can be one of `pil`, `np`, `pt`, `latent`. - do_denormalize (`list[bool]`, *optional*, defaults to `None`): - Whether to denormalize the image to [0,1]. If `None`, will use the value of `do_normalize` in the - `VaeImageProcessor` config. - - Returns: - `PIL.Image.Image`, `np.ndarray` or `torch.Tensor`: - The postprocessed image. - """ - if not isinstance(image, torch.Tensor): - raise ValueError( - f"Input for postprocessing is in incorrect format: {type(image)}. We only support pytorch tensor" - ) - if output_type not in ["latent", "pt", "np", "pil"]: - deprecation_message = ( - 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: " - "`pil`, `np`, `pt`, `latent`" - ) - deprecate("Unsupported output_type", "1.0.0", deprecation_message, standard_warn=False) - output_type = "np" - - image = self._denormalize_conditionally(image, do_denormalize) - - image = self.pt_to_numpy(image) - - if output_type == "np": - if image.shape[-1] == 6: - image_depth = np.stack([self.rgblike_to_depthmap(im[:, :, 3:]) for im in image], axis=0) - else: - image_depth = image[:, :, :, 3:] - return image[:, :, :, :3], image_depth - - if output_type == "pil": - return self.numpy_to_pil(image), self.numpy_to_depth(image) - else: - raise Exception(f"This type {output_type} is not supported") - - def preprocess( - self, - rgb: torch.Tensor | PIL.Image.Image | np.ndarray, - depth: torch.Tensor | PIL.Image.Image | np.ndarray, - height: int | None = None, - width: int | None = None, - target_res: int | None = None, - ) -> torch.Tensor: - r""" - Preprocess the image input. Accepted formats are PIL images, NumPy arrays, or PyTorch tensors. - - Args: - rgb (`torch.Tensor | PIL.Image.Image | np.ndarray`): - The RGB input image, which can be a single image or a batch. - depth (`torch.Tensor | PIL.Image.Image | np.ndarray`): - The depth input image, which can be a single image or a batch. - height (`int | None`, *optional*, defaults to `None`): - The desired height of the processed image. If `None`, defaults to the height of the input image. - width (`int | None`, *optional*, defaults to `None`): - The desired width of the processed image. If `None`, defaults to the width of the input image. - target_res (`int | None`, *optional*, defaults to `None`): - Target resolution for resizing the images. If specified, overrides height and width. - - Returns: - `tuple[torch.Tensor, torch.Tensor]`: - A tuple containing the processed RGB and depth images as PyTorch tensors. - """ - supported_formats = (PIL.Image.Image, np.ndarray, torch.Tensor) - - # Expand the missing dimension for 3-dimensional pytorch tensor or numpy array that represents grayscale image - if self.config.do_convert_grayscale and isinstance(rgb, (torch.Tensor, np.ndarray)) and rgb.ndim == 3: - raise Exception("This is not yet supported") - - if isinstance(rgb, supported_formats): - rgb = [rgb] - depth = [depth] - elif not (isinstance(rgb, list) and all(isinstance(i, supported_formats) for i in rgb)): - raise ValueError( - f"Input is in incorrect format: {[type(i) for i in rgb]}. Currently, we only support {', '.join(supported_formats)}" - ) - - if isinstance(rgb[0], PIL.Image.Image): - if self.config.do_convert_rgb: - raise Exception("This is not yet supported") - # rgb = [self.convert_to_rgb(i) for i in rgb] - # depth = [self.convert_to_depth(i) for i in depth] #TODO define convert_to_depth - if self.config.do_resize or target_res: - height, width = self.get_default_height_width(rgb[0], height, width) if not target_res else target_res - rgb = [self.resize(i, height, width) for i in rgb] - depth = [self.resize(i, height, width) for i in depth] - rgb = self.pil_to_numpy(rgb) # to np - rgb = self.numpy_to_pt(rgb) # to pt - - depth = self.depth_pil_to_numpy(depth) # to np - depth = self.numpy_to_pt(depth) # to pt - - elif isinstance(rgb[0], np.ndarray): - rgb = np.concatenate(rgb, axis=0) if rgb[0].ndim == 4 else np.stack(rgb, axis=0) - rgb = self.numpy_to_pt(rgb) - height, width = self.get_default_height_width(rgb, height, width) - if self.config.do_resize: - rgb = self.resize(rgb, height, width) - - depth = np.concatenate(depth, axis=0) if rgb[0].ndim == 4 else np.stack(depth, axis=0) - depth = self.numpy_to_pt(depth) - height, width = self.get_default_height_width(depth, height, width) - if self.config.do_resize: - depth = self.resize(depth, height, width) - - elif isinstance(rgb[0], torch.Tensor): - raise Exception("This is not yet supported") - # rgb = torch.cat(rgb, axis=0) if rgb[0].ndim == 4 else torch.stack(rgb, axis=0) - - # if self.config.do_convert_grayscale and rgb.ndim == 3: - # rgb = rgb.unsqueeze(1) - - # channel = rgb.shape[1] - - # height, width = self.get_default_height_width(rgb, height, width) - # if self.config.do_resize: - # rgb = self.resize(rgb, height, width) - - # depth = torch.cat(depth, axis=0) if depth[0].ndim == 4 else torch.stack(depth, axis=0) - - # if self.config.do_convert_grayscale and depth.ndim == 3: - # depth = depth.unsqueeze(1) - - # channel = depth.shape[1] - # # don't need any preprocess if the image is latents - # if depth == 4: - # return rgb, depth - - # height, width = self.get_default_height_width(depth, height, width) - # if self.config.do_resize: - # depth = self.resize(depth, height, width) - # expected range [0,1], normalize to [-1,1] - do_normalize = self.config.do_normalize - if rgb.min() < 0 and do_normalize: - warnings.warn( - "Passing `image` as torch tensor with value range in [-1,1] is deprecated. The expected value range for image tensor is [0,1] " - f"when passing as pytorch tensor or numpy Array. You passed `image` with value range [{rgb.min()},{rgb.max()}]", - FutureWarning, - ) - do_normalize = False - - if do_normalize: - rgb = self.normalize(rgb) - depth = self.normalize(depth) - - if self.config.do_binarize: - rgb = self.binarize(rgb) - depth = self.binarize(depth) - - return rgb, depth - - -class IPAdapterMaskProcessor(VaeImageProcessor): - """ - Image processor for IP Adapter image masks. - - Args: - do_resize (`bool`, *optional*, defaults to `True`): - Whether to downscale the image's (height, width) dimensions to multiples of `vae_scale_factor`. - vae_scale_factor (`int`, *optional*, defaults to `8`): - VAE scale factor. If `do_resize` is `True`, the image is automatically resized to multiples of this factor. - resample (`str`, *optional*, defaults to `lanczos`): - Resampling filter to use when resizing the image. - do_normalize (`bool`, *optional*, defaults to `False`): - Whether to normalize the image to [-1,1]. - do_binarize (`bool`, *optional*, defaults to `True`): - Whether to binarize the image to 0/1. - do_convert_grayscale (`bool`, *optional*, defaults to be `True`): - Whether to convert the images to grayscale format. - - """ - - config_name = CONFIG_NAME - - @register_to_config - def __init__( - self, - do_resize: bool = True, - vae_scale_factor: int = 8, - resample: str = "lanczos", - do_normalize: bool = False, - do_binarize: bool = True, - do_convert_grayscale: bool = True, - ): - super().__init__( - do_resize=do_resize, - vae_scale_factor=vae_scale_factor, - resample=resample, - do_normalize=do_normalize, - do_binarize=do_binarize, - do_convert_grayscale=do_convert_grayscale, - ) - - @staticmethod - def downsample(mask: torch.Tensor, batch_size: int, num_queries: int, value_embed_dim: int): - """ - Downsamples the provided mask tensor to match the expected dimensions for scaled dot-product attention. If the - aspect ratio of the mask does not match the aspect ratio of the output image, a warning is issued. - - Args: - mask (`torch.Tensor`): - The input mask tensor generated with `IPAdapterMaskProcessor.preprocess()`. - batch_size (`int`): - The batch size. - num_queries (`int`): - The number of queries. - value_embed_dim (`int`): - The dimensionality of the value embeddings. - - Returns: - `torch.Tensor`: - The downsampled mask tensor. - - """ - o_h = mask.shape[1] - o_w = mask.shape[2] - ratio = o_w / o_h - mask_h = int(math.sqrt(num_queries / ratio)) - mask_h = int(mask_h) + int((num_queries % int(mask_h)) != 0) - mask_w = num_queries // mask_h - - mask_downsample = F.interpolate(mask.unsqueeze(0), size=(mask_h, mask_w), mode="bicubic").squeeze(0) - - # Repeat batch_size times - if mask_downsample.shape[0] < batch_size: - mask_downsample = mask_downsample.repeat(batch_size, 1, 1) - - mask_downsample = mask_downsample.view(mask_downsample.shape[0], -1) - - downsampled_area = mask_h * mask_w - # If the output image and the mask do not have the same aspect ratio, tensor shapes will not match - # Pad tensor if downsampled_mask.shape[1] is smaller than num_queries - if downsampled_area < num_queries: - warnings.warn( - "The aspect ratio of the mask does not match the aspect ratio of the output image. " - "Please update your masks or adjust the output size for optimal performance.", - UserWarning, - ) - mask_downsample = F.pad(mask_downsample, (0, num_queries - mask_downsample.shape[1]), value=0.0) - # Discard last embeddings if downsampled_mask.shape[1] is bigger than num_queries - if downsampled_area > num_queries: - warnings.warn( - "The aspect ratio of the mask does not match the aspect ratio of the output image. " - "Please update your masks or adjust the output size for optimal performance.", - UserWarning, - ) - mask_downsample = mask_downsample[:, :num_queries] - - # Repeat last dimension to match SDPA output shape - mask_downsample = mask_downsample.view(mask_downsample.shape[0], mask_downsample.shape[1], 1).repeat( - 1, 1, value_embed_dim - ) - - return mask_downsample - - -class PixArtImageProcessor(VaeImageProcessor): - """ - Image processor for PixArt image resize and crop. - - Args: - do_resize (`bool`, *optional*, defaults to `True`): - Whether to downscale the image's (height, width) dimensions to multiples of `vae_scale_factor`. Can accept - `height` and `width` arguments from [`image_processor.VaeImageProcessor.preprocess`] method. - vae_scale_factor (`int`, *optional*, defaults to `8`): - VAE scale factor. If `do_resize` is `True`, the image is automatically resized to multiples of this factor. - resample (`str`, *optional*, defaults to `lanczos`): - Resampling filter to use when resizing the image. - do_normalize (`bool`, *optional*, defaults to `True`): - Whether to normalize the image to [-1,1]. - do_binarize (`bool`, *optional*, defaults to `False`): - Whether to binarize the image to 0/1. - do_convert_rgb (`bool`, *optional*, defaults to be `False`): - Whether to convert the images to RGB format. - do_convert_grayscale (`bool`, *optional*, defaults to be `False`): - Whether to convert the images to grayscale format. - """ - - @register_to_config - def __init__( - self, - do_resize: bool = True, - vae_scale_factor: int = 8, - resample: str = "lanczos", - do_normalize: bool = True, - do_binarize: bool = False, - do_convert_grayscale: bool = False, - ): - super().__init__( - do_resize=do_resize, - vae_scale_factor=vae_scale_factor, - resample=resample, - do_normalize=do_normalize, - do_binarize=do_binarize, - do_convert_grayscale=do_convert_grayscale, - ) - - @staticmethod - def classify_height_width_bin(height: int, width: int, ratios: dict) -> tuple[int, int]: - r""" - Returns the binned height and width based on the aspect ratio. - - Args: - height (`int`): The height of the image. - width (`int`): The width of the image. - ratios (`dict`): A dictionary where keys are aspect ratios and values are tuples of (height, width). - - Returns: - `tuple[int, int]`: The closest binned height and width. - """ - ar = float(height / width) - closest_ratio = min(ratios.keys(), key=lambda ratio: abs(float(ratio) - ar)) - default_hw = ratios[closest_ratio] - return int(default_hw[0]), int(default_hw[1]) - - @staticmethod - def resize_and_crop_tensor(samples: torch.Tensor, new_width: int, new_height: int) -> torch.Tensor: - r""" - Resizes and crops a tensor of images to the specified dimensions. - - Args: - samples (`torch.Tensor`): - A tensor of shape (N, C, H, W) where N is the batch size, C is the number of channels, H is the height, - and W is the width. - new_width (`int`): The desired width of the output images. - new_height (`int`): The desired height of the output images. - - Returns: - `torch.Tensor`: A tensor containing the resized and cropped images. - """ - orig_height, orig_width = samples.shape[2], samples.shape[3] - - # Check if resizing is needed - if orig_height != new_height or orig_width != new_width: - ratio = max(new_height / orig_height, new_width / orig_width) - resized_width = int(orig_width * ratio) - resized_height = int(orig_height * ratio) - - # Resize - samples = F.interpolate( - samples, size=(resized_height, resized_width), mode="bilinear", align_corners=False - ) - - # Center Crop - start_x = (resized_width - new_width) // 2 - end_x = start_x + new_width - start_y = (resized_height - new_height) // 2 - end_y = start_y + new_height - samples = samples[:, :, start_y:end_y, start_x:end_x] - - return samples diff --git a/diffusers/loaders/__init__.py b/diffusers/loaders/__init__.py deleted file mode 100644 index 1c6693bd0c0808607a628a2ce534082d18174420..0000000000000000000000000000000000000000 --- a/diffusers/loaders/__init__.py +++ /dev/null @@ -1,159 +0,0 @@ -from typing import TYPE_CHECKING - -from ..utils import DIFFUSERS_SLOW_IMPORT, _LazyModule, deprecate -from ..utils.import_utils import is_peft_available, is_torch_available, is_transformers_available - - -def text_encoder_lora_state_dict(text_encoder): - deprecate( - "text_encoder_load_state_dict in `models`", - "0.27.0", - "`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.", - ) - state_dict = {} - - for name, module in text_encoder_attn_modules(text_encoder): - for k, v in module.q_proj.lora_linear_layer.state_dict().items(): - state_dict[f"{name}.q_proj.lora_linear_layer.{k}"] = v - - for k, v in module.k_proj.lora_linear_layer.state_dict().items(): - state_dict[f"{name}.k_proj.lora_linear_layer.{k}"] = v - - for k, v in module.v_proj.lora_linear_layer.state_dict().items(): - state_dict[f"{name}.v_proj.lora_linear_layer.{k}"] = v - - for k, v in module.out_proj.lora_linear_layer.state_dict().items(): - state_dict[f"{name}.out_proj.lora_linear_layer.{k}"] = v - - return state_dict - - -if is_transformers_available(): - - def text_encoder_attn_modules(text_encoder): - deprecate( - "text_encoder_attn_modules in `models`", - "0.27.0", - "`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.", - ) - from transformers import CLIPTextModel, CLIPTextModelWithProjection - - attn_modules = [] - - if isinstance(text_encoder, (CLIPTextModel, CLIPTextModelWithProjection)): - for i, layer in enumerate(text_encoder.text_model.encoder.layers): - name = f"text_model.encoder.layers.{i}.self_attn" - mod = layer.self_attn - attn_modules.append((name, mod)) - else: - raise ValueError(f"do not know how to get attention modules for: {text_encoder.__class__.__name__}") - - return attn_modules - - -_import_structure = {} - -if is_torch_available(): - _import_structure["single_file_model"] = ["FromOriginalModelMixin"] - _import_structure["transformer_flux"] = ["FluxTransformer2DLoadersMixin"] - _import_structure["transformer_sd3"] = ["SD3Transformer2DLoadersMixin"] - _import_structure["unet"] = ["UNet2DConditionLoadersMixin"] - _import_structure["utils"] = ["AttnProcsLayers"] - if is_transformers_available(): - _import_structure["single_file"] = ["FromSingleFileMixin"] - _import_structure["lora_pipeline"] = [ - "AceStepLoraLoaderMixin", - "AmusedLoraLoaderMixin", - "AnimaLoraLoaderMixin", - "StableDiffusionLoraLoaderMixin", - "SD3LoraLoaderMixin", - "AuraFlowLoraLoaderMixin", - "StableDiffusionXLLoraLoaderMixin", - "LTX2LoraLoaderMixin", - "LTXVideoLoraLoaderMixin", - "LoraLoaderMixin", - "FluxLoraLoaderMixin", - "CogVideoXLoraLoaderMixin", - "CogView4LoraLoaderMixin", - "Mochi1LoraLoaderMixin", - "HunyuanVideoLoraLoaderMixin", - "SanaLoraLoaderMixin", - "Lumina2LoraLoaderMixin", - "WanLoraLoaderMixin", - "HeliosLoraLoaderMixin", - "KandinskyLoraLoaderMixin", - "HiDreamImageLoraLoaderMixin", - "SkyReelsV2LoraLoaderMixin", - "QwenImageLoraLoaderMixin", - "Krea2LoraLoaderMixin", - "ZImageLoraLoaderMixin", - "Flux2LoraLoaderMixin", - "Ideogram4LoraLoaderMixin", - "ErnieImageLoraLoaderMixin", - "CosmosLoraLoaderMixin", - ] - _import_structure["textual_inversion"] = ["TextualInversionLoaderMixin"] - _import_structure["ip_adapter"] = [ - "IPAdapterMixin", - "FluxIPAdapterMixin", - "SD3IPAdapterMixin", - "ModularIPAdapterMixin", - ] - -_import_structure["peft"] = ["PeftAdapterMixin"] - - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - if is_torch_available(): - from .single_file_model import FromOriginalModelMixin - from .transformer_flux import FluxTransformer2DLoadersMixin - from .transformer_sd3 import SD3Transformer2DLoadersMixin - from .unet import UNet2DConditionLoadersMixin - from .utils import AttnProcsLayers - - if is_transformers_available(): - from .ip_adapter import ( - FluxIPAdapterMixin, - IPAdapterMixin, - ModularIPAdapterMixin, - SD3IPAdapterMixin, - ) - from .lora_pipeline import ( - AceStepLoraLoaderMixin, - AmusedLoraLoaderMixin, - AnimaLoraLoaderMixin, - AuraFlowLoraLoaderMixin, - CogVideoXLoraLoaderMixin, - CogView4LoraLoaderMixin, - CosmosLoraLoaderMixin, - ErnieImageLoraLoaderMixin, - Flux2LoraLoaderMixin, - FluxLoraLoaderMixin, - HeliosLoraLoaderMixin, - HiDreamImageLoraLoaderMixin, - HunyuanVideoLoraLoaderMixin, - Ideogram4LoraLoaderMixin, - KandinskyLoraLoaderMixin, - Krea2LoraLoaderMixin, - LoraLoaderMixin, - LTX2LoraLoaderMixin, - LTXVideoLoraLoaderMixin, - Lumina2LoraLoaderMixin, - Mochi1LoraLoaderMixin, - QwenImageLoraLoaderMixin, - SanaLoraLoaderMixin, - SD3LoraLoaderMixin, - SkyReelsV2LoraLoaderMixin, - StableDiffusionLoraLoaderMixin, - StableDiffusionXLLoraLoaderMixin, - WanLoraLoaderMixin, - ZImageLoraLoaderMixin, - ) - from .single_file import FromSingleFileMixin - from .textual_inversion import TextualInversionLoaderMixin - - from .peft import PeftAdapterMixin -else: - import sys - - sys.modules[__name__] = _LazyModule(__name__, globals()["__file__"], _import_structure, module_spec=__spec__) diff --git a/diffusers/loaders/ip_adapter.py b/diffusers/loaders/ip_adapter.py deleted file mode 100644 index 5f8d3f48c99755239b7c268b0cff38483239bc12..0000000000000000000000000000000000000000 --- a/diffusers/loaders/ip_adapter.py +++ /dev/null @@ -1,1134 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from pathlib import Path -from typing import List, Union - -import torch -import torch.nn.functional as F -from huggingface_hub.utils import validate_hf_hub_args -from safetensors import safe_open - -from ..models.modeling_utils import _LOW_CPU_MEM_USAGE_DEFAULT, load_state_dict -from ..utils import ( - USE_PEFT_BACKEND, - _get_detailed_type, - _get_model_file, - _is_valid_type, - is_accelerate_available, - is_torch_version, - is_transformers_available, - is_transformers_version, - logging, -) -from .unet_loader_utils import _maybe_expand_lora_scales - - -if is_transformers_available(): - from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection, SiglipImageProcessor, SiglipVisionModel - -from ..models.attention_processor import ( - AttnProcessor, - AttnProcessor2_0, - IPAdapterAttnProcessor, - IPAdapterAttnProcessor2_0, - IPAdapterXFormersAttnProcessor, - JointAttnProcessor2_0, - SD3IPAdapterJointAttnProcessor2_0, -) - - -logger = logging.get_logger(__name__) - - -class IPAdapterMixin: - """Mixin for handling IP Adapters.""" - - @validate_hf_hub_args - def load_ip_adapter( - self, - pretrained_model_name_or_path_or_dict: str | list[str] | dict[str, torch.Tensor], - subfolder: str | list[str], - weight_name: str | list[str], - image_encoder_folder: str | None = "image_encoder", - **kwargs, - ): - """ - Parameters: - pretrained_model_name_or_path_or_dict (`str` or `list[str]` or `os.PathLike` or `list[os.PathLike]` or `dict` or `list[dict]`): - Can be either: - - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved - with [`ModelMixin.save_pretrained`]. - - A [torch state - dict](https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict). - subfolder (`str` or `list[str]`): - The subfolder location of a model file within a larger model repository on the Hub or locally. If a - list is passed, it should have the same length as `weight_name`. - weight_name (`str` or `list[str]`): - The name of the weight file to load. If a list is passed, it should have the same length as - `subfolder`. - image_encoder_folder (`str`, *optional*, defaults to `image_encoder`): - The subfolder location of the image encoder within a larger model repository on the Hub or locally. - Pass `None` to not load the image encoder. If the image encoder is located in a folder inside - `subfolder`, you only need to pass the name of the folder that contains image encoder weights, e.g. - `image_encoder_folder="image_encoder"`. If the image encoder is located in a folder other than - `subfolder`, you should pass the path to the folder that contains image encoder weights, for example, - `image_encoder_folder="different_subfolder/image_encoder"`. - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 else `False`): - Speed up model loading only loading the pretrained weights and not initializing the weights. This also - tries to not use more than 1x model size in CPU memory (including peak memory) while loading the model. - Only supported for PyTorch >= 1.9.0. If you are using an older version of PyTorch, setting this - argument to `True` will raise an error. - """ - - # handle the list inputs for multiple IP Adapters - if not isinstance(weight_name, list): - weight_name = [weight_name] - - if not isinstance(pretrained_model_name_or_path_or_dict, list): - pretrained_model_name_or_path_or_dict = [pretrained_model_name_or_path_or_dict] - if len(pretrained_model_name_or_path_or_dict) == 1: - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict * len(weight_name) - - if not isinstance(subfolder, list): - subfolder = [subfolder] - if len(subfolder) == 1: - subfolder = subfolder * len(weight_name) - - if len(weight_name) != len(pretrained_model_name_or_path_or_dict): - raise ValueError("`weight_name` and `pretrained_model_name_or_path_or_dict` must have the same length.") - - if len(weight_name) != len(subfolder): - raise ValueError("`weight_name` and `subfolder` must have the same length.") - - # Load the main state dict first. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT) - - if low_cpu_mem_usage and not is_accelerate_available(): - low_cpu_mem_usage = False - logger.warning( - "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the" - " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install" - " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip" - " install accelerate\n```\n." - ) - - if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `low_cpu_mem_usage=False`." - ) - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - state_dicts = [] - for pretrained_model_name_or_path_or_dict, weight_name, subfolder in zip( - pretrained_model_name_or_path_or_dict, weight_name, subfolder - ): - if not isinstance(pretrained_model_name_or_path_or_dict, dict): - model_file = _get_model_file( - pretrained_model_name_or_path_or_dict, - weights_name=weight_name, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - ) - if weight_name.endswith(".safetensors"): - state_dict = {"image_proj": {}, "ip_adapter": {}} - with safe_open(model_file, framework="pt", device="cpu") as f: - for key in f.keys(): - if key.startswith("image_proj."): - state_dict["image_proj"][key.replace("image_proj.", "")] = f.get_tensor(key) - elif key.startswith("ip_adapter."): - state_dict["ip_adapter"][key.replace("ip_adapter.", "")] = f.get_tensor(key) - else: - state_dict = load_state_dict(model_file) - else: - state_dict = pretrained_model_name_or_path_or_dict - - keys = list(state_dict.keys()) - if "image_proj" not in keys and "ip_adapter" not in keys: - raise ValueError("Required keys are (`image_proj` and `ip_adapter`) missing from the state dict.") - - state_dicts.append(state_dict) - - # load CLIP image encoder here if it has not been registered to the pipeline yet - if hasattr(self, "image_encoder") and getattr(self, "image_encoder", None) is None: - if image_encoder_folder is not None: - if not isinstance(pretrained_model_name_or_path_or_dict, dict): - logger.info(f"loading image_encoder from {pretrained_model_name_or_path_or_dict}") - if image_encoder_folder.count("/") == 0: - image_encoder_subfolder = Path(subfolder, image_encoder_folder).as_posix() - else: - image_encoder_subfolder = Path(image_encoder_folder).as_posix() - - # transformers renamed `torch_dtype` to `dtype` in 4.56.0. - dtype_kwarg = ( - {"dtype": self.dtype} - if is_transformers_version(">=", "4.56.0") - else {"torch_dtype": self.dtype} - ) - image_encoder = CLIPVisionModelWithProjection.from_pretrained( - pretrained_model_name_or_path_or_dict, - subfolder=image_encoder_subfolder, - low_cpu_mem_usage=low_cpu_mem_usage, - cache_dir=cache_dir, - local_files_only=local_files_only, - **dtype_kwarg, - ).to(self.device) - self.register_modules(image_encoder=image_encoder) - else: - raise ValueError( - "`image_encoder` cannot be loaded because `pretrained_model_name_or_path_or_dict` is a state dict." - ) - else: - logger.warning( - "image_encoder is not loaded since `image_encoder_folder=None` passed. You will not be able to use `ip_adapter_image` when calling the pipeline with IP-Adapter." - "Use `ip_adapter_image_embeds` to pass pre-generated image embedding instead." - ) - - # create feature extractor if it has not been registered to the pipeline yet - if hasattr(self, "feature_extractor") and getattr(self, "feature_extractor", None) is None: - # FaceID IP adapters don't need the image encoder so it's not present, in this case we default to 224 - default_clip_size = 224 - clip_image_size = ( - self.image_encoder.config.image_size if self.image_encoder is not None else default_clip_size - ) - feature_extractor = CLIPImageProcessor(size=clip_image_size, crop_size=clip_image_size) - self.register_modules(feature_extractor=feature_extractor) - - # load ip-adapter into unet - unet = getattr(self, self.unet_name) if not hasattr(self, "unet") else self.unet - unet._load_ip_adapter_weights(state_dicts, low_cpu_mem_usage=low_cpu_mem_usage) - - extra_loras = unet._load_ip_adapter_loras(state_dicts) - if extra_loras != {}: - if not USE_PEFT_BACKEND: - logger.warning("PEFT backend is required to load these weights.") - else: - # apply the IP Adapter Face ID LoRA weights - peft_config = getattr(unet, "peft_config", {}) - for k, lora in extra_loras.items(): - if f"faceid_{k}" not in peft_config: - self.load_lora_weights(lora, adapter_name=f"faceid_{k}") - self.set_adapters([f"faceid_{k}"], adapter_weights=[1.0]) - - def set_ip_adapter_scale(self, scale): - """ - Set IP-Adapter scales per-transformer block. Input `scale` could be a single config or a list of configs for - granular control over each IP-Adapter behavior. A config can be a float or a dictionary. - - Example: - - ```py - # To use original IP-Adapter - scale = 1.0 - pipeline.set_ip_adapter_scale(scale) - - # To use style block only - scale = { - "up": {"block_0": [0.0, 1.0, 0.0]}, - } - pipeline.set_ip_adapter_scale(scale) - - # To use style+layout blocks - scale = { - "down": {"block_2": [0.0, 1.0]}, - "up": {"block_0": [0.0, 1.0, 0.0]}, - } - pipeline.set_ip_adapter_scale(scale) - - # To use style and layout from 2 reference images - scales = [{"down": {"block_2": [0.0, 1.0]}}, {"up": {"block_0": [0.0, 1.0, 0.0]}}] - pipeline.set_ip_adapter_scale(scales) - ``` - """ - unet = getattr(self, self.unet_name) if not hasattr(self, "unet") else self.unet - if not isinstance(scale, list): - scale = [scale] - scale_configs = _maybe_expand_lora_scales(unet, scale, default_scale=0.0) - - for attn_name, attn_processor in unet.attn_processors.items(): - if isinstance( - attn_processor, (IPAdapterAttnProcessor, IPAdapterAttnProcessor2_0, IPAdapterXFormersAttnProcessor) - ): - if len(scale_configs) != len(attn_processor.scale): - raise ValueError( - f"Cannot assign {len(scale_configs)} scale_configs to {len(attn_processor.scale)} IP-Adapter." - ) - elif len(scale_configs) == 1: - scale_configs = scale_configs * len(attn_processor.scale) - for i, scale_config in enumerate(scale_configs): - if isinstance(scale_config, dict): - for k, s in scale_config.items(): - if attn_name.startswith(k): - attn_processor.scale[i] = s - else: - attn_processor.scale[i] = scale_config - - def unload_ip_adapter(self): - """ - Unloads the IP Adapter weights - - Examples: - - ```python - >>> # Assuming `pipeline` is already loaded with the IP Adapter weights. - >>> pipeline.unload_ip_adapter() - >>> ... - ``` - """ - # remove CLIP image encoder - if hasattr(self, "image_encoder") and getattr(self, "image_encoder", None) is not None: - self.image_encoder = None - self.register_to_config(image_encoder=[None, None]) - - # remove feature extractor only when safety_checker is None as safety_checker uses - # the feature_extractor later - if not hasattr(self, "safety_checker"): - if hasattr(self, "feature_extractor") and getattr(self, "feature_extractor", None) is not None: - self.feature_extractor = None - self.register_to_config(feature_extractor=[None, None]) - - # remove hidden encoder - self.unet.encoder_hid_proj = None - self.unet.config.encoder_hid_dim_type = None - - # Kolors: restore `encoder_hid_proj` with `text_encoder_hid_proj` - if hasattr(self.unet, "text_encoder_hid_proj") and self.unet.text_encoder_hid_proj is not None: - self.unet.encoder_hid_proj = self.unet.text_encoder_hid_proj - self.unet.text_encoder_hid_proj = None - self.unet.config.encoder_hid_dim_type = "text_proj" - - # restore original Unet attention processors layers - attn_procs = {} - for name, value in self.unet.attn_processors.items(): - attn_processor_class = ( - AttnProcessor2_0() if hasattr(F, "scaled_dot_product_attention") else AttnProcessor() - ) - attn_procs[name] = ( - attn_processor_class - if isinstance( - value, (IPAdapterAttnProcessor, IPAdapterAttnProcessor2_0, IPAdapterXFormersAttnProcessor) - ) - else value.__class__() - ) - self.unet.set_attn_processor(attn_procs) - - -class ModularIPAdapterMixin: - """Mixin for handling IP Adapters.""" - - @validate_hf_hub_args - def load_ip_adapter( - self, - pretrained_model_name_or_path_or_dict: str | list[str] | dict[str, torch.Tensor], - subfolder: str | list[str], - weight_name: str | list[str], - **kwargs, - ): - """ - Parameters: - pretrained_model_name_or_path_or_dict (`str` or `list[str]` or `os.PathLike` or `list[os.PathLike]` or `dict` or `list[dict]`): - Can be either: - - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved - with [`ModelMixin.save_pretrained`]. - - A [torch state - dict](https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict). - subfolder (`str` or `list[str]`): - The subfolder location of a model file within a larger model repository on the Hub or locally. If a - list is passed, it should have the same length as `weight_name`. - weight_name (`str` or `list[str]`): - The name of the weight file to load. If a list is passed, it should have the same length as - `subfolder`. - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 else `False`): - Speed up model loading only loading the pretrained weights and not initializing the weights. This also - tries to not use more than 1x model size in CPU memory (including peak memory) while loading the model. - Only supported for PyTorch >= 1.9.0. If you are using an older version of PyTorch, setting this - argument to `True` will raise an error. - """ - - # handle the list inputs for multiple IP Adapters - if not isinstance(weight_name, list): - weight_name = [weight_name] - - if not isinstance(pretrained_model_name_or_path_or_dict, list): - pretrained_model_name_or_path_or_dict = [pretrained_model_name_or_path_or_dict] - if len(pretrained_model_name_or_path_or_dict) == 1: - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict * len(weight_name) - - if not isinstance(subfolder, list): - subfolder = [subfolder] - if len(subfolder) == 1: - subfolder = subfolder * len(weight_name) - - if len(weight_name) != len(pretrained_model_name_or_path_or_dict): - raise ValueError("`weight_name` and `pretrained_model_name_or_path_or_dict` must have the same length.") - - if len(weight_name) != len(subfolder): - raise ValueError("`weight_name` and `subfolder` must have the same length.") - - # Load the main state dict first. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT) - - if low_cpu_mem_usage and not is_accelerate_available(): - low_cpu_mem_usage = False - logger.warning( - "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the" - " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install" - " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip" - " install accelerate\n```\n." - ) - - if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `low_cpu_mem_usage=False`." - ) - - user_agent = { - "file_type": "attn_procs_weights", - "framework": "pytorch", - } - state_dicts = [] - for pretrained_model_name_or_path_or_dict, weight_name, subfolder in zip( - pretrained_model_name_or_path_or_dict, weight_name, subfolder - ): - if not isinstance(pretrained_model_name_or_path_or_dict, dict): - model_file = _get_model_file( - pretrained_model_name_or_path_or_dict, - weights_name=weight_name, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - ) - if weight_name.endswith(".safetensors"): - state_dict = {"image_proj": {}, "ip_adapter": {}} - with safe_open(model_file, framework="pt", device="cpu") as f: - for key in f.keys(): - if key.startswith("image_proj."): - state_dict["image_proj"][key.replace("image_proj.", "")] = f.get_tensor(key) - elif key.startswith("ip_adapter."): - state_dict["ip_adapter"][key.replace("ip_adapter.", "")] = f.get_tensor(key) - else: - state_dict = load_state_dict(model_file) - else: - state_dict = pretrained_model_name_or_path_or_dict - - keys = list(state_dict.keys()) - if "image_proj" not in keys and "ip_adapter" not in keys: - raise ValueError("Required keys are (`image_proj` and `ip_adapter`) missing from the state dict.") - - state_dicts.append(state_dict) - - unet_name = getattr(self, "unet_name", "unet") - unet = getattr(self, unet_name) - unet._load_ip_adapter_weights(state_dicts, low_cpu_mem_usage=low_cpu_mem_usage) - - extra_loras = unet._load_ip_adapter_loras(state_dicts) - if extra_loras != {}: - if not USE_PEFT_BACKEND: - logger.warning("PEFT backend is required to load these weights.") - else: - # apply the IP Adapter Face ID LoRA weights - peft_config = getattr(unet, "peft_config", {}) - for k, lora in extra_loras.items(): - if f"faceid_{k}" not in peft_config: - self.load_lora_weights(lora, adapter_name=f"faceid_{k}") - self.set_adapters([f"faceid_{k}"], adapter_weights=[1.0]) - - def set_ip_adapter_scale(self, scale): - """ - Set IP-Adapter scales per-transformer block. Input `scale` could be a single config or a list of configs for - granular control over each IP-Adapter behavior. A config can be a float or a dictionary. - - Example: - - ```py - # To use original IP-Adapter - scale = 1.0 - pipeline.set_ip_adapter_scale(scale) - - # To use style block only - scale = { - "up": {"block_0": [0.0, 1.0, 0.0]}, - } - pipeline.set_ip_adapter_scale(scale) - - # To use style+layout blocks - scale = { - "down": {"block_2": [0.0, 1.0]}, - "up": {"block_0": [0.0, 1.0, 0.0]}, - } - pipeline.set_ip_adapter_scale(scale) - - # To use style and layout from 2 reference images - scales = [{"down": {"block_2": [0.0, 1.0]}}, {"up": {"block_0": [0.0, 1.0, 0.0]}}] - pipeline.set_ip_adapter_scale(scales) - ``` - """ - unet_name = getattr(self, "unet_name", "unet") - unet = getattr(self, unet_name) - if not isinstance(scale, list): - scale = [scale] - scale_configs = _maybe_expand_lora_scales(unet, scale, default_scale=0.0) - - for attn_name, attn_processor in unet.attn_processors.items(): - if isinstance( - attn_processor, (IPAdapterAttnProcessor, IPAdapterAttnProcessor2_0, IPAdapterXFormersAttnProcessor) - ): - if len(scale_configs) != len(attn_processor.scale): - raise ValueError( - f"Cannot assign {len(scale_configs)} scale_configs to {len(attn_processor.scale)} IP-Adapter." - ) - elif len(scale_configs) == 1: - scale_configs = scale_configs * len(attn_processor.scale) - for i, scale_config in enumerate(scale_configs): - if isinstance(scale_config, dict): - for k, s in scale_config.items(): - if attn_name.startswith(k): - attn_processor.scale[i] = s - else: - attn_processor.scale[i] = scale_config - - def unload_ip_adapter(self): - """ - Unloads the IP Adapter weights - - Examples: - - ```python - >>> # Assuming `pipeline` is already loaded with the IP Adapter weights. - >>> pipeline.unload_ip_adapter() - >>> ... - ``` - """ - - # remove hidden encoder - if self.unet is None: - return - - self.unet.encoder_hid_proj = None - self.unet.config.encoder_hid_dim_type = None - - # Kolors: restore `encoder_hid_proj` with `text_encoder_hid_proj` - if hasattr(self.unet, "text_encoder_hid_proj") and self.unet.text_encoder_hid_proj is not None: - self.unet.encoder_hid_proj = self.unet.text_encoder_hid_proj - self.unet.text_encoder_hid_proj = None - self.unet.config.encoder_hid_dim_type = "text_proj" - - # restore original Unet attention processors layers - attn_procs = {} - for name, value in self.unet.attn_processors.items(): - attn_processor_class = ( - AttnProcessor2_0() if hasattr(F, "scaled_dot_product_attention") else AttnProcessor() - ) - attn_procs[name] = ( - attn_processor_class - if isinstance( - value, (IPAdapterAttnProcessor, IPAdapterAttnProcessor2_0, IPAdapterXFormersAttnProcessor) - ) - else value.__class__() - ) - self.unet.set_attn_processor(attn_procs) - - -class FluxIPAdapterMixin: - """Mixin for handling Flux IP Adapters.""" - - @validate_hf_hub_args - def load_ip_adapter( - self, - pretrained_model_name_or_path_or_dict: str | list[str] | dict[str, torch.Tensor], - weight_name: str | list[str], - subfolder: str | list[str] | None = "", - image_encoder_pretrained_model_name_or_path: str | None = "image_encoder", - image_encoder_subfolder: str | None = "", - image_encoder_dtype: torch.dtype = torch.float16, - **kwargs, - ): - """ - Parameters: - pretrained_model_name_or_path_or_dict (`str` or `list[str]` or `os.PathLike` or `list[os.PathLike]` or `dict` or `list[dict]`): - Can be either: - - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved - with [`ModelMixin.save_pretrained`]. - - A [torch state - dict](https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict). - subfolder (`str` or `list[str]`): - The subfolder location of a model file within a larger model repository on the Hub or locally. If a - list is passed, it should have the same length as `weight_name`. - weight_name (`str` or `list[str]`): - The name of the weight file to load. If a list is passed, it should have the same length as - `weight_name`. - image_encoder_pretrained_model_name_or_path (`str`, *optional*, defaults to `./image_encoder`): - Can be either: - - - A string, the *model id* (for example `openai/clip-vit-large-patch14`) of a pretrained model - hosted on the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved - with [`ModelMixin.save_pretrained`]. - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 else `False`): - Speed up model loading only loading the pretrained weights and not initializing the weights. This also - tries to not use more than 1x model size in CPU memory (including peak memory) while loading the model. - Only supported for PyTorch >= 1.9.0. If you are using an older version of PyTorch, setting this - argument to `True` will raise an error. - """ - - # handle the list inputs for multiple IP Adapters - if not isinstance(weight_name, list): - weight_name = [weight_name] - - if not isinstance(pretrained_model_name_or_path_or_dict, list): - pretrained_model_name_or_path_or_dict = [pretrained_model_name_or_path_or_dict] - if len(pretrained_model_name_or_path_or_dict) == 1: - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict * len(weight_name) - - if not isinstance(subfolder, list): - subfolder = [subfolder] - if len(subfolder) == 1: - subfolder = subfolder * len(weight_name) - - if len(weight_name) != len(pretrained_model_name_or_path_or_dict): - raise ValueError("`weight_name` and `pretrained_model_name_or_path_or_dict` must have the same length.") - - if len(weight_name) != len(subfolder): - raise ValueError("`weight_name` and `subfolder` must have the same length.") - - # Load the main state dict first. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT) - - if low_cpu_mem_usage and not is_accelerate_available(): - low_cpu_mem_usage = False - logger.warning( - "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the" - " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install" - " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip" - " install accelerate\n```\n." - ) - - if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `low_cpu_mem_usage=False`." - ) - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - state_dicts = [] - for pretrained_model_name_or_path_or_dict, weight_name, subfolder in zip( - pretrained_model_name_or_path_or_dict, weight_name, subfolder - ): - if not isinstance(pretrained_model_name_or_path_or_dict, dict): - model_file = _get_model_file( - pretrained_model_name_or_path_or_dict, - weights_name=weight_name, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - ) - if weight_name.endswith(".safetensors"): - state_dict = {"image_proj": {}, "ip_adapter": {}} - with safe_open(model_file, framework="pt", device="cpu") as f: - image_proj_keys = ["ip_adapter_proj_model.", "image_proj."] - ip_adapter_keys = ["double_blocks.", "ip_adapter."] - for key in f.keys(): - if any(key.startswith(prefix) for prefix in image_proj_keys): - diffusers_name = ".".join(key.split(".")[1:]) - state_dict["image_proj"][diffusers_name] = f.get_tensor(key) - elif any(key.startswith(prefix) for prefix in ip_adapter_keys): - diffusers_name = ( - ".".join(key.split(".")[1:]) - .replace("ip_adapter_double_stream_k_proj", "to_k_ip") - .replace("ip_adapter_double_stream_v_proj", "to_v_ip") - .replace("processor.", "") - ) - state_dict["ip_adapter"][diffusers_name] = f.get_tensor(key) - else: - state_dict = load_state_dict(model_file) - else: - state_dict = pretrained_model_name_or_path_or_dict - - keys = list(state_dict.keys()) - if keys != ["image_proj", "ip_adapter"]: - raise ValueError("Required keys are (`image_proj` and `ip_adapter`) missing from the state dict.") - - state_dicts.append(state_dict) - - # load CLIP image encoder here if it has not been registered to the pipeline yet - if hasattr(self, "image_encoder") and getattr(self, "image_encoder", None) is None: - if image_encoder_pretrained_model_name_or_path is not None: - if not isinstance(pretrained_model_name_or_path_or_dict, dict): - logger.info(f"loading image_encoder from {image_encoder_pretrained_model_name_or_path}") - image_encoder = ( - CLIPVisionModelWithProjection.from_pretrained( - image_encoder_pretrained_model_name_or_path, - subfolder=image_encoder_subfolder, - low_cpu_mem_usage=low_cpu_mem_usage, - cache_dir=cache_dir, - local_files_only=local_files_only, - torch_dtype=image_encoder_dtype, - ) - .to(self.device) - .eval() - ) - self.register_modules(image_encoder=image_encoder) - else: - raise ValueError( - "`image_encoder` cannot be loaded because `pretrained_model_name_or_path_or_dict` is a state dict." - ) - else: - logger.warning( - "image_encoder is not loaded since `image_encoder_folder=None` passed. You will not be able to use `ip_adapter_image` when calling the pipeline with IP-Adapter." - "Use `ip_adapter_image_embeds` to pass pre-generated image embedding instead." - ) - - # create feature extractor if it has not been registered to the pipeline yet - if hasattr(self, "feature_extractor") and getattr(self, "feature_extractor", None) is None: - # FaceID IP adapters don't need the image encoder so it's not present, in this case we default to 224 - default_clip_size = 224 - clip_image_size = ( - self.image_encoder.config.image_size if self.image_encoder is not None else default_clip_size - ) - feature_extractor = CLIPImageProcessor(size=clip_image_size, crop_size=clip_image_size) - self.register_modules(feature_extractor=feature_extractor) - - # load ip-adapter into transformer - self.transformer._load_ip_adapter_weights(state_dicts, low_cpu_mem_usage=low_cpu_mem_usage) - - def set_ip_adapter_scale(self, scale: float | list[float] | list[list[float]]): - """ - Set IP-Adapter scales per-transformer block. Input `scale` could be a single config or a list of configs for - granular control over each IP-Adapter behavior. A config can be a float or a list. - - `float` is converted to list and repeated for the number of blocks and the number of IP adapters. `list[float]` - length match the number of blocks, it is repeated for each IP adapter. `list[list[float]]` must match the - number of IP adapters and each must match the number of blocks. - - Example: - - ```py - # To use original IP-Adapter - scale = 1.0 - pipeline.set_ip_adapter_scale(scale) - - - def LinearStrengthModel(start, finish, size): - return [(start + (finish - start) * (i / (size - 1))) for i in range(size)] - - - ip_strengths = LinearStrengthModel(0.3, 0.92, 19) - pipeline.set_ip_adapter_scale(ip_strengths) - ``` - """ - - scale_type = Union[int, float] - num_ip_adapters = self.transformer.encoder_hid_proj.num_ip_adapters - num_layers = self.transformer.config.num_layers - - # Single value for all layers of all IP-Adapters - if isinstance(scale, scale_type): - scale = [scale for _ in range(num_ip_adapters)] - # List of per-layer scales for a single IP-Adapter - elif _is_valid_type(scale, List[scale_type]) and num_ip_adapters == 1: - scale = [scale] - # Invalid scale type - elif not _is_valid_type(scale, List[Union[scale_type, List[scale_type]]]): - raise TypeError(f"Unexpected type {_get_detailed_type(scale)} for scale.") - - if len(scale) != num_ip_adapters: - raise ValueError(f"Cannot assign {len(scale)} scales to {num_ip_adapters} IP-Adapters.") - - if any(len(s) != num_layers for s in scale if isinstance(s, list)): - invalid_scale_sizes = {len(s) for s in scale if isinstance(s, list)} - {num_layers} - raise ValueError( - f"Expected list of {num_layers} scales, got {', '.join(str(x) for x in invalid_scale_sizes)}." - ) - - # Scalars are transformed to lists with length num_layers - scale_configs = [[s] * num_layers if isinstance(s, scale_type) else s for s in scale] - - # Set scales. zip over scale_configs prevents going into single transformer layers - for attn_processor, *scale in zip(self.transformer.attn_processors.values(), *scale_configs): - attn_processor.scale = scale - - def unload_ip_adapter(self): - """ - Unloads the IP Adapter weights - - Examples: - - ```python - >>> # Assuming `pipeline` is already loaded with the IP Adapter weights. - >>> pipeline.unload_ip_adapter() - >>> ... - ``` - """ - # TODO: once the 1.0.0 deprecations are in, we can move the imports to top-level - from ..models.transformers.transformer_flux import FluxAttnProcessor, FluxIPAdapterAttnProcessor - - # remove CLIP image encoder - if hasattr(self, "image_encoder") and getattr(self, "image_encoder", None) is not None: - self.image_encoder = None - self.register_to_config(image_encoder=[None, None]) - - # remove feature extractor only when safety_checker is None as safety_checker uses - # the feature_extractor later - if not hasattr(self, "safety_checker"): - if hasattr(self, "feature_extractor") and getattr(self, "feature_extractor", None) is not None: - self.feature_extractor = None - self.register_to_config(feature_extractor=[None, None]) - - # remove hidden encoder - self.transformer.encoder_hid_proj = None - self.transformer.config.encoder_hid_dim_type = None - - # restore original Transformer attention processors layers - attn_procs = {} - for name, value in self.transformer.attn_processors.items(): - attn_processor_class = FluxAttnProcessor() - attn_procs[name] = ( - attn_processor_class if isinstance(value, FluxIPAdapterAttnProcessor) else value.__class__() - ) - self.transformer.set_attn_processor(attn_procs) - - -class SD3IPAdapterMixin: - """Mixin for handling StableDiffusion 3 IP Adapters.""" - - @property - def is_ip_adapter_active(self) -> bool: - """Checks if IP-Adapter is loaded and scale > 0. - - IP-Adapter scale controls the influence of the image prompt versus text prompt. When this value is set to 0, - the image context is irrelevant. - - Returns: - `bool`: True when IP-Adapter is loaded and any layer has scale > 0. - """ - scales = [ - attn_proc.scale - for attn_proc in self.transformer.attn_processors.values() - if isinstance(attn_proc, SD3IPAdapterJointAttnProcessor2_0) - ] - - return len(scales) > 0 and any(scale > 0 for scale in scales) - - @validate_hf_hub_args - def load_ip_adapter( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - weight_name: str = "ip-adapter.safetensors", - subfolder: str | None = None, - image_encoder_folder: str | None = "image_encoder", - **kwargs, - ) -> None: - """ - Parameters: - pretrained_model_name_or_path_or_dict (`str` or `os.PathLike` or `dict`): - Can be either: - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved - with [`ModelMixin.save_pretrained`]. - - A [torch state - dict](https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict). - weight_name (`str`, defaults to "ip-adapter.safetensors"): - The name of the weight file to load. If a list is passed, it should have the same length as - `subfolder`. - subfolder (`str`, *optional*): - The subfolder location of a model file within a larger model repository on the Hub or locally. If a - list is passed, it should have the same length as `weight_name`. - image_encoder_folder (`str`, *optional*, defaults to `image_encoder`): - The subfolder location of the image encoder within a larger model repository on the Hub or locally. - Pass `None` to not load the image encoder. If the image encoder is located in a folder inside - `subfolder`, you only need to pass the name of the folder that contains image encoder weights, e.g. - `image_encoder_folder="image_encoder"`. If the image encoder is located in a folder other than - `subfolder`, you should pass the path to the folder that contains image encoder weights, for example, - `image_encoder_folder="different_subfolder/image_encoder"`. - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 else `False`): - Speed up model loading only loading the pretrained weights and not initializing the weights. This also - tries to not use more than 1x model size in CPU memory (including peak memory) while loading the model. - Only supported for PyTorch >= 1.9.0. If you are using an older version of PyTorch, setting this - argument to `True` will raise an error. - """ - # Load the main state dict first - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT) - - if low_cpu_mem_usage and not is_accelerate_available(): - low_cpu_mem_usage = False - logger.warning( - "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the" - " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install" - " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip" - " install accelerate\n```\n." - ) - - if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `low_cpu_mem_usage=False`." - ) - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - if not isinstance(pretrained_model_name_or_path_or_dict, dict): - model_file = _get_model_file( - pretrained_model_name_or_path_or_dict, - weights_name=weight_name, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - ) - if weight_name.endswith(".safetensors"): - state_dict = {"image_proj": {}, "ip_adapter": {}} - with safe_open(model_file, framework="pt", device="cpu") as f: - for key in f.keys(): - if key.startswith("image_proj."): - state_dict["image_proj"][key.replace("image_proj.", "")] = f.get_tensor(key) - elif key.startswith("ip_adapter."): - state_dict["ip_adapter"][key.replace("ip_adapter.", "")] = f.get_tensor(key) - else: - state_dict = load_state_dict(model_file) - else: - state_dict = pretrained_model_name_or_path_or_dict - - keys = list(state_dict.keys()) - if "image_proj" not in keys and "ip_adapter" not in keys: - raise ValueError("Required keys are (`image_proj` and `ip_adapter`) missing from the state dict.") - - # Load image_encoder and feature_extractor here if they haven't been registered to the pipeline yet - if hasattr(self, "image_encoder") and getattr(self, "image_encoder", None) is None: - if image_encoder_folder is not None: - if not isinstance(pretrained_model_name_or_path_or_dict, dict): - logger.info(f"loading image_encoder from {pretrained_model_name_or_path_or_dict}") - if image_encoder_folder.count("/") == 0: - image_encoder_subfolder = Path(subfolder, image_encoder_folder).as_posix() - else: - image_encoder_subfolder = Path(image_encoder_folder).as_posix() - - # Commons args for loading image encoder and image processor - kwargs = { - "low_cpu_mem_usage": low_cpu_mem_usage, - "cache_dir": cache_dir, - "local_files_only": local_files_only, - } - # transformers renamed `torch_dtype` to `dtype` in 4.56.0. - dtype_kwarg = ( - {"dtype": self.dtype} - if is_transformers_version(">=", "4.56.0") - else {"torch_dtype": self.dtype} - ) - - self.register_modules( - feature_extractor=SiglipImageProcessor.from_pretrained(image_encoder_subfolder, **kwargs), - image_encoder=SiglipVisionModel.from_pretrained( - image_encoder_subfolder, **dtype_kwarg, **kwargs - ).to(self.device), - ) - else: - raise ValueError( - "`image_encoder` cannot be loaded because `pretrained_model_name_or_path_or_dict` is a state dict." - ) - else: - logger.warning( - "image_encoder is not loaded since `image_encoder_folder=None` passed. You will not be able to use `ip_adapter_image` when calling the pipeline with IP-Adapter." - "Use `ip_adapter_image_embeds` to pass pre-generated image embedding instead." - ) - - # Load IP-Adapter into transformer - self.transformer._load_ip_adapter_weights(state_dict, low_cpu_mem_usage=low_cpu_mem_usage) - - def set_ip_adapter_scale(self, scale: float) -> None: - """ - Set IP-Adapter scale, which controls image prompt conditioning. A value of 1.0 means the model is only - conditioned on the image prompt, and 0.0 only conditioned by the text prompt. Lowering this value encourages - the model to produce more diverse images, but they may not be as aligned with the image prompt. - - Example: - - ```python - >>> # Assuming `pipeline` is already loaded with the IP Adapter weights. - >>> pipeline.set_ip_adapter_scale(0.6) - >>> ... - ``` - - Args: - scale (float): - IP-Adapter scale to be set. - - """ - for attn_processor in self.transformer.attn_processors.values(): - if isinstance(attn_processor, SD3IPAdapterJointAttnProcessor2_0): - attn_processor.scale = scale - - def unload_ip_adapter(self) -> None: - """ - Unloads the IP Adapter weights. - - Example: - - ```python - >>> # Assuming `pipeline` is already loaded with the IP Adapter weights. - >>> pipeline.unload_ip_adapter() - >>> ... - ``` - """ - # Remove image encoder - if hasattr(self, "image_encoder") and getattr(self, "image_encoder", None) is not None: - self.image_encoder = None - self.register_to_config(image_encoder=None) - - # Remove feature extractor - if hasattr(self, "feature_extractor") and getattr(self, "feature_extractor", None) is not None: - self.feature_extractor = None - self.register_to_config(feature_extractor=None) - - # Remove image projection - self.transformer.image_proj = None - - # Restore original attention processors layers - attn_procs = { - name: ( - JointAttnProcessor2_0() if isinstance(value, SD3IPAdapterJointAttnProcessor2_0) else value.__class__() - ) - for name, value in self.transformer.attn_processors.items() - } - self.transformer.set_attn_processor(attn_procs) diff --git a/diffusers/loaders/lora_base.py b/diffusers/loaders/lora_base.py deleted file mode 100644 index d4c88d35924f71eb090defa9889586aa7edb008a..0000000000000000000000000000000000000000 --- a/diffusers/loaders/lora_base.py +++ /dev/null @@ -1,1098 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from __future__ import annotations - -import copy -import inspect -import json -import os -from pathlib import Path -from typing import Callable - -import safetensors -import torch -import torch.nn as nn -from huggingface_hub import model_info -from huggingface_hub.constants import HF_HUB_OFFLINE - -from ..models.modeling_utils import ModelMixin, load_state_dict -from ..utils import ( - USE_PEFT_BACKEND, - _get_model_file, - convert_state_dict_to_diffusers, - convert_state_dict_to_peft, - delete_adapter_layers, - deprecate, - get_adapter_name, - is_accelerate_available, - is_peft_available, - is_peft_version, - is_transformers_available, - is_transformers_version, - logging, - recurse_remove_peft_layers, - scale_lora_layers, - set_adapter_layers, - set_weights_and_activate_adapters, -) -from ..utils.peft_utils import _create_lora_config -from ..utils.state_dict_utils import _load_sft_state_dict_metadata - - -if is_transformers_available(): - from transformers import PreTrainedModel - -if is_peft_available(): - from peft.tuners.tuners_utils import BaseTunerLayer - -if is_accelerate_available(): - from accelerate.hooks import AlignDevicesHook, CpuOffload, remove_hook_from_module - -logger = logging.get_logger(__name__) - -LORA_WEIGHT_NAME = "pytorch_lora_weights.bin" -LORA_WEIGHT_NAME_SAFE = "pytorch_lora_weights.safetensors" -LORA_ADAPTER_METADATA_KEY = "lora_adapter_metadata" - - -def fuse_text_encoder_lora(text_encoder, lora_scale=1.0, safe_fusing=False, adapter_names=None): - """ - Fuses LoRAs for the text encoder. - - Args: - text_encoder (`torch.nn.Module`): - The text encoder module to set the adapter layers for. If `None`, it will try to get the `text_encoder` - attribute. - lora_scale (`float`, defaults to 1.0): - Controls how much to influence the outputs with the LoRA parameters. - safe_fusing (`bool`, defaults to `False`): - Whether to check fused weights for NaN values before fusing and if values are NaN not fusing them. - adapter_names (`list[str]` or `str`): - The names of the adapters to use. - """ - merge_kwargs = {"safe_merge": safe_fusing} - - for module in text_encoder.modules(): - if isinstance(module, BaseTunerLayer): - if lora_scale != 1.0: - module.scale_layer(lora_scale) - - # For BC with previous PEFT versions, we need to check the signature - # of the `merge` method to see if it supports the `adapter_names` argument. - supported_merge_kwargs = list(inspect.signature(module.merge).parameters) - if "adapter_names" in supported_merge_kwargs: - merge_kwargs["adapter_names"] = adapter_names - elif "adapter_names" not in supported_merge_kwargs and adapter_names is not None: - raise ValueError( - "The `adapter_names` argument is not supported with your PEFT version. " - "Please upgrade to the latest version of PEFT. `pip install -U peft`" - ) - - module.merge(**merge_kwargs) - - -def unfuse_text_encoder_lora(text_encoder): - """ - Unfuses LoRAs for the text encoder. - - Args: - text_encoder (`torch.nn.Module`): - The text encoder module to set the adapter layers for. If `None`, it will try to get the `text_encoder` - attribute. - """ - for module in text_encoder.modules(): - if isinstance(module, BaseTunerLayer): - module.unmerge() - - -def set_adapters_for_text_encoder( - adapter_names: list[str] | str, - text_encoder: "PreTrainedModel" | None = None, # noqa: F821 - text_encoder_weights: float | list[float] | list[None] | None = None, -): - """ - Sets the adapter layers for the text encoder. - - Args: - adapter_names (`list[str]` or `str`): - The names of the adapters to use. - text_encoder (`torch.nn.Module`, *optional*): - The text encoder module to set the adapter layers for. If `None`, it will try to get the `text_encoder` - attribute. - text_encoder_weights (`list[float]`, *optional*): - The weights to use for the text encoder. If `None`, the weights are set to `1.0` for all the adapters. - """ - if text_encoder is None: - raise ValueError( - "The pipeline does not have a default `pipe.text_encoder` class. Please make sure to pass a `text_encoder` instead." - ) - - def process_weights(adapter_names, weights): - # Expand weights into a list, one entry per adapter - # e.g. for 2 adapters: 7 -> [7,7] ; [3, None] -> [3, None] - if not isinstance(weights, list): - weights = [weights] * len(adapter_names) - - if len(adapter_names) != len(weights): - raise ValueError( - f"Length of adapter names {len(adapter_names)} is not equal to the length of the weights {len(weights)}" - ) - - # Set None values to default of 1.0 - # e.g. [7,7] -> [7,7] ; [3, None] -> [3,1] - weights = [w if w is not None else 1.0 for w in weights] - - return weights - - adapter_names = [adapter_names] if isinstance(adapter_names, str) else adapter_names - text_encoder_weights = process_weights(adapter_names, text_encoder_weights) - set_weights_and_activate_adapters(text_encoder, adapter_names, text_encoder_weights) - - -def disable_lora_for_text_encoder(text_encoder: "PreTrainedModel" | None = None): - """ - Disables the LoRA layers for the text encoder. - - Args: - text_encoder (`torch.nn.Module`, *optional*): - The text encoder module to disable the LoRA layers for. If `None`, it will try to get the `text_encoder` - attribute. - """ - if text_encoder is None: - raise ValueError("Text Encoder not found.") - set_adapter_layers(text_encoder, enabled=False) - - -def enable_lora_for_text_encoder(text_encoder: "PreTrainedModel" | None = None): - """ - Enables the LoRA layers for the text encoder. - - Args: - text_encoder (`torch.nn.Module`, *optional*): - The text encoder module to enable the LoRA layers for. If `None`, it will try to get the `text_encoder` - attribute. - """ - if text_encoder is None: - raise ValueError("Text Encoder not found.") - set_adapter_layers(text_encoder, enabled=True) - - -def _remove_text_encoder_monkey_patch(text_encoder): - recurse_remove_peft_layers(text_encoder) - if getattr(text_encoder, "peft_config", None) is not None: - del text_encoder.peft_config - text_encoder._hf_peft_config_loaded = None - - -def _fetch_state_dict( - pretrained_model_name_or_path_or_dict, - weight_name, - use_safetensors, - local_files_only, - cache_dir, - force_download, - proxies, - token, - revision, - subfolder, - user_agent, - allow_pickle, - metadata=None, -): - model_file = None - if not isinstance(pretrained_model_name_or_path_or_dict, dict): - # Let's first try to load .safetensors weights - if (use_safetensors and weight_name is None) or ( - weight_name is not None and weight_name.endswith(".safetensors") - ): - try: - # Here we're relaxing the loading check to enable more Inference API - # friendliness where sometimes, it's not at all possible to automatically - # determine `weight_name`. - if weight_name is None: - weight_name = _best_guess_weight_name( - pretrained_model_name_or_path_or_dict, - file_extension=".safetensors", - local_files_only=local_files_only, - ) - model_file = _get_model_file( - pretrained_model_name_or_path_or_dict, - weights_name=weight_name or LORA_WEIGHT_NAME_SAFE, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - ) - state_dict = safetensors.torch.load_file(model_file, device="cpu") - metadata = _load_sft_state_dict_metadata(model_file) - - except (IOError, safetensors.SafetensorError) as e: - if not allow_pickle: - raise e - # try loading non-safetensors weights - model_file = None - metadata = None - pass - - if model_file is None: - if weight_name is None: - weight_name = _best_guess_weight_name( - pretrained_model_name_or_path_or_dict, file_extension=".bin", local_files_only=local_files_only - ) - model_file = _get_model_file( - pretrained_model_name_or_path_or_dict, - weights_name=weight_name or LORA_WEIGHT_NAME, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - ) - state_dict = load_state_dict(model_file) - metadata = None - else: - state_dict = pretrained_model_name_or_path_or_dict - - return state_dict, metadata - - -def _best_guess_weight_name( - pretrained_model_name_or_path_or_dict, file_extension=".safetensors", local_files_only=False -): - targeted_files = [] - - if os.path.isfile(pretrained_model_name_or_path_or_dict): - return - elif os.path.isdir(pretrained_model_name_or_path_or_dict): - targeted_files = [f for f in os.listdir(pretrained_model_name_or_path_or_dict) if f.endswith(file_extension)] - elif local_files_only or HF_HUB_OFFLINE: - raise ValueError("When using the offline mode, you must specify a `weight_name`.") - else: - files_in_repo = model_info(pretrained_model_name_or_path_or_dict).siblings - targeted_files = [f.rfilename for f in files_in_repo if f.rfilename.endswith(file_extension)] - if len(targeted_files) == 0: - return - - # "scheduler" does not correspond to a LoRA checkpoint. - # "optimizer" does not correspond to a LoRA checkpoint - # only top-level checkpoints are considered and not the other ones, hence "checkpoint". - unallowed_substrings = {"scheduler", "optimizer", "checkpoint"} - targeted_files = list( - filter(lambda x: all(substring not in x for substring in unallowed_substrings), targeted_files) - ) - - if any(f.endswith(LORA_WEIGHT_NAME) for f in targeted_files): - targeted_files = list(filter(lambda x: x.endswith(LORA_WEIGHT_NAME), targeted_files)) - elif any(f.endswith(LORA_WEIGHT_NAME_SAFE) for f in targeted_files): - targeted_files = list(filter(lambda x: x.endswith(LORA_WEIGHT_NAME_SAFE), targeted_files)) - - if len(targeted_files) > 1: - logger.warning( - f"Provided path contains more than one weights file in the {file_extension} format. `{targeted_files[0]}` is going to be loaded, for precise control, specify a `weight_name` in `load_lora_weights`." - ) - weight_name = targeted_files[0] - return weight_name - - -def _pack_dict_with_prefix(state_dict, prefix): - sd_with_prefix = {f"{prefix}.{key}": value for key, value in state_dict.items()} - return sd_with_prefix - - -def _load_lora_into_text_encoder( - state_dict, - network_alphas, - text_encoder, - prefix=None, - lora_scale=1.0, - text_encoder_name="text_encoder", - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, -): - from ..hooks.group_offloading import _maybe_remove_and_reapply_group_offloading - - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - if network_alphas and metadata: - raise ValueError("`network_alphas` and `metadata` cannot be specified both at the same time.") - - peft_kwargs = {} - if low_cpu_mem_usage: - if not is_peft_version(">=", "0.13.1"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - if not is_transformers_version(">", "4.45.2"): - # Note from sayakpaul: It's not in `transformers` stable yet. - # https://github.com/huggingface/transformers/pull/33725/ - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `transformers` version. Please update it with `pip install -U transformers`." - ) - peft_kwargs["low_cpu_mem_usage"] = low_cpu_mem_usage - - # If the serialization format is new (introduced in https://github.com/huggingface/diffusers/pull/2918), - # then the `state_dict` keys should have `unet_name` and/or `text_encoder_name` as - # their prefixes. - prefix = text_encoder_name if prefix is None else prefix - - # Safe prefix to check with. - if hotswap and any(text_encoder_name in key for key in state_dict.keys()): - raise ValueError("At the moment, hotswapping is not supported for text encoders, please pass `hotswap=False`.") - - # Load the layers corresponding to text encoder and make necessary adjustments. - if prefix is not None: - state_dict = {k.removeprefix(f"{prefix}."): v for k, v in state_dict.items() if k.startswith(f"{prefix}.")} - if metadata is not None: - metadata = {k.removeprefix(f"{prefix}."): v for k, v in metadata.items() if k.startswith(f"{prefix}.")} - - if len(state_dict) > 0: - logger.info(f"Loading {prefix}.") - rank = {} - state_dict = convert_state_dict_to_diffusers(state_dict) - - # convert state dict - state_dict = convert_state_dict_to_peft(state_dict) - - for name, _ in text_encoder.named_modules(): - if name.endswith((".q_proj", ".k_proj", ".v_proj", ".out_proj", ".fc1", ".fc2")): - rank_key = f"{name}.lora_B.weight" - if rank_key in state_dict: - rank[rank_key] = state_dict[rank_key].shape[1] - - if network_alphas is not None: - alpha_keys = [k for k in network_alphas.keys() if k.startswith(prefix) and k.split(".")[0] == prefix] - network_alphas = {k.removeprefix(f"{prefix}."): v for k, v in network_alphas.items() if k in alpha_keys} - - # create `LoraConfig` - lora_config = _create_lora_config(state_dict, network_alphas, metadata, rank, is_unet=False) - - # adapter_name - if adapter_name is None: - adapter_name = get_adapter_name(text_encoder) - - # - - if prefix is not None and not state_dict: - model_class_name = text_encoder.__class__.__name__ - logger.warning( - f"No LoRA keys associated to {model_class_name} found with the {prefix=}. " - "This is safe to ignore if LoRA state dict didn't originally have any " - f"{model_class_name} related params. You can also try specifying `prefix=None` " - "to resolve the warning. Otherwise, open an issue if you think it's unexpected: " - "https://github.com/huggingface/diffusers/issues/new" - ) - - -def _func_optionally_disable_offloading(_pipeline): - """ - Optionally removes offloading in case the pipeline has been already sequentially offloaded to CPU. - - Args: - _pipeline (`DiffusionPipeline`): - The pipeline to disable offloading for. - - Returns: - tuple: - A tuple indicating if `is_model_cpu_offload` or `is_sequential_cpu_offload` or `is_group_offload` is True. - """ - from ..hooks.group_offloading import _is_group_offload_enabled - - is_model_cpu_offload = False - is_sequential_cpu_offload = False - is_group_offload = False - - if _pipeline is not None and _pipeline.hf_device_map is None: - for _, component in _pipeline.components.items(): - if not isinstance(component, nn.Module): - continue - is_group_offload = is_group_offload or _is_group_offload_enabled(component) - if not hasattr(component, "_hf_hook"): - continue - is_model_cpu_offload = is_model_cpu_offload or isinstance(component._hf_hook, CpuOffload) - is_sequential_cpu_offload = is_sequential_cpu_offload or ( - isinstance(component._hf_hook, AlignDevicesHook) - or hasattr(component._hf_hook, "hooks") - and isinstance(component._hf_hook.hooks[0], AlignDevicesHook) - ) - - if is_sequential_cpu_offload or is_model_cpu_offload: - logger.info( - "Accelerate hooks detected. Since you have called `load_lora_weights()`, the previous hooks will be first removed. Then the LoRA parameters will be loaded and the hooks will be applied again." - ) - for _, component in _pipeline.components.items(): - if not isinstance(component, nn.Module) or not hasattr(component, "_hf_hook"): - continue - remove_hook_from_module(component, recurse=is_sequential_cpu_offload) - - return (is_model_cpu_offload, is_sequential_cpu_offload, is_group_offload) - - -class LoraBaseMixin: - """Utility class for handling LoRAs.""" - - _lora_loadable_modules = [] - _merged_adapters = set() - - @property - def lora_scale(self) -> float: - """ - Returns the lora scale which can be set at run time by the pipeline. # if `_lora_scale` has not been set, - return 1. - """ - return self._lora_scale if hasattr(self, "_lora_scale") else 1.0 - - @property - def num_fused_loras(self): - """Returns the number of LoRAs that have been fused.""" - return len(self._merged_adapters) - - @property - def fused_loras(self): - """Returns names of the LoRAs that have been fused.""" - return self._merged_adapters - - def load_lora_weights(self, **kwargs): - raise NotImplementedError("`load_lora_weights()` is not implemented.") - - @classmethod - def save_lora_weights(cls, **kwargs): - raise NotImplementedError("`save_lora_weights()` not implemented.") - - @classmethod - def lora_state_dict(cls, **kwargs): - raise NotImplementedError("`lora_state_dict()` is not implemented.") - - def unload_lora_weights(self): - """ - Unloads the LoRA parameters. - - Examples: - - ```python - >>> # Assuming `pipeline` is already loaded with the LoRA parameters. - >>> pipeline.unload_lora_weights() - >>> ... - ``` - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - for component in self._lora_loadable_modules: - model = getattr(self, component, None) - if model is not None: - if issubclass(model.__class__, ModelMixin): - model.unload_lora() - elif issubclass(model.__class__, PreTrainedModel): - _remove_text_encoder_monkey_patch(model) - - def fuse_lora( - self, - components: list[str] | None = None, - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - Fuses the LoRA parameters into the original parameters of the corresponding blocks. - - Args: - components: (`list[str]`): list of LoRA-injectable components to fuse the LoRAs into. - lora_scale (`float`, defaults to 1.0): - Controls how much to influence the outputs with the LoRA parameters. - safe_fusing (`bool`, defaults to `False`): - Whether to check fused weights for NaN values before fusing and if values are NaN not fusing them. - adapter_names (`list[str]`, *optional*): - Adapter names to be used for fusing. If nothing is passed, all active adapters will be fused. - - Example: - - ```py - from diffusers import DiffusionPipeline - import torch - - pipeline = DiffusionPipeline.from_pretrained( - "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16 - ).to("cuda") - pipeline.load_lora_weights("nerijs/pixel-art-xl", weight_name="pixel-art-xl.safetensors", adapter_name="pixel") - pipeline.fuse_lora(lora_scale=0.7) - ``` - """ - if components is None: - components = [] - - if "fuse_unet" in kwargs: - depr_message = "Passing `fuse_unet` to `fuse_lora()` is deprecated and will be ignored. Please use the `components` argument and provide a list of the components whose LoRAs are to be fused. `fuse_unet` will be removed in a future version." - deprecate( - "fuse_unet", - "1.0.0", - depr_message, - ) - if "fuse_transformer" in kwargs: - depr_message = "Passing `fuse_transformer` to `fuse_lora()` is deprecated and will be ignored. Please use the `components` argument and provide a list of the components whose LoRAs are to be fused. `fuse_transformer` will be removed in a future version." - deprecate( - "fuse_transformer", - "1.0.0", - depr_message, - ) - if "fuse_text_encoder" in kwargs: - depr_message = "Passing `fuse_text_encoder` to `fuse_lora()` is deprecated and will be ignored. Please use the `components` argument and provide a list of the components whose LoRAs are to be fused. `fuse_text_encoder` will be removed in a future version." - deprecate( - "fuse_text_encoder", - "1.0.0", - depr_message, - ) - - if len(components) == 0: - raise ValueError("`components` cannot be an empty list.") - - # Need to retrieve the names as `adapter_names` can be None. So we cannot directly use it - # in `self._merged_adapters = self._merged_adapters | merged_adapter_names`. - merged_adapter_names = set() - for fuse_component in components: - if fuse_component not in self._lora_loadable_modules: - raise ValueError(f"{fuse_component} is not found in {self._lora_loadable_modules=}.") - - model = getattr(self, fuse_component, None) - if model is not None: - # check if diffusers model - if issubclass(model.__class__, ModelMixin): - model.fuse_lora(lora_scale, safe_fusing=safe_fusing, adapter_names=adapter_names) - for module in model.modules(): - if isinstance(module, BaseTunerLayer): - merged_adapter_names.update(set(module.merged_adapters)) - # handle transformers models. - if issubclass(model.__class__, PreTrainedModel): - fuse_text_encoder_lora( - model, lora_scale=lora_scale, safe_fusing=safe_fusing, adapter_names=adapter_names - ) - for module in model.modules(): - if isinstance(module, BaseTunerLayer): - merged_adapter_names.update(set(module.merged_adapters)) - - self._merged_adapters = self._merged_adapters | merged_adapter_names - - def unfuse_lora(self, components: list[str] | None = None, **kwargs): - r""" - Reverses the effect of - [`pipe.fuse_lora()`](https://huggingface.co/docs/diffusers/main/en/api/loaders#diffusers.loaders.LoraBaseMixin.fuse_lora). - - Args: - components (`list[str]`): list of LoRA-injectable components to unfuse LoRA from. - unfuse_unet (`bool`, defaults to `True`): Whether to unfuse the UNet LoRA parameters. - unfuse_text_encoder (`bool`, defaults to `True`): - Whether to unfuse the text encoder LoRA parameters. If the text encoder wasn't monkey-patched with the - LoRA parameters then it won't have any effect. - """ - if components is None: - components = [] - - if "unfuse_unet" in kwargs: - depr_message = "Passing `unfuse_unet` to `unfuse_lora()` is deprecated and will be ignored. Please use the `components` argument. `unfuse_unet` will be removed in a future version." - deprecate( - "unfuse_unet", - "1.0.0", - depr_message, - ) - if "unfuse_transformer" in kwargs: - depr_message = "Passing `unfuse_transformer` to `unfuse_lora()` is deprecated and will be ignored. Please use the `components` argument. `unfuse_transformer` will be removed in a future version." - deprecate( - "unfuse_transformer", - "1.0.0", - depr_message, - ) - if "unfuse_text_encoder" in kwargs: - depr_message = "Passing `unfuse_text_encoder` to `unfuse_lora()` is deprecated and will be ignored. Please use the `components` argument. `unfuse_text_encoder` will be removed in a future version." - deprecate( - "unfuse_text_encoder", - "1.0.0", - depr_message, - ) - - if len(components) == 0: - raise ValueError("`components` cannot be an empty list.") - - for fuse_component in components: - if fuse_component not in self._lora_loadable_modules: - raise ValueError(f"{fuse_component} is not found in {self._lora_loadable_modules=}.") - - model = getattr(self, fuse_component, None) - if model is not None: - if issubclass(model.__class__, (ModelMixin, PreTrainedModel)): - for module in model.modules(): - if isinstance(module, BaseTunerLayer): - for adapter in set(module.merged_adapters): - if adapter and adapter in self._merged_adapters: - self._merged_adapters = self._merged_adapters - {adapter} - module.unmerge() - - def set_adapters( - self, - adapter_names: list[str] | str, - adapter_weights: float | dict | list[float] | list[dict] | None = None, - ): - """ - Set the currently active adapters for use in the pipeline. - - Args: - adapter_names (`list[str]` or `str`): - The names of the adapters to use. - adapter_weights (`list[float, float]`, *optional*): - The adapter(s) weights to use with the UNet. If `None`, the weights are set to `1.0` for all the - adapters. - - Example: - - ```py - from diffusers import AutoPipelineForText2Image - import torch - - pipeline = AutoPipelineForText2Image.from_pretrained( - "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16 - ).to("cuda") - pipeline.load_lora_weights( - "jbilcke-hf/sdxl-cinematic-1", weight_name="pytorch_lora_weights.safetensors", adapter_name="cinematic" - ) - pipeline.load_lora_weights("nerijs/pixel-art-xl", weight_name="pixel-art-xl.safetensors", adapter_name="pixel") - pipeline.set_adapters(["cinematic", "pixel"], adapter_weights=[0.5, 0.5]) - ``` - """ - if isinstance(adapter_weights, dict): - components_passed = set(adapter_weights.keys()) - lora_components = set(self._lora_loadable_modules) - - invalid_components = sorted(components_passed - lora_components) - if invalid_components: - logger.warning( - f"The following components in `adapter_weights` are not part of the pipeline: {invalid_components}. " - f"Available components that are LoRA-compatible: {self._lora_loadable_modules}. So, weights belonging " - "to the invalid components will be removed and ignored." - ) - adapter_weights = {k: v for k, v in adapter_weights.items() if k not in invalid_components} - - adapter_names = [adapter_names] if isinstance(adapter_names, str) else adapter_names - adapter_weights = copy.deepcopy(adapter_weights) - - # Expand weights into a list, one entry per adapter - if not isinstance(adapter_weights, list): - adapter_weights = [adapter_weights] * len(adapter_names) - - if len(adapter_names) != len(adapter_weights): - raise ValueError( - f"Length of adapter names {len(adapter_names)} is not equal to the length of the weights {len(adapter_weights)}" - ) - - list_adapters = self.get_list_adapters() # eg {"unet": ["adapter1", "adapter2"], "text_encoder": ["adapter2"]} - # eg ["adapter1", "adapter2"] - all_adapters = {adapter for adapters in list_adapters.values() for adapter in adapters} - missing_adapters = set(adapter_names) - all_adapters - if len(missing_adapters) > 0: - raise ValueError( - f"Adapter name(s) {missing_adapters} not in the list of present adapters: {all_adapters}." - ) - - # eg {"adapter1": ["unet"], "adapter2": ["unet", "text_encoder"]} - invert_list_adapters = { - adapter: [part for part, adapters in list_adapters.items() if adapter in adapters] - for adapter in all_adapters - } - - # Decompose weights into weights for denoiser and text encoders. - _component_adapter_weights = {} - for component in self._lora_loadable_modules: - model = getattr(self, component, None) - # To guard for cases like Wan. In Wan2.1 and WanVace, we have a single denoiser. - # Whereas in Wan 2.2, we have two denoisers. - if model is None: - continue - - for adapter_name, weights in zip(adapter_names, adapter_weights): - if isinstance(weights, dict): - component_adapter_weights = weights.pop(component, None) - if component_adapter_weights is not None and component not in invert_list_adapters[adapter_name]: - logger.warning( - ( - f"Lora weight dict for adapter '{adapter_name}' contains {component}," - f"but this will be ignored because {adapter_name} does not contain weights for {component}." - f"Valid parts for {adapter_name} are: {invert_list_adapters[adapter_name]}." - ) - ) - - else: - component_adapter_weights = weights - - _component_adapter_weights.setdefault(component, []) - _component_adapter_weights[component].append(component_adapter_weights) - - if issubclass(model.__class__, ModelMixin): - model.set_adapters(adapter_names, _component_adapter_weights[component]) - elif issubclass(model.__class__, PreTrainedModel): - set_adapters_for_text_encoder(adapter_names, model, _component_adapter_weights[component]) - - def disable_lora(self): - """ - Disables the active LoRA layers of the pipeline. - - Example: - - ```py - from diffusers import AutoPipelineForText2Image - import torch - - pipeline = AutoPipelineForText2Image.from_pretrained( - "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16 - ).to("cuda") - pipeline.load_lora_weights( - "jbilcke-hf/sdxl-cinematic-1", weight_name="pytorch_lora_weights.safetensors", adapter_name="cinematic" - ) - pipeline.disable_lora() - ``` - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - for component in self._lora_loadable_modules: - model = getattr(self, component, None) - if model is not None: - if issubclass(model.__class__, ModelMixin): - model.disable_lora() - elif issubclass(model.__class__, PreTrainedModel): - disable_lora_for_text_encoder(model) - - def enable_lora(self): - """ - Enables the active LoRA layers of the pipeline. - - Example: - - ```py - from diffusers import AutoPipelineForText2Image - import torch - - pipeline = AutoPipelineForText2Image.from_pretrained( - "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16 - ).to("cuda") - pipeline.load_lora_weights( - "jbilcke-hf/sdxl-cinematic-1", weight_name="pytorch_lora_weights.safetensors", adapter_name="cinematic" - ) - pipeline.enable_lora() - ``` - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - for component in self._lora_loadable_modules: - model = getattr(self, component, None) - if model is not None: - if issubclass(model.__class__, ModelMixin): - model.enable_lora() - elif issubclass(model.__class__, PreTrainedModel): - enable_lora_for_text_encoder(model) - - def delete_adapters(self, adapter_names: list[str] | str): - """ - Delete an adapter's LoRA layers from the pipeline. - - Args: - adapter_names (`list[str, str]`): - The names of the adapters to delete. - - Example: - - ```py - from diffusers import AutoPipelineForText2Image - import torch - - pipeline = AutoPipelineForText2Image.from_pretrained( - "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16 - ).to("cuda") - pipeline.load_lora_weights( - "jbilcke-hf/sdxl-cinematic-1", weight_name="pytorch_lora_weights.safetensors", adapter_names="cinematic" - ) - pipeline.delete_adapters("cinematic") - ``` - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - if isinstance(adapter_names, str): - adapter_names = [adapter_names] - - for component in self._lora_loadable_modules: - model = getattr(self, component, None) - if model is not None: - if issubclass(model.__class__, ModelMixin): - model.delete_adapters(adapter_names) - elif issubclass(model.__class__, PreTrainedModel): - for adapter_name in adapter_names: - delete_adapter_layers(model, adapter_name) - - def get_active_adapters(self) -> list[str]: - """ - Gets the list of the current active adapters. - - Example: - - ```python - from diffusers import DiffusionPipeline - - pipeline = DiffusionPipeline.from_pretrained( - "stabilityai/stable-diffusion-xl-base-1.0", - ).to("cuda") - pipeline.load_lora_weights("CiroN2022/toy-face", weight_name="toy_face_sdxl.safetensors", adapter_name="toy") - pipeline.get_active_adapters() - ``` - """ - if not USE_PEFT_BACKEND: - raise ValueError( - "PEFT backend is required for this method. Please install the latest version of PEFT `pip install -U peft`" - ) - - active_adapters = [] - - for component in self._lora_loadable_modules: - model = getattr(self, component, None) - if model is not None and issubclass(model.__class__, ModelMixin): - for module in model.modules(): - if isinstance(module, BaseTunerLayer): - active_adapters = module.active_adapters - break - - return active_adapters - - def get_list_adapters(self) -> dict[str, list[str]]: - """ - Gets the current list of all available adapters in the pipeline. - """ - if not USE_PEFT_BACKEND: - raise ValueError( - "PEFT backend is required for this method. Please install the latest version of PEFT `pip install -U peft`" - ) - - set_adapters = {} - - for component in self._lora_loadable_modules: - model = getattr(self, component, None) - if ( - model is not None - and issubclass(model.__class__, (ModelMixin, PreTrainedModel)) - and hasattr(model, "peft_config") - ): - set_adapters[component] = list(model.peft_config.keys()) - - return set_adapters - - def set_lora_device(self, adapter_names: list[str], device: torch.device | str | int) -> None: - """ - Moves the LoRAs listed in `adapter_names` to a target device. Useful for offloading the LoRA to the CPU in case - you want to load multiple adapters and free some GPU memory. - - After offloading the LoRA adapters to CPU, as long as the rest of the model is still on GPU, the LoRA adapters - can no longer be used for inference, as that would cause a device mismatch. Remember to set the device back to - GPU before using those LoRA adapters for inference. - - ```python - >>> pipe.load_lora_weights(path_1, adapter_name="adapter-1") - >>> pipe.load_lora_weights(path_2, adapter_name="adapter-2") - >>> pipe.set_adapters("adapter-1") - >>> image_1 = pipe(**kwargs) - >>> # switch to adapter-2, offload adapter-1 - >>> pipeline.set_lora_device(adapter_names=["adapter-1"], device="cpu") - >>> pipeline.set_lora_device(adapter_names=["adapter-2"], device="cuda:0") - >>> pipe.set_adapters("adapter-2") - >>> image_2 = pipe(**kwargs) - >>> # switch back to adapter-1, offload adapter-2 - >>> pipeline.set_lora_device(adapter_names=["adapter-2"], device="cpu") - >>> pipeline.set_lora_device(adapter_names=["adapter-1"], device="cuda:0") - >>> pipe.set_adapters("adapter-1") - >>> ... - ``` - - Args: - adapter_names (`list[str]`): - list of adapters to send device to. - device (`torch.device | str | int`): - Device to send the adapters to. Can be either a torch device, a str or an integer. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - for component in self._lora_loadable_modules: - model = getattr(self, component, None) - if model is not None: - for module in model.modules(): - if isinstance(module, BaseTunerLayer): - for adapter_name in adapter_names: - if adapter_name not in module.lora_A: - # it is sufficient to check lora_A - continue - - module.lora_A[adapter_name].to(device) - module.lora_B[adapter_name].to(device) - # this is a param, not a module, so device placement is not in-place -> re-assign - if hasattr(module, "lora_magnitude_vector") and module.lora_magnitude_vector is not None: - if adapter_name in module.lora_magnitude_vector: - module.lora_magnitude_vector[adapter_name] = module.lora_magnitude_vector[ - adapter_name - ].to(device) - - def enable_lora_hotswap(self, **kwargs) -> None: - """ - Hotswap adapters without triggering recompilation of a model or if the ranks of the loaded adapters are - different. - - Args: - target_rank (`int`): - The highest rank among all the adapters that will be loaded. - check_compiled (`str`, *optional*, defaults to `"error"`): - How to handle a model that is already compiled. The check can return the following messages: - - "error" (default): raise an error - - "warn": issue a warning - - "ignore": do nothing - """ - for key, component in self.components.items(): - if hasattr(component, "enable_lora_hotswap") and (key in self._lora_loadable_modules): - component.enable_lora_hotswap(**kwargs) - - @staticmethod - def pack_weights(layers, prefix): - layers_weights = layers.state_dict() if isinstance(layers, torch.nn.Module) else layers - return _pack_dict_with_prefix(layers_weights, prefix) - - @staticmethod - def write_lora_layers( - state_dict: dict[str, torch.Tensor], - save_directory: str, - is_main_process: bool, - weight_name: str, - save_function: Callable, - safe_serialization: bool, - lora_adapter_metadata: dict | None = None, - ): - """Writes the state dict of the LoRA layers (optionally with metadata) to disk.""" - if os.path.isfile(save_directory): - logger.error(f"Provided path ({save_directory}) should be a directory, not a file") - return - - if lora_adapter_metadata and not safe_serialization: - raise ValueError("`lora_adapter_metadata` cannot be specified when not using `safe_serialization`.") - if lora_adapter_metadata and not isinstance(lora_adapter_metadata, dict): - raise TypeError("`lora_adapter_metadata` must be of type `dict`.") - - if save_function is None: - if safe_serialization: - - def save_function(weights, filename): - # Inject framework format. - metadata = {"format": "pt"} - if lora_adapter_metadata: - for key, value in lora_adapter_metadata.items(): - if isinstance(value, set): - lora_adapter_metadata[key] = list(value) - metadata[LORA_ADAPTER_METADATA_KEY] = json.dumps( - lora_adapter_metadata, indent=2, sort_keys=True - ) - - return safetensors.torch.save_file(weights, filename, metadata=metadata) - - else: - save_function = torch.save - - os.makedirs(save_directory, exist_ok=True) - - if weight_name is None: - if safe_serialization: - weight_name = LORA_WEIGHT_NAME_SAFE - else: - weight_name = LORA_WEIGHT_NAME - - save_path = Path(save_directory, weight_name).as_posix() - save_function(state_dict, save_path) - logger.info(f"Model weights saved in {save_path}") - - @classmethod - def _save_lora_weights( - cls, - save_directory: str | os.PathLike, - lora_layers: dict[str, dict[str, torch.nn.Module | torch.Tensor]], - lora_metadata: dict[str, dict | None], - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - ): - """ - Helper method to pack and save LoRA weights and metadata. This method centralizes the saving logic for all - pipeline types. - """ - state_dict = {} - final_lora_adapter_metadata = {} - - for prefix, layers in lora_layers.items(): - state_dict.update(cls.pack_weights(layers, prefix)) - - for prefix, metadata in lora_metadata.items(): - if metadata: - final_lora_adapter_metadata.update(_pack_dict_with_prefix(metadata, prefix)) - - cls.write_lora_layers( - state_dict=state_dict, - save_directory=save_directory, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - lora_adapter_metadata=final_lora_adapter_metadata if final_lora_adapter_metadata else None, - ) - - @classmethod - def _optionally_disable_offloading(cls, _pipeline): - return _func_optionally_disable_offloading(_pipeline=_pipeline) diff --git a/diffusers/loaders/lora_conversion_utils.py b/diffusers/loaders/lora_conversion_utils.py deleted file mode 100644 index 07e3351685e8ff424e72e7ffb8bd4e73e708b18c..0000000000000000000000000000000000000000 --- a/diffusers/loaders/lora_conversion_utils.py +++ /dev/null @@ -1,3124 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import re - -import torch - -from ..utils import is_peft_version, logging, state_dict_all_zero - - -logger = logging.get_logger(__name__) - - -def swap_scale_shift(weight): - shift, scale = weight.chunk(2, dim=0) - new_weight = torch.cat([scale, shift], dim=0) - return new_weight - - -def _maybe_map_sgm_blocks_to_diffusers(state_dict, unet_config, delimiter="_", block_slice_pos=5): - # 1. get all state_dict_keys - all_keys = list(state_dict.keys()) - sgm_patterns = ["input_blocks", "middle_block", "output_blocks"] - not_sgm_patterns = ["down_blocks", "mid_block", "up_blocks"] - - # check if state_dict contains both patterns - contains_sgm_patterns = False - contains_not_sgm_patterns = False - for key in all_keys: - if any(p in key for p in sgm_patterns): - contains_sgm_patterns = True - elif any(p in key for p in not_sgm_patterns): - contains_not_sgm_patterns = True - - # if state_dict contains both patterns, remove sgm - # we can then return state_dict immediately - if contains_sgm_patterns and contains_not_sgm_patterns: - for key in all_keys: - if any(p in key for p in sgm_patterns): - state_dict.pop(key) - return state_dict - - # 2. check if needs remapping, if not return original dict - is_in_sgm_format = False - for key in all_keys: - if any(p in key for p in sgm_patterns): - is_in_sgm_format = True - break - - if not is_in_sgm_format: - return state_dict - - # 3. Else remap from SGM patterns - new_state_dict = {} - inner_block_map = ["resnets", "attentions", "upsamplers"] - - # Retrieves # of down, mid and up blocks - input_block_ids, middle_block_ids, output_block_ids = set(), set(), set() - - for layer in all_keys: - if "text" in layer: - new_state_dict[layer] = state_dict.pop(layer) - elif not any(p in layer for p in sgm_patterns) or f"input_blocks{delimiter}0{delimiter}0" in layer: - # SDXL's sgm UNet has modules outside the input/middle/output block structure that - # _convert_unet_lora_key maps directly: time_embed, label_emb, out (out.2 = conv_out) - # and input_blocks.0.0 (= conv_in). Pass these through instead of block-remapping - # (conv_in's input_blocks.0 would otherwise be parsed as a down-block) or raising. - new_state_dict[layer] = state_dict.pop(layer) - else: - layer_id = int(layer.split(delimiter)[:block_slice_pos][-1]) - if sgm_patterns[0] in layer: - input_block_ids.add(layer_id) - elif sgm_patterns[1] in layer: - middle_block_ids.add(layer_id) - elif sgm_patterns[2] in layer: - output_block_ids.add(layer_id) - else: - raise ValueError(f"Checkpoint not supported because layer {layer} not supported.") - - input_blocks = { - layer_id: [key for key in state_dict if f"input_blocks{delimiter}{layer_id}" in key] - for layer_id in input_block_ids - } - middle_blocks = { - layer_id: [key for key in state_dict if f"middle_block{delimiter}{layer_id}" in key] - for layer_id in middle_block_ids - } - output_blocks = { - layer_id: [key for key in state_dict if f"output_blocks{delimiter}{layer_id}" in key] - for layer_id in output_block_ids - } - - # Rename keys accordingly - for i in input_block_ids: - block_id = (i - 1) // (unet_config.layers_per_block + 1) - layer_in_block_id = (i - 1) % (unet_config.layers_per_block + 1) - - for key in input_blocks[i]: - inner_block_id = int(key.split(delimiter)[block_slice_pos]) - inner_block_key = inner_block_map[inner_block_id] if "op" not in key else "downsamplers" - inner_layers_in_block = str(layer_in_block_id) if "op" not in key else "0" - new_key = delimiter.join( - key.split(delimiter)[: block_slice_pos - 1] - + [str(block_id), inner_block_key, inner_layers_in_block] - + key.split(delimiter)[block_slice_pos + 1 :] - ) - new_state_dict[new_key] = state_dict.pop(key) - - for i in middle_block_ids: - key_part = None - if i == 0: - key_part = [inner_block_map[0], "0"] - elif i == 1: - key_part = [inner_block_map[1], "0"] - elif i == 2: - key_part = [inner_block_map[0], "1"] - else: - raise ValueError(f"Invalid middle block id {i}.") - - for key in middle_blocks[i]: - new_key = delimiter.join( - key.split(delimiter)[: block_slice_pos - 1] + key_part + key.split(delimiter)[block_slice_pos:] - ) - new_state_dict[new_key] = state_dict.pop(key) - - for i in output_block_ids: - block_id = i // (unet_config.layers_per_block + 1) - layer_in_block_id = i % (unet_config.layers_per_block + 1) - - for key in output_blocks[i]: - inner_block_id = int(key.split(delimiter)[block_slice_pos]) - inner_block_key = inner_block_map[inner_block_id] - inner_layers_in_block = str(layer_in_block_id) if inner_block_id < 2 else "0" - new_key = delimiter.join( - key.split(delimiter)[: block_slice_pos - 1] - + [str(block_id), inner_block_key, inner_layers_in_block] - + key.split(delimiter)[block_slice_pos + 1 :] - ) - new_state_dict[new_key] = state_dict.pop(key) - - if state_dict: - raise ValueError("At this point all state dict entries have to be converted.") - - return new_state_dict - - -def _convert_non_diffusers_lora_to_diffusers(state_dict, unet_name="unet", text_encoder_name="text_encoder"): - """ - Converts a non-Diffusers LoRA state dict to a Diffusers compatible state dict. - - Args: - state_dict (`dict`): The state dict to convert. - unet_name (`str`, optional): The name of the U-Net module in the Diffusers model. Defaults to "unet". - text_encoder_name (`str`, optional): The name of the text encoder module in the Diffusers model. Defaults to - "text_encoder". - - Returns: - `tuple`: A tuple containing the converted state dict and a dictionary of alphas. - """ - unet_state_dict = {} - te_state_dict = {} - te2_state_dict = {} - network_alphas = {} - - # Check for DoRA-enabled LoRAs. - dora_present_in_unet = any("dora_scale" in k and "lora_unet_" in k for k in state_dict) - dora_present_in_te = any("dora_scale" in k and ("lora_te_" in k or "lora_te1_" in k) for k in state_dict) - dora_present_in_te2 = any("dora_scale" in k and "lora_te2_" in k for k in state_dict) - if dora_present_in_unet or dora_present_in_te or dora_present_in_te2: - if is_peft_version("<", "0.9.0"): - raise ValueError( - "You need `peft` 0.9.0 at least to use DoRA-enabled LoRAs. Please upgrade your installation of `peft`." - ) - - # Iterate over all LoRA weights. - all_lora_keys = list(state_dict.keys()) - for key in all_lora_keys: - if not key.endswith("lora_down.weight"): - continue - - # Extract LoRA name. - lora_name = key.split(".")[0] - - # Find corresponding up weight and alpha. - lora_name_up = lora_name + ".lora_up.weight" - lora_name_alpha = lora_name + ".alpha" - - # Handle U-Net LoRAs. - if lora_name.startswith("lora_unet_"): - diffusers_name = _convert_unet_lora_key(key) - - # Store down and up weights. - unet_state_dict[diffusers_name] = state_dict.pop(key) - unet_state_dict[diffusers_name.replace(".down.", ".up.")] = state_dict.pop(lora_name_up) - - # Store DoRA scale if present. - if dora_present_in_unet: - dora_scale_key_to_replace = "_lora.down." if "_lora.down." in diffusers_name else ".lora.down." - unet_state_dict[diffusers_name.replace(dora_scale_key_to_replace, ".lora_magnitude_vector.")] = ( - state_dict.pop(key.replace("lora_down.weight", "dora_scale")) - ) - - # Handle text encoder LoRAs. - elif lora_name.startswith(("lora_te_", "lora_te1_", "lora_te2_")): - diffusers_name = _convert_text_encoder_lora_key(key, lora_name) - - # Store down and up weights for te or te2. - if lora_name.startswith(("lora_te_", "lora_te1_")): - te_state_dict[diffusers_name] = state_dict.pop(key) - te_state_dict[diffusers_name.replace(".down.", ".up.")] = state_dict.pop(lora_name_up) - else: - te2_state_dict[diffusers_name] = state_dict.pop(key) - te2_state_dict[diffusers_name.replace(".down.", ".up.")] = state_dict.pop(lora_name_up) - - # Store DoRA scale if present. - if dora_present_in_te or dora_present_in_te2: - dora_scale_key_to_replace_te = ( - "_lora.down." if "_lora.down." in diffusers_name else ".lora_linear_layer." - ) - if lora_name.startswith(("lora_te_", "lora_te1_")): - te_state_dict[diffusers_name.replace(dora_scale_key_to_replace_te, ".lora_magnitude_vector.")] = ( - state_dict.pop(key.replace("lora_down.weight", "dora_scale")) - ) - elif lora_name.startswith("lora_te2_"): - te2_state_dict[diffusers_name.replace(dora_scale_key_to_replace_te, ".lora_magnitude_vector.")] = ( - state_dict.pop(key.replace("lora_down.weight", "dora_scale")) - ) - - # Store alpha if present. - if lora_name_alpha in state_dict: - alpha = state_dict.pop(lora_name_alpha).item() - network_alphas.update(_get_alpha_name(lora_name_alpha, diffusers_name, alpha)) - - # Check if any keys remain. - if len(state_dict) > 0: - raise ValueError(f"The following keys have not been correctly renamed: \n\n {', '.join(state_dict.keys())}") - - logger.info("Non-diffusers checkpoint detected.") - - # Construct final state dict. - unet_state_dict = {f"{unet_name}.{module_name}": params for module_name, params in unet_state_dict.items()} - te_state_dict = {f"{text_encoder_name}.{module_name}": params for module_name, params in te_state_dict.items()} - te2_state_dict = ( - {f"text_encoder_2.{module_name}": params for module_name, params in te2_state_dict.items()} - if len(te2_state_dict) > 0 - else None - ) - if te2_state_dict is not None: - te_state_dict.update(te2_state_dict) - - new_state_dict = {**unet_state_dict, **te_state_dict} - return new_state_dict, network_alphas - - -def _convert_unet_lora_key(key): - """ - Converts a U-Net LoRA key to a Diffusers compatible key. - """ - diffusers_name = key.replace("lora_unet_", "").replace("_", ".") - - # kohya-ss trains SDXL on its own sgm/LDM UNet, so conv_in / conv_out arrive as - # input_blocks.0.0 / out.2. Map these before the block renames below, otherwise - # input_blocks.0.0 would become down_blocks.0.0 instead of conv_in. - diffusers_name = diffusers_name.replace("input.blocks.0.0", "conv_in") - diffusers_name = diffusers_name.replace("out.2", "conv_out") - - # Replace common U-Net naming patterns. - diffusers_name = diffusers_name.replace("input.blocks", "down_blocks") - diffusers_name = diffusers_name.replace("down.blocks", "down_blocks") - diffusers_name = diffusers_name.replace("middle.block", "mid_block") - diffusers_name = diffusers_name.replace("mid.block", "mid_block") - diffusers_name = diffusers_name.replace("output.blocks", "up_blocks") - diffusers_name = diffusers_name.replace("up.blocks", "up_blocks") - diffusers_name = diffusers_name.replace("transformer.blocks", "transformer_blocks") - diffusers_name = diffusers_name.replace("to.q.lora", "to_q_lora") - diffusers_name = diffusers_name.replace("to.k.lora", "to_k_lora") - diffusers_name = diffusers_name.replace("to.v.lora", "to_v_lora") - diffusers_name = diffusers_name.replace("to.out.0.lora", "to_out_lora") - diffusers_name = diffusers_name.replace("proj.in", "proj_in") - diffusers_name = diffusers_name.replace("proj.out", "proj_out") - diffusers_name = diffusers_name.replace("emb.layers", "time_emb_proj") - diffusers_name = diffusers_name.replace("conv.in", "conv_in") - diffusers_name = diffusers_name.replace("conv.out", "conv_out") - diffusers_name = diffusers_name.replace("time.embed.0", "time_embedding.linear_1") - diffusers_name = diffusers_name.replace("time.embed.2", "time_embedding.linear_2") - # sgm label_emb (SDXL added-conditioning MLP) -> diffusers add_embedding. Map before the - # SDXL index-strip heuristic below, which would otherwise collapse the layer index. - diffusers_name = diffusers_name.replace("label.emb.0.0", "add_embedding.linear_1") - diffusers_name = diffusers_name.replace("label.emb.0.2", "add_embedding.linear_2") - # kohya-ss trains SD 1.x on the diffusers UNet (not the sgm UNet it uses for SDXL), - # so the time-embedding MLP keeps the diffusers spelling time_embedding.linear_N - # rather than the sgm time_embed.N handled above. - diffusers_name = diffusers_name.replace("time.embedding.linear.1", "time_embedding.linear_1") - diffusers_name = diffusers_name.replace("time.embedding.linear.2", "time_embedding.linear_2") - - # SDXL specific conversions. - if "emb" in diffusers_name and "time.emb.proj" not in diffusers_name: - pattern = r"\.\d+(?=\D*$)" - diffusers_name = re.sub(pattern, "", diffusers_name, count=1) - if ".in." in diffusers_name: - diffusers_name = diffusers_name.replace("in.layers.2", "conv1") - if ".out." in diffusers_name: - diffusers_name = diffusers_name.replace("out.layers.3", "conv2") - if "downsamplers" in diffusers_name or "upsamplers" in diffusers_name: - diffusers_name = diffusers_name.replace("op", "conv") - if "skip" in diffusers_name: - diffusers_name = diffusers_name.replace("skip.connection", "conv_shortcut") - - # LyCORIS specific conversions. - if "time.emb.proj" in diffusers_name: - diffusers_name = diffusers_name.replace("time.emb.proj", "time_emb_proj") - if "conv.shortcut" in diffusers_name: - diffusers_name = diffusers_name.replace("conv.shortcut", "conv_shortcut") - - # General conversions. - if "transformer_blocks" in diffusers_name: - if "attn1" in diffusers_name or "attn2" in diffusers_name: - diffusers_name = diffusers_name.replace("attn1", "attn1.processor") - diffusers_name = diffusers_name.replace("attn2", "attn2.processor") - elif "ff" in diffusers_name: - pass - elif any(key in diffusers_name for key in ("proj_in", "proj_out")): - pass - else: - pass - - return diffusers_name - - -def _convert_text_encoder_lora_key(key, lora_name): - """ - Converts a text encoder LoRA key to a Diffusers compatible key. - """ - if lora_name.startswith(("lora_te_", "lora_te1_")): - key_to_replace = "lora_te_" if lora_name.startswith("lora_te_") else "lora_te1_" - else: - key_to_replace = "lora_te2_" - - diffusers_name = key.replace(key_to_replace, "").replace("_", ".") - diffusers_name = diffusers_name.replace("text.model", "text_model") - diffusers_name = diffusers_name.replace("self.attn", "self_attn") - diffusers_name = diffusers_name.replace("q.proj.lora", "to_q_lora") - diffusers_name = diffusers_name.replace("k.proj.lora", "to_k_lora") - diffusers_name = diffusers_name.replace("v.proj.lora", "to_v_lora") - diffusers_name = diffusers_name.replace("out.proj.lora", "to_out_lora") - diffusers_name = diffusers_name.replace("text.projection", "text_projection") - - if "self_attn" in diffusers_name or "text_projection" in diffusers_name: - pass - elif "mlp" in diffusers_name: - # Be aware that this is the new diffusers convention and the rest of the code might - # not utilize it yet. - diffusers_name = diffusers_name.replace(".lora.", ".lora_linear_layer.") - - return diffusers_name - - -def _get_alpha_name(lora_name_alpha, diffusers_name, alpha): - """ - Gets the correct alpha name for the Diffusers model. - """ - if lora_name_alpha.startswith("lora_unet_"): - prefix = "unet." - elif lora_name_alpha.startswith(("lora_te_", "lora_te1_")): - prefix = "text_encoder." - else: - prefix = "text_encoder_2." - new_name = prefix + diffusers_name.split(".lora.")[0] + ".alpha" - return {new_name: alpha} - - -# The utilities under `_convert_kohya_flux_lora_to_diffusers()` -# are adapted from https://github.com/kohya-ss/sd-scripts/blob/a61cf73a5cb5209c3f4d1a3688dd276a4dfd1ecb/networks/convert_flux_lora.py -def _convert_kohya_flux_lora_to_diffusers(state_dict): - def _convert_to_ai_toolkit(sds_sd, ait_sd, sds_key, ait_key): - if sds_key + ".lora_down.weight" not in sds_sd: - return - down_weight = sds_sd.pop(sds_key + ".lora_down.weight") - - # scale weight by alpha and dim - rank = down_weight.shape[0] - default_alpha = torch.tensor(rank, dtype=down_weight.dtype, device=down_weight.device, requires_grad=False) - alpha = sds_sd.pop(sds_key + ".alpha", default_alpha).item() # alpha is scalar - scale = alpha / rank # LoRA is scaled by 'alpha / rank' in forward pass, so we need to scale it back here - - # calculate scale_down and scale_up to keep the same value. if scale is 4, scale_down is 2 and scale_up is 2 - scale_down = scale - scale_up = 1.0 - while scale_down * 2 < scale_up: - scale_down *= 2 - scale_up /= 2 - - ait_sd[ait_key + ".lora_A.weight"] = down_weight * scale_down - ait_sd[ait_key + ".lora_B.weight"] = sds_sd.pop(sds_key + ".lora_up.weight") * scale_up - - def _convert_to_ai_toolkit_cat(sds_sd, ait_sd, sds_key, ait_keys, dims=None): - if sds_key + ".lora_down.weight" not in sds_sd: - return - down_weight = sds_sd.pop(sds_key + ".lora_down.weight") - up_weight = sds_sd.pop(sds_key + ".lora_up.weight") - sd_lora_rank = down_weight.shape[0] - - # scale weight by alpha and dim - default_alpha = torch.tensor( - sd_lora_rank, dtype=down_weight.dtype, device=down_weight.device, requires_grad=False - ) - alpha = sds_sd.pop(sds_key + ".alpha", default_alpha) - scale = alpha / sd_lora_rank - - # calculate scale_down and scale_up - scale_down = scale - scale_up = 1.0 - while scale_down * 2 < scale_up: - scale_down *= 2 - scale_up /= 2 - - down_weight = down_weight * scale_down - up_weight = up_weight * scale_up - - # calculate dims if not provided - num_splits = len(ait_keys) - if dims is None: - dims = [up_weight.shape[0] // num_splits] * num_splits - else: - assert sum(dims) == up_weight.shape[0] - - # check upweight is sparse or not - is_sparse = False - if sd_lora_rank % num_splits == 0: - ait_rank = sd_lora_rank // num_splits - is_sparse = True - i = 0 - for j in range(len(dims)): - for k in range(len(dims)): - if j == k: - continue - is_sparse = is_sparse and torch.all( - up_weight[i : i + dims[j], k * ait_rank : (k + 1) * ait_rank] == 0 - ) - i += dims[j] - if is_sparse: - logger.info(f"weight is sparse: {sds_key}") - - # make ai-toolkit weight - ait_down_keys = [k + ".lora_A.weight" for k in ait_keys] - ait_up_keys = [k + ".lora_B.weight" for k in ait_keys] - if not is_sparse: - # down_weight is copied to each split - ait_sd.update(dict.fromkeys(ait_down_keys, down_weight)) - - # up_weight is split to each split - ait_sd.update({k: v for k, v in zip(ait_up_keys, torch.split(up_weight, dims, dim=0))}) # noqa: C416 - else: - # down_weight is chunked to each split - ait_sd.update({k: v for k, v in zip(ait_down_keys, torch.chunk(down_weight, num_splits, dim=0))}) # noqa: C416 - - # up_weight is sparse: only non-zero values are copied to each split - i = 0 - for j in range(len(dims)): - ait_sd[ait_up_keys[j]] = up_weight[i : i + dims[j], j * ait_rank : (j + 1) * ait_rank].contiguous() - i += dims[j] - - def _convert_sd_scripts_to_ai_toolkit(sds_sd): - ait_sd = {} - for i in range(19): - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - f"lora_unet_double_blocks_{i}_img_attn_proj", - f"transformer.transformer_blocks.{i}.attn.to_out.0", - ) - _convert_to_ai_toolkit_cat( - sds_sd, - ait_sd, - f"lora_unet_double_blocks_{i}_img_attn_qkv", - [ - f"transformer.transformer_blocks.{i}.attn.to_q", - f"transformer.transformer_blocks.{i}.attn.to_k", - f"transformer.transformer_blocks.{i}.attn.to_v", - ], - ) - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - f"lora_unet_double_blocks_{i}_img_mlp_0", - f"transformer.transformer_blocks.{i}.ff.net.0.proj", - ) - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - f"lora_unet_double_blocks_{i}_img_mlp_2", - f"transformer.transformer_blocks.{i}.ff.net.2", - ) - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - f"lora_unet_double_blocks_{i}_img_mod_lin", - f"transformer.transformer_blocks.{i}.norm1.linear", - ) - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - f"lora_unet_double_blocks_{i}_txt_attn_proj", - f"transformer.transformer_blocks.{i}.attn.to_add_out", - ) - _convert_to_ai_toolkit_cat( - sds_sd, - ait_sd, - f"lora_unet_double_blocks_{i}_txt_attn_qkv", - [ - f"transformer.transformer_blocks.{i}.attn.add_q_proj", - f"transformer.transformer_blocks.{i}.attn.add_k_proj", - f"transformer.transformer_blocks.{i}.attn.add_v_proj", - ], - ) - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - f"lora_unet_double_blocks_{i}_txt_mlp_0", - f"transformer.transformer_blocks.{i}.ff_context.net.0.proj", - ) - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - f"lora_unet_double_blocks_{i}_txt_mlp_2", - f"transformer.transformer_blocks.{i}.ff_context.net.2", - ) - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - f"lora_unet_double_blocks_{i}_txt_mod_lin", - f"transformer.transformer_blocks.{i}.norm1_context.linear", - ) - - for i in range(38): - _convert_to_ai_toolkit_cat( - sds_sd, - ait_sd, - f"lora_unet_single_blocks_{i}_linear1", - [ - f"transformer.single_transformer_blocks.{i}.attn.to_q", - f"transformer.single_transformer_blocks.{i}.attn.to_k", - f"transformer.single_transformer_blocks.{i}.attn.to_v", - f"transformer.single_transformer_blocks.{i}.proj_mlp", - ], - dims=[3072, 3072, 3072, 12288], - ) - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - f"lora_unet_single_blocks_{i}_linear2", - f"transformer.single_transformer_blocks.{i}.proj_out", - ) - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - f"lora_unet_single_blocks_{i}_modulation_lin", - f"transformer.single_transformer_blocks.{i}.norm.linear", - ) - - # TODO: alphas. - def assign_remaining_weights(assignments, source): - for lora_key in ["lora_A", "lora_B"]: - orig_lora_key = "lora_down" if lora_key == "lora_A" else "lora_up" - for target_fmt, source_fmt, transform in assignments: - target_key = target_fmt.format(lora_key=lora_key) - source_key = source_fmt.format(orig_lora_key=orig_lora_key) - value = source.pop(source_key, None) - if value is None: - continue - if transform and lora_key == "lora_B": - value = transform(value) - ait_sd[target_key] = value - - # Consume any leftover final_layer alpha keys so they don't - # reach the remaining_keys guard and cause a false "Incompatible keys" error. - for key in list(source.keys()): - if "final_layer" in key and key.endswith(".alpha"): - source.pop(key) - - if any("guidance_in" in k for k in sds_sd): - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - "lora_unet_guidance_in_in_layer", - "time_text_embed.guidance_embedder.linear_1", - ) - - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - "lora_unet_guidance_in_out_layer", - "time_text_embed.guidance_embedder.linear_2", - ) - - if any("img_in" in k for k in sds_sd): - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - "lora_unet_img_in", - "x_embedder", - ) - - if any("txt_in" in k for k in sds_sd): - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - "lora_unet_txt_in", - "context_embedder", - ) - - if any("time_in" in k for k in sds_sd): - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - "lora_unet_time_in_in_layer", - "time_text_embed.timestep_embedder.linear_1", - ) - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - "lora_unet_time_in_out_layer", - "time_text_embed.timestep_embedder.linear_2", - ) - - if any("vector_in" in k for k in sds_sd): - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - "lora_unet_vector_in_in_layer", - "time_text_embed.text_embedder.linear_1", - ) - _convert_to_ai_toolkit( - sds_sd, - ait_sd, - "lora_unet_vector_in_out_layer", - "time_text_embed.text_embedder.linear_2", - ) - - if any("final_layer" in k for k in sds_sd): - # Notice the swap in processing for "final_layer". - assign_remaining_weights( - [ - ( - "norm_out.linear.{lora_key}.weight", - "lora_unet_final_layer_adaLN_modulation_1.{orig_lora_key}.weight", - swap_scale_shift, - ), - ("proj_out.{lora_key}.weight", "lora_unet_final_layer_linear.{orig_lora_key}.weight", None), - ], - sds_sd, - ) - - remaining_keys = list(sds_sd.keys()) - te_state_dict = {} - if remaining_keys: - if not all(k.startswith(("lora_te", "lora_te1")) for k in remaining_keys): - raise ValueError(f"Incompatible keys detected: \n\n {', '.join(remaining_keys)}") - for key in remaining_keys: - if not key.endswith("lora_down.weight"): - continue - - lora_name = key.split(".")[0] - lora_name_up = f"{lora_name}.lora_up.weight" - lora_name_alpha = f"{lora_name}.alpha" - diffusers_name = _convert_text_encoder_lora_key(key, lora_name) - - if lora_name.startswith(("lora_te_", "lora_te1_")): - down_weight = sds_sd.pop(key) - sd_lora_rank = down_weight.shape[0] - te_state_dict[diffusers_name] = down_weight - te_state_dict[diffusers_name.replace(".down.", ".up.")] = sds_sd.pop(lora_name_up) - - if lora_name_alpha in sds_sd: - alpha = sds_sd.pop(lora_name_alpha).item() - scale = alpha / sd_lora_rank - - scale_down = scale - scale_up = 1.0 - while scale_down * 2 < scale_up: - scale_down *= 2 - scale_up /= 2 - - te_state_dict[diffusers_name] *= scale_down - te_state_dict[diffusers_name.replace(".down.", ".up.")] *= scale_up - - if len(sds_sd) > 0: - logger.warning(f"Unsupported keys for ai-toolkit: {sds_sd.keys()}") - - if te_state_dict: - te_state_dict = {f"text_encoder.{module_name}": params for module_name, params in te_state_dict.items()} - - new_state_dict = {**ait_sd, **te_state_dict} - return new_state_dict - - def _convert_mixture_state_dict_to_diffusers(state_dict): - new_state_dict = {} - - def _convert(original_key, diffusers_key, state_dict, new_state_dict): - down_key = f"{original_key}.lora_down.weight" - down_weight = state_dict.pop(down_key) - lora_rank = down_weight.shape[0] - - up_weight_key = f"{original_key}.lora_up.weight" - up_weight = state_dict.pop(up_weight_key) - - alpha_key = f"{original_key}.alpha" - alpha = state_dict.pop(alpha_key) - - # scale weight by alpha and dim - scale = alpha / lora_rank - # calculate scale_down and scale_up - scale_down = scale - scale_up = 1.0 - while scale_down * 2 < scale_up: - scale_down *= 2 - scale_up /= 2 - down_weight = down_weight * scale_down - up_weight = up_weight * scale_up - - diffusers_down_key = f"{diffusers_key}.lora_A.weight" - new_state_dict[diffusers_down_key] = down_weight - new_state_dict[diffusers_down_key.replace(".lora_A.", ".lora_B.")] = up_weight - - all_unique_keys = { - k.replace(".lora_down.weight", "").replace(".lora_up.weight", "").replace(".alpha", "") - for k in state_dict - if not k.startswith(("lora_unet_")) - } - assert all(k.startswith(("lora_transformer_", "lora_te1_")) for k in all_unique_keys), f"{all_unique_keys=}" - - has_te_keys = False - for k in all_unique_keys: - if k.startswith("lora_transformer_single_transformer_blocks_"): - i = int(k.split("lora_transformer_single_transformer_blocks_")[-1].split("_")[0]) - diffusers_key = f"single_transformer_blocks.{i}" - elif k.startswith("lora_transformer_transformer_blocks_"): - i = int(k.split("lora_transformer_transformer_blocks_")[-1].split("_")[0]) - diffusers_key = f"transformer_blocks.{i}" - elif k.startswith("lora_te1_"): - has_te_keys = True - continue - elif k.startswith("lora_transformer_context_embedder"): - diffusers_key = "context_embedder" - elif k.startswith("lora_transformer_norm_out_linear"): - diffusers_key = "norm_out.linear" - elif k.startswith("lora_transformer_proj_out"): - diffusers_key = "proj_out" - elif k.startswith("lora_transformer_x_embedder"): - diffusers_key = "x_embedder" - elif k.startswith("lora_transformer_time_text_embed_guidance_embedder_linear_"): - i = int(k.split("lora_transformer_time_text_embed_guidance_embedder_linear_")[-1]) - diffusers_key = f"time_text_embed.guidance_embedder.linear_{i}" - elif k.startswith("lora_transformer_time_text_embed_text_embedder_linear_"): - i = int(k.split("lora_transformer_time_text_embed_text_embedder_linear_")[-1]) - diffusers_key = f"time_text_embed.text_embedder.linear_{i}" - elif k.startswith("lora_transformer_time_text_embed_timestep_embedder_linear_"): - i = int(k.split("lora_transformer_time_text_embed_timestep_embedder_linear_")[-1]) - diffusers_key = f"time_text_embed.timestep_embedder.linear_{i}" - else: - raise NotImplementedError(f"Handling for key ({k}) is not implemented.") - - if "attn_" in k: - if "_to_out_0" in k: - diffusers_key += ".attn.to_out.0" - elif "_to_add_out" in k: - diffusers_key += ".attn.to_add_out" - elif any(qkv in k for qkv in ["to_q", "to_k", "to_v"]): - remaining = k.split("attn_")[-1] - diffusers_key += f".attn.{remaining}" - elif any(add_qkv in k for add_qkv in ["add_q_proj", "add_k_proj", "add_v_proj"]): - remaining = k.split("attn_")[-1] - diffusers_key += f".attn.{remaining}" - - _convert(k, diffusers_key, state_dict, new_state_dict) - - if has_te_keys: - layer_pattern = re.compile(r"lora_te1_text_model_encoder_layers_(\d+)") - attn_mapping = { - "q_proj": ".self_attn.q_proj", - "k_proj": ".self_attn.k_proj", - "v_proj": ".self_attn.v_proj", - "out_proj": ".self_attn.out_proj", - } - mlp_mapping = {"fc1": ".mlp.fc1", "fc2": ".mlp.fc2"} - for k in all_unique_keys: - if not k.startswith("lora_te1_"): - continue - - match = layer_pattern.search(k) - if not match: - continue - i = int(match.group(1)) - diffusers_key = f"text_model.encoder.layers.{i}" - - if "attn" in k: - for key_fragment, suffix in attn_mapping.items(): - if key_fragment in k: - diffusers_key += suffix - break - elif "mlp" in k: - for key_fragment, suffix in mlp_mapping.items(): - if key_fragment in k: - diffusers_key += suffix - break - - _convert(k, diffusers_key, state_dict, new_state_dict) - - remaining_all_unet = False - if state_dict: - remaining_all_unet = all(k.startswith("lora_unet_") for k in state_dict) - if remaining_all_unet: - keys = list(state_dict.keys()) - for k in keys: - state_dict.pop(k) - - if len(state_dict) > 0: - raise ValueError( - f"Expected an empty state dict at this point but its has these keys which couldn't be parsed: {list(state_dict.keys())}." - ) - - transformer_state_dict = { - f"transformer.{k}": v for k, v in new_state_dict.items() if not k.startswith("text_model.") - } - te_state_dict = {f"text_encoder.{k}": v for k, v in new_state_dict.items() if k.startswith("text_model.")} - return {**transformer_state_dict, **te_state_dict} - - # This is weird. - # https://huggingface.co/sayakpaul/different-lora-from-civitai/tree/main?show_file_info=sharp_detailed_foot.safetensors - # has both `peft` and non-peft state dict. - has_peft_state_dict = any(k.startswith("transformer.") for k in state_dict) - if has_peft_state_dict: - state_dict = { - k.replace("lora_down.weight", "lora_A.weight").replace("lora_up.weight", "lora_B.weight"): v - for k, v in state_dict.items() - if k.startswith("transformer.") - } - return state_dict - - # Another weird one. - has_mixture = any( - k.startswith("lora_transformer_") and ("lora_down" in k or "lora_up" in k or "alpha" in k) for k in state_dict - ) - - # ComfyUI. - if not has_mixture: - state_dict = {k.replace("diffusion_model.", "lora_unet_"): v for k, v in state_dict.items()} - state_dict = {k.replace("text_encoders.clip_l.transformer.", "lora_te_"): v for k, v in state_dict.items()} - - has_position_embedding = any("position_embedding" in k for k in state_dict) - if has_position_embedding: - zero_status_pe = state_dict_all_zero(state_dict, "position_embedding") - if zero_status_pe: - logger.info( - "The `position_embedding` LoRA params are all zeros which make them ineffective. " - "So, we will purge them out of the current state dict to make loading possible." - ) - - else: - logger.info( - "The state_dict has position_embedding LoRA params and we currently do not support them. " - "Open an issue if you need this supported - https://github.com/huggingface/diffusers/issues/new." - ) - state_dict = {k: v for k, v in state_dict.items() if "position_embedding" not in k} - - has_t5xxl = any(k.startswith("text_encoders.t5xxl.transformer.") for k in state_dict) - if has_t5xxl: - zero_status_t5 = state_dict_all_zero(state_dict, "text_encoders.t5xxl") - if zero_status_t5: - logger.info( - "The `t5xxl` LoRA params are all zeros which make them ineffective. " - "So, we will purge them out of the current state dict to make loading possible." - ) - else: - logger.info( - "T5-xxl keys found in the state dict, which are currently unsupported. We will filter them out." - "Open an issue if this is a problem - https://github.com/huggingface/diffusers/issues/new." - ) - state_dict = {k: v for k, v in state_dict.items() if not k.startswith("text_encoders.t5xxl.transformer.")} - - has_diffb = any("diff_b" in k and k.startswith(("lora_unet_", "lora_te_", "lora_te1_")) for k in state_dict) - if has_diffb: - zero_status_diff_b = state_dict_all_zero(state_dict, ".diff_b") - if zero_status_diff_b: - logger.info( - "The `diff_b` LoRA params are all zeros which make them ineffective. " - "So, we will purge them out of the current state dict to make loading possible." - ) - else: - logger.info( - "`diff_b` keys found in the state dict which are currently unsupported. " - "So, we will filter out those keys. Open an issue if this is a problem - " - "https://github.com/huggingface/diffusers/issues/new." - ) - state_dict = {k: v for k, v in state_dict.items() if ".diff_b" not in k} - - has_norm_diff = any(".norm" in k and ".diff" in k for k in state_dict) - if has_norm_diff: - zero_status_diff = state_dict_all_zero(state_dict, ".diff") - if zero_status_diff: - logger.info( - "The `diff` LoRA params are all zeros which make them ineffective. " - "So, we will purge them out of the current state dict to make loading possible." - ) - else: - logger.info( - "Normalization diff keys found in the state dict which are currently unsupported. " - "So, we will filter out those keys. Open an issue if this is a problem - " - "https://github.com/huggingface/diffusers/issues/new." - ) - state_dict = {k: v for k, v in state_dict.items() if ".norm" not in k and ".diff" not in k} - - limit_substrings = ["lora_down", "lora_up"] - if any("alpha" in k for k in state_dict): - limit_substrings.append("alpha") - - state_dict = { - _custom_replace(k, limit_substrings): v - for k, v in state_dict.items() - if k.startswith(("lora_unet_", "lora_te_", "lora_te1_")) - } - - if any("text_projection" in k for k in state_dict): - logger.info( - "`text_projection` keys found in the `state_dict` which are unexpected. " - "So, we will filter out those keys. Open an issue if this is a problem - " - "https://github.com/huggingface/diffusers/issues/new." - ) - state_dict = {k: v for k, v in state_dict.items() if "text_projection" not in k} - - if has_mixture: - return _convert_mixture_state_dict_to_diffusers(state_dict) - - return _convert_sd_scripts_to_ai_toolkit(state_dict) - - -# Adapted from https://gist.github.com/Leommm-byte/6b331a1e9bd53271210b26543a7065d6 -# Some utilities were reused from -# https://github.com/kohya-ss/sd-scripts/blob/a61cf73a5cb5209c3f4d1a3688dd276a4dfd1ecb/networks/convert_flux_lora.py -def _convert_xlabs_flux_lora_to_diffusers(old_state_dict): - new_state_dict = {} - orig_keys = list(old_state_dict.keys()) - - def handle_qkv(sds_sd, ait_sd, sds_key, ait_keys, dims=None): - down_weight = sds_sd.pop(sds_key) - up_weight = sds_sd.pop(sds_key.replace(".down.weight", ".up.weight")) - - # calculate dims if not provided - num_splits = len(ait_keys) - if dims is None: - dims = [up_weight.shape[0] // num_splits] * num_splits - else: - assert sum(dims) == up_weight.shape[0] - - # make ai-toolkit weight - ait_down_keys = [k + ".lora_A.weight" for k in ait_keys] - ait_up_keys = [k + ".lora_B.weight" for k in ait_keys] - - # down_weight is copied to each split - ait_sd.update(dict.fromkeys(ait_down_keys, down_weight)) - - # up_weight is split to each split - ait_sd.update({k: v for k, v in zip(ait_up_keys, torch.split(up_weight, dims, dim=0))}) # noqa: C416 - - for old_key in orig_keys: - # Handle double_blocks - if old_key.startswith(("diffusion_model.double_blocks", "double_blocks")): - block_num = re.search(r"double_blocks\.(\d+)", old_key).group(1) - new_key = f"transformer.transformer_blocks.{block_num}" - - if "processor.proj_lora1" in old_key: - new_key += ".attn.to_out.0" - elif "processor.proj_lora2" in old_key: - new_key += ".attn.to_add_out" - # Handle text latents. - elif "processor.qkv_lora2" in old_key and "up" not in old_key: - handle_qkv( - old_state_dict, - new_state_dict, - old_key, - [ - f"transformer.transformer_blocks.{block_num}.attn.add_q_proj", - f"transformer.transformer_blocks.{block_num}.attn.add_k_proj", - f"transformer.transformer_blocks.{block_num}.attn.add_v_proj", - ], - ) - # continue - # Handle image latents. - elif "processor.qkv_lora1" in old_key and "up" not in old_key: - handle_qkv( - old_state_dict, - new_state_dict, - old_key, - [ - f"transformer.transformer_blocks.{block_num}.attn.to_q", - f"transformer.transformer_blocks.{block_num}.attn.to_k", - f"transformer.transformer_blocks.{block_num}.attn.to_v", - ], - ) - # continue - - if "down" in old_key: - new_key += ".lora_A.weight" - elif "up" in old_key: - new_key += ".lora_B.weight" - - # Handle single_blocks - elif old_key.startswith(("diffusion_model.single_blocks", "single_blocks")): - block_num = re.search(r"single_blocks\.(\d+)", old_key).group(1) - new_key = f"transformer.single_transformer_blocks.{block_num}" - - if "proj_lora" in old_key: - new_key += ".proj_out" - elif "qkv_lora" in old_key and "up" not in old_key: - handle_qkv( - old_state_dict, - new_state_dict, - old_key, - [ - f"transformer.single_transformer_blocks.{block_num}.attn.to_q", - f"transformer.single_transformer_blocks.{block_num}.attn.to_k", - f"transformer.single_transformer_blocks.{block_num}.attn.to_v", - ], - ) - - if "down" in old_key: - new_key += ".lora_A.weight" - elif "up" in old_key: - new_key += ".lora_B.weight" - - else: - # Handle other potential key patterns here - new_key = old_key - - # Since we already handle qkv above. - if "qkv" not in old_key: - new_state_dict[new_key] = old_state_dict.pop(old_key) - - if len(old_state_dict) > 0: - raise ValueError(f"`old_state_dict` should be at this point but has: {list(old_state_dict.keys())}.") - - return new_state_dict - - -def _custom_replace(key: str, substrings: list[str]) -> str: - # Replaces the "."s with "_"s upto the `substrings`. - # Example: - # lora_unet.foo.bar.lora_A.weight -> lora_unet_foo_bar.lora_A.weight - pattern = "(" + "|".join(re.escape(sub) for sub in substrings) + ")" - - match = re.search(pattern, key) - if match: - start_sub = match.start() - if start_sub > 0 and key[start_sub - 1] == ".": - boundary = start_sub - 1 - else: - boundary = start_sub - left = key[:boundary].replace(".", "_") - right = key[boundary:] - return left + right - else: - return key.replace(".", "_") - - -def _convert_bfl_flux_control_lora_to_diffusers(original_state_dict): - converted_state_dict = {} - original_state_dict_keys = list(original_state_dict.keys()) - num_layers = 19 - num_single_layers = 38 - inner_dim = 3072 - mlp_ratio = 4.0 - - for lora_key in ["lora_A", "lora_B"]: - ## time_text_embed.timestep_embedder <- time_in - converted_state_dict[f"time_text_embed.timestep_embedder.linear_1.{lora_key}.weight"] = ( - original_state_dict.pop(f"time_in.in_layer.{lora_key}.weight") - ) - if f"time_in.in_layer.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"time_text_embed.timestep_embedder.linear_1.{lora_key}.bias"] = ( - original_state_dict.pop(f"time_in.in_layer.{lora_key}.bias") - ) - - converted_state_dict[f"time_text_embed.timestep_embedder.linear_2.{lora_key}.weight"] = ( - original_state_dict.pop(f"time_in.out_layer.{lora_key}.weight") - ) - if f"time_in.out_layer.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"time_text_embed.timestep_embedder.linear_2.{lora_key}.bias"] = ( - original_state_dict.pop(f"time_in.out_layer.{lora_key}.bias") - ) - - ## time_text_embed.text_embedder <- vector_in - converted_state_dict[f"time_text_embed.text_embedder.linear_1.{lora_key}.weight"] = original_state_dict.pop( - f"vector_in.in_layer.{lora_key}.weight" - ) - if f"vector_in.in_layer.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"time_text_embed.text_embedder.linear_1.{lora_key}.bias"] = original_state_dict.pop( - f"vector_in.in_layer.{lora_key}.bias" - ) - - converted_state_dict[f"time_text_embed.text_embedder.linear_2.{lora_key}.weight"] = original_state_dict.pop( - f"vector_in.out_layer.{lora_key}.weight" - ) - if f"vector_in.out_layer.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"time_text_embed.text_embedder.linear_2.{lora_key}.bias"] = original_state_dict.pop( - f"vector_in.out_layer.{lora_key}.bias" - ) - - # guidance - has_guidance = any("guidance" in k for k in original_state_dict) - if has_guidance: - converted_state_dict[f"time_text_embed.guidance_embedder.linear_1.{lora_key}.weight"] = ( - original_state_dict.pop(f"guidance_in.in_layer.{lora_key}.weight") - ) - if f"guidance_in.in_layer.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"time_text_embed.guidance_embedder.linear_1.{lora_key}.bias"] = ( - original_state_dict.pop(f"guidance_in.in_layer.{lora_key}.bias") - ) - - converted_state_dict[f"time_text_embed.guidance_embedder.linear_2.{lora_key}.weight"] = ( - original_state_dict.pop(f"guidance_in.out_layer.{lora_key}.weight") - ) - if f"guidance_in.out_layer.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"time_text_embed.guidance_embedder.linear_2.{lora_key}.bias"] = ( - original_state_dict.pop(f"guidance_in.out_layer.{lora_key}.bias") - ) - - # context_embedder - converted_state_dict[f"context_embedder.{lora_key}.weight"] = original_state_dict.pop( - f"txt_in.{lora_key}.weight" - ) - if f"txt_in.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"context_embedder.{lora_key}.bias"] = original_state_dict.pop( - f"txt_in.{lora_key}.bias" - ) - - # x_embedder - converted_state_dict[f"x_embedder.{lora_key}.weight"] = original_state_dict.pop(f"img_in.{lora_key}.weight") - if f"img_in.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"x_embedder.{lora_key}.bias"] = original_state_dict.pop(f"img_in.{lora_key}.bias") - - # double transformer blocks - for i in range(num_layers): - block_prefix = f"transformer_blocks.{i}." - - for lora_key in ["lora_A", "lora_B"]: - # norms - converted_state_dict[f"{block_prefix}norm1.linear.{lora_key}.weight"] = original_state_dict.pop( - f"double_blocks.{i}.img_mod.lin.{lora_key}.weight" - ) - if f"double_blocks.{i}.img_mod.lin.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}norm1.linear.{lora_key}.bias"] = original_state_dict.pop( - f"double_blocks.{i}.img_mod.lin.{lora_key}.bias" - ) - - converted_state_dict[f"{block_prefix}norm1_context.linear.{lora_key}.weight"] = original_state_dict.pop( - f"double_blocks.{i}.txt_mod.lin.{lora_key}.weight" - ) - if f"double_blocks.{i}.txt_mod.lin.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}norm1_context.linear.{lora_key}.bias"] = original_state_dict.pop( - f"double_blocks.{i}.txt_mod.lin.{lora_key}.bias" - ) - - # Q, K, V - if lora_key == "lora_A": - sample_lora_weight = original_state_dict.pop(f"double_blocks.{i}.img_attn.qkv.{lora_key}.weight") - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.weight"] = torch.cat([sample_lora_weight]) - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.weight"] = torch.cat([sample_lora_weight]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.weight"] = torch.cat([sample_lora_weight]) - - context_lora_weight = original_state_dict.pop(f"double_blocks.{i}.txt_attn.qkv.{lora_key}.weight") - converted_state_dict[f"{block_prefix}attn.add_q_proj.{lora_key}.weight"] = torch.cat( - [context_lora_weight] - ) - converted_state_dict[f"{block_prefix}attn.add_k_proj.{lora_key}.weight"] = torch.cat( - [context_lora_weight] - ) - converted_state_dict[f"{block_prefix}attn.add_v_proj.{lora_key}.weight"] = torch.cat( - [context_lora_weight] - ) - else: - sample_q, sample_k, sample_v = torch.chunk( - original_state_dict.pop(f"double_blocks.{i}.img_attn.qkv.{lora_key}.weight"), 3, dim=0 - ) - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.weight"] = torch.cat([sample_q]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.weight"] = torch.cat([sample_k]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.weight"] = torch.cat([sample_v]) - - context_q, context_k, context_v = torch.chunk( - original_state_dict.pop(f"double_blocks.{i}.txt_attn.qkv.{lora_key}.weight"), 3, dim=0 - ) - converted_state_dict[f"{block_prefix}attn.add_q_proj.{lora_key}.weight"] = torch.cat([context_q]) - converted_state_dict[f"{block_prefix}attn.add_k_proj.{lora_key}.weight"] = torch.cat([context_k]) - converted_state_dict[f"{block_prefix}attn.add_v_proj.{lora_key}.weight"] = torch.cat([context_v]) - - if f"double_blocks.{i}.img_attn.qkv.{lora_key}.bias" in original_state_dict_keys: - sample_q_bias, sample_k_bias, sample_v_bias = torch.chunk( - original_state_dict.pop(f"double_blocks.{i}.img_attn.qkv.{lora_key}.bias"), 3, dim=0 - ) - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.bias"] = torch.cat([sample_q_bias]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.bias"] = torch.cat([sample_k_bias]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.bias"] = torch.cat([sample_v_bias]) - - if f"double_blocks.{i}.txt_attn.qkv.{lora_key}.bias" in original_state_dict_keys: - context_q_bias, context_k_bias, context_v_bias = torch.chunk( - original_state_dict.pop(f"double_blocks.{i}.txt_attn.qkv.{lora_key}.bias"), 3, dim=0 - ) - converted_state_dict[f"{block_prefix}attn.add_q_proj.{lora_key}.bias"] = torch.cat([context_q_bias]) - converted_state_dict[f"{block_prefix}attn.add_k_proj.{lora_key}.bias"] = torch.cat([context_k_bias]) - converted_state_dict[f"{block_prefix}attn.add_v_proj.{lora_key}.bias"] = torch.cat([context_v_bias]) - - # ff img_mlp - converted_state_dict[f"{block_prefix}ff.net.0.proj.{lora_key}.weight"] = original_state_dict.pop( - f"double_blocks.{i}.img_mlp.0.{lora_key}.weight" - ) - if f"double_blocks.{i}.img_mlp.0.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}ff.net.0.proj.{lora_key}.bias"] = original_state_dict.pop( - f"double_blocks.{i}.img_mlp.0.{lora_key}.bias" - ) - - converted_state_dict[f"{block_prefix}ff.net.2.{lora_key}.weight"] = original_state_dict.pop( - f"double_blocks.{i}.img_mlp.2.{lora_key}.weight" - ) - if f"double_blocks.{i}.img_mlp.2.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}ff.net.2.{lora_key}.bias"] = original_state_dict.pop( - f"double_blocks.{i}.img_mlp.2.{lora_key}.bias" - ) - - converted_state_dict[f"{block_prefix}ff_context.net.0.proj.{lora_key}.weight"] = original_state_dict.pop( - f"double_blocks.{i}.txt_mlp.0.{lora_key}.weight" - ) - if f"double_blocks.{i}.txt_mlp.0.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}ff_context.net.0.proj.{lora_key}.bias"] = original_state_dict.pop( - f"double_blocks.{i}.txt_mlp.0.{lora_key}.bias" - ) - - converted_state_dict[f"{block_prefix}ff_context.net.2.{lora_key}.weight"] = original_state_dict.pop( - f"double_blocks.{i}.txt_mlp.2.{lora_key}.weight" - ) - if f"double_blocks.{i}.txt_mlp.2.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}ff_context.net.2.{lora_key}.bias"] = original_state_dict.pop( - f"double_blocks.{i}.txt_mlp.2.{lora_key}.bias" - ) - - # output projections. - converted_state_dict[f"{block_prefix}attn.to_out.0.{lora_key}.weight"] = original_state_dict.pop( - f"double_blocks.{i}.img_attn.proj.{lora_key}.weight" - ) - if f"double_blocks.{i}.img_attn.proj.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}attn.to_out.0.{lora_key}.bias"] = original_state_dict.pop( - f"double_blocks.{i}.img_attn.proj.{lora_key}.bias" - ) - converted_state_dict[f"{block_prefix}attn.to_add_out.{lora_key}.weight"] = original_state_dict.pop( - f"double_blocks.{i}.txt_attn.proj.{lora_key}.weight" - ) - if f"double_blocks.{i}.txt_attn.proj.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}attn.to_add_out.{lora_key}.bias"] = original_state_dict.pop( - f"double_blocks.{i}.txt_attn.proj.{lora_key}.bias" - ) - - # qk_norm - converted_state_dict[f"{block_prefix}attn.norm_q.weight"] = original_state_dict.pop( - f"double_blocks.{i}.img_attn.norm.query_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_k.weight"] = original_state_dict.pop( - f"double_blocks.{i}.img_attn.norm.key_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_added_q.weight"] = original_state_dict.pop( - f"double_blocks.{i}.txt_attn.norm.query_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_added_k.weight"] = original_state_dict.pop( - f"double_blocks.{i}.txt_attn.norm.key_norm.scale" - ) - - # single transformer blocks - for i in range(num_single_layers): - block_prefix = f"single_transformer_blocks.{i}." - - for lora_key in ["lora_A", "lora_B"]: - # norm.linear <- single_blocks.0.modulation.lin - converted_state_dict[f"{block_prefix}norm.linear.{lora_key}.weight"] = original_state_dict.pop( - f"single_blocks.{i}.modulation.lin.{lora_key}.weight" - ) - if f"single_blocks.{i}.modulation.lin.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}norm.linear.{lora_key}.bias"] = original_state_dict.pop( - f"single_blocks.{i}.modulation.lin.{lora_key}.bias" - ) - - # Q, K, V, mlp - mlp_hidden_dim = int(inner_dim * mlp_ratio) - split_size = (inner_dim, inner_dim, inner_dim, mlp_hidden_dim) - - if lora_key == "lora_A": - lora_weight = original_state_dict.pop(f"single_blocks.{i}.linear1.{lora_key}.weight") - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.weight"] = torch.cat([lora_weight]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.weight"] = torch.cat([lora_weight]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.weight"] = torch.cat([lora_weight]) - converted_state_dict[f"{block_prefix}proj_mlp.{lora_key}.weight"] = torch.cat([lora_weight]) - - if f"single_blocks.{i}.linear1.{lora_key}.bias" in original_state_dict_keys: - lora_bias = original_state_dict.pop(f"single_blocks.{i}.linear1.{lora_key}.bias") - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.bias"] = torch.cat([lora_bias]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.bias"] = torch.cat([lora_bias]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.bias"] = torch.cat([lora_bias]) - converted_state_dict[f"{block_prefix}proj_mlp.{lora_key}.bias"] = torch.cat([lora_bias]) - else: - q, k, v, mlp = torch.split( - original_state_dict.pop(f"single_blocks.{i}.linear1.{lora_key}.weight"), split_size, dim=0 - ) - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.weight"] = torch.cat([q]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.weight"] = torch.cat([k]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.weight"] = torch.cat([v]) - converted_state_dict[f"{block_prefix}proj_mlp.{lora_key}.weight"] = torch.cat([mlp]) - - if f"single_blocks.{i}.linear1.{lora_key}.bias" in original_state_dict_keys: - q_bias, k_bias, v_bias, mlp_bias = torch.split( - original_state_dict.pop(f"single_blocks.{i}.linear1.{lora_key}.bias"), split_size, dim=0 - ) - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.bias"] = torch.cat([q_bias]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.bias"] = torch.cat([k_bias]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.bias"] = torch.cat([v_bias]) - converted_state_dict[f"{block_prefix}proj_mlp.{lora_key}.bias"] = torch.cat([mlp_bias]) - - # output projections. - converted_state_dict[f"{block_prefix}proj_out.{lora_key}.weight"] = original_state_dict.pop( - f"single_blocks.{i}.linear2.{lora_key}.weight" - ) - if f"single_blocks.{i}.linear2.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}proj_out.{lora_key}.bias"] = original_state_dict.pop( - f"single_blocks.{i}.linear2.{lora_key}.bias" - ) - - # qk norm - converted_state_dict[f"{block_prefix}attn.norm_q.weight"] = original_state_dict.pop( - f"single_blocks.{i}.norm.query_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_k.weight"] = original_state_dict.pop( - f"single_blocks.{i}.norm.key_norm.scale" - ) - - for lora_key in ["lora_A", "lora_B"]: - converted_state_dict[f"proj_out.{lora_key}.weight"] = original_state_dict.pop( - f"final_layer.linear.{lora_key}.weight" - ) - if f"final_layer.linear.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"proj_out.{lora_key}.bias"] = original_state_dict.pop( - f"final_layer.linear.{lora_key}.bias" - ) - - converted_state_dict[f"norm_out.linear.{lora_key}.weight"] = swap_scale_shift( - original_state_dict.pop(f"final_layer.adaLN_modulation.1.{lora_key}.weight") - ) - if f"final_layer.adaLN_modulation.1.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"norm_out.linear.{lora_key}.bias"] = swap_scale_shift( - original_state_dict.pop(f"final_layer.adaLN_modulation.1.{lora_key}.bias") - ) - - if len(original_state_dict) > 0: - raise ValueError(f"`original_state_dict` should be empty at this point but has {original_state_dict.keys()=}.") - - for key in list(converted_state_dict.keys()): - converted_state_dict[f"transformer.{key}"] = converted_state_dict.pop(key) - - return converted_state_dict - - -def _convert_fal_kontext_lora_to_diffusers(original_state_dict): - converted_state_dict = {} - original_state_dict_keys = list(original_state_dict.keys()) - num_layers = 19 - num_single_layers = 38 - inner_dim = 3072 - mlp_ratio = 4.0 - - # double transformer blocks - for i in range(num_layers): - block_prefix = f"transformer_blocks.{i}." - original_block_prefix = "base_model.model." - - for lora_key in ["lora_A", "lora_B"]: - # norms - converted_state_dict[f"{block_prefix}norm1.linear.{lora_key}.weight"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.img_mod.lin.{lora_key}.weight" - ) - if f"double_blocks.{i}.img_mod.lin.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}norm1.linear.{lora_key}.bias"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.img_mod.lin.{lora_key}.bias" - ) - - converted_state_dict[f"{block_prefix}norm1_context.linear.{lora_key}.weight"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.txt_mod.lin.{lora_key}.weight" - ) - - # Q, K, V - if lora_key == "lora_A": - sample_lora_weight = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.img_attn.qkv.{lora_key}.weight" - ) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.weight"] = torch.cat([sample_lora_weight]) - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.weight"] = torch.cat([sample_lora_weight]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.weight"] = torch.cat([sample_lora_weight]) - - context_lora_weight = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.txt_attn.qkv.{lora_key}.weight" - ) - converted_state_dict[f"{block_prefix}attn.add_q_proj.{lora_key}.weight"] = torch.cat( - [context_lora_weight] - ) - converted_state_dict[f"{block_prefix}attn.add_k_proj.{lora_key}.weight"] = torch.cat( - [context_lora_weight] - ) - converted_state_dict[f"{block_prefix}attn.add_v_proj.{lora_key}.weight"] = torch.cat( - [context_lora_weight] - ) - else: - sample_q, sample_k, sample_v = torch.chunk( - original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.img_attn.qkv.{lora_key}.weight" - ), - 3, - dim=0, - ) - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.weight"] = torch.cat([sample_q]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.weight"] = torch.cat([sample_k]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.weight"] = torch.cat([sample_v]) - - context_q, context_k, context_v = torch.chunk( - original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.txt_attn.qkv.{lora_key}.weight" - ), - 3, - dim=0, - ) - converted_state_dict[f"{block_prefix}attn.add_q_proj.{lora_key}.weight"] = torch.cat([context_q]) - converted_state_dict[f"{block_prefix}attn.add_k_proj.{lora_key}.weight"] = torch.cat([context_k]) - converted_state_dict[f"{block_prefix}attn.add_v_proj.{lora_key}.weight"] = torch.cat([context_v]) - - if f"double_blocks.{i}.img_attn.qkv.{lora_key}.bias" in original_state_dict_keys: - sample_q_bias, sample_k_bias, sample_v_bias = torch.chunk( - original_state_dict.pop(f"{original_block_prefix}double_blocks.{i}.img_attn.qkv.{lora_key}.bias"), - 3, - dim=0, - ) - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.bias"] = torch.cat([sample_q_bias]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.bias"] = torch.cat([sample_k_bias]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.bias"] = torch.cat([sample_v_bias]) - - if f"double_blocks.{i}.txt_attn.qkv.{lora_key}.bias" in original_state_dict_keys: - context_q_bias, context_k_bias, context_v_bias = torch.chunk( - original_state_dict.pop(f"{original_block_prefix}double_blocks.{i}.txt_attn.qkv.{lora_key}.bias"), - 3, - dim=0, - ) - converted_state_dict[f"{block_prefix}attn.add_q_proj.{lora_key}.bias"] = torch.cat([context_q_bias]) - converted_state_dict[f"{block_prefix}attn.add_k_proj.{lora_key}.bias"] = torch.cat([context_k_bias]) - converted_state_dict[f"{block_prefix}attn.add_v_proj.{lora_key}.bias"] = torch.cat([context_v_bias]) - - # ff img_mlp - converted_state_dict[f"{block_prefix}ff.net.0.proj.{lora_key}.weight"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.img_mlp.0.{lora_key}.weight" - ) - if f"{original_block_prefix}double_blocks.{i}.img_mlp.0.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}ff.net.0.proj.{lora_key}.bias"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.img_mlp.0.{lora_key}.bias" - ) - - converted_state_dict[f"{block_prefix}ff.net.2.{lora_key}.weight"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.img_mlp.2.{lora_key}.weight" - ) - if f"{original_block_prefix}double_blocks.{i}.img_mlp.2.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}ff.net.2.{lora_key}.bias"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.img_mlp.2.{lora_key}.bias" - ) - - converted_state_dict[f"{block_prefix}ff_context.net.0.proj.{lora_key}.weight"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.txt_mlp.0.{lora_key}.weight" - ) - if f"{original_block_prefix}double_blocks.{i}.txt_mlp.0.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}ff_context.net.0.proj.{lora_key}.bias"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.txt_mlp.0.{lora_key}.bias" - ) - - converted_state_dict[f"{block_prefix}ff_context.net.2.{lora_key}.weight"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.txt_mlp.2.{lora_key}.weight" - ) - if f"{original_block_prefix}double_blocks.{i}.txt_mlp.2.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}ff_context.net.2.{lora_key}.bias"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.txt_mlp.2.{lora_key}.bias" - ) - - # output projections. - converted_state_dict[f"{block_prefix}attn.to_out.0.{lora_key}.weight"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.img_attn.proj.{lora_key}.weight" - ) - if f"{original_block_prefix}double_blocks.{i}.img_attn.proj.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}attn.to_out.0.{lora_key}.bias"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.img_attn.proj.{lora_key}.bias" - ) - converted_state_dict[f"{block_prefix}attn.to_add_out.{lora_key}.weight"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.txt_attn.proj.{lora_key}.weight" - ) - if f"{original_block_prefix}double_blocks.{i}.txt_attn.proj.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}attn.to_add_out.{lora_key}.bias"] = original_state_dict.pop( - f"{original_block_prefix}double_blocks.{i}.txt_attn.proj.{lora_key}.bias" - ) - - # single transformer blocks - for i in range(num_single_layers): - block_prefix = f"single_transformer_blocks.{i}." - - for lora_key in ["lora_A", "lora_B"]: - # norm.linear <- single_blocks.0.modulation.lin - converted_state_dict[f"{block_prefix}norm.linear.{lora_key}.weight"] = original_state_dict.pop( - f"{original_block_prefix}single_blocks.{i}.modulation.lin.{lora_key}.weight" - ) - if f"{original_block_prefix}single_blocks.{i}.modulation.lin.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}norm.linear.{lora_key}.bias"] = original_state_dict.pop( - f"{original_block_prefix}single_blocks.{i}.modulation.lin.{lora_key}.bias" - ) - - # Q, K, V, mlp - mlp_hidden_dim = int(inner_dim * mlp_ratio) - split_size = (inner_dim, inner_dim, inner_dim, mlp_hidden_dim) - - if lora_key == "lora_A": - lora_weight = original_state_dict.pop( - f"{original_block_prefix}single_blocks.{i}.linear1.{lora_key}.weight" - ) - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.weight"] = torch.cat([lora_weight]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.weight"] = torch.cat([lora_weight]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.weight"] = torch.cat([lora_weight]) - converted_state_dict[f"{block_prefix}proj_mlp.{lora_key}.weight"] = torch.cat([lora_weight]) - - if f"{original_block_prefix}single_blocks.{i}.linear1.{lora_key}.bias" in original_state_dict_keys: - lora_bias = original_state_dict.pop(f"single_blocks.{i}.linear1.{lora_key}.bias") - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.bias"] = torch.cat([lora_bias]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.bias"] = torch.cat([lora_bias]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.bias"] = torch.cat([lora_bias]) - converted_state_dict[f"{block_prefix}proj_mlp.{lora_key}.bias"] = torch.cat([lora_bias]) - else: - q, k, v, mlp = torch.split( - original_state_dict.pop(f"{original_block_prefix}single_blocks.{i}.linear1.{lora_key}.weight"), - split_size, - dim=0, - ) - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.weight"] = torch.cat([q]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.weight"] = torch.cat([k]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.weight"] = torch.cat([v]) - converted_state_dict[f"{block_prefix}proj_mlp.{lora_key}.weight"] = torch.cat([mlp]) - - if f"{original_block_prefix}single_blocks.{i}.linear1.{lora_key}.bias" in original_state_dict_keys: - q_bias, k_bias, v_bias, mlp_bias = torch.split( - original_state_dict.pop(f"{original_block_prefix}single_blocks.{i}.linear1.{lora_key}.bias"), - split_size, - dim=0, - ) - converted_state_dict[f"{block_prefix}attn.to_q.{lora_key}.bias"] = torch.cat([q_bias]) - converted_state_dict[f"{block_prefix}attn.to_k.{lora_key}.bias"] = torch.cat([k_bias]) - converted_state_dict[f"{block_prefix}attn.to_v.{lora_key}.bias"] = torch.cat([v_bias]) - converted_state_dict[f"{block_prefix}proj_mlp.{lora_key}.bias"] = torch.cat([mlp_bias]) - - # output projections. - converted_state_dict[f"{block_prefix}proj_out.{lora_key}.weight"] = original_state_dict.pop( - f"{original_block_prefix}single_blocks.{i}.linear2.{lora_key}.weight" - ) - if f"{original_block_prefix}single_blocks.{i}.linear2.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"{block_prefix}proj_out.{lora_key}.bias"] = original_state_dict.pop( - f"{original_block_prefix}single_blocks.{i}.linear2.{lora_key}.bias" - ) - - for lora_key in ["lora_A", "lora_B"]: - converted_state_dict[f"proj_out.{lora_key}.weight"] = original_state_dict.pop( - f"{original_block_prefix}final_layer.linear.{lora_key}.weight" - ) - if f"{original_block_prefix}final_layer.linear.{lora_key}.bias" in original_state_dict_keys: - converted_state_dict[f"proj_out.{lora_key}.bias"] = original_state_dict.pop( - f"{original_block_prefix}final_layer.linear.{lora_key}.bias" - ) - - if len(original_state_dict) > 0: - raise ValueError(f"`original_state_dict` should be empty at this point but has {original_state_dict.keys()=}.") - - for key in list(converted_state_dict.keys()): - converted_state_dict[f"transformer.{key}"] = converted_state_dict.pop(key) - - return converted_state_dict - - -def _convert_hunyuan_video_lora_to_diffusers(original_state_dict): - converted_state_dict = {k: original_state_dict.pop(k) for k in list(original_state_dict.keys())} - - def remap_norm_scale_shift_(key, state_dict): - weight = state_dict.pop(key) - shift, scale = weight.chunk(2, dim=0) - new_weight = torch.cat([scale, shift], dim=0) - state_dict[key.replace("final_layer.adaLN_modulation.1", "norm_out.linear")] = new_weight - - def remap_txt_in_(key, state_dict): - def rename_key(key): - new_key = key.replace("individual_token_refiner.blocks", "token_refiner.refiner_blocks") - new_key = new_key.replace("adaLN_modulation.1", "norm_out.linear") - new_key = new_key.replace("txt_in", "context_embedder") - new_key = new_key.replace("t_embedder.mlp.0", "time_text_embed.timestep_embedder.linear_1") - new_key = new_key.replace("t_embedder.mlp.2", "time_text_embed.timestep_embedder.linear_2") - new_key = new_key.replace("c_embedder", "time_text_embed.text_embedder") - new_key = new_key.replace("mlp", "ff") - return new_key - - if "self_attn_qkv" in key: - weight = state_dict.pop(key) - to_q, to_k, to_v = weight.chunk(3, dim=0) - state_dict[rename_key(key.replace("self_attn_qkv", "attn.to_q"))] = to_q - state_dict[rename_key(key.replace("self_attn_qkv", "attn.to_k"))] = to_k - state_dict[rename_key(key.replace("self_attn_qkv", "attn.to_v"))] = to_v - else: - state_dict[rename_key(key)] = state_dict.pop(key) - - def remap_img_attn_qkv_(key, state_dict): - weight = state_dict.pop(key) - if "lora_A" in key: - state_dict[key.replace("img_attn_qkv", "attn.to_q")] = weight - state_dict[key.replace("img_attn_qkv", "attn.to_k")] = weight - state_dict[key.replace("img_attn_qkv", "attn.to_v")] = weight - else: - to_q, to_k, to_v = weight.chunk(3, dim=0) - state_dict[key.replace("img_attn_qkv", "attn.to_q")] = to_q - state_dict[key.replace("img_attn_qkv", "attn.to_k")] = to_k - state_dict[key.replace("img_attn_qkv", "attn.to_v")] = to_v - - def remap_txt_attn_qkv_(key, state_dict): - weight = state_dict.pop(key) - if "lora_A" in key: - state_dict[key.replace("txt_attn_qkv", "attn.add_q_proj")] = weight - state_dict[key.replace("txt_attn_qkv", "attn.add_k_proj")] = weight - state_dict[key.replace("txt_attn_qkv", "attn.add_v_proj")] = weight - else: - to_q, to_k, to_v = weight.chunk(3, dim=0) - state_dict[key.replace("txt_attn_qkv", "attn.add_q_proj")] = to_q - state_dict[key.replace("txt_attn_qkv", "attn.add_k_proj")] = to_k - state_dict[key.replace("txt_attn_qkv", "attn.add_v_proj")] = to_v - - def remap_single_transformer_blocks_(key, state_dict): - hidden_size = 3072 - - if "linear1.lora_A.weight" in key or "linear1.lora_B.weight" in key: - linear1_weight = state_dict.pop(key) - if "lora_A" in key: - new_key = key.replace("single_blocks", "single_transformer_blocks").removesuffix( - ".linear1.lora_A.weight" - ) - state_dict[f"{new_key}.attn.to_q.lora_A.weight"] = linear1_weight - state_dict[f"{new_key}.attn.to_k.lora_A.weight"] = linear1_weight - state_dict[f"{new_key}.attn.to_v.lora_A.weight"] = linear1_weight - state_dict[f"{new_key}.proj_mlp.lora_A.weight"] = linear1_weight - else: - split_size = (hidden_size, hidden_size, hidden_size, linear1_weight.size(0) - 3 * hidden_size) - q, k, v, mlp = torch.split(linear1_weight, split_size, dim=0) - new_key = key.replace("single_blocks", "single_transformer_blocks").removesuffix( - ".linear1.lora_B.weight" - ) - state_dict[f"{new_key}.attn.to_q.lora_B.weight"] = q - state_dict[f"{new_key}.attn.to_k.lora_B.weight"] = k - state_dict[f"{new_key}.attn.to_v.lora_B.weight"] = v - state_dict[f"{new_key}.proj_mlp.lora_B.weight"] = mlp - - elif "linear1.lora_A.bias" in key or "linear1.lora_B.bias" in key: - linear1_bias = state_dict.pop(key) - if "lora_A" in key: - new_key = key.replace("single_blocks", "single_transformer_blocks").removesuffix( - ".linear1.lora_A.bias" - ) - state_dict[f"{new_key}.attn.to_q.lora_A.bias"] = linear1_bias - state_dict[f"{new_key}.attn.to_k.lora_A.bias"] = linear1_bias - state_dict[f"{new_key}.attn.to_v.lora_A.bias"] = linear1_bias - state_dict[f"{new_key}.proj_mlp.lora_A.bias"] = linear1_bias - else: - split_size = (hidden_size, hidden_size, hidden_size, linear1_bias.size(0) - 3 * hidden_size) - q_bias, k_bias, v_bias, mlp_bias = torch.split(linear1_bias, split_size, dim=0) - new_key = key.replace("single_blocks", "single_transformer_blocks").removesuffix( - ".linear1.lora_B.bias" - ) - state_dict[f"{new_key}.attn.to_q.lora_B.bias"] = q_bias - state_dict[f"{new_key}.attn.to_k.lora_B.bias"] = k_bias - state_dict[f"{new_key}.attn.to_v.lora_B.bias"] = v_bias - state_dict[f"{new_key}.proj_mlp.lora_B.bias"] = mlp_bias - - else: - new_key = key.replace("single_blocks", "single_transformer_blocks") - new_key = new_key.replace("linear2", "proj_out") - new_key = new_key.replace("q_norm", "attn.norm_q") - new_key = new_key.replace("k_norm", "attn.norm_k") - state_dict[new_key] = state_dict.pop(key) - - TRANSFORMER_KEYS_RENAME_DICT = { - "img_in": "x_embedder", - "time_in.mlp.0": "time_text_embed.timestep_embedder.linear_1", - "time_in.mlp.2": "time_text_embed.timestep_embedder.linear_2", - "guidance_in.mlp.0": "time_text_embed.guidance_embedder.linear_1", - "guidance_in.mlp.2": "time_text_embed.guidance_embedder.linear_2", - "vector_in.in_layer": "time_text_embed.text_embedder.linear_1", - "vector_in.out_layer": "time_text_embed.text_embedder.linear_2", - "double_blocks": "transformer_blocks", - "img_attn_q_norm": "attn.norm_q", - "img_attn_k_norm": "attn.norm_k", - "img_attn_proj": "attn.to_out.0", - "txt_attn_q_norm": "attn.norm_added_q", - "txt_attn_k_norm": "attn.norm_added_k", - "txt_attn_proj": "attn.to_add_out", - "img_mod.linear": "norm1.linear", - "img_norm1": "norm1.norm", - "img_norm2": "norm2", - "img_mlp": "ff", - "txt_mod.linear": "norm1_context.linear", - "txt_norm1": "norm1.norm", - "txt_norm2": "norm2_context", - "txt_mlp": "ff_context", - "self_attn_proj": "attn.to_out.0", - "modulation.linear": "norm.linear", - "pre_norm": "norm.norm", - "final_layer.norm_final": "norm_out.norm", - "final_layer.linear": "proj_out", - "fc1": "net.0.proj", - "fc2": "net.2", - "input_embedder": "proj_in", - } - - TRANSFORMER_SPECIAL_KEYS_REMAP = { - "txt_in": remap_txt_in_, - "img_attn_qkv": remap_img_attn_qkv_, - "txt_attn_qkv": remap_txt_attn_qkv_, - "single_blocks": remap_single_transformer_blocks_, - "final_layer.adaLN_modulation.1": remap_norm_scale_shift_, - } - - # Some folks attempt to make their state dict compatible with diffusers by adding "transformer." prefix to all keys - # and use their custom code. To make sure both "original" and "attempted diffusers" loras work as expected, we make - # sure that both follow the same initial format by stripping off the "transformer." prefix. - for key in list(converted_state_dict.keys()): - if key.startswith("transformer."): - converted_state_dict[key[len("transformer.") :]] = converted_state_dict.pop(key) - if key.startswith("diffusion_model."): - converted_state_dict[key[len("diffusion_model.") :]] = converted_state_dict.pop(key) - - # Rename and remap the state dict keys - for key in list(converted_state_dict.keys()): - new_key = key[:] - for replace_key, rename_key in TRANSFORMER_KEYS_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - converted_state_dict[new_key] = converted_state_dict.pop(key) - - for key in list(converted_state_dict.keys()): - for special_key, handler_fn_inplace in TRANSFORMER_SPECIAL_KEYS_REMAP.items(): - if special_key not in key: - continue - handler_fn_inplace(key, converted_state_dict) - - # Add back the "transformer." prefix - for key in list(converted_state_dict.keys()): - converted_state_dict[f"transformer.{key}"] = converted_state_dict.pop(key) - - return converted_state_dict - - -def _convert_non_diffusers_lumina2_lora_to_diffusers(state_dict): - # Remove "diffusion_model." prefix from keys. - state_dict = {k[len("diffusion_model.") :]: v for k, v in state_dict.items()} - converted_state_dict = {} - - def get_num_layers(keys, pattern): - layers = set() - for key in keys: - match = re.search(pattern, key) - if match: - layers.add(int(match.group(1))) - return len(layers) - - def process_block(prefix, index, convert_norm): - # Process attention qkv: pop lora_A and lora_B weights. - lora_down = state_dict.pop(f"{prefix}.{index}.attention.qkv.lora_A.weight") - lora_up = state_dict.pop(f"{prefix}.{index}.attention.qkv.lora_B.weight") - for attn_key in ["to_q", "to_k", "to_v"]: - converted_state_dict[f"{prefix}.{index}.attn.{attn_key}.lora_A.weight"] = lora_down - for attn_key, weight in zip(["to_q", "to_k", "to_v"], torch.split(lora_up, [2304, 768, 768], dim=0)): - converted_state_dict[f"{prefix}.{index}.attn.{attn_key}.lora_B.weight"] = weight - - # Process attention out weights. - converted_state_dict[f"{prefix}.{index}.attn.to_out.0.lora_A.weight"] = state_dict.pop( - f"{prefix}.{index}.attention.out.lora_A.weight" - ) - converted_state_dict[f"{prefix}.{index}.attn.to_out.0.lora_B.weight"] = state_dict.pop( - f"{prefix}.{index}.attention.out.lora_B.weight" - ) - - # Process feed-forward weights for layers 1, 2, and 3. - for layer in range(1, 4): - converted_state_dict[f"{prefix}.{index}.feed_forward.linear_{layer}.lora_A.weight"] = state_dict.pop( - f"{prefix}.{index}.feed_forward.w{layer}.lora_A.weight" - ) - converted_state_dict[f"{prefix}.{index}.feed_forward.linear_{layer}.lora_B.weight"] = state_dict.pop( - f"{prefix}.{index}.feed_forward.w{layer}.lora_B.weight" - ) - - if convert_norm: - converted_state_dict[f"{prefix}.{index}.norm1.linear.lora_A.weight"] = state_dict.pop( - f"{prefix}.{index}.adaLN_modulation.1.lora_A.weight" - ) - converted_state_dict[f"{prefix}.{index}.norm1.linear.lora_B.weight"] = state_dict.pop( - f"{prefix}.{index}.adaLN_modulation.1.lora_B.weight" - ) - - noise_refiner_pattern = r"noise_refiner\.(\d+)\." - num_noise_refiner_layers = get_num_layers(state_dict.keys(), noise_refiner_pattern) - for i in range(num_noise_refiner_layers): - process_block("noise_refiner", i, convert_norm=True) - - context_refiner_pattern = r"context_refiner\.(\d+)\." - num_context_refiner_layers = get_num_layers(state_dict.keys(), context_refiner_pattern) - for i in range(num_context_refiner_layers): - process_block("context_refiner", i, convert_norm=False) - - core_transformer_pattern = r"layers\.(\d+)\." - num_core_transformer_layers = get_num_layers(state_dict.keys(), core_transformer_pattern) - for i in range(num_core_transformer_layers): - process_block("layers", i, convert_norm=True) - - if len(state_dict) > 0: - raise ValueError(f"`state_dict` should be empty at this point but has {state_dict.keys()=}") - - for key in list(converted_state_dict.keys()): - converted_state_dict[f"transformer.{key}"] = converted_state_dict.pop(key) - - return converted_state_dict - - -def _convert_non_diffusers_wan_lora_to_diffusers(state_dict): - converted_state_dict = {} - original_state_dict = {k[len("diffusion_model.") :]: v for k, v in state_dict.items()} - - block_numbers = {int(k.split(".")[1]) for k in original_state_dict if k.startswith("blocks.")} - min_block = min(block_numbers) - max_block = max(block_numbers) - - is_i2v_lora = any("k_img" in k for k in original_state_dict) and any("v_img" in k for k in original_state_dict) - lora_down_key = "lora_A" if any("lora_A" in k for k in original_state_dict) else "lora_down" - lora_up_key = "lora_B" if any("lora_B" in k for k in original_state_dict) else "lora_up" - has_time_projection_weight = any( - k.startswith("time_projection") and k.endswith(".weight") for k in original_state_dict - ) - - def get_alpha_scales(down_weight, alpha_key): - rank = down_weight.shape[0] - alpha = original_state_dict.pop(alpha_key).item() - scale = alpha / rank # LoRA is scaled by 'alpha / rank' in forward pass, so we need to scale it back here - scale_down = scale - scale_up = 1.0 - while scale_down * 2 < scale_up: - scale_down *= 2 - scale_up /= 2 - return scale_down, scale_up - - for key in list(original_state_dict.keys()): - if key.endswith((".diff", ".diff_b")) and "norm" in key: - # NOTE: we don't support this because norm layer diff keys are just zeroed values. We can support it - # in future if needed and they are not zeroed. - original_state_dict.pop(key) - logger.debug(f"Removing {key} key from the state dict as it is a norm diff key. This is unsupported.") - - if "time_projection" in key and not has_time_projection_weight: - # AccVideo lora has diff bias keys but not the weight keys. This causes a weird problem where - # our lora config adds the time proj lora layers, but we don't have the weights for them. - # CausVid lora has the weight keys and the bias keys. - original_state_dict.pop(key) - - # For the `diff_b` keys, we treat them as lora_bias. - # https://huggingface.co/docs/peft/main/en/package_reference/lora#peft.LoraConfig.lora_bias - - for i in range(min_block, max_block + 1): - # Self-attention - for o, c in zip(["q", "k", "v", "o"], ["to_q", "to_k", "to_v", "to_out.0"]): - alpha_key = f"blocks.{i}.self_attn.{o}.alpha" - has_alpha = alpha_key in original_state_dict - original_key_A = f"blocks.{i}.self_attn.{o}.{lora_down_key}.weight" - converted_key_A = f"blocks.{i}.attn1.{c}.lora_A.weight" - - original_key_B = f"blocks.{i}.self_attn.{o}.{lora_up_key}.weight" - converted_key_B = f"blocks.{i}.attn1.{c}.lora_B.weight" - - if has_alpha: - down_weight = original_state_dict.pop(original_key_A) - up_weight = original_state_dict.pop(original_key_B) - scale_down, scale_up = get_alpha_scales(down_weight, alpha_key) - converted_state_dict[converted_key_A] = down_weight * scale_down - converted_state_dict[converted_key_B] = up_weight * scale_up - - else: - if original_key_A in original_state_dict: - converted_state_dict[converted_key_A] = original_state_dict.pop(original_key_A) - if original_key_B in original_state_dict: - converted_state_dict[converted_key_B] = original_state_dict.pop(original_key_B) - - original_key = f"blocks.{i}.self_attn.{o}.diff_b" - converted_key = f"blocks.{i}.attn1.{c}.lora_B.bias" - if original_key in original_state_dict: - converted_state_dict[converted_key] = original_state_dict.pop(original_key) - - # Cross-attention - for o, c in zip(["q", "k", "v", "o"], ["to_q", "to_k", "to_v", "to_out.0"]): - alpha_key = f"blocks.{i}.cross_attn.{o}.alpha" - has_alpha = alpha_key in original_state_dict - original_key_A = f"blocks.{i}.cross_attn.{o}.{lora_down_key}.weight" - converted_key_A = f"blocks.{i}.attn2.{c}.lora_A.weight" - - original_key_B = f"blocks.{i}.cross_attn.{o}.{lora_up_key}.weight" - converted_key_B = f"blocks.{i}.attn2.{c}.lora_B.weight" - - if original_key_A in original_state_dict: - down_weight = original_state_dict.pop(original_key_A) - converted_state_dict[converted_key_A] = down_weight - if original_key_B in original_state_dict: - up_weight = original_state_dict.pop(original_key_B) - converted_state_dict[converted_key_B] = up_weight - if has_alpha: - scale_down, scale_up = get_alpha_scales(down_weight, alpha_key) - converted_state_dict[converted_key_A] *= scale_down - converted_state_dict[converted_key_B] *= scale_up - - original_key = f"blocks.{i}.cross_attn.{o}.diff_b" - converted_key = f"blocks.{i}.attn2.{c}.lora_B.bias" - if original_key in original_state_dict: - converted_state_dict[converted_key] = original_state_dict.pop(original_key) - - if is_i2v_lora: - for o, c in zip(["k_img", "v_img"], ["add_k_proj", "add_v_proj"]): - alpha_key = f"blocks.{i}.cross_attn.{o}.alpha" - has_alpha = alpha_key in original_state_dict - original_key_A = f"blocks.{i}.cross_attn.{o}.{lora_down_key}.weight" - converted_key_A = f"blocks.{i}.attn2.{c}.lora_A.weight" - - original_key_B = f"blocks.{i}.cross_attn.{o}.{lora_up_key}.weight" - converted_key_B = f"blocks.{i}.attn2.{c}.lora_B.weight" - - if original_key_A in original_state_dict: - down_weight = original_state_dict.pop(original_key_A) - converted_state_dict[converted_key_A] = down_weight - if original_key_B in original_state_dict: - up_weight = original_state_dict.pop(original_key_B) - converted_state_dict[converted_key_B] = up_weight - if has_alpha: - scale_down, scale_up = get_alpha_scales(down_weight, alpha_key) - converted_state_dict[converted_key_A] *= scale_down - converted_state_dict[converted_key_B] *= scale_up - - original_key = f"blocks.{i}.cross_attn.{o}.diff_b" - converted_key = f"blocks.{i}.attn2.{c}.lora_B.bias" - if original_key in original_state_dict: - converted_state_dict[converted_key] = original_state_dict.pop(original_key) - - # FFN - for o, c in zip(["ffn.0", "ffn.2"], ["net.0.proj", "net.2"]): - alpha_key = f"blocks.{i}.{o}.alpha" - has_alpha = alpha_key in original_state_dict - original_key_A = f"blocks.{i}.{o}.{lora_down_key}.weight" - converted_key_A = f"blocks.{i}.ffn.{c}.lora_A.weight" - - original_key_B = f"blocks.{i}.{o}.{lora_up_key}.weight" - converted_key_B = f"blocks.{i}.ffn.{c}.lora_B.weight" - - if original_key_A in original_state_dict: - down_weight = original_state_dict.pop(original_key_A) - converted_state_dict[converted_key_A] = down_weight - if original_key_B in original_state_dict: - up_weight = original_state_dict.pop(original_key_B) - converted_state_dict[converted_key_B] = up_weight - if has_alpha: - scale_down, scale_up = get_alpha_scales(down_weight, alpha_key) - converted_state_dict[converted_key_A] *= scale_down - converted_state_dict[converted_key_B] *= scale_up - - original_key = f"blocks.{i}.{o}.diff_b" - converted_key = f"blocks.{i}.ffn.{c}.lora_B.bias" - if original_key in original_state_dict: - converted_state_dict[converted_key] = original_state_dict.pop(original_key) - - # Remaining. - if original_state_dict: - if any("time_projection" in k for k in original_state_dict): - original_key = f"time_projection.1.{lora_down_key}.weight" - converted_key = "condition_embedder.time_proj.lora_A.weight" - if original_key in original_state_dict: - converted_state_dict[converted_key] = original_state_dict.pop(original_key) - - original_key = f"time_projection.1.{lora_up_key}.weight" - converted_key = "condition_embedder.time_proj.lora_B.weight" - if original_key in original_state_dict: - converted_state_dict[converted_key] = original_state_dict.pop(original_key) - - if "time_projection.1.diff_b" in original_state_dict: - converted_state_dict["condition_embedder.time_proj.lora_B.bias"] = original_state_dict.pop( - "time_projection.1.diff_b" - ) - - if any("head.head" in k for k in original_state_dict): - if any(f"head.head.{lora_down_key}.weight" in k for k in state_dict): - converted_state_dict["proj_out.lora_A.weight"] = original_state_dict.pop( - f"head.head.{lora_down_key}.weight" - ) - if any(f"head.head.{lora_up_key}.weight" in k for k in state_dict): - converted_state_dict["proj_out.lora_B.weight"] = original_state_dict.pop( - f"head.head.{lora_up_key}.weight" - ) - if "head.head.diff_b" in original_state_dict: - converted_state_dict["proj_out.lora_B.bias"] = original_state_dict.pop("head.head.diff_b") - - # Notes: https://huggingface.co/lightx2v/Wan2.2-Distill-Loras - # This is my (sayakpaul) assumption that this particular key belongs to the down matrix. - # Since for this particular LoRA, we don't have the corresponding up matrix, I will use - # an identity. - if any("head.head" in k and k.endswith(".diff") for k in state_dict): - if f"head.head.{lora_down_key}.weight" in state_dict: - logger.info( - f"The state dict seems to be have both `head.head.diff` and `head.head.{lora_down_key}.weight` keys, which is unexpected." - ) - converted_state_dict["proj_out.lora_A.weight"] = original_state_dict.pop("head.head.diff") - down_matrix_head = converted_state_dict["proj_out.lora_A.weight"] - up_matrix_shape = (down_matrix_head.shape[0], converted_state_dict["proj_out.lora_B.bias"].shape[0]) - converted_state_dict["proj_out.lora_B.weight"] = torch.eye( - *up_matrix_shape, dtype=down_matrix_head.dtype, device=down_matrix_head.device - ).T - - for text_time in ["text_embedding", "time_embedding"]: - if any(text_time in k for k in original_state_dict): - for b_n in [0, 2]: - diffusers_b_n = 1 if b_n == 0 else 2 - diffusers_name = ( - "condition_embedder.text_embedder" - if text_time == "text_embedding" - else "condition_embedder.time_embedder" - ) - if any(f"{text_time}.{b_n}" in k for k in original_state_dict): - converted_state_dict[f"{diffusers_name}.linear_{diffusers_b_n}.lora_A.weight"] = ( - original_state_dict.pop(f"{text_time}.{b_n}.{lora_down_key}.weight") - ) - converted_state_dict[f"{diffusers_name}.linear_{diffusers_b_n}.lora_B.weight"] = ( - original_state_dict.pop(f"{text_time}.{b_n}.{lora_up_key}.weight") - ) - if f"{text_time}.{b_n}.diff_b" in original_state_dict: - converted_state_dict[f"{diffusers_name}.linear_{diffusers_b_n}.lora_B.bias"] = ( - original_state_dict.pop(f"{text_time}.{b_n}.diff_b") - ) - - for img_ours, img_theirs in [ - ("ff.net.0.proj", "img_emb.proj.1"), - ("ff.net.2", "img_emb.proj.3"), - ]: - original_key = f"{img_theirs}.{lora_down_key}.weight" - converted_key = f"condition_embedder.image_embedder.{img_ours}.lora_A.weight" - if original_key in original_state_dict: - converted_state_dict[converted_key] = original_state_dict.pop(original_key) - - original_key = f"{img_theirs}.{lora_up_key}.weight" - converted_key = f"condition_embedder.image_embedder.{img_ours}.lora_B.weight" - if original_key in original_state_dict: - converted_state_dict[converted_key] = original_state_dict.pop(original_key) - bias_key_theirs = original_key.removesuffix(f".{lora_up_key}.weight") + ".diff_b" - if bias_key_theirs in original_state_dict: - bias_key = converted_key.removesuffix(".weight") + ".bias" - converted_state_dict[bias_key] = original_state_dict.pop(bias_key_theirs) - - if len(original_state_dict) > 0: - diff = all(".diff" in k for k in original_state_dict) - if diff: - diff_keys = {k for k in original_state_dict if k.endswith(".diff")} - if not all("lora" not in k for k in diff_keys): - raise ValueError - logger.info( - "The remaining `state_dict` contains `diff` keys which we do not handle yet. If you see performance issues, please file an issue: " - "https://github.com/huggingface/diffusers//issues/new" - ) - else: - raise ValueError(f"`state_dict` should be empty at this point but has {original_state_dict.keys()=}") - - for key in list(converted_state_dict.keys()): - converted_state_dict[f"transformer.{key}"] = converted_state_dict.pop(key) - - return converted_state_dict - - -def _convert_musubi_wan_lora_to_diffusers(state_dict): - # https://github.com/kohya-ss/musubi-tuner - converted_state_dict = {} - original_state_dict = {k[len("lora_unet_") :]: v for k, v in state_dict.items()} - - num_blocks = len({k.split("blocks_")[1].split("_")[0] for k in original_state_dict}) - is_i2v_lora = any("k_img" in k for k in original_state_dict) and any("v_img" in k for k in original_state_dict) - - def get_alpha_scales(down_weight, key): - rank = down_weight.shape[0] - alpha = original_state_dict.pop(key + ".alpha").item() - scale = alpha / rank # LoRA is scaled by 'alpha / rank' in forward pass, so we need to scale it back here - scale_down = scale - scale_up = 1.0 - while scale_down * 2 < scale_up: - scale_down *= 2 - scale_up /= 2 - return scale_down, scale_up - - for i in range(num_blocks): - # Self-attention - for o, c in zip(["q", "k", "v", "o"], ["to_q", "to_k", "to_v", "to_out.0"]): - down_weight = original_state_dict.pop(f"blocks_{i}_self_attn_{o}.lora_down.weight") - up_weight = original_state_dict.pop(f"blocks_{i}_self_attn_{o}.lora_up.weight") - scale_down, scale_up = get_alpha_scales(down_weight, f"blocks_{i}_self_attn_{o}") - converted_state_dict[f"blocks.{i}.attn1.{c}.lora_A.weight"] = down_weight * scale_down - converted_state_dict[f"blocks.{i}.attn1.{c}.lora_B.weight"] = up_weight * scale_up - - # Cross-attention - for o, c in zip(["q", "k", "v", "o"], ["to_q", "to_k", "to_v", "to_out.0"]): - down_weight = original_state_dict.pop(f"blocks_{i}_cross_attn_{o}.lora_down.weight") - up_weight = original_state_dict.pop(f"blocks_{i}_cross_attn_{o}.lora_up.weight") - scale_down, scale_up = get_alpha_scales(down_weight, f"blocks_{i}_cross_attn_{o}") - converted_state_dict[f"blocks.{i}.attn2.{c}.lora_A.weight"] = down_weight * scale_down - converted_state_dict[f"blocks.{i}.attn2.{c}.lora_B.weight"] = up_weight * scale_up - - if is_i2v_lora: - for o, c in zip(["k_img", "v_img"], ["add_k_proj", "add_v_proj"]): - down_weight = original_state_dict.pop(f"blocks_{i}_cross_attn_{o}.lora_down.weight") - up_weight = original_state_dict.pop(f"blocks_{i}_cross_attn_{o}.lora_up.weight") - scale_down, scale_up = get_alpha_scales(down_weight, f"blocks_{i}_cross_attn_{o}") - converted_state_dict[f"blocks.{i}.attn2.{c}.lora_A.weight"] = down_weight * scale_down - converted_state_dict[f"blocks.{i}.attn2.{c}.lora_B.weight"] = up_weight * scale_up - - # FFN - for o, c in zip(["ffn_0", "ffn_2"], ["net.0.proj", "net.2"]): - down_weight = original_state_dict.pop(f"blocks_{i}_{o}.lora_down.weight") - up_weight = original_state_dict.pop(f"blocks_{i}_{o}.lora_up.weight") - scale_down, scale_up = get_alpha_scales(down_weight, f"blocks_{i}_{o}") - converted_state_dict[f"blocks.{i}.ffn.{c}.lora_A.weight"] = down_weight * scale_down - converted_state_dict[f"blocks.{i}.ffn.{c}.lora_B.weight"] = up_weight * scale_up - - if len(original_state_dict) > 0: - raise ValueError(f"`state_dict` should be empty at this point but has {original_state_dict.keys()=}") - - for key in list(converted_state_dict.keys()): - converted_state_dict[f"transformer.{key}"] = converted_state_dict.pop(key) - - return converted_state_dict - - -def _convert_non_diffusers_hidream_lora_to_diffusers(state_dict, non_diffusers_prefix="diffusion_model"): - if not all(k.startswith(non_diffusers_prefix) for k in state_dict): - raise ValueError("Invalid LoRA state dict for HiDream.") - converted_state_dict = {k.removeprefix(f"{non_diffusers_prefix}."): v for k, v in state_dict.items()} - converted_state_dict = {f"transformer.{k}": v for k, v in converted_state_dict.items()} - return converted_state_dict - - -def _convert_non_diffusers_ltxv_lora_to_diffusers(state_dict, non_diffusers_prefix="diffusion_model"): - if not all(k.startswith(f"{non_diffusers_prefix}.") for k in state_dict): - raise ValueError("Invalid LoRA state dict for LTX-Video.") - converted_state_dict = {k.removeprefix(f"{non_diffusers_prefix}."): v for k, v in state_dict.items()} - converted_state_dict = {f"transformer.{k}": v for k, v in converted_state_dict.items()} - return converted_state_dict - - -def _convert_non_diffusers_ltx2_lora_to_diffusers(state_dict, non_diffusers_prefix="diffusion_model"): - # Remove the prefix - state_dict = {k: v for k, v in state_dict.items() if k.startswith(f"{non_diffusers_prefix}.")} - converted_state_dict = {k.removeprefix(f"{non_diffusers_prefix}."): v for k, v in state_dict.items()} - - if non_diffusers_prefix == "diffusion_model": - rename_dict = { - "patchify_proj": "proj_in", - "audio_patchify_proj": "audio_proj_in", - "av_ca_video_scale_shift_adaln_single": "av_cross_attn_video_scale_shift", - "av_ca_a2v_gate_adaln_single": "av_cross_attn_video_a2v_gate", - "av_ca_audio_scale_shift_adaln_single": "av_cross_attn_audio_scale_shift", - "av_ca_v2a_gate_adaln_single": "av_cross_attn_audio_v2a_gate", - "scale_shift_table_a2v_ca_video": "video_a2v_cross_attn_scale_shift_table", - "scale_shift_table_a2v_ca_audio": "audio_a2v_cross_attn_scale_shift_table", - "q_norm": "norm_q", - "k_norm": "norm_k", - # LTX-2.3 - "audio_prompt_adaln_single": "audio_prompt_adaln", - "prompt_adaln_single": "prompt_adaln", - } - else: - rename_dict = {"aggregate_embed": "text_proj_in"} - - # Apply renaming - renamed_state_dict = {} - for key, value in converted_state_dict.items(): - new_key = key[:] - for old_pattern, new_pattern in rename_dict.items(): - new_key = new_key.replace(old_pattern, new_pattern) - renamed_state_dict[new_key] = value - - # Handle adaln_single -> time_embed and audio_adaln_single -> audio_time_embed - final_state_dict = {} - for key, value in renamed_state_dict.items(): - if key.startswith("adaln_single."): - new_key = key.replace("adaln_single.", "time_embed.") - final_state_dict[new_key] = value - elif key.startswith("audio_adaln_single."): - new_key = key.replace("audio_adaln_single.", "audio_time_embed.") - final_state_dict[new_key] = value - else: - final_state_dict[key] = value - - # Add transformer prefix - prefix = "transformer" if non_diffusers_prefix == "diffusion_model" else "connectors" - final_state_dict = {f"{prefix}.{k}": v for k, v in final_state_dict.items()} - - return final_state_dict - - -def _convert_non_diffusers_qwen_lora_to_diffusers(state_dict): - has_diffusion_model = any(k.startswith("diffusion_model.") for k in state_dict) - if has_diffusion_model: - state_dict = {k.removeprefix("diffusion_model."): v for k, v in state_dict.items()} - - has_lora_unet = any(k.startswith("lora_unet_") for k in state_dict) - if has_lora_unet: - state_dict = {k.removeprefix("lora_unet_"): v for k, v in state_dict.items()} - - # Top-level (non-block) modules: convert_key below assumes every key lives under - # transformer_blocks_ and blindly strips/re-prepends that prefix, which collapses - # these module names onto each other. Map them explicitly before that logic runs. - # The flattened name -> dotted diffusers name is fixed, and the .lora_down/.lora_up/ - # .alpha suffix is preserved. - top_level_modules = { - "img_in": "img_in", - "txt_in": "txt_in", - "proj_out": "proj_out", - "norm_out_linear": "norm_out.linear", - "time_text_embed_timestep_embedder_linear_1": "time_text_embed.timestep_embedder.linear_1", - "time_text_embed_timestep_embedder_linear_2": "time_text_embed.timestep_embedder.linear_2", - } - - def convert_key(key: str) -> str: - prefix = "transformer_blocks" - for flat, dotted in top_level_modules.items(): - if key == flat or key.startswith(flat + "."): - return dotted + key[len(flat) :] - - if "." in key: - base, suffix = key.rsplit(".", 1) - else: - base, suffix = key, "" - - start = f"{prefix}_" - rest = base[len(start) :] - - if "." in rest: - head, tail = rest.split(".", 1) - tail = "." + tail - else: - head, tail = rest, "" - - # Protected n-grams that must keep their internal underscores - protected = { - # pairs - ("to", "q"), - ("to", "k"), - ("to", "v"), - ("to", "out"), - ("add", "q"), - ("add", "k"), - ("add", "v"), - ("txt", "mlp"), - ("img", "mlp"), - ("txt", "mod"), - ("img", "mod"), - # triplets - ("add", "q", "proj"), - ("add", "k", "proj"), - ("add", "v", "proj"), - ("to", "add", "out"), - } - - prot_by_len = {} - for ng in protected: - prot_by_len.setdefault(len(ng), set()).add(ng) - - parts = head.split("_") - merged = [] - i = 0 - lengths_desc = sorted(prot_by_len.keys(), reverse=True) - - while i < len(parts): - matched = False - for L in lengths_desc: - if i + L <= len(parts) and tuple(parts[i : i + L]) in prot_by_len[L]: - merged.append("_".join(parts[i : i + L])) - i += L - matched = True - break - if not matched: - merged.append(parts[i]) - i += 1 - - head_converted = ".".join(merged) - converted_base = f"{prefix}.{head_converted}{tail}" - return converted_base + (("." + suffix) if suffix else "") - - state_dict = {convert_key(k): v for k, v in state_dict.items()} - - has_default = any("default." in k for k in state_dict) - if has_default: - state_dict = {k.replace("default.", ""): v for k, v in state_dict.items()} - - converted_state_dict = {} - all_keys = list(state_dict.keys()) - down_key = ".lora_down.weight" - up_key = ".lora_up.weight" - a_key = ".lora_A.weight" - b_key = ".lora_B.weight" - - has_non_diffusers_lora_id = any(down_key in k or up_key in k for k in all_keys) - has_diffusers_lora_id = any(a_key in k or b_key in k for k in all_keys) - - if has_non_diffusers_lora_id: - - def get_alpha_scales(down_weight, alpha_key): - rank = down_weight.shape[0] - alpha = state_dict.pop(alpha_key).item() - scale = alpha / rank # LoRA is scaled by 'alpha / rank' in forward pass, so we need to scale it back here - scale_down = scale - scale_up = 1.0 - while scale_down * 2 < scale_up: - scale_down *= 2 - scale_up /= 2 - return scale_down, scale_up - - for k in all_keys: - if k.endswith(down_key): - diffusers_down_key = k.replace(down_key, ".lora_A.weight") - diffusers_up_key = k.replace(down_key, up_key).replace(up_key, ".lora_B.weight") - alpha_key = k.replace(down_key, ".alpha") - - down_weight = state_dict.pop(k) - up_weight = state_dict.pop(k.replace(down_key, up_key)) - scale_down, scale_up = get_alpha_scales(down_weight, alpha_key) - converted_state_dict[diffusers_down_key] = down_weight * scale_down - converted_state_dict[diffusers_up_key] = up_weight * scale_up - - # Already in diffusers format (lora_A/lora_B), just pop - elif has_diffusers_lora_id: - for k in all_keys: - if a_key in k or b_key in k: - converted_state_dict[k] = state_dict.pop(k) - elif ".alpha" in k: - state_dict.pop(k) - - if len(state_dict) > 0: - raise ValueError(f"`state_dict` should be empty at this point but has {state_dict.keys()=}") - - converted_state_dict = {f"transformer.{k}": v for k, v in converted_state_dict.items()} - return converted_state_dict - - -def _convert_non_diffusers_anima_lora_to_diffusers(state_dict): - rename_dict = { - "blocks.": "transformer_blocks.", - "adaln_modulation_self_attn.1": "norm1.linear_1", - "adaln_modulation_self_attn.2": "norm1.linear_2", - "adaln_modulation_cross_attn.1": "norm2.linear_1", - "adaln_modulation_cross_attn.2": "norm2.linear_2", - "adaln_modulation_mlp.1": "norm3.linear_1", - "adaln_modulation_mlp.2": "norm3.linear_2", - "self_attn.q_proj": "attn1.to_q", - "self_attn.k_proj": "attn1.to_k", - "self_attn.v_proj": "attn1.to_v", - "self_attn.output_proj": "attn1.to_out.0", - "cross_attn.q_proj": "attn2.to_q", - "cross_attn.k_proj": "attn2.to_k", - "cross_attn.v_proj": "attn2.to_v", - "cross_attn.output_proj": "attn2.to_out.0", - "mlp.layer1": "ff.net.0.proj", - "mlp.layer2": "ff.net.2", - "final_layer.adaln_modulation.1": "norm_out.linear_1", - "final_layer.adaln_modulation.2": "norm_out.linear_2", - "final_layer.linear": "proj_out", - "t_embedder.1": "time_embed.t_embedder", - "t_embedding_norm": "time_embed.norm", - "x_embedder.proj.1": "patch_embed.proj", - } - - converted_state_dict = {} - for key, value in state_dict.items(): - if not key.startswith("diffusion_model."): - converted_state_dict[key] = value - continue - - new_key = key.removeprefix("diffusion_model.") - if new_key.startswith("llm_adapter."): - new_key = f"text_conditioner.{new_key.removeprefix('llm_adapter.')}" - else: - for old_key, new_key_part in rename_dict.items(): - new_key = new_key.replace(old_key, new_key_part) - new_key = f"transformer.{new_key}" - - converted_state_dict[new_key] = value - - return converted_state_dict - - -def _convert_non_diffusers_flux2_lora_to_diffusers(state_dict): - converted_state_dict = {} - - prefix = "diffusion_model." - original_state_dict = {k[len(prefix) :]: v for k, v in state_dict.items()} - - has_lora_down_up = any("lora_down" in k or "lora_up" in k for k in original_state_dict.keys()) - if has_lora_down_up: - temp_state_dict = {} - for k, v in original_state_dict.items(): - new_key = k.replace("lora_down", "lora_A").replace("lora_up", "lora_B") - temp_state_dict[new_key] = v - original_state_dict = temp_state_dict - - # Some Flux2 checkpoints skip the ai-toolkit `single_blocks` / `double_blocks` - # layout and already store expanded diffusers block names. Accept those - # directly, and normalize the legacy `sformer_blocks` alias used by some exports. - possible_expanded_block_prefixes = { - "single_transformer_blocks.": "single_transformer_blocks.", - "transformer_blocks.": "transformer_blocks.", - "sformer_blocks.": "transformer_blocks.", - } - for key in list(original_state_dict.keys()): - for source_prefix, target_prefix in possible_expanded_block_prefixes.items(): - if key.startswith(source_prefix): - converted_state_dict[target_prefix + key[len(source_prefix) :]] = original_state_dict.pop(key) - break - - num_double_layers = 0 - num_single_layers = 0 - for key in original_state_dict.keys(): - if key.startswith("single_blocks."): - num_single_layers = max(num_single_layers, int(key.split(".")[1]) + 1) - elif key.startswith("double_blocks."): - num_double_layers = max(num_double_layers, int(key.split(".")[1]) + 1) - - lora_keys = ("lora_A", "lora_B") - attn_types = ("img_attn", "txt_attn") - - for sl in range(num_single_layers): - single_block_prefix = f"single_blocks.{sl}" - attn_prefix = f"single_transformer_blocks.{sl}.attn" - - for lora_key in lora_keys: - linear1_key = f"{single_block_prefix}.linear1.{lora_key}.weight" - if linear1_key in original_state_dict: - converted_state_dict[f"{attn_prefix}.to_qkv_mlp_proj.{lora_key}.weight"] = original_state_dict.pop( - linear1_key - ) - - linear2_key = f"{single_block_prefix}.linear2.{lora_key}.weight" - if linear2_key in original_state_dict: - converted_state_dict[f"{attn_prefix}.to_out.{lora_key}.weight"] = original_state_dict.pop(linear2_key) - - for dl in range(num_double_layers): - transformer_block_prefix = f"transformer_blocks.{dl}" - - for lora_key in lora_keys: - for attn_type in attn_types: - attn_prefix = f"{transformer_block_prefix}.attn" - qkv_key = f"double_blocks.{dl}.{attn_type}.qkv.{lora_key}.weight" - - if qkv_key not in original_state_dict: - continue - - fused_qkv_weight = original_state_dict.pop(qkv_key) - - if lora_key == "lora_A": - diff_attn_proj_keys = ( - ["to_q", "to_k", "to_v"] - if attn_type == "img_attn" - else ["add_q_proj", "add_k_proj", "add_v_proj"] - ) - for proj_key in diff_attn_proj_keys: - converted_state_dict[f"{attn_prefix}.{proj_key}.{lora_key}.weight"] = torch.cat( - [fused_qkv_weight] - ) - else: - sample_q, sample_k, sample_v = torch.chunk(fused_qkv_weight, 3, dim=0) - - if attn_type == "img_attn": - converted_state_dict[f"{attn_prefix}.to_q.{lora_key}.weight"] = torch.cat([sample_q]) - converted_state_dict[f"{attn_prefix}.to_k.{lora_key}.weight"] = torch.cat([sample_k]) - converted_state_dict[f"{attn_prefix}.to_v.{lora_key}.weight"] = torch.cat([sample_v]) - else: - converted_state_dict[f"{attn_prefix}.add_q_proj.{lora_key}.weight"] = torch.cat([sample_q]) - converted_state_dict[f"{attn_prefix}.add_k_proj.{lora_key}.weight"] = torch.cat([sample_k]) - converted_state_dict[f"{attn_prefix}.add_v_proj.{lora_key}.weight"] = torch.cat([sample_v]) - - proj_mappings = [ - ("img_attn.proj", "attn.to_out.0"), - ("txt_attn.proj", "attn.to_add_out"), - ] - for org_proj, diff_proj in proj_mappings: - for lora_key in lora_keys: - original_key = f"double_blocks.{dl}.{org_proj}.{lora_key}.weight" - if original_key in original_state_dict: - diffusers_key = f"{transformer_block_prefix}.{diff_proj}.{lora_key}.weight" - converted_state_dict[diffusers_key] = original_state_dict.pop(original_key) - - mlp_mappings = [ - ("img_mlp.0", "ff.linear_in"), - ("img_mlp.2", "ff.linear_out"), - ("txt_mlp.0", "ff_context.linear_in"), - ("txt_mlp.2", "ff_context.linear_out"), - ] - for org_mlp, diff_mlp in mlp_mappings: - for lora_key in lora_keys: - original_key = f"double_blocks.{dl}.{org_mlp}.{lora_key}.weight" - if original_key in original_state_dict: - diffusers_key = f"{transformer_block_prefix}.{diff_mlp}.{lora_key}.weight" - converted_state_dict[diffusers_key] = original_state_dict.pop(original_key) - - extra_mappings = { - "img_in": "x_embedder", - "txt_in": "context_embedder", - "time_in.in_layer": "time_guidance_embed.timestep_embedder.linear_1", - "time_in.out_layer": "time_guidance_embed.timestep_embedder.linear_2", - "guidance_in.in_layer": "time_guidance_embed.guidance_embedder.linear_1", - "guidance_in.out_layer": "time_guidance_embed.guidance_embedder.linear_2", - "final_layer.linear": "proj_out", - "final_layer.adaLN_modulation.1": "norm_out.linear", - "single_stream_modulation.lin": "single_stream_modulation.linear", - "double_stream_modulation_img.lin": "double_stream_modulation_img.linear", - "double_stream_modulation_txt.lin": "double_stream_modulation_txt.linear", - } - - for org_key, diff_key in extra_mappings.items(): - for lora_key in lora_keys: - original_key = f"{org_key}.{lora_key}.weight" - if original_key in original_state_dict: - converted_state_dict[f"{diff_key}.{lora_key}.weight"] = original_state_dict.pop(original_key) - - if len(original_state_dict) > 0: - raise ValueError(f"`original_state_dict` should be empty at this point but has {original_state_dict.keys()=}.") - - for key in list(converted_state_dict.keys()): - converted_state_dict[f"transformer.{key}"] = converted_state_dict.pop(key) - - return converted_state_dict - - -def _convert_kohya_flux2_lora_to_diffusers(state_dict): - def _convert_to_ai_toolkit(sds_sd, ait_sd, sds_key, ait_key): - if sds_key + ".lora_down.weight" not in sds_sd: - return - down_weight = sds_sd.pop(sds_key + ".lora_down.weight") - - # scale weight by alpha and dim - rank = down_weight.shape[0] - default_alpha = torch.tensor(rank, dtype=down_weight.dtype, device=down_weight.device, requires_grad=False) - alpha = sds_sd.pop(sds_key + ".alpha", default_alpha).item() - scale = alpha / rank - - scale_down = scale - scale_up = 1.0 - while scale_down * 2 < scale_up: - scale_down *= 2 - scale_up /= 2 - - ait_sd[ait_key + ".lora_A.weight"] = down_weight * scale_down - ait_sd[ait_key + ".lora_B.weight"] = sds_sd.pop(sds_key + ".lora_up.weight") * scale_up - - def _convert_to_ai_toolkit_cat(sds_sd, ait_sd, sds_key, ait_keys, dims=None): - if sds_key + ".lora_down.weight" not in sds_sd: - return - down_weight = sds_sd.pop(sds_key + ".lora_down.weight") - up_weight = sds_sd.pop(sds_key + ".lora_up.weight") - sd_lora_rank = down_weight.shape[0] - - default_alpha = torch.tensor( - sd_lora_rank, dtype=down_weight.dtype, device=down_weight.device, requires_grad=False - ) - alpha = sds_sd.pop(sds_key + ".alpha", default_alpha) - scale = alpha / sd_lora_rank - - scale_down = scale - scale_up = 1.0 - while scale_down * 2 < scale_up: - scale_down *= 2 - scale_up /= 2 - - down_weight = down_weight * scale_down - up_weight = up_weight * scale_up - - num_splits = len(ait_keys) - if dims is None: - dims = [up_weight.shape[0] // num_splits] * num_splits - else: - assert sum(dims) == up_weight.shape[0] - - # check if upweight is sparse - is_sparse = False - if sd_lora_rank % num_splits == 0: - ait_rank = sd_lora_rank // num_splits - is_sparse = True - i = 0 - for j in range(len(dims)): - for k in range(len(dims)): - if j == k: - continue - is_sparse = is_sparse and torch.all( - up_weight[i : i + dims[j], k * ait_rank : (k + 1) * ait_rank] == 0 - ) - i += dims[j] - if is_sparse: - logger.info(f"weight is sparse: {sds_key}") - - ait_down_keys = [k + ".lora_A.weight" for k in ait_keys] - ait_up_keys = [k + ".lora_B.weight" for k in ait_keys] - if not is_sparse: - ait_sd.update(dict.fromkeys(ait_down_keys, down_weight)) - ait_sd.update({k: v for k, v in zip(ait_up_keys, torch.split(up_weight, dims, dim=0))}) # noqa: C416 - else: - ait_sd.update({k: v for k, v in zip(ait_down_keys, torch.chunk(down_weight, num_splits, dim=0))}) # noqa: C416 - i = 0 - for j in range(len(dims)): - ait_sd[ait_up_keys[j]] = up_weight[i : i + dims[j], j * ait_rank : (j + 1) * ait_rank].contiguous() - i += dims[j] - - # Detect number of blocks from keys - num_double_layers = 0 - num_single_layers = 0 - for key in state_dict.keys(): - if key.startswith("lora_unet_double_blocks_"): - block_idx = int(key.split("_")[4]) - num_double_layers = max(num_double_layers, block_idx + 1) - elif key.startswith("lora_unet_single_blocks_"): - block_idx = int(key.split("_")[4]) - num_single_layers = max(num_single_layers, block_idx + 1) - - ait_sd = {} - - for i in range(num_double_layers): - # Attention projections - _convert_to_ai_toolkit( - state_dict, - ait_sd, - f"lora_unet_double_blocks_{i}_img_attn_proj", - f"transformer.transformer_blocks.{i}.attn.to_out.0", - ) - _convert_to_ai_toolkit_cat( - state_dict, - ait_sd, - f"lora_unet_double_blocks_{i}_img_attn_qkv", - [ - f"transformer.transformer_blocks.{i}.attn.to_q", - f"transformer.transformer_blocks.{i}.attn.to_k", - f"transformer.transformer_blocks.{i}.attn.to_v", - ], - ) - _convert_to_ai_toolkit( - state_dict, - ait_sd, - f"lora_unet_double_blocks_{i}_txt_attn_proj", - f"transformer.transformer_blocks.{i}.attn.to_add_out", - ) - _convert_to_ai_toolkit_cat( - state_dict, - ait_sd, - f"lora_unet_double_blocks_{i}_txt_attn_qkv", - [ - f"transformer.transformer_blocks.{i}.attn.add_q_proj", - f"transformer.transformer_blocks.{i}.attn.add_k_proj", - f"transformer.transformer_blocks.{i}.attn.add_v_proj", - ], - ) - # MLP layers (Flux2 uses ff.linear_in/linear_out) - _convert_to_ai_toolkit( - state_dict, - ait_sd, - f"lora_unet_double_blocks_{i}_img_mlp_0", - f"transformer.transformer_blocks.{i}.ff.linear_in", - ) - _convert_to_ai_toolkit( - state_dict, - ait_sd, - f"lora_unet_double_blocks_{i}_img_mlp_2", - f"transformer.transformer_blocks.{i}.ff.linear_out", - ) - _convert_to_ai_toolkit( - state_dict, - ait_sd, - f"lora_unet_double_blocks_{i}_txt_mlp_0", - f"transformer.transformer_blocks.{i}.ff_context.linear_in", - ) - _convert_to_ai_toolkit( - state_dict, - ait_sd, - f"lora_unet_double_blocks_{i}_txt_mlp_2", - f"transformer.transformer_blocks.{i}.ff_context.linear_out", - ) - - for i in range(num_single_layers): - # Single blocks: linear1 -> attn.to_qkv_mlp_proj (fused, no split needed) - _convert_to_ai_toolkit( - state_dict, - ait_sd, - f"lora_unet_single_blocks_{i}_linear1", - f"transformer.single_transformer_blocks.{i}.attn.to_qkv_mlp_proj", - ) - # Single blocks: linear2 -> attn.to_out - _convert_to_ai_toolkit( - state_dict, - ait_sd, - f"lora_unet_single_blocks_{i}_linear2", - f"transformer.single_transformer_blocks.{i}.attn.to_out", - ) - - # Handle optional extra keys - extra_mappings = { - "lora_unet_img_in": "transformer.x_embedder", - "lora_unet_txt_in": "transformer.context_embedder", - "lora_unet_time_in_in_layer": "transformer.time_guidance_embed.timestep_embedder.linear_1", - "lora_unet_time_in_out_layer": "transformer.time_guidance_embed.timestep_embedder.linear_2", - "lora_unet_final_layer_linear": "transformer.proj_out", - } - for sds_key, ait_key in extra_mappings.items(): - _convert_to_ai_toolkit(state_dict, ait_sd, sds_key, ait_key) - - remaining_keys = list(state_dict.keys()) - if remaining_keys: - logger.warning(f"Unsupported keys for Kohya Flux2 LoRA conversion: {remaining_keys}") - - return ait_sd - - -def _convert_non_diffusers_z_image_lora_to_diffusers(state_dict): - """ - Convert non-diffusers ZImage LoRA state dict to diffusers format. - - Handles: - - `diffusion_model.` prefix removal - - `lora_unet_` prefix conversion with key mapping - - `default.` prefix removal - - `.lora_down.weight`/`.lora_up.weight` → `.lora_A.weight`/`.lora_B.weight` conversion with alpha scaling - """ - has_diffusion_model = any(k.startswith("diffusion_model.") for k in state_dict) - if has_diffusion_model: - state_dict = {k.removeprefix("diffusion_model."): v for k, v in state_dict.items()} - - has_lora_unet = any(k.startswith("lora_unet_") for k in state_dict) - if has_lora_unet: - state_dict = {k.removeprefix("lora_unet_"): v for k, v in state_dict.items()} - - def convert_key(key: str) -> str: - # ZImage has: layers, noise_refiner, context_refiner blocks - # Keys may be like: layers_0_attention_to_q.lora_down.weight - - if "." in key: - base, suffix = key.rsplit(".", 1) - else: - base, suffix = key, "" - - # Protected n-grams that must keep their internal underscores - protected = { - # pairs for attention - ("to", "q"), - ("to", "k"), - ("to", "v"), - ("to", "out"), - # feed_forward - ("feed", "forward"), - } - - prot_by_len = {} - for ng in protected: - prot_by_len.setdefault(len(ng), set()).add(ng) - - parts = base.split("_") - merged = [] - i = 0 - lengths_desc = sorted(prot_by_len.keys(), reverse=True) - - while i < len(parts): - matched = False - for L in lengths_desc: - if i + L <= len(parts) and tuple(parts[i : i + L]) in prot_by_len[L]: - merged.append("_".join(parts[i : i + L])) - i += L - matched = True - break - if not matched: - merged.append(parts[i]) - i += 1 - - converted_base = ".".join(merged) - return converted_base + (("." + suffix) if suffix else "") - - state_dict = {convert_key(k): v for k, v in state_dict.items()} - - def normalize_out_key(k: str) -> str: - if ".to_out" in k: - return k - return re.sub( - r"\.out(?=\.(?:lora_down|lora_up)\.weight$|\.alpha$)", - ".to_out.0", - k, - ) - - state_dict = {normalize_out_key(k): v for k, v in state_dict.items()} - - has_default = any("default." in k for k in state_dict) - if has_default: - state_dict = {k.replace("default.", ""): v for k, v in state_dict.items()} - - # Normalize ZImage-specific dot-separated module names to underscore form so they - # match the diffusers model parameter names. convert_key blindly split every "_", - # so module names whose own names contain underscores (and aren't protected as the - # attention/feed_forward n-grams are) come out over-split here. This runs on the full - # key (before the weight/alpha handlers below) so it fixes .lora_A/B and .alpha alike. - zimage_module_name_fixups = { - "context.refiner.": "context_refiner.", - "noise.refiner.": "noise_refiner.", - "adaLN.modulation.": "adaLN_modulation.", - "all.final.layer.": "all_final_layer.", - "all.x.embedder.": "all_x_embedder.", - "cap.embedder.": "cap_embedder.", - "t.embedder.": "t_embedder.", - } - - def fixup_module_names(k: str) -> str: - for dotted, underscored in zimage_module_name_fixups.items(): - k = k.replace(dotted, underscored) - return k - - state_dict = {fixup_module_names(k): v for k, v in state_dict.items()} - - converted_state_dict = {} - all_keys = list(state_dict.keys()) - down_key = ".lora_down.weight" - up_key = ".lora_up.weight" - a_key = ".lora_A.weight" - b_key = ".lora_B.weight" - - has_non_diffusers_lora_id = any(down_key in k or up_key in k for k in all_keys) - has_diffusers_lora_id = any(a_key in k or b_key in k for k in all_keys) - - def get_alpha_scales(down_weight, alpha_key): - rank = down_weight.shape[0] - alpha_tensor = state_dict.pop(alpha_key, None) - if alpha_tensor is None: - return 1.0, 1.0 - scale = ( - alpha_tensor.item() / rank - ) # LoRA is scaled by 'alpha / rank' in forward pass, so we need to scale it back here - scale_down = scale - scale_up = 1.0 - while scale_down * 2 < scale_up: - scale_down *= 2 - scale_up /= 2 - return scale_down, scale_up - - if has_non_diffusers_lora_id: - for k in all_keys: - if k.endswith(down_key): - diffusers_down_key = k.replace(down_key, ".lora_A.weight") - diffusers_up_key = k.replace(down_key, up_key).replace(up_key, ".lora_B.weight") - alpha_key = k.replace(down_key, ".alpha") - - down_weight = state_dict.pop(k) - up_weight = state_dict.pop(k.replace(down_key, up_key)) - scale_down, scale_up = get_alpha_scales(down_weight, alpha_key) - converted_state_dict[diffusers_down_key] = down_weight * scale_down - converted_state_dict[diffusers_up_key] = up_weight * scale_up - - # Already in diffusers format (lora_A/lora_B), apply alpha scaling and pop. - elif has_diffusers_lora_id: - for k in all_keys: - if k.endswith(a_key): - diffusers_up_key = k.replace(a_key, b_key) - alpha_key = k.replace(a_key, ".alpha") - - down_weight = state_dict.pop(k) - up_weight = state_dict.pop(diffusers_up_key) - scale_down, scale_up = get_alpha_scales(down_weight, alpha_key) - converted_state_dict[k] = down_weight * scale_down - converted_state_dict[diffusers_up_key] = up_weight * scale_up - - # Handle dot-format LoRA keys: ".lora.down.weight" / ".lora.up.weight". - # Some external ZImage trainers (e.g. Anime-Z) use dots instead of underscores in - # lora weight names and also include redundant keys: - # - "qkv.lora.*" duplicates individual "to.q/k/v.lora.*" keys → skip qkv - # - "out.lora.*" duplicates "to_out.0.lora.*" keys → skip bare out - # - "to.q/k/v.lora.*" → normalise to "to_q/k/v.lora_A/B.weight" - lora_dot_down_key = ".lora.down.weight" - lora_dot_up_key = ".lora.up.weight" - has_lora_dot_format = any(lora_dot_down_key in k for k in state_dict) - - if has_lora_dot_format: - dot_keys = list(state_dict.keys()) - for k in dot_keys: - if lora_dot_down_key not in k: - continue - if k not in state_dict: - continue # already popped by a prior iteration - - base = k[: -len(lora_dot_down_key)] - - # Skip combined "qkv" projection — individual to.q/k/v keys are also present. - if base.endswith(".qkv"): - state_dict.pop(k) - state_dict.pop(k.replace(lora_dot_down_key, lora_dot_up_key), None) - state_dict.pop(base + ".alpha", None) - continue - - # Skip bare "out.lora.*" — "to_out.0.lora.*" covers the same projection. - if re.search(r"\.out$", base) and ".to_out" not in base: - state_dict.pop(k) - state_dict.pop(k.replace(lora_dot_down_key, lora_dot_up_key), None) - continue - - # Normalise "to.q/k/v" → "to_q/k/v" for the diffusers output key. - norm_k = re.sub( - r"\.to\.([qkv])" + re.escape(lora_dot_down_key) + r"$", - r".to_\1" + lora_dot_down_key, - k, - ) - norm_base = norm_k[: -len(lora_dot_down_key)] - alpha_key = norm_base + ".alpha" - - diffusers_down = norm_k.replace(lora_dot_down_key, ".lora_A.weight") - diffusers_up = norm_k.replace(lora_dot_down_key, ".lora_B.weight") - - down_weight = state_dict.pop(k) - up_weight = state_dict.pop(k.replace(lora_dot_down_key, lora_dot_up_key)) - scale_down, scale_up = get_alpha_scales(down_weight, alpha_key) - converted_state_dict[diffusers_down] = down_weight * scale_down - converted_state_dict[diffusers_up] = up_weight * scale_up - - if len(state_dict) > 0: - raise ValueError(f"`state_dict` should be empty at this point but has {state_dict.keys()=}") - - converted_state_dict = {f"transformer.{k}": v for k, v in converted_state_dict.items()} - return converted_state_dict - - -def _convert_non_diffusers_ideogram4_lora_to_diffusers(state_dict): - """ - Convert non-diffusers Ideogram4 LoRA state dict to diffusers format. - - Handles: - - `diffusion_model.` / `conditional_transformer.` prefix removal - - `lora_down`/`lora_up` (kohya) -> `lora_A`/`lora_B`, with `.alpha` folded into the weights - - fused `attention.qkv` -> split `to_q`/`to_k`/`to_v`; `attention.o` -> `to_out.0` - - `feed_forward.w1`/`w2`/`w3` and `adaln_modulation` map one-to-one - """ - for prefix in ("diffusion_model.", "conditional_transformer."): - if any(k.startswith(prefix) for k in state_dict): - state_dict = {k.removeprefix(prefix): v for k, v in state_dict.items()} - break - - is_kohya = any(".lora_down.weight" in k for k in state_dict) - down_suffix = ".lora_down.weight" if is_kohya else ".lora_A.weight" - up_suffix = ".lora_up.weight" if is_kohya else ".lora_B.weight" - - def get_alpha_scales(down_weight, alpha_key): - rank = down_weight.shape[0] - alpha_tensor = state_dict.pop(alpha_key, None) - if alpha_tensor is None: - return 1.0, 1.0 - # LoRA is scaled by `alpha / rank` in the forward pass; split the factor between down and up. - scale = alpha_tensor.item() / rank - scale_down, scale_up = scale, 1.0 - while scale_down * 2 < scale_up: - scale_down *= 2 - scale_up /= 2 - return scale_down, scale_up - - def pull(base): - """Pop the scaled (lora_A, lora_B) pair for a module path, or return None if absent.""" - down_key = base + down_suffix - if down_key not in state_dict: - return None - down = state_dict.pop(down_key) - up = state_dict.pop(base + up_suffix) - scale_down, scale_up = get_alpha_scales(down, base + ".alpha") - return down * scale_down, up * scale_up - - num_layers = 0 - for k in state_dict: - match = re.match(r"layers\.(\d+)\.", k) - if match: - num_layers = max(num_layers, int(match.group(1)) + 1) - - converted_state_dict = {} - for i in range(num_layers): - layer_prefix = f"layers.{i}" - - # Fused qkv -> split to_q / to_k / to_v (shared down/lora_A, chunk up/lora_B in thirds). - qkv = pull(f"{layer_prefix}.attention.qkv") - if qkv is not None: - down, up = qkv - up_q, up_k, up_v = torch.chunk(up, 3, dim=0) - for proj, up_proj in (("to_q", up_q), ("to_k", up_k), ("to_v", up_v)): - converted_state_dict[f"{layer_prefix}.attention.{proj}.lora_A.weight"] = down.clone() - converted_state_dict[f"{layer_prefix}.attention.{proj}.lora_B.weight"] = up_proj.contiguous() - - # attention.o -> attention.to_out.0 - out = pull(f"{layer_prefix}.attention.o") - if out is not None: - down, up = out - converted_state_dict[f"{layer_prefix}.attention.to_out.0.lora_A.weight"] = down - converted_state_dict[f"{layer_prefix}.attention.to_out.0.lora_B.weight"] = up - - # feed_forward.{w1,w2,w3} and adaln_modulation map one-to-one. - for module in ("feed_forward.w1", "feed_forward.w2", "feed_forward.w3", "adaln_modulation"): - pair = pull(f"{layer_prefix}.{module}") - if pair is not None: - down, up = pair - converted_state_dict[f"{layer_prefix}.{module}.lora_A.weight"] = down - converted_state_dict[f"{layer_prefix}.{module}.lora_B.weight"] = up - - if len(state_dict) > 0: - raise ValueError( - f"`state_dict` should be empty at this point but has {sorted(state_dict.keys())}. " - "This may be an unsupported Ideogram4 LoRA layout." - ) - - return {f"transformer.{k}": v for k, v in converted_state_dict.items()} - - -def _convert_non_diffusers_krea2_lora_to_diffusers(state_dict): - """ - Convert a non-diffusers Krea 2 LoRA state dict to the diffusers format. - - Maps the original `krea-ai/krea-2` module names onto `Krea2Transformer2DModel`. Handles both the `diffusion_model.` - prefix (Krea 2 reference trainer / ComfyUI) and the `base_model.model.` prefix (Ostris AI-Toolkit). - """ - state_dict = { - k.removeprefix("base_model.model.").removeprefix("diffusion_model."): v for k, v in state_dict.items() - } - - attn_map = {"wq": "to_q", "wk": "to_k", "wv": "to_v", "wo": "to_out.0", "gate": "to_gate"} - ff_map = {"gate": "ff.gate", "up": "ff.up", "down": "ff.down"} - # AI-Toolkit stores these standalone modules under abbreviated `nn.Sequential`-style names. - standalone_map = { - "first": "img_in", - "last.linear": "final_layer.linear", - "tmlp.0": "time_embed.linear_1", - "tmlp.2": "time_embed.linear_2", - "tproj.1": "time_mod_proj", - "txtmlp.1": "txt_in.linear_1", - "txtmlp.3": "txt_in.linear_2", - "txtfusion.projector": "text_fusion.projector", - } - - def convert_module(module): - m = re.match(r"blocks\.(\d+)\.(attn|mlp)\.(\w+)$", module) - if m: - idx, kind, sub = m.groups() - if kind == "attn" and sub in attn_map: - return f"transformer_blocks.{idx}.attn.{attn_map[sub]}" - if kind == "mlp" and sub in ff_map: - return f"transformer_blocks.{idx}.{ff_map[sub]}" - return None - m = re.match(r"txtfusion\.(layerwise_blocks|refiner_blocks)\.(\d+)\.(attn|mlp)\.(\w+)$", module) - if m: - block, idx, kind, sub = m.groups() - if kind == "attn" and sub in attn_map: - return f"text_fusion.{block}.{idx}.attn.{attn_map[sub]}" - if kind == "mlp" and sub in ff_map: - return f"text_fusion.{block}.{idx}.{ff_map[sub]}" - return None - return standalone_map.get(module) - - converted_state_dict = {} - for key in list(state_dict): - match = re.search(r"\.(?:lora_[AB])\.weight$", key) - if match is None: - continue - diffusers_module = convert_module(key[: match.start()]) - if diffusers_module is None: - continue - converted_state_dict[f"transformer.{diffusers_module}{key[match.start() :]}"] = state_dict.pop(key) - - if len(state_dict) > 0: - raise ValueError(f"`state_dict` should be empty at this point but has {state_dict.keys()=}") - - return converted_state_dict - - -def _convert_non_diffusers_ace_step_lora_to_diffusers(state_dict): - """Convert an ACE-Step-1.5 (PEFT format) LoRA state dict to diffusers key names. - - The original ACE-Step repo targets ``q_proj``, ``k_proj``, ``v_proj``, ``o_proj`` on the DiT decoder while - diffusers renames them to ``to_q``, ``to_k``, ``to_v``, ``to_out.0``. Keys arrive as - ``base_model.model.layers.{i}.{self_attn|cross_attn}.{proj}.lora_{A|B}.weight`` and are mapped to - ``transformer.layers.{i}.{self_attn|cross_attn}.{proj_diffusers}.lora_{A|B}.weight``. - """ - _PROJ_RENAMES = { - ".q_proj.": ".to_q.", - ".k_proj.": ".to_k.", - ".v_proj.": ".to_v.", - ".o_proj.": ".to_out.0.", - } - - converted_state_dict = {} - for key in list(state_dict.keys()): - new_key = key - if new_key.startswith("base_model.model."): - new_key = new_key[len("base_model.model.") :] - for old, new in _PROJ_RENAMES.items(): - new_key = new_key.replace(old, new) - new_key = f"transformer.{new_key}" - converted_state_dict[new_key] = state_dict.pop(key) - - return converted_state_dict diff --git a/diffusers/loaders/lora_pipeline.py b/diffusers/loaders/lora_pipeline.py deleted file mode 100644 index 8de23d81528ce2943042f63444cab2030ef5cf65..0000000000000000000000000000000000000000 --- a/diffusers/loaders/lora_pipeline.py +++ /dev/null @@ -1,7048 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import os -from typing import Callable - -import torch -from huggingface_hub.utils import validate_hf_hub_args - -from ..utils import ( - USE_PEFT_BACKEND, - deprecate, - get_submodule_by_name, - is_bitsandbytes_available, - is_gguf_available, - is_peft_available, - is_peft_version, - is_torch_version, - is_transformers_available, - is_transformers_version, - logging, -) -from .lora_base import ( # noqa - LORA_WEIGHT_NAME, - LORA_WEIGHT_NAME_SAFE, - LoraBaseMixin, - _fetch_state_dict, - _load_lora_into_text_encoder, - _pack_dict_with_prefix, -) -from .lora_conversion_utils import ( - _convert_bfl_flux_control_lora_to_diffusers, - _convert_fal_kontext_lora_to_diffusers, - _convert_hunyuan_video_lora_to_diffusers, - _convert_kohya_flux2_lora_to_diffusers, - _convert_kohya_flux_lora_to_diffusers, - _convert_musubi_wan_lora_to_diffusers, - _convert_non_diffusers_ace_step_lora_to_diffusers, - _convert_non_diffusers_anima_lora_to_diffusers, - _convert_non_diffusers_flux2_lora_to_diffusers, - _convert_non_diffusers_hidream_lora_to_diffusers, - _convert_non_diffusers_ideogram4_lora_to_diffusers, - _convert_non_diffusers_krea2_lora_to_diffusers, - _convert_non_diffusers_lora_to_diffusers, - _convert_non_diffusers_ltx2_lora_to_diffusers, - _convert_non_diffusers_ltxv_lora_to_diffusers, - _convert_non_diffusers_lumina2_lora_to_diffusers, - _convert_non_diffusers_qwen_lora_to_diffusers, - _convert_non_diffusers_wan_lora_to_diffusers, - _convert_non_diffusers_z_image_lora_to_diffusers, - _convert_xlabs_flux_lora_to_diffusers, - _maybe_map_sgm_blocks_to_diffusers, -) - - -_LOW_CPU_MEM_USAGE_DEFAULT_LORA = False -if is_torch_version(">=", "1.9.0"): - if ( - is_peft_available() - and is_peft_version(">=", "0.13.1") - and is_transformers_available() - and is_transformers_version(">", "4.45.2") - ): - _LOW_CPU_MEM_USAGE_DEFAULT_LORA = True - - -logger = logging.get_logger(__name__) - -TEXT_ENCODER_NAME = "text_encoder" -UNET_NAME = "unet" -TRANSFORMER_NAME = "transformer" -LTX2_CONNECTOR_NAME = "connectors" - -_MODULE_NAME_TO_ATTRIBUTE_MAP_FLUX = {"x_embedder": "in_channels"} - - -def _maybe_dequantize_weight_for_expanded_lora(model, module): - if is_bitsandbytes_available(): - from ..quantizers.bitsandbytes import dequantize_bnb_weight - - if is_gguf_available(): - from ..quantizers.gguf.utils import dequantize_gguf_tensor - - is_bnb_4bit_quantized = module.weight.__class__.__name__ == "Params4bit" - is_bnb_8bit_quantized = module.weight.__class__.__name__ == "Int8Params" - is_gguf_quantized = module.weight.__class__.__name__ == "GGUFParameter" - - if is_bnb_4bit_quantized and not is_bitsandbytes_available(): - raise ValueError( - "The checkpoint seems to have been quantized with `bitsandbytes` (4bits). Install `bitsandbytes` to load quantized checkpoints." - ) - if is_bnb_8bit_quantized and not is_bitsandbytes_available(): - raise ValueError( - "The checkpoint seems to have been quantized with `bitsandbytes` (8bits). Install `bitsandbytes` to load quantized checkpoints." - ) - if is_gguf_quantized and not is_gguf_available(): - raise ValueError( - "The checkpoint seems to have been quantized with `gguf`. Install `gguf` to load quantized checkpoints." - ) - - weight_on_cpu = False - if module.weight.device.type == "cpu": - weight_on_cpu = True - - device = torch.accelerator.current_accelerator().type if hasattr(torch, "accelerator") else "cuda" - if is_bnb_4bit_quantized or is_bnb_8bit_quantized: - module_weight = dequantize_bnb_weight( - module.weight.to(device) if weight_on_cpu else module.weight, - state=module.weight.quant_state if is_bnb_4bit_quantized else module.state, - dtype=model.dtype, - ).data - elif is_gguf_quantized: - module_weight = dequantize_gguf_tensor( - module.weight.to(device) if weight_on_cpu else module.weight, - ) - module_weight = module_weight.to(model.dtype) - else: - module_weight = module.weight.data - - if weight_on_cpu: - module_weight = module_weight.cpu() - - return module_weight - - -class StableDiffusionLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into Stable Diffusion [`UNet2DConditionModel`] and - [`CLIPTextModel`](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel). - """ - - _lora_loadable_modules = ["unet", "text_encoder"] - unet_name = UNET_NAME - text_encoder_name = TEXT_ENCODER_NAME - - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """Load LoRA weights specified in `pretrained_model_name_or_path_or_dict` into `self.unet` and - `self.text_encoder`. - - All kwargs are forwarded to `self.lora_state_dict`. - - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details on how the state dict is - loaded. - - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details on how the state dict is - loaded into `self.unet`. - - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_text_encoder`] for more details on how the state - dict is loaded into `self.text_encoder`. - - Parameters: - pretrained_model_name_or_path_or_dict (`str` or `os.PathLike` or `dict`): - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`]. - adapter_name (`str`, *optional*): - Adapter name to be used for referencing the loaded adapter model. If not specified, it will use - `default_{i}` where i is the total number of adapters being loaded. - low_cpu_mem_usage (`bool`, *optional*): - Speed up model loading by only loading the pretrained LoRA weights and not initializing the random - weights. - hotswap (`bool`, *optional*): - Defaults to `False`. Whether to substitute an existing (LoRA) adapter with the newly loaded adapter - in-place. This means that, instead of loading an additional adapter, this will take the existing - adapter weights and replace them with the weights of the new adapter. This can be faster and more - memory efficient. However, the main advantage of hotswapping is that when the model is compiled with - torch.compile, loading the new adapter does not require recompilation of the model. When using - hotswapping, the passed `adapter_name` should be the name of an already loaded adapter. - - If the new adapter and the old adapter have different ranks and/or LoRA alphas (i.e. scaling), you need - to call an additional method before loading the adapter: - - ```py - pipeline = ... # load diffusers pipeline - max_rank = ... # the highest rank among all LoRAs that you want to load - # call *before* compiling and loading the LoRA adapter - pipeline.enable_lora_hotswap(target_rank=max_rank) - pipeline.load_lora_weights(file_name) - # optionally compile the model now - ``` - - Note that hotswapping adapters of the text encoder is not yet supported. There are some further - limitations to this technique, which are documented here: - https://huggingface.co/docs/peft/main/en/package_reference/hotswap - kwargs (`dict`, *optional*): - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`]. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and not is_peft_version(">=", "0.13.1"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, network_alphas, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_unet( - state_dict, - network_alphas=network_alphas, - unet=getattr(self, self.unet_name) if not hasattr(self, "unet") else self.unet, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - self.load_lora_into_text_encoder( - state_dict, - network_alphas=network_alphas, - text_encoder=getattr(self, self.text_encoder_name) - if not hasattr(self, "text_encoder") - else self.text_encoder, - lora_scale=self.lora_scale, - adapter_name=adapter_name, - _pipeline=self, - metadata=metadata, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - Return state dict for lora weights and the network alphas. - - > [!WARNING] > We support loading A1111 formatted LoRA checkpoints in a limited capacity. > > This function is - experimental and might change in the future. - - Parameters: - pretrained_model_name_or_path_or_dict (`str` or `os.PathLike` or `dict`): - Can be either: - - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved - with [`ModelMixin.save_pretrained`]. - - A [torch state - dict](https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict). - - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - subfolder (`str`, *optional*, defaults to `""`): - The subfolder location of a model file within a larger model repository on the Hub or locally. - weight_name (`str`, *optional*, defaults to None): - Name of the serialized state dict file. - return_lora_metadata (`bool`, *optional*, defaults to False): - When enabled, additionally return the LoRA adapter metadata, typically found in the state dict. - """ - # Load the main state dict first which has the LoRA layers for either of - # UNet and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - unet_config = kwargs.pop("unet_config", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - network_alphas = None - # TODO: replace it with a method from `state_dict_utils` - if all( - ( - k.startswith("lora_te_") - or k.startswith("lora_unet_") - or k.startswith("lora_te1_") - or k.startswith("lora_te2_") - ) - for k in state_dict.keys() - ): - # Map SDXL blocks correctly. - if unet_config is not None: - # use unet config to remap block numbers - state_dict = _maybe_map_sgm_blocks_to_diffusers(state_dict, unet_config) - state_dict, network_alphas = _convert_non_diffusers_lora_to_diffusers(state_dict) - - out = (state_dict, network_alphas, metadata) if return_lora_metadata else (state_dict, network_alphas) - return out - - @classmethod - def load_lora_into_unet( - cls, - state_dict, - network_alphas, - unet, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - This will load the LoRA layers specified in `state_dict` into `unet`. - - Parameters: - state_dict (`dict`): - A standard state dict containing the lora layer parameters. The keys can either be indexed directly - into the unet or prefixed with an additional `unet` which can be used to distinguish between text - encoder lora layers. - network_alphas (`dict[str, float]`): - The value of the network alpha used for stable learning and preventing underflow. This value has the - same meaning as the `--network_alpha` option in the kohya-ss trainer script. Refer to [this - link](https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning). - unet (`UNet2DConditionModel`): - The UNet model to load the LoRA layers into. - adapter_name (`str`, *optional*): - Adapter name to be used for referencing the loaded adapter model. If not specified, it will use - `default_{i}` where i is the total number of adapters being loaded. - low_cpu_mem_usage (`bool`, *optional*): - Speed up model loading only loading the pretrained LoRA weights and not initializing the random - weights. - hotswap (`bool`, *optional*): - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`]. - metadata (`dict`): - Optional LoRA adapter metadata. When supplied, the `LoraConfig` arguments of `peft` won't be derived - from the state dict. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - if low_cpu_mem_usage and not is_peft_version(">=", "0.13.1"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # If the serialization format is new (introduced in https://github.com/huggingface/diffusers/pull/2918), - # then the `state_dict` keys should have `cls.unet_name` and/or `cls.text_encoder_name` as - # their prefixes. - logger.info(f"Loading {cls.unet_name}.") - unet.load_lora_adapter( - state_dict, - prefix=cls.unet_name, - network_alphas=network_alphas, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - def load_lora_into_text_encoder( - cls, - state_dict, - network_alphas, - text_encoder, - prefix=None, - lora_scale=1.0, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - This will load the LoRA layers specified in `state_dict` into `text_encoder` - - Parameters: - state_dict (`dict`): - A standard state dict containing the lora layer parameters. The key should be prefixed with an - additional `text_encoder` to distinguish between unet lora layers. - network_alphas (`dict[str, float]`): - The value of the network alpha used for stable learning and preventing underflow. This value has the - same meaning as the `--network_alpha` option in the kohya-ss trainer script. Refer to [this - link](https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning). - text_encoder (`CLIPTextModel`): - The text encoder model to load the LoRA layers into. - prefix (`str`): - Expected prefix of the `text_encoder` in the `state_dict`. - lora_scale (`float`): - How much to scale the output of the lora linear layer before it is added with the output of the regular - lora layer. - adapter_name (`str`, *optional*): - Adapter name to be used for referencing the loaded adapter model. If not specified, it will use - `default_{i}` where i is the total number of adapters being loaded. - low_cpu_mem_usage (`bool`, *optional*): - Speed up model loading by only loading the pretrained LoRA weights and not initializing the random - weights. - hotswap (`bool`, *optional*): - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`]. - metadata (`dict`): - Optional LoRA adapter metadata. When supplied, the `LoraConfig` arguments of `peft` won't be derived - from the state dict. - """ - _load_lora_into_text_encoder( - state_dict=state_dict, - network_alphas=network_alphas, - lora_scale=lora_scale, - text_encoder=text_encoder, - prefix=prefix, - text_encoder_name=cls.text_encoder_name, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - unet_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - text_encoder_lora_layers: dict[str, torch.nn.Module] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - unet_lora_adapter_metadata=None, - text_encoder_lora_adapter_metadata=None, - ): - r""" - Save the LoRA parameters corresponding to the UNet and text encoder. - - Arguments: - save_directory (`str` or `os.PathLike`): - Directory to save LoRA parameters to. Will be created if it doesn't exist. - unet_lora_layers (`dict[str, torch.nn.Module]` or `dict[str, torch.Tensor]`): - State dict of the LoRA layers corresponding to the `unet`. - text_encoder_lora_layers (`dict[str, torch.nn.Module]` or `dict[str, torch.Tensor]`): - State dict of the LoRA layers corresponding to the `text_encoder`. Must explicitly pass the text - encoder LoRA state dict because it comes from 🤗 Transformers. - is_main_process (`bool`, *optional*, defaults to `True`): - Whether the process calling this is the main process or not. Useful during distributed training and you - need to call this function on all processes. In this case, set `is_main_process=True` only on the main - process to avoid race conditions. - save_function (`Callable`): - The function to use to save the state dictionary. Useful during distributed training when you need to - replace `torch.save` with another method. Can be configured with the environment variable - `DIFFUSERS_SAVE_MODE`. - safe_serialization (`bool`, *optional*, defaults to `True`): - Whether to save the model using `safetensors` or the traditional PyTorch way with `pickle`. - unet_lora_adapter_metadata: - LoRA adapter metadata associated with the unet to be serialized with the state dict. - text_encoder_lora_adapter_metadata: - LoRA adapter metadata associated with the text encoder to be serialized with the state dict. - """ - lora_layers = {} - lora_metadata = {} - - if unet_lora_layers: - lora_layers[cls.unet_name] = unet_lora_layers - lora_metadata[cls.unet_name] = unet_lora_adapter_metadata - - if text_encoder_lora_layers: - lora_layers[cls.text_encoder_name] = text_encoder_lora_layers - lora_metadata[cls.text_encoder_name] = text_encoder_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `unet_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - def fuse_lora( - self, - components: list[str] = ["unet", "text_encoder"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - Fuses the LoRA parameters into the original parameters of the corresponding blocks. - - Args: - components: (`list[str]`): list of LoRA-injectable components to fuse the LoRAs into. - lora_scale (`float`, defaults to 1.0): - Controls how much to influence the outputs with the LoRA parameters. - safe_fusing (`bool`, defaults to `False`): - Whether to check fused weights for NaN values before fusing and if values are NaN not fusing them. - adapter_names (`list[str]`, *optional*): - Adapter names to be used for fusing. If nothing is passed, all active adapters will be fused. - - Example: - - ```py - from diffusers import DiffusionPipeline - import torch - - pipeline = DiffusionPipeline.from_pretrained( - "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16 - ).to("cuda") - pipeline.load_lora_weights("nerijs/pixel-art-xl", weight_name="pixel-art-xl.safetensors", adapter_name="pixel") - pipeline.fuse_lora(lora_scale=0.7) - ``` - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - def unfuse_lora(self, components: list[str] = ["unet", "text_encoder"], **kwargs): - r""" - Reverses the effect of - [`pipe.fuse_lora()`](https://huggingface.co/docs/diffusers/main/en/api/loaders#diffusers.loaders.LoraBaseMixin.fuse_lora). - - Args: - components (`list[str]`): list of LoRA-injectable components to unfuse LoRA from. - unfuse_unet (`bool`, defaults to `True`): Whether to unfuse the UNet LoRA parameters. - unfuse_text_encoder (`bool`, defaults to `True`): - Whether to unfuse the text encoder LoRA parameters. If the text encoder wasn't monkey-patched with the - LoRA parameters then it won't have any effect. - """ - super().unfuse_lora(components=components, **kwargs) - - -class StableDiffusionXLLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into Stable Diffusion XL [`UNet2DConditionModel`], - [`CLIPTextModel`](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), and - [`CLIPTextModelWithProjection`](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModelWithProjection). - """ - - _lora_loadable_modules = ["unet", "text_encoder", "text_encoder_2"] - unet_name = UNET_NAME - text_encoder_name = TEXT_ENCODER_NAME - - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and not is_peft_version(">=", "0.13.1"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # We could have accessed the unet config from `lora_state_dict()` too. We pass - # it here explicitly to be able to tell that it's coming from an SDXL - # pipeline. - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, network_alphas, metadata = self.lora_state_dict( - pretrained_model_name_or_path_or_dict, - unet_config=self.unet.config, - **kwargs, - ) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_unet( - state_dict, - network_alphas=network_alphas, - unet=self.unet, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - self.load_lora_into_text_encoder( - state_dict, - network_alphas=network_alphas, - text_encoder=self.text_encoder, - prefix=self.text_encoder_name, - lora_scale=self.lora_scale, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - self.load_lora_into_text_encoder( - state_dict, - network_alphas=network_alphas, - text_encoder=self.text_encoder_2, - prefix=f"{self.text_encoder_name}_2", - lora_scale=self.lora_scale, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - @validate_hf_hub_args - # Copied from diffusers.loaders.lora_pipeline.StableDiffusionLoraLoaderMixin.lora_state_dict - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - Return state dict for lora weights and the network alphas. - - > [!WARNING] > We support loading A1111 formatted LoRA checkpoints in a limited capacity. > > This function is - experimental and might change in the future. - - Parameters: - pretrained_model_name_or_path_or_dict (`str` or `os.PathLike` or `dict`): - Can be either: - - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved - with [`ModelMixin.save_pretrained`]. - - A [torch state - dict](https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict). - - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - subfolder (`str`, *optional*, defaults to `""`): - The subfolder location of a model file within a larger model repository on the Hub or locally. - weight_name (`str`, *optional*, defaults to None): - Name of the serialized state dict file. - return_lora_metadata (`bool`, *optional*, defaults to False): - When enabled, additionally return the LoRA adapter metadata, typically found in the state dict. - """ - # Load the main state dict first which has the LoRA layers for either of - # UNet and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - unet_config = kwargs.pop("unet_config", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - network_alphas = None - # TODO: replace it with a method from `state_dict_utils` - if all( - ( - k.startswith("lora_te_") - or k.startswith("lora_unet_") - or k.startswith("lora_te1_") - or k.startswith("lora_te2_") - ) - for k in state_dict.keys() - ): - # Map SDXL blocks correctly. - if unet_config is not None: - # use unet config to remap block numbers - state_dict = _maybe_map_sgm_blocks_to_diffusers(state_dict, unet_config) - state_dict, network_alphas = _convert_non_diffusers_lora_to_diffusers(state_dict) - - out = (state_dict, network_alphas, metadata) if return_lora_metadata else (state_dict, network_alphas) - return out - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.StableDiffusionLoraLoaderMixin.load_lora_into_unet - def load_lora_into_unet( - cls, - state_dict, - network_alphas, - unet, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - This will load the LoRA layers specified in `state_dict` into `unet`. - - Parameters: - state_dict (`dict`): - A standard state dict containing the lora layer parameters. The keys can either be indexed directly - into the unet or prefixed with an additional `unet` which can be used to distinguish between text - encoder lora layers. - network_alphas (`dict[str, float]`): - The value of the network alpha used for stable learning and preventing underflow. This value has the - same meaning as the `--network_alpha` option in the kohya-ss trainer script. Refer to [this - link](https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning). - unet (`UNet2DConditionModel`): - The UNet model to load the LoRA layers into. - adapter_name (`str`, *optional*): - Adapter name to be used for referencing the loaded adapter model. If not specified, it will use - `default_{i}` where i is the total number of adapters being loaded. - low_cpu_mem_usage (`bool`, *optional*): - Speed up model loading only loading the pretrained LoRA weights and not initializing the random - weights. - hotswap (`bool`, *optional*): - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`]. - metadata (`dict`): - Optional LoRA adapter metadata. When supplied, the `LoraConfig` arguments of `peft` won't be derived - from the state dict. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - if low_cpu_mem_usage and not is_peft_version(">=", "0.13.1"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # If the serialization format is new (introduced in https://github.com/huggingface/diffusers/pull/2918), - # then the `state_dict` keys should have `cls.unet_name` and/or `cls.text_encoder_name` as - # their prefixes. - logger.info(f"Loading {cls.unet_name}.") - unet.load_lora_adapter( - state_dict, - prefix=cls.unet_name, - network_alphas=network_alphas, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.StableDiffusionLoraLoaderMixin.load_lora_into_text_encoder - def load_lora_into_text_encoder( - cls, - state_dict, - network_alphas, - text_encoder, - prefix=None, - lora_scale=1.0, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - This will load the LoRA layers specified in `state_dict` into `text_encoder` - - Parameters: - state_dict (`dict`): - A standard state dict containing the lora layer parameters. The key should be prefixed with an - additional `text_encoder` to distinguish between unet lora layers. - network_alphas (`dict[str, float]`): - The value of the network alpha used for stable learning and preventing underflow. This value has the - same meaning as the `--network_alpha` option in the kohya-ss trainer script. Refer to [this - link](https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning). - text_encoder (`CLIPTextModel`): - The text encoder model to load the LoRA layers into. - prefix (`str`): - Expected prefix of the `text_encoder` in the `state_dict`. - lora_scale (`float`): - How much to scale the output of the lora linear layer before it is added with the output of the regular - lora layer. - adapter_name (`str`, *optional*): - Adapter name to be used for referencing the loaded adapter model. If not specified, it will use - `default_{i}` where i is the total number of adapters being loaded. - low_cpu_mem_usage (`bool`, *optional*): - Speed up model loading by only loading the pretrained LoRA weights and not initializing the random - weights. - hotswap (`bool`, *optional*): - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`]. - metadata (`dict`): - Optional LoRA adapter metadata. When supplied, the `LoraConfig` arguments of `peft` won't be derived - from the state dict. - """ - _load_lora_into_text_encoder( - state_dict=state_dict, - network_alphas=network_alphas, - lora_scale=lora_scale, - text_encoder=text_encoder, - prefix=prefix, - text_encoder_name=cls.text_encoder_name, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - unet_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - text_encoder_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - text_encoder_2_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - unet_lora_adapter_metadata=None, - text_encoder_lora_adapter_metadata=None, - text_encoder_2_lora_adapter_metadata=None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if unet_lora_layers: - lora_layers[cls.unet_name] = unet_lora_layers - lora_metadata[cls.unet_name] = unet_lora_adapter_metadata - - if text_encoder_lora_layers: - lora_layers["text_encoder"] = text_encoder_lora_layers - lora_metadata["text_encoder"] = text_encoder_lora_adapter_metadata - - if text_encoder_2_lora_layers: - lora_layers["text_encoder_2"] = text_encoder_2_lora_layers - lora_metadata["text_encoder_2"] = text_encoder_2_lora_adapter_metadata - - if not lora_layers: - raise ValueError( - "You must pass at least one of `unet_lora_layers`, `text_encoder_lora_layers`, or `text_encoder_2_lora_layers`." - ) - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - def fuse_lora( - self, - components: list[str] = ["unet", "text_encoder", "text_encoder_2"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - def unfuse_lora(self, components: list[str] = ["unet", "text_encoder", "text_encoder_2"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class SD3LoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`SD3Transformer2DModel`], - [`CLIPTextModel`](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), and - [`CLIPTextModelWithProjection`](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModelWithProjection). - - Specific to [`StableDiffusion3Pipeline`]. - """ - - _lora_loadable_modules = ["transformer", "text_encoder", "text_encoder_2"] - transformer_name = TRANSFORMER_NAME - text_encoder_name = TEXT_ENCODER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name=None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - self.load_lora_into_text_encoder( - state_dict, - network_alphas=None, - text_encoder=self.text_encoder, - prefix=self.text_encoder_name, - lora_scale=self.lora_scale, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - self.load_lora_into_text_encoder( - state_dict, - network_alphas=None, - text_encoder=self.text_encoder_2, - prefix=f"{self.text_encoder_name}_2", - lora_scale=self.lora_scale, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.StableDiffusionLoraLoaderMixin.load_lora_into_text_encoder - def load_lora_into_text_encoder( - cls, - state_dict, - network_alphas, - text_encoder, - prefix=None, - lora_scale=1.0, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - This will load the LoRA layers specified in `state_dict` into `text_encoder` - - Parameters: - state_dict (`dict`): - A standard state dict containing the lora layer parameters. The key should be prefixed with an - additional `text_encoder` to distinguish between unet lora layers. - network_alphas (`dict[str, float]`): - The value of the network alpha used for stable learning and preventing underflow. This value has the - same meaning as the `--network_alpha` option in the kohya-ss trainer script. Refer to [this - link](https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning). - text_encoder (`CLIPTextModel`): - The text encoder model to load the LoRA layers into. - prefix (`str`): - Expected prefix of the `text_encoder` in the `state_dict`. - lora_scale (`float`): - How much to scale the output of the lora linear layer before it is added with the output of the regular - lora layer. - adapter_name (`str`, *optional*): - Adapter name to be used for referencing the loaded adapter model. If not specified, it will use - `default_{i}` where i is the total number of adapters being loaded. - low_cpu_mem_usage (`bool`, *optional*): - Speed up model loading by only loading the pretrained LoRA weights and not initializing the random - weights. - hotswap (`bool`, *optional*): - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`]. - metadata (`dict`): - Optional LoRA adapter metadata. When supplied, the `LoraConfig` arguments of `peft` won't be derived - from the state dict. - """ - _load_lora_into_text_encoder( - state_dict=state_dict, - network_alphas=network_alphas, - lora_scale=lora_scale, - text_encoder=text_encoder, - prefix=prefix, - text_encoder_name=cls.text_encoder_name, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.StableDiffusionXLLoraLoaderMixin.save_lora_weights with unet->transformer - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - text_encoder_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - text_encoder_2_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata=None, - text_encoder_lora_adapter_metadata=None, - text_encoder_2_lora_adapter_metadata=None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if text_encoder_lora_layers: - lora_layers["text_encoder"] = text_encoder_lora_layers - lora_metadata["text_encoder"] = text_encoder_lora_adapter_metadata - - if text_encoder_2_lora_layers: - lora_layers["text_encoder_2"] = text_encoder_2_lora_layers - lora_metadata["text_encoder_2"] = text_encoder_2_lora_adapter_metadata - - if not lora_layers: - raise ValueError( - "You must pass at least one of `transformer_lora_layers`, `text_encoder_lora_layers`, or `text_encoder_2_lora_layers`." - ) - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.StableDiffusionXLLoraLoaderMixin.fuse_lora with unet->transformer - def fuse_lora( - self, - components: list[str] = ["transformer", "text_encoder", "text_encoder_2"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.StableDiffusionXLLoraLoaderMixin.unfuse_lora with unet->transformer - def unfuse_lora(self, components: list[str] = ["transformer", "text_encoder", "text_encoder_2"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class AuraFlowLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`AuraFlowTransformer2DModel`] Specific to [`AuraFlowPipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.lora_state_dict - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->AuraFlowTransformer2DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.SanaLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.SanaLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer", "text_encoder"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class FluxLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`FluxTransformer2DModel`], - [`CLIPTextModel`](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel). - - Specific to [`FluxPipeline`]. - """ - - _lora_loadable_modules = ["transformer", "text_encoder"] - transformer_name = TRANSFORMER_NAME - text_encoder_name = TEXT_ENCODER_NAME - _control_lora_supported_norm_keys = ["norm_q", "norm_k", "norm_added_q", "norm_added_k"] - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - return_alphas: bool = False, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - # TODO (sayakpaul): to a follow-up to clean and try to unify the conditions. - is_kohya = any(".lora_down.weight" in k for k in state_dict) - if is_kohya: - state_dict = _convert_kohya_flux_lora_to_diffusers(state_dict) - # Kohya already takes care of scaling the LoRA parameters with alpha. - return cls._prepare_outputs( - state_dict, - metadata=metadata, - alphas=None, - return_alphas=return_alphas, - return_metadata=return_lora_metadata, - ) - - is_xlabs = any("processor" in k for k in state_dict) - if is_xlabs: - state_dict = _convert_xlabs_flux_lora_to_diffusers(state_dict) - # xlabs doesn't use `alpha`. - return cls._prepare_outputs( - state_dict, - metadata=metadata, - alphas=None, - return_alphas=return_alphas, - return_metadata=return_lora_metadata, - ) - - is_bfl_control = any("query_norm.scale" in k for k in state_dict) - if is_bfl_control: - state_dict = _convert_bfl_flux_control_lora_to_diffusers(state_dict) - return cls._prepare_outputs( - state_dict, - metadata=metadata, - alphas=None, - return_alphas=return_alphas, - return_metadata=return_lora_metadata, - ) - - is_fal_kontext = any("base_model" in k for k in state_dict) - if is_fal_kontext: - state_dict = _convert_fal_kontext_lora_to_diffusers(state_dict) - return cls._prepare_outputs( - state_dict, - metadata=metadata, - alphas=None, - return_alphas=return_alphas, - return_metadata=return_lora_metadata, - ) - - # For state dicts like - # https://huggingface.co/TheLastBen/Jon_Snow_Flux_LoRA - keys = list(state_dict.keys()) - network_alphas = {} - for k in keys: - if "alpha" in k: - alpha_value = state_dict.get(k) - if (torch.is_tensor(alpha_value) and torch.is_floating_point(alpha_value)) or isinstance( - alpha_value, float - ): - network_alphas[k] = state_dict.pop(k) - else: - raise ValueError( - f"The alpha key ({k}) seems to be incorrect. If you think this error is unexpected, please open as issue." - ) - - if return_alphas or return_lora_metadata: - return cls._prepare_outputs( - state_dict, - metadata=metadata, - alphas=network_alphas, - return_alphas=return_alphas, - return_metadata=return_lora_metadata, - ) - else: - return state_dict - - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and not is_peft_version(">=", "0.13.1"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, network_alphas, metadata = self.lora_state_dict( - pretrained_model_name_or_path_or_dict, return_alphas=True, **kwargs - ) - - has_lora_keys = any("lora" in key for key in state_dict.keys()) - - # Flux Control LoRAs also have norm keys - has_norm_keys = any( - norm_key in key for key in state_dict.keys() for norm_key in self._control_lora_supported_norm_keys - ) - - if not (has_lora_keys or has_norm_keys): - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - transformer_lora_state_dict = { - k: state_dict.get(k) - for k in list(state_dict.keys()) - if k.startswith(f"{self.transformer_name}.") and "lora" in k - } - transformer_norm_state_dict = { - k: state_dict.pop(k) - for k in list(state_dict.keys()) - if k.startswith(f"{self.transformer_name}.") - and any(norm_key in k for norm_key in self._control_lora_supported_norm_keys) - } - - transformer = getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer - has_param_with_expanded_shape = False - if len(transformer_lora_state_dict) > 0: - has_param_with_expanded_shape = self._maybe_expand_transformer_param_shape_or_error_( - transformer, transformer_lora_state_dict, transformer_norm_state_dict - ) - - if has_param_with_expanded_shape: - logger.info( - "The LoRA weights contain parameters that have different shapes that expected by the transformer. " - "As a result, the state_dict of the transformer has been expanded to match the LoRA parameter shapes. " - "To get a comprehensive list of parameter names that were modified, enable debug logging." - ) - if len(transformer_lora_state_dict) > 0: - transformer_lora_state_dict = self._maybe_expand_lora_state_dict( - transformer=transformer, lora_state_dict=transformer_lora_state_dict - ) - for k in transformer_lora_state_dict: - state_dict.update({k: transformer_lora_state_dict[k]}) - - self.load_lora_into_transformer( - state_dict, - network_alphas=network_alphas, - transformer=transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - if len(transformer_norm_state_dict) > 0: - transformer._transformer_norm_layers = self._load_norm_into_transformer( - transformer_norm_state_dict, - transformer=transformer, - discard_original_layers=False, - ) - - self.load_lora_into_text_encoder( - state_dict, - network_alphas=network_alphas, - text_encoder=self.text_encoder, - prefix=self.text_encoder_name, - lora_scale=self.lora_scale, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - def load_lora_into_transformer( - cls, - state_dict, - network_alphas, - transformer, - adapter_name=None, - metadata=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and not is_peft_version(">=", "0.13.1"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=network_alphas, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - def _load_norm_into_transformer( - cls, - state_dict, - transformer, - prefix=None, - discard_original_layers=False, - ) -> dict[str, torch.Tensor]: - # Remove prefix if present - prefix = prefix or cls.transformer_name - for key in list(state_dict.keys()): - if key.split(".")[0] == prefix: - state_dict[key.removeprefix(f"{prefix}.")] = state_dict.pop(key) - - # Find invalid keys - transformer_state_dict = transformer.state_dict() - transformer_keys = set(transformer_state_dict.keys()) - state_dict_keys = set(state_dict.keys()) - extra_keys = list(state_dict_keys - transformer_keys) - - if extra_keys: - logger.warning( - f"Unsupported keys found in state dict when trying to load normalization layers into the transformer. The following keys will be ignored:\n{extra_keys}." - ) - - for key in extra_keys: - state_dict.pop(key) - - # Save the layers that are going to be overwritten so that unload_lora_weights can work as expected - overwritten_layers_state_dict = {} - if not discard_original_layers: - for key in state_dict.keys(): - overwritten_layers_state_dict[key] = transformer_state_dict[key].clone() - - logger.info( - "The provided state dict contains normalization layers in addition to LoRA layers. The normalization layers will directly update the state_dict of the transformer " - 'as opposed to the LoRA layers that will co-exist separately until the "fuse_lora()" method is called. That is to say, the normalization layers will always be directly ' - "fused into the transformer and can only be unfused if `discard_original_layers=True` is passed. This might also have implications when dealing with multiple LoRAs. " - "If you notice something unexpected, please open an issue: https://github.com/huggingface/diffusers/issues." - ) - - # We can't load with strict=True because the current state_dict does not contain all the transformer keys - incompatible_keys = transformer.load_state_dict(state_dict, strict=False) - unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None) - - # We shouldn't expect to see the supported norm keys here being present in the unexpected keys. - if unexpected_keys: - if any(norm_key in k for k in unexpected_keys for norm_key in cls._control_lora_supported_norm_keys): - raise ValueError( - f"Found {unexpected_keys} as unexpected keys while trying to load norm layers into the transformer." - ) - - return overwritten_layers_state_dict - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.StableDiffusionLoraLoaderMixin.load_lora_into_text_encoder - def load_lora_into_text_encoder( - cls, - state_dict, - network_alphas, - text_encoder, - prefix=None, - lora_scale=1.0, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - This will load the LoRA layers specified in `state_dict` into `text_encoder` - - Parameters: - state_dict (`dict`): - A standard state dict containing the lora layer parameters. The key should be prefixed with an - additional `text_encoder` to distinguish between unet lora layers. - network_alphas (`dict[str, float]`): - The value of the network alpha used for stable learning and preventing underflow. This value has the - same meaning as the `--network_alpha` option in the kohya-ss trainer script. Refer to [this - link](https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning). - text_encoder (`CLIPTextModel`): - The text encoder model to load the LoRA layers into. - prefix (`str`): - Expected prefix of the `text_encoder` in the `state_dict`. - lora_scale (`float`): - How much to scale the output of the lora linear layer before it is added with the output of the regular - lora layer. - adapter_name (`str`, *optional*): - Adapter name to be used for referencing the loaded adapter model. If not specified, it will use - `default_{i}` where i is the total number of adapters being loaded. - low_cpu_mem_usage (`bool`, *optional*): - Speed up model loading by only loading the pretrained LoRA weights and not initializing the random - weights. - hotswap (`bool`, *optional*): - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`]. - metadata (`dict`): - Optional LoRA adapter metadata. When supplied, the `LoraConfig` arguments of `peft` won't be derived - from the state dict. - """ - _load_lora_into_text_encoder( - state_dict=state_dict, - network_alphas=network_alphas, - lora_scale=lora_scale, - text_encoder=text_encoder, - prefix=prefix, - text_encoder_name=cls.text_encoder_name, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.StableDiffusionLoraLoaderMixin.save_lora_weights with unet->transformer - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - text_encoder_lora_layers: dict[str, torch.nn.Module] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata=None, - text_encoder_lora_adapter_metadata=None, - ): - r""" - Save the LoRA parameters corresponding to the UNet and text encoder. - - Arguments: - save_directory (`str` or `os.PathLike`): - Directory to save LoRA parameters to. Will be created if it doesn't exist. - transformer_lora_layers (`dict[str, torch.nn.Module]` or `dict[str, torch.Tensor]`): - State dict of the LoRA layers corresponding to the `transformer`. - text_encoder_lora_layers (`dict[str, torch.nn.Module]` or `dict[str, torch.Tensor]`): - State dict of the LoRA layers corresponding to the `text_encoder`. Must explicitly pass the text - encoder LoRA state dict because it comes from 🤗 Transformers. - is_main_process (`bool`, *optional*, defaults to `True`): - Whether the process calling this is the main process or not. Useful during distributed training and you - need to call this function on all processes. In this case, set `is_main_process=True` only on the main - process to avoid race conditions. - save_function (`Callable`): - The function to use to save the state dictionary. Useful during distributed training when you need to - replace `torch.save` with another method. Can be configured with the environment variable - `DIFFUSERS_SAVE_MODE`. - safe_serialization (`bool`, *optional*, defaults to `True`): - Whether to save the model using `safetensors` or the traditional PyTorch way with `pickle`. - transformer_lora_adapter_metadata: - LoRA adapter metadata associated with the transformer to be serialized with the state dict. - text_encoder_lora_adapter_metadata: - LoRA adapter metadata associated with the text encoder to be serialized with the state dict. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if text_encoder_lora_layers: - lora_layers[cls.text_encoder_name] = text_encoder_lora_layers - lora_metadata[cls.text_encoder_name] = text_encoder_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - - transformer = getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer - if ( - hasattr(transformer, "_transformer_norm_layers") - and isinstance(transformer._transformer_norm_layers, dict) - and len(transformer._transformer_norm_layers.keys()) > 0 - ): - logger.info( - "The provided state dict contains normalization layers in addition to LoRA layers. The normalization layers will be directly updated the state_dict of the transformer " - "as opposed to the LoRA layers that will co-exist separately until the 'fuse_lora()' method is called. That is to say, the normalization layers will always be directly " - "fused into the transformer and can only be unfused if `discard_original_layers=True` is passed." - ) - - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - def unfuse_lora(self, components: list[str] = ["transformer", "text_encoder"], **kwargs): - r""" - Reverses the effect of - [`pipe.fuse_lora()`](https://huggingface.co/docs/diffusers/main/en/api/loaders#diffusers.loaders.LoraBaseMixin.fuse_lora). - - Args: - components (`list[str]`): list of LoRA-injectable components to unfuse LoRA from. - """ - transformer = getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer - if hasattr(transformer, "_transformer_norm_layers") and transformer._transformer_norm_layers: - transformer.load_state_dict(transformer._transformer_norm_layers, strict=False) - - super().unfuse_lora(components=components, **kwargs) - - # We override this here account for `_transformer_norm_layers` and `_overwritten_params`. - def unload_lora_weights(self, reset_to_overwritten_params=False): - """ - Unloads the LoRA parameters. - - Args: - reset_to_overwritten_params (`bool`, defaults to `False`): Whether to reset the LoRA-loaded modules - to their original params. Refer to the [Flux - documentation](https://huggingface.co/docs/diffusers/main/en/api/pipelines/flux) to learn more. - - Examples: - - ```python - >>> # Assuming `pipeline` is already loaded with the LoRA parameters. - >>> pipeline.unload_lora_weights() - >>> ... - ``` - """ - super().unload_lora_weights() - - transformer = getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer - if hasattr(transformer, "_transformer_norm_layers") and transformer._transformer_norm_layers: - transformer.load_state_dict(transformer._transformer_norm_layers, strict=False) - transformer._transformer_norm_layers = None - - if reset_to_overwritten_params and getattr(transformer, "_overwritten_params", None) is not None: - overwritten_params = transformer._overwritten_params - module_names = set() - - for param_name in overwritten_params: - if param_name.endswith(".weight"): - module_names.add(param_name.replace(".weight", "")) - - for name, module in transformer.named_modules(): - if isinstance(module, torch.nn.Linear) and name in module_names: - module_weight = module.weight.data - module_bias = module.bias.data if module.bias is not None else None - bias = module_bias is not None - - parent_module_name, _, current_module_name = name.rpartition(".") - parent_module = transformer.get_submodule(parent_module_name) - - current_param_weight = overwritten_params[f"{name}.weight"] - in_features, out_features = current_param_weight.shape[1], current_param_weight.shape[0] - with torch.device("meta"): - original_module = torch.nn.Linear( - in_features, - out_features, - bias=bias, - dtype=module_weight.dtype, - ) - - tmp_state_dict = {"weight": current_param_weight} - if module_bias is not None: - tmp_state_dict.update({"bias": overwritten_params[f"{name}.bias"]}) - original_module.load_state_dict(tmp_state_dict, assign=True, strict=True) - setattr(parent_module, current_module_name, original_module) - - del tmp_state_dict - - if current_module_name in _MODULE_NAME_TO_ATTRIBUTE_MAP_FLUX: - attribute_name = _MODULE_NAME_TO_ATTRIBUTE_MAP_FLUX[current_module_name] - new_value = int(current_param_weight.shape[1]) - old_value = getattr(transformer.config, attribute_name) - setattr(transformer.config, attribute_name, new_value) - logger.info( - f"Set the {attribute_name} attribute of the model to {new_value} from {old_value}." - ) - - @classmethod - def _maybe_expand_transformer_param_shape_or_error_( - cls, - transformer: torch.nn.Module, - lora_state_dict=None, - norm_state_dict=None, - prefix=None, - ) -> bool: - """ - Control LoRA expands the shape of the input layer from (3072, 64) to (3072, 128). This method handles that and - generalizes things a bit so that any parameter that needs expansion receives appropriate treatment. - """ - state_dict = {} - if lora_state_dict is not None: - state_dict.update(lora_state_dict) - if norm_state_dict is not None: - state_dict.update(norm_state_dict) - - # Remove prefix if present - prefix = prefix or cls.transformer_name - for key in list(state_dict.keys()): - if key.split(".")[0] == prefix: - state_dict[key.removeprefix(f"{prefix}.")] = state_dict.pop(key) - - # Expand transformer parameter shapes if they don't match lora - has_param_with_shape_update = False - overwritten_params = {} - - is_peft_loaded = getattr(transformer, "peft_config", None) is not None - is_quantized = hasattr(transformer, "hf_quantizer") - for name, module in transformer.named_modules(): - if isinstance(module, torch.nn.Linear): - module_weight = module.weight.data - module_bias = module.bias.data if module.bias is not None else None - bias = module_bias is not None - - lora_base_name = name.replace(".base_layer", "") if is_peft_loaded else name - lora_A_weight_name = f"{lora_base_name}.lora_A.weight" - lora_B_weight_name = f"{lora_base_name}.lora_B.weight" - if lora_A_weight_name not in state_dict: - continue - - in_features = state_dict[lora_A_weight_name].shape[1] - out_features = state_dict[lora_B_weight_name].shape[0] - - # Model maybe loaded with different quantization schemes which may flatten the params. - # `bitsandbytes`, for example, flatten the weights when using 4bit. 8bit bnb models - # preserve weight shape. - module_weight_shape = cls._calculate_module_shape(model=transformer, base_module=module) - - # This means there's no need for an expansion in the params, so we simply skip. - if tuple(module_weight_shape) == (out_features, in_features): - continue - - module_out_features, module_in_features = module_weight_shape - debug_message = "" - if in_features > module_in_features: - debug_message += ( - f'Expanding the nn.Linear input/output features for module="{name}" because the provided LoRA ' - f"checkpoint contains higher number of features than expected. The number of input_features will be " - f"expanded from {module_in_features} to {in_features}" - ) - if out_features > module_out_features: - debug_message += ( - ", and the number of output features will be " - f"expanded from {module_out_features} to {out_features}." - ) - else: - debug_message += "." - if debug_message: - logger.debug(debug_message) - - if out_features > module_out_features or in_features > module_in_features: - has_param_with_shape_update = True - parent_module_name, _, current_module_name = name.rpartition(".") - parent_module = transformer.get_submodule(parent_module_name) - - if is_quantized: - module_weight = _maybe_dequantize_weight_for_expanded_lora(transformer, module) - - # TODO: consider if this layer needs to be a quantized layer as well if `is_quantized` is True. - with torch.device("meta"): - expanded_module = torch.nn.Linear( - in_features, out_features, bias=bias, dtype=module_weight.dtype - ) - # Only weights are expanded and biases are not. This is because only the input dimensions - # are changed while the output dimensions remain the same. The shape of the weight tensor - # is (out_features, in_features), while the shape of bias tensor is (out_features,), which - # explains the reason why only weights are expanded. - new_weight = torch.zeros_like( - expanded_module.weight.data, device=module_weight.device, dtype=module_weight.dtype - ) - slices = tuple(slice(0, dim) for dim in module_weight_shape) - new_weight[slices] = module_weight - tmp_state_dict = {"weight": new_weight} - if module_bias is not None: - tmp_state_dict["bias"] = module_bias - expanded_module.load_state_dict(tmp_state_dict, strict=True, assign=True) - - setattr(parent_module, current_module_name, expanded_module) - - del tmp_state_dict - - if current_module_name in _MODULE_NAME_TO_ATTRIBUTE_MAP_FLUX: - attribute_name = _MODULE_NAME_TO_ATTRIBUTE_MAP_FLUX[current_module_name] - new_value = int(expanded_module.weight.data.shape[1]) - old_value = getattr(transformer.config, attribute_name) - setattr(transformer.config, attribute_name, new_value) - logger.info( - f"Set the {attribute_name} attribute of the model to {new_value} from {old_value}." - ) - - # For `unload_lora_weights()`. - # TODO: this could lead to more memory overhead if the number of overwritten params - # are large. Should be revisited later and tackled through a `discard_original_layers` arg. - overwritten_params[f"{current_module_name}.weight"] = module_weight - if module_bias is not None: - overwritten_params[f"{current_module_name}.bias"] = module_bias - - if len(overwritten_params) > 0: - transformer._overwritten_params = overwritten_params - - return has_param_with_shape_update - - @classmethod - def _maybe_expand_lora_state_dict(cls, transformer, lora_state_dict): - expanded_module_names = set() - transformer_state_dict = transformer.state_dict() - prefix = f"{cls.transformer_name}." - - lora_module_names = [ - key[: -len(".lora_A.weight")] for key in lora_state_dict if key.endswith(".lora_A.weight") - ] - lora_module_names = [name[len(prefix) :] for name in lora_module_names if name.startswith(prefix)] - lora_module_names = sorted(set(lora_module_names)) - transformer_module_names = sorted({name for name, _ in transformer.named_modules()}) - unexpected_modules = set(lora_module_names) - set(transformer_module_names) - if unexpected_modules: - logger.debug(f"Found unexpected modules: {unexpected_modules}. These will be ignored.") - - for k in lora_module_names: - if k in unexpected_modules: - continue - - base_param_name = ( - f"{k.replace(prefix, '')}.base_layer.weight" - if f"{k.replace(prefix, '')}.base_layer.weight" in transformer_state_dict - else f"{k.replace(prefix, '')}.weight" - ) - base_weight_param = transformer_state_dict[base_param_name] - lora_A_param = lora_state_dict[f"{prefix}{k}.lora_A.weight"] - - # TODO (sayakpaul): Handle the cases when we actually need to expand when using quantization. - base_module_shape = cls._calculate_module_shape(model=transformer, base_weight_param_name=base_param_name) - - if base_module_shape[1] > lora_A_param.shape[1]: - shape = (lora_A_param.shape[0], base_weight_param.shape[1]) - expanded_state_dict_weight = torch.zeros(shape, device=base_weight_param.device) - expanded_state_dict_weight[:, : lora_A_param.shape[1]].copy_(lora_A_param) - lora_state_dict[f"{prefix}{k}.lora_A.weight"] = expanded_state_dict_weight - expanded_module_names.add(k) - elif base_module_shape[1] < lora_A_param.shape[1]: - raise NotImplementedError( - f"This LoRA param ({k}.lora_A.weight) has an incompatible shape {lora_A_param.shape}. Please open an issue to file for a feature request - https://github.com/huggingface/diffusers/issues/new." - ) - - if expanded_module_names: - logger.info( - f"The following LoRA modules were zero padded to match the state dict of {cls.transformer_name}: {expanded_module_names}. Please open an issue if you think this was unexpected - https://github.com/huggingface/diffusers/issues/new." - ) - - return lora_state_dict - - @staticmethod - def _calculate_module_shape( - model: "torch.nn.Module", - base_module: "torch.nn.Linear" = None, - base_weight_param_name: str = None, - ) -> "torch.Size": - def _get_weight_shape(weight: torch.Tensor): - if weight.__class__.__name__ == "Params4bit": - return weight.quant_state.shape - elif weight.__class__.__name__ == "GGUFParameter": - return weight.quant_shape - else: - return weight.shape - - if base_module is not None: - return _get_weight_shape(base_module.weight) - elif base_weight_param_name is not None: - if not base_weight_param_name.endswith(".weight"): - raise ValueError( - f"Invalid `base_weight_param_name` passed as it does not end with '.weight' {base_weight_param_name=}." - ) - module_path = base_weight_param_name.rsplit(".weight", 1)[0] - submodule = get_submodule_by_name(model, module_path) - return _get_weight_shape(submodule.weight) - - raise ValueError("Either `base_module` or `base_weight_param_name` must be provided.") - - @staticmethod - def _prepare_outputs(state_dict, metadata, alphas=None, return_alphas=False, return_metadata=False): - outputs = [state_dict] - if return_alphas: - outputs.append(alphas) - if return_metadata: - outputs.append(metadata) - return tuple(outputs) if (return_alphas or return_metadata) else state_dict - - -# The reason why we subclass from `StableDiffusionLoraLoaderMixin` here is because Amused initially -# relied on `StableDiffusionLoraLoaderMixin` for its LoRA support. -class AmusedLoraLoaderMixin(StableDiffusionLoraLoaderMixin): - _lora_loadable_modules = ["transformer", "text_encoder"] - transformer_name = TRANSFORMER_NAME - text_encoder_name = TEXT_ENCODER_NAME - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.FluxLoraLoaderMixin.load_lora_into_transformer with FluxTransformer2DModel->UVit2DModel - def load_lora_into_transformer( - cls, - state_dict, - network_alphas, - transformer, - adapter_name=None, - metadata=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and not is_peft_version(">=", "0.13.1"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=network_alphas, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.StableDiffusionLoraLoaderMixin.load_lora_into_text_encoder - def load_lora_into_text_encoder( - cls, - state_dict, - network_alphas, - text_encoder, - prefix=None, - lora_scale=1.0, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - This will load the LoRA layers specified in `state_dict` into `text_encoder` - - Parameters: - state_dict (`dict`): - A standard state dict containing the lora layer parameters. The key should be prefixed with an - additional `text_encoder` to distinguish between unet lora layers. - network_alphas (`dict[str, float]`): - The value of the network alpha used for stable learning and preventing underflow. This value has the - same meaning as the `--network_alpha` option in the kohya-ss trainer script. Refer to [this - link](https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning). - text_encoder (`CLIPTextModel`): - The text encoder model to load the LoRA layers into. - prefix (`str`): - Expected prefix of the `text_encoder` in the `state_dict`. - lora_scale (`float`): - How much to scale the output of the lora linear layer before it is added with the output of the regular - lora layer. - adapter_name (`str`, *optional*): - Adapter name to be used for referencing the loaded adapter model. If not specified, it will use - `default_{i}` where i is the total number of adapters being loaded. - low_cpu_mem_usage (`bool`, *optional*): - Speed up model loading by only loading the pretrained LoRA weights and not initializing the random - weights. - hotswap (`bool`, *optional*): - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`]. - metadata (`dict`): - Optional LoRA adapter metadata. When supplied, the `LoraConfig` arguments of `peft` won't be derived - from the state dict. - """ - _load_lora_into_text_encoder( - state_dict=state_dict, - network_alphas=network_alphas, - lora_scale=lora_scale, - text_encoder=text_encoder, - prefix=prefix, - text_encoder_name=cls.text_encoder_name, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - text_encoder_lora_layers: dict[str, torch.nn.Module] = None, - transformer_lora_layers: dict[str, torch.nn.Module] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - ): - r""" - Save the LoRA parameters corresponding to the UNet and text encoder. - - Arguments: - save_directory (`str` or `os.PathLike`): - Directory to save LoRA parameters to. Will be created if it doesn't exist. - unet_lora_layers (`dict[str, torch.nn.Module]` or `dict[str, torch.Tensor]`): - State dict of the LoRA layers corresponding to the `unet`. - text_encoder_lora_layers (`dict[str, torch.nn.Module]` or `dict[str, torch.Tensor]`): - State dict of the LoRA layers corresponding to the `text_encoder`. Must explicitly pass the text - encoder LoRA state dict because it comes from 🤗 Transformers. - is_main_process (`bool`, *optional*, defaults to `True`): - Whether the process calling this is the main process or not. Useful during distributed training and you - need to call this function on all processes. In this case, set `is_main_process=True` only on the main - process to avoid race conditions. - save_function (`Callable`): - The function to use to save the state dictionary. Useful during distributed training when you need to - replace `torch.save` with another method. Can be configured with the environment variable - `DIFFUSERS_SAVE_MODE`. - safe_serialization (`bool`, *optional*, defaults to `True`): - Whether to save the model using `safetensors` or the traditional PyTorch way with `pickle`. - """ - state_dict = {} - - if not (transformer_lora_layers or text_encoder_lora_layers): - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - if transformer_lora_layers: - state_dict.update(cls.pack_weights(transformer_lora_layers, cls.transformer_name)) - - if text_encoder_lora_layers: - state_dict.update(cls.pack_weights(text_encoder_lora_layers, cls.text_encoder_name)) - - # Save the model - cls.write_lora_layers( - state_dict=state_dict, - save_directory=save_directory, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - -class CogVideoXLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`CogVideoXTransformer3DModel`]. Specific to [`CogVideoXPipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.lora_state_dict - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->CogVideoXTransformer3DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class Mochi1LoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`MochiTransformer3DModel`]. Specific to [`MochiPipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.lora_state_dict - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->MochiTransformer3DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class LTXVideoLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`LTXVideoTransformer3DModel`]. Specific to [`LTXPipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - is_non_diffusers_format = any(k.startswith("diffusion_model.") for k in state_dict) - if is_non_diffusers_format: - state_dict = _convert_non_diffusers_ltxv_lora_to_diffusers(state_dict) - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->LTXVideoTransformer3DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class LTX2LoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`LTX2VideoTransformer3DModel`]. Specific to [`LTX2Pipeline`]. - """ - - _lora_loadable_modules = ["transformer", "connectors"] - transformer_name = TRANSFORMER_NAME - connectors_name = LTX2_CONNECTOR_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - final_state_dict = state_dict - is_non_diffusers_format = any(k.startswith("diffusion_model.") for k in state_dict) - has_connector = any(k.startswith("text_embedding_projection.") for k in state_dict) - if is_non_diffusers_format: - final_state_dict = _convert_non_diffusers_ltx2_lora_to_diffusers(state_dict) - if has_connector: - connectors_state_dict = _convert_non_diffusers_ltx2_lora_to_diffusers( - state_dict, "text_embedding_projection" - ) - final_state_dict.update(connectors_state_dict) - out = (final_state_dict, metadata) if return_lora_metadata else final_state_dict - return out - - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - transformer_peft_state_dict = { - k: v for k, v in state_dict.items() if k.startswith(f"{self.transformer_name}.") - } - connectors_peft_state_dict = {k: v for k, v in state_dict.items() if k.startswith(f"{self.connectors_name}.")} - self.load_lora_into_transformer( - transformer_peft_state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - if connectors_peft_state_dict: - self.load_lora_into_transformer( - connectors_peft_state_dict, - transformer=getattr(self, self.connectors_name) - if not hasattr(self, "connectors") - else self.connectors, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - prefix=self.connectors_name, - ) - - @classmethod - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - prefix: str = "transformer", - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {prefix}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - prefix=prefix, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class SanaLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`SanaTransformer2DModel`]. Specific to [`SanaPipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.lora_state_dict - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->SanaTransformer2DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class HeliosLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`HeliosTransformer3DModel`]. Specific to [`HeliosPipeline`] and [`HeliosPyramidPipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - if any(k.startswith("diffusion_model.") for k in state_dict): - state_dict = _convert_non_diffusers_wan_lora_to_diffusers(state_dict) - elif any(k.startswith("lora_unet_") for k in state_dict): - state_dict = _convert_musubi_wan_lora_to_diffusers(state_dict) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->WanTransformer3DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class HunyuanVideoLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`HunyuanVideoTransformer3DModel`]. Specific to [`HunyuanVideoPipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - is_original_hunyuan_video = any("img_attn_qkv" in k for k in state_dict) - if is_original_hunyuan_video: - state_dict = _convert_hunyuan_video_lora_to_diffusers(state_dict) - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->HunyuanVideoTransformer3DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class Lumina2LoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`Lumina2Transformer2DModel`]. Specific to [`Lumina2Text2ImgPipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - # conversion. - non_diffusers = any(k.startswith("diffusion_model.") for k in state_dict) - if non_diffusers: - state_dict = _convert_non_diffusers_lumina2_lora_to_diffusers(state_dict) - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->Lumina2Transformer2DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.SanaLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.SanaLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class KandinskyLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`Kandinsky5Transformer3DModel`], - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.lora_state_dict - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class WanLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`WanTransformer3DModel`]. Specific to [`WanPipeline`] and `[WanImageToVideoPipeline`]. - """ - - _lora_loadable_modules = ["transformer", "transformer_2"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - if any(k.startswith("diffusion_model.") for k in state_dict): - state_dict = _convert_non_diffusers_wan_lora_to_diffusers(state_dict) - elif any(k.startswith("lora_unet_") for k in state_dict): - state_dict = _convert_musubi_wan_lora_to_diffusers(state_dict) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - @classmethod - def _maybe_expand_t2v_lora_for_i2v( - cls, - transformer: torch.nn.Module, - state_dict, - ): - if transformer.config.image_dim is None: - return state_dict - - target_device = transformer.device - - if any(k.startswith("transformer.blocks.") for k in state_dict): - num_blocks = len({k.split("blocks.")[1].split(".")[0] for k in state_dict if "blocks." in k}) - is_i2v_lora = any("add_k_proj" in k for k in state_dict) and any("add_v_proj" in k for k in state_dict) - has_bias = any(".lora_B.bias" in k for k in state_dict) - - if is_i2v_lora: - return state_dict - - for i in range(num_blocks): - for o, c in zip(["k_img", "v_img"], ["add_k_proj", "add_v_proj"]): - # These keys should exist if the block `i` was part of the T2V LoRA. - ref_key_lora_A = f"transformer.blocks.{i}.attn2.to_k.lora_A.weight" - ref_key_lora_B = f"transformer.blocks.{i}.attn2.to_k.lora_B.weight" - - if ref_key_lora_A not in state_dict or ref_key_lora_B not in state_dict: - continue - - state_dict[f"transformer.blocks.{i}.attn2.{c}.lora_A.weight"] = torch.zeros_like( - state_dict[f"transformer.blocks.{i}.attn2.to_k.lora_A.weight"], device=target_device - ) - state_dict[f"transformer.blocks.{i}.attn2.{c}.lora_B.weight"] = torch.zeros_like( - state_dict[f"transformer.blocks.{i}.attn2.to_k.lora_B.weight"], device=target_device - ) - - # If the original LoRA had biases (indicated by has_bias) - # AND the specific reference bias key exists for this block. - - ref_key_lora_B_bias = f"transformer.blocks.{i}.attn2.to_k.lora_B.bias" - if has_bias and ref_key_lora_B_bias in state_dict: - ref_lora_B_bias_tensor = state_dict[ref_key_lora_B_bias] - state_dict[f"transformer.blocks.{i}.attn2.{c}.lora_B.bias"] = torch.zeros_like( - ref_lora_B_bias_tensor, - device=target_device, - ) - - return state_dict - - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - # convert T2V LoRA to I2V LoRA (when loaded to Wan I2V) by adding zeros for the additional (missing) _img layers - state_dict = self._maybe_expand_t2v_lora_for_i2v( - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - state_dict=state_dict, - ) - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - load_into_transformer_2 = kwargs.pop("load_into_transformer_2", False) - if load_into_transformer_2: - if not hasattr(self, "transformer_2"): - raise AttributeError( - f"'{type(self).__name__}' object has no attribute transformer_2" - "Note that Wan2.1 models do not have a transformer_2 component." - "Ensure the model has a transformer_2 component before setting load_into_transformer_2=True." - ) - self.load_lora_into_transformer( - state_dict, - transformer=self.transformer_2, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - else: - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) - if not hasattr(self, "transformer") - else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->WanTransformer3DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class SkyReelsV2LoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`SkyReelsV2Transformer3DModel`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - # Copied from diffusers.loaders.lora_pipeline.WanLoraLoaderMixin.lora_state_dict - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - if any(k.startswith("diffusion_model.") for k in state_dict): - state_dict = _convert_non_diffusers_wan_lora_to_diffusers(state_dict) - elif any(k.startswith("lora_unet_") for k in state_dict): - state_dict = _convert_musubi_wan_lora_to_diffusers(state_dict) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.WanLoraLoaderMixin._maybe_expand_t2v_lora_for_i2v - def _maybe_expand_t2v_lora_for_i2v( - cls, - transformer: torch.nn.Module, - state_dict, - ): - if transformer.config.image_dim is None: - return state_dict - - target_device = transformer.device - - if any(k.startswith("transformer.blocks.") for k in state_dict): - num_blocks = len({k.split("blocks.")[1].split(".")[0] for k in state_dict if "blocks." in k}) - is_i2v_lora = any("add_k_proj" in k for k in state_dict) and any("add_v_proj" in k for k in state_dict) - has_bias = any(".lora_B.bias" in k for k in state_dict) - - if is_i2v_lora: - return state_dict - - for i in range(num_blocks): - for o, c in zip(["k_img", "v_img"], ["add_k_proj", "add_v_proj"]): - # These keys should exist if the block `i` was part of the T2V LoRA. - ref_key_lora_A = f"transformer.blocks.{i}.attn2.to_k.lora_A.weight" - ref_key_lora_B = f"transformer.blocks.{i}.attn2.to_k.lora_B.weight" - - if ref_key_lora_A not in state_dict or ref_key_lora_B not in state_dict: - continue - - state_dict[f"transformer.blocks.{i}.attn2.{c}.lora_A.weight"] = torch.zeros_like( - state_dict[f"transformer.blocks.{i}.attn2.to_k.lora_A.weight"], device=target_device - ) - state_dict[f"transformer.blocks.{i}.attn2.{c}.lora_B.weight"] = torch.zeros_like( - state_dict[f"transformer.blocks.{i}.attn2.to_k.lora_B.weight"], device=target_device - ) - - # If the original LoRA had biases (indicated by has_bias) - # AND the specific reference bias key exists for this block. - - ref_key_lora_B_bias = f"transformer.blocks.{i}.attn2.to_k.lora_B.bias" - if has_bias and ref_key_lora_B_bias in state_dict: - ref_lora_B_bias_tensor = state_dict[ref_key_lora_B_bias] - state_dict[f"transformer.blocks.{i}.attn2.{c}.lora_B.bias"] = torch.zeros_like( - ref_lora_B_bias_tensor, - device=target_device, - ) - - return state_dict - - # Copied from diffusers.loaders.lora_pipeline.WanLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - # convert T2V LoRA to I2V LoRA (when loaded to Wan I2V) by adding zeros for the additional (missing) _img layers - state_dict = self._maybe_expand_t2v_lora_for_i2v( - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - state_dict=state_dict, - ) - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - load_into_transformer_2 = kwargs.pop("load_into_transformer_2", False) - if load_into_transformer_2: - if not hasattr(self, "transformer_2"): - raise AttributeError( - f"'{type(self).__name__}' object has no attribute transformer_2" - "Note that Wan2.1 models do not have a transformer_2 component." - "Ensure the model has a transformer_2 component before setting load_into_transformer_2=True." - ) - self.load_lora_into_transformer( - state_dict, - transformer=self.transformer_2, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - else: - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) - if not hasattr(self, "transformer") - else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->SkyReelsV2Transformer3DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class CogView4LoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`WanTransformer3DModel`]. Specific to [`CogView4Pipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.lora_state_dict - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->CogView4Transformer2DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class HiDreamImageLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`HiDreamImageTransformer2DModel`]. Specific to [`HiDreamImagePipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - is_non_diffusers_format = any("diffusion_model" in k for k in state_dict) - if is_non_diffusers_format: - state_dict = _convert_non_diffusers_hidream_lora_to_diffusers(state_dict) - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->HiDreamImageTransformer2DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.SanaLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.SanaLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class QwenImageLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`QwenImageTransformer2DModel`]. Specific to [`QwenImagePipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - has_alphas_in_sd = any(k.endswith(".alpha") for k in state_dict) - has_lora_unet = any(k.startswith("lora_unet_") for k in state_dict) - has_diffusion_model = any(k.startswith("diffusion_model.") for k in state_dict) - has_default = any("default." in k for k in state_dict) - if has_alphas_in_sd or has_lora_unet or has_diffusion_model or has_default: - state_dict = _convert_non_diffusers_qwen_lora_to_diffusers(state_dict) - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->QwenImageTransformer2DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class Krea2LoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`Krea2Transformer2DModel`]. Specific to [`Krea2Pipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - is_non_diffusers_format = any( - k.startswith("diffusion_model.") or k.startswith("base_model.model.") for k in state_dict - ) - if is_non_diffusers_format: - state_dict = _convert_non_diffusers_krea2_lora_to_diffusers(state_dict) - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->Krea2Transformer2DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class ZImageLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`ZImageTransformer2DModel`]. Specific to [`ZImagePipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - has_alphas_in_sd = any(k.endswith(".alpha") for k in state_dict) - has_lora_unet = any(k.startswith("lora_unet_") for k in state_dict) - has_diffusion_model = any(k.startswith("diffusion_model.") for k in state_dict) - has_default = any("default." in k for k in state_dict) - if has_alphas_in_sd or has_lora_unet or has_diffusion_model or has_default: - state_dict = _convert_non_diffusers_z_image_lora_to_diffusers(state_dict) - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->ZImageTransformer2DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class AnimaLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`CosmosTransformer3DModel`] and [`AnimaTextConditioner`]. - """ - - _lora_loadable_modules = ["transformer", "text_conditioner"] - transformer_name = TRANSFORMER_NAME - text_conditioner_name = "text_conditioner" - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - has_diffusion_model = any(k.startswith("diffusion_model.") for k in state_dict) - if has_diffusion_model: - state_dict = _convert_non_diffusers_anima_lora_to_diffusers(state_dict) - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - transformer_state_dict = {k: v for k, v in state_dict.items() if k.startswith(f"{self.transformer_name}.")} - text_conditioner_state_dict = { - k: v for k, v in state_dict.items() if k.startswith(f"{self.text_conditioner_name}.") - } - - if transformer_state_dict: - self.load_lora_into_transformer( - transformer_state_dict, - transformer=self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - if text_conditioner_state_dict: - self.load_lora_into_text_conditioner( - text_conditioner_state_dict, - text_conditioner=self.text_conditioner, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - def load_lora_into_text_conditioner( - cls, - state_dict, - text_conditioner, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - logger.info(f"Loading {cls.text_conditioner_name}.") - text_conditioner.load_lora_adapter( - state_dict, - prefix=cls.text_conditioner_name, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - def fuse_lora( - self, - components: list[str] = ["transformer", "text_conditioner"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - def unfuse_lora(self, components: list[str] = ["transformer", "text_conditioner"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class Flux2LoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`Flux2Transformer2DModel`]. Specific to [`Flux2Pipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - is_kohya = any(".lora_down.weight" in k for k in state_dict) - if is_kohya: - state_dict = _convert_kohya_flux2_lora_to_diffusers(state_dict) - # Kohya already takes care of scaling the LoRA parameters with alpha. - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - is_peft_format = any(k.startswith("base_model.model.") for k in state_dict) - if is_peft_format: - state_dict = {k.replace("base_model.model.", "diffusion_model."): v for k, v in state_dict.items()} - - is_ai_toolkit = any(k.startswith("diffusion_model.") for k in state_dict) - if is_ai_toolkit: - state_dict = _convert_non_diffusers_flux2_lora_to_diffusers(state_dict) - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->CogView4Transformer2DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class Ideogram4LoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`Ideogram4Transformer2DModel`]. Specific to [`Ideogram4Pipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - # ai-toolkit (ostris) saves Ideogram4 LoRAs under a `diffusion_model.` prefix with a fused - # `attention.qkv` projection; convert those to the diffusers layout before loading. - is_non_diffusers_format = any(k.startswith("diffusion_model.") for k in state_dict) or any( - ".attention.qkv." in k for k in state_dict - ) - if is_non_diffusers_format: - state_dict = _convert_non_diffusers_ideogram4_lora_to_diffusers(state_dict) - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->CogView4Transformer2DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class ErnieImageLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`ErnieImageTransformer2DModel`]. Specific to [`ErnieImagePipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - # PEFT format -> normalize to diffusion_model.* prefix - is_peft_format = any(k.startswith("base_model.model.") for k in state_dict) - if is_peft_format: - state_dict = {k.replace("base_model.model.", "diffusion_model."): v for k, v in state_dict.items()} - - # AI-Toolkit / diffusion_model.* prefix -> swap to transformer.* - # The Ernie LoRA naming under diffusion_model.* already matches diffusers module - # paths (layers.X.self_attention.to_q etc.), so only the prefix needs to change. - is_diffusion_model_prefix = any(k.startswith("diffusion_model.") for k in state_dict) - if is_diffusion_model_prefix: - state_dict = {k.replace("diffusion_model.", "transformer."): v for k, v in state_dict.items()} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->ErnieImageTransformer2DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class CosmosLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`CosmosTransformer3DModel`], Specific to [`Cosmos2_5_PredictBasePipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - text_encoder_name = TEXT_ENCODER_NAME - - @classmethod - @validate_hf_hub_args - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.lora_state_dict - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - # Load the main state dict first which has the LoRA layers for either of - # transformer and text encoder or both. - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->CosmosTransformer3DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class AceStepLoraLoaderMixin(LoraBaseMixin): - r""" - Load LoRA layers into [`AceStepTransformer1DModel`]. Specific to [`AceStepPipeline`]. - """ - - _lora_loadable_modules = ["transformer"] - transformer_name = TRANSFORMER_NAME - - @classmethod - @validate_hf_hub_args - def lora_state_dict( - cls, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details. - """ - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - return_lora_metadata = kwargs.pop("return_lora_metadata", False) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - ) - - is_dora_scale_present = any("dora_scale" in k for k in state_dict) - if is_dora_scale_present: - warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." - logger.warning(warn_msg) - state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} - - # Detect original ACE-Step-1.5 PEFT format (q_proj/k_proj naming). - is_original_ace_step = any("q_proj" in k or "k_proj" in k for k in state_dict) - if is_original_ace_step: - state_dict = _convert_non_diffusers_ace_step_lora_to_diffusers(state_dict) - - out = (state_dict, metadata) if return_lora_metadata else state_dict - return out - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights - def load_lora_weights( - self, - pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], - adapter_name: str | None = None, - hotswap: bool = False, - **kwargs, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_weights`] for more details. - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # if a dict is passed, copy it instead of modifying it inplace - if isinstance(pretrained_model_name_or_path_or_dict, dict): - pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() - - # First, ensure that the checkpoint is a compatible one and can be successfully loaded. - kwargs["return_lora_metadata"] = True - state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) - - is_correct_format = all("lora" in key for key in state_dict.keys()) - if not is_correct_format: - raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") - - self.load_lora_into_transformer( - state_dict, - transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=self, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->AceStepTransformer1DModel - def load_lora_into_transformer( - cls, - state_dict, - transformer, - adapter_name=None, - _pipeline=None, - low_cpu_mem_usage=False, - hotswap: bool = False, - metadata=None, - ): - """ - See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_unet`] for more details. - """ - if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - # Load the layers corresponding to transformer. - logger.info(f"Loading {cls.transformer_name}.") - transformer.load_lora_adapter( - state_dict, - network_alphas=None, - adapter_name=adapter_name, - metadata=metadata, - _pipeline=_pipeline, - low_cpu_mem_usage=low_cpu_mem_usage, - hotswap=hotswap, - ) - - @classmethod - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights - def save_lora_weights( - cls, - save_directory: str | os.PathLike, - transformer_lora_layers: dict[str, torch.nn.Module | torch.Tensor] = None, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - transformer_lora_adapter_metadata: dict | None = None, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.save_lora_weights`] for more information. - """ - lora_layers = {} - lora_metadata = {} - - if transformer_lora_layers: - lora_layers[cls.transformer_name] = transformer_lora_layers - lora_metadata[cls.transformer_name] = transformer_lora_adapter_metadata - - if not lora_layers: - raise ValueError("You must pass at least one of `transformer_lora_layers` or `text_encoder_lora_layers`.") - - cls._save_lora_weights( - save_directory=save_directory, - lora_layers=lora_layers, - lora_metadata=lora_metadata, - is_main_process=is_main_process, - weight_name=weight_name, - save_function=save_function, - safe_serialization=safe_serialization, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.fuse_lora - def fuse_lora( - self, - components: list[str] = ["transformer"], - lora_scale: float = 1.0, - safe_fusing: bool = False, - adapter_names: list[str] | None = None, - **kwargs, - ): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.fuse_lora`] for more details. - """ - super().fuse_lora( - components=components, - lora_scale=lora_scale, - safe_fusing=safe_fusing, - adapter_names=adapter_names, - **kwargs, - ) - - # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.unfuse_lora - def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): - r""" - See [`~loaders.StableDiffusionLoraLoaderMixin.unfuse_lora`] for more details. - """ - super().unfuse_lora(components=components, **kwargs) - - -class LoraLoaderMixin(StableDiffusionLoraLoaderMixin): - def __init__(self, *args, **kwargs): - deprecation_message = "LoraLoaderMixin is deprecated and this will be removed in a future version. Please use `StableDiffusionLoraLoaderMixin`, instead." - deprecate("LoraLoaderMixin", "1.0.0", deprecation_message) - super().__init__(*args, **kwargs) diff --git a/diffusers/loaders/peft.py b/diffusers/loaders/peft.py deleted file mode 100644 index daa078bc25d51b177a8744a61d30a03748d6840e..0000000000000000000000000000000000000000 --- a/diffusers/loaders/peft.py +++ /dev/null @@ -1,832 +0,0 @@ -# coding=utf-8 -# Copyright 2025 The HuggingFace Inc. team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import inspect -import json -import os -from collections import defaultdict -from functools import partial -from pathlib import Path -from typing import Literal - -import safetensors -import torch - -from ..hooks.group_offloading import _maybe_remove_and_reapply_group_offloading -from ..utils import ( - MIN_PEFT_VERSION, - USE_PEFT_BACKEND, - check_peft_version, - convert_sai_sd_control_lora_state_dict_to_peft, - convert_unet_state_dict_to_peft, - delete_adapter_layers, - get_adapter_name, - is_peft_available, - is_peft_version, - logging, - set_adapter_layers, - set_weights_and_activate_adapters, -) -from ..utils.peft_utils import _create_lora_config, _maybe_warn_for_unhandled_keys -from .lora_base import _fetch_state_dict, _func_optionally_disable_offloading -from .unet_loader_utils import _maybe_expand_lora_scales - - -logger = logging.get_logger(__name__) - -_SET_ADAPTER_SCALE_FN_MAPPING = defaultdict( - lambda: (lambda model_cls, weights: weights), - { - "UNet2DConditionModel": _maybe_expand_lora_scales, - "UNetMotionModel": _maybe_expand_lora_scales, - }, -) - - -class PeftAdapterMixin: - """ - A class containing all functions for loading and using adapters weights that are supported in PEFT library. For - more details about adapters and injecting them in a base model, check out the PEFT - [documentation](https://huggingface.co/docs/peft/index). - - Install the latest version of PEFT, and use this mixin to: - - - Attach new adapters in the model. - - Attach multiple adapters and iteratively activate/deactivate them. - - Activate/deactivate all adapters from the model. - - Get a list of the active adapters. - """ - - _hf_peft_config_loaded = False - # kwargs for prepare_model_for_compiled_hotswap, if required - _prepare_lora_hotswap_kwargs: dict | None = None - - @classmethod - # Copied from diffusers.loaders.lora_base.LoraBaseMixin._optionally_disable_offloading - def _optionally_disable_offloading(cls, _pipeline): - return _func_optionally_disable_offloading(_pipeline=_pipeline) - - def load_lora_adapter( - self, pretrained_model_name_or_path_or_dict, prefix="transformer", hotswap: bool = False, **kwargs - ): - r""" - Loads a LoRA adapter into the underlying model. - - Parameters: - pretrained_model_name_or_path_or_dict (`str` or `os.PathLike` or `dict`): - Can be either: - - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved - with [`ModelMixin.save_pretrained`]. - - A [torch state - dict](https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict). - - prefix (`str`, *optional*): Prefix to filter the state dict. - - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - subfolder (`str`, *optional*, defaults to `""`): - The subfolder location of a model file within a larger model repository on the Hub or locally. - network_alphas (`dict[str, float]`): - The value of the network alpha used for stable learning and preventing underflow. This value has the - same meaning as the `--network_alpha` option in the kohya-ss trainer script. Refer to [this - link](https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning). - low_cpu_mem_usage (`bool`, *optional*): - Speed up model loading by only loading the pretrained LoRA weights and not initializing the random - weights. - hotswap : (`bool`, *optional*) - Defaults to `False`. Whether to substitute an existing (LoRA) adapter with the newly loaded adapter - in-place. This means that, instead of loading an additional adapter, this will take the existing - adapter weights and replace them with the weights of the new adapter. This can be faster and more - memory efficient. However, the main advantage of hotswapping is that when the model is compiled with - torch.compile, loading the new adapter does not require recompilation of the model. When using - hotswapping, the passed `adapter_name` should be the name of an already loaded adapter. - - If the new adapter and the old adapter have different ranks and/or LoRA alphas (i.e. scaling), you need - to call an additional method before loading the adapter: - - ```py - pipeline = ... # load diffusers pipeline - max_rank = ... # the highest rank among all LoRAs that you want to load - # call *before* compiling and loading the LoRA adapter - pipeline.enable_lora_hotswap(target_rank=max_rank) - pipeline.load_lora_weights(file_name) - # optionally compile the model now - ``` - - Note that hotswapping adapters of the text encoder is not yet supported. There are some further - limitations to this technique, which are documented here: - https://huggingface.co/docs/peft/main/en/package_reference/hotswap - metadata: - LoRA adapter metadata. When supplied, the metadata inferred through the state dict isn't used to - initialize `LoraConfig`. - """ - from peft import inject_adapter_in_model, set_peft_model_state_dict - from peft.tuners.tuners_utils import BaseTunerLayer - - from ..hooks.group_offloading import _maybe_remove_and_reapply_group_offloading - - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - adapter_name = kwargs.pop("adapter_name", None) - network_alphas = kwargs.pop("network_alphas", None) - _pipeline = kwargs.pop("_pipeline", None) - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", False) - metadata = kwargs.pop("metadata", None) - allow_pickle = False - - if low_cpu_mem_usage and is_peft_version("<=", "0.13.0"): - raise ValueError( - "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." - ) - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - state_dict, metadata = _fetch_state_dict( - pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, - weight_name=weight_name, - use_safetensors=use_safetensors, - local_files_only=local_files_only, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - allow_pickle=allow_pickle, - metadata=metadata, - ) - if network_alphas is not None and prefix is None: - raise ValueError("`network_alphas` cannot be None when `prefix` is None.") - if network_alphas and metadata: - raise ValueError("Both `network_alphas` and `metadata` cannot be specified.") - - if prefix is not None: - state_dict = {k.removeprefix(f"{prefix}."): v for k, v in state_dict.items() if k.startswith(f"{prefix}.")} - if metadata is not None: - metadata = {k.removeprefix(f"{prefix}."): v for k, v in metadata.items() if k.startswith(f"{prefix}.")} - - if len(state_dict) > 0: - if adapter_name in getattr(self, "peft_config", {}) and not hotswap: - raise ValueError( - f"Adapter name {adapter_name} already in use in the model - please select a new adapter name." - ) - elif adapter_name not in getattr(self, "peft_config", {}) and hotswap: - raise ValueError( - f"Trying to hotswap LoRA adapter '{adapter_name}' but there is no existing adapter by that name. " - "Please choose an existing adapter name or set `hotswap=False` to prevent hotswapping." - ) - - # check with first key if is not in peft format - first_key = next(iter(state_dict.keys())) - if "lora_A" not in first_key: - state_dict = convert_unet_state_dict_to_peft(state_dict) - - # Control LoRA from SAI is different from BFL Control LoRA - # https://huggingface.co/stabilityai/control-lora - # https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors - is_sai_sd_control_lora = "lora_controlnet" in state_dict - if is_sai_sd_control_lora: - state_dict = convert_sai_sd_control_lora_state_dict_to_peft(state_dict) - - rank = {} - for key, val in state_dict.items(): - # Cannot figure out rank from lora layers that don't have at least 2 dimensions. - # Bias layers in LoRA only have a single dimension - if "lora_B" in key and val.ndim > 1: - # Check out https://github.com/huggingface/peft/pull/2419 for the `^` symbol. - # We may run into some ambiguous configuration values when a model has module - # names, sharing a common prefix (`proj_out.weight` and `blocks.transformer.proj_out.weight`, - # for example) and they have different LoRA ranks. - rank[f"^{key}"] = val.shape[1] - - if network_alphas is not None and len(network_alphas) >= 1: - alpha_keys = [k for k in network_alphas.keys() if k.startswith(f"{prefix}.")] - network_alphas = { - k.removeprefix(f"{prefix}."): v for k, v in network_alphas.items() if k in alpha_keys - } - - # adapter_name - if adapter_name is None: - adapter_name = get_adapter_name(self) - - # create LoraConfig - lora_config = _create_lora_config( - state_dict, - network_alphas, - metadata, - rank, - model_state_dict=self.state_dict(), - adapter_name=adapter_name, - ) - - # Adjust LoRA config for Control LoRA - if is_sai_sd_control_lora: - lora_config.lora_alpha = lora_config.r - lora_config.alpha_pattern = lora_config.rank_pattern - lora_config.bias = "all" - lora_config.modules_to_save = lora_config.exclude_modules - lora_config.exclude_modules = None - - # =", "0.13.1"): - peft_kwargs["low_cpu_mem_usage"] = low_cpu_mem_usage - - if hotswap or (self._prepare_lora_hotswap_kwargs is not None): - if is_peft_version(">", "0.14.0"): - from peft.utils.hotswap import ( - check_hotswap_configs_compatible, - hotswap_adapter_from_state_dict, - prepare_model_for_compiled_hotswap, - ) - else: - msg = ( - "Hotswapping requires PEFT > v0.14. Please upgrade PEFT to a higher version or install it " - "from source." - ) - raise ImportError(msg) - - if hotswap: - - def map_state_dict_for_hotswap(sd): - # For hotswapping, we need the adapter name to be present in the state dict keys - new_sd = {} - for k, v in sd.items(): - if k.endswith("lora_A.weight") or k.endswith("lora_B.weight"): - k = k[: -len(".weight")] + f".{adapter_name}.weight" - elif k.endswith("lora_B.bias"): # lora_bias=True option - k = k[: -len(".bias")] + f".{adapter_name}.bias" - new_sd[k] = v - return new_sd - - # To handle scenarios where we cannot successfully set state dict. If it's unsuccessful, - # we should also delete the `peft_config` associated to the `adapter_name`. - try: - if hotswap: - state_dict = map_state_dict_for_hotswap(state_dict) - check_hotswap_configs_compatible(self.peft_config[adapter_name], lora_config) - try: - hotswap_adapter_from_state_dict( - model=self, - state_dict=state_dict, - adapter_name=adapter_name, - config=lora_config, - ) - except Exception as e: - logger.error(f"Hotswapping {adapter_name} was unsuccessful with the following error: \n{e}") - raise - # the hotswap function raises if there are incompatible keys, so if we reach this point we can set - # it to None - incompatible_keys = None - else: - inject_adapter_in_model( - lora_config, self, adapter_name=adapter_name, state_dict=state_dict, **peft_kwargs - ) - incompatible_keys = set_peft_model_state_dict(self, state_dict, adapter_name, **peft_kwargs) - - if self._prepare_lora_hotswap_kwargs is not None: - # For hotswapping of compiled models or adapters with different ranks. - # If the user called enable_lora_hotswap, we need to ensure it is called: - # - after the first adapter was loaded - # - before the model is compiled and the 2nd adapter is being hotswapped in - # Therefore, it needs to be called here - prepare_model_for_compiled_hotswap( - self, config=lora_config, **self._prepare_lora_hotswap_kwargs - ) - # We only want to call prepare_model_for_compiled_hotswap once - self._prepare_lora_hotswap_kwargs = None - - # Set peft config loaded flag to True if module has been successfully injected and incompatible keys retrieved - if not self._hf_peft_config_loaded: - self._hf_peft_config_loaded = True - except Exception as e: - # In case `inject_adapter_in_model()` was unsuccessful even before injecting the `peft_config`. - if hasattr(self, "peft_config"): - for module in self.modules(): - if isinstance(module, BaseTunerLayer): - active_adapters = module.active_adapters - for active_adapter in active_adapters: - if adapter_name in active_adapter: - module.delete_adapter(adapter_name) - - self.peft_config.pop(adapter_name) - logger.error(f"Loading {adapter_name} was unsuccessful with the following error: \n{e}") - raise - - _maybe_warn_for_unhandled_keys(incompatible_keys, adapter_name) - - # Offload back. - if is_model_cpu_offload: - _pipeline.enable_model_cpu_offload() - elif is_sequential_cpu_offload: - _pipeline.enable_sequential_cpu_offload() - elif is_group_offload: - for component in _pipeline.components.values(): - if isinstance(component, torch.nn.Module): - _maybe_remove_and_reapply_group_offloading(component) - # Unsafe code /> - - if prefix is not None and not state_dict: - model_class_name = self.__class__.__name__ - logger.warning( - f"No LoRA keys associated to {model_class_name} found with the {prefix=}. " - "This is safe to ignore if LoRA state dict didn't originally have any " - f"{model_class_name} related params. You can also try specifying `prefix=None` " - "to resolve the warning. Otherwise, open an issue if you think it's unexpected: " - "https://github.com/huggingface/diffusers/issues/new" - ) - - def save_lora_adapter( - self, - save_directory, - adapter_name: str = "default", - upcast_before_saving: bool = False, - safe_serialization: bool = True, - weight_name: str | None = None, - ): - """ - Save the LoRA parameters corresponding to the underlying model. - - Arguments: - save_directory (`str` or `os.PathLike`): - Directory to save LoRA parameters to. Will be created if it doesn't exist. - adapter_name: (`str`, defaults to "default"): The name of the adapter to serialize. Useful when the - underlying model has multiple adapters loaded. - upcast_before_saving (`bool`, defaults to `False`): - Whether to cast the underlying model to `torch.float32` before serialization. - safe_serialization (`bool`, *optional*, defaults to `True`): - Whether to save the model using `safetensors` or the traditional PyTorch way with `pickle`. - weight_name: (`str`, *optional*, defaults to `None`): Name of the file to serialize the state dict with. - """ - from peft.utils import get_peft_model_state_dict - - from .lora_base import LORA_ADAPTER_METADATA_KEY, LORA_WEIGHT_NAME, LORA_WEIGHT_NAME_SAFE - - if adapter_name is None: - adapter_name = get_adapter_name(self) - - if adapter_name not in getattr(self, "peft_config", {}): - raise ValueError(f"Adapter name {adapter_name} not found in the model.") - - lora_adapter_metadata = self.peft_config[adapter_name].to_dict() - - lora_layers_to_save = get_peft_model_state_dict( - self.to(dtype=torch.float32 if upcast_before_saving else None), adapter_name=adapter_name - ) - if os.path.isfile(save_directory): - raise ValueError(f"Provided path ({save_directory}) should be a directory, not a file") - - if safe_serialization: - - def save_function(weights, filename): - # Inject framework format. - metadata = {"format": "pt"} - if lora_adapter_metadata is not None: - for key, value in lora_adapter_metadata.items(): - if isinstance(value, set): - lora_adapter_metadata[key] = list(value) - metadata[LORA_ADAPTER_METADATA_KEY] = json.dumps(lora_adapter_metadata, indent=2, sort_keys=True) - - return safetensors.torch.save_file(weights, filename, metadata=metadata) - - else: - save_function = torch.save - - os.makedirs(save_directory, exist_ok=True) - - if weight_name is None: - if safe_serialization: - weight_name = LORA_WEIGHT_NAME_SAFE - else: - weight_name = LORA_WEIGHT_NAME - - save_path = Path(save_directory, weight_name).as_posix() - save_function(lora_layers_to_save, save_path) - logger.info(f"Model weights saved in {save_path}") - - def set_adapters( - self, - adapter_names: list[str] | str, - weights: float | dict | list[float] | list[dict] | list[None] | None = None, - ): - """ - Set the currently active adapters for use in the diffusion network (e.g. unet, transformer, etc.). - - Args: - adapter_names (`list[str]` or `str`): - The names of the adapters to use. - weights (`Union[List[float], float]`, *optional*): - The adapter(s) weights to use with the UNet. If `None`, the weights are set to `1.0` for all the - adapters. - - Example: - - ```py - from diffusers import AutoPipelineForText2Image - import torch - - pipeline = AutoPipelineForText2Image.from_pretrained( - "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16 - ).to("cuda") - pipeline.load_lora_weights( - "jbilcke-hf/sdxl-cinematic-1", weight_name="pytorch_lora_weights.safetensors", adapter_name="cinematic" - ) - pipeline.load_lora_weights("nerijs/pixel-art-xl", weight_name="pixel-art-xl.safetensors", adapter_name="pixel") - pipeline.unet.set_adapters(["cinematic", "pixel"], weights=[0.5, 0.5]) - ``` - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for `set_adapters()`.") - - adapter_names = [adapter_names] if isinstance(adapter_names, str) else adapter_names - - # Expand weights into a list, one entry per adapter - # examples for e.g. 2 adapters: [{...}, 7] -> [7,7] ; None -> [None, None] - if not isinstance(weights, list): - weights = [weights] * len(adapter_names) - - if len(adapter_names) != len(weights): - raise ValueError( - f"Length of adapter names {len(adapter_names)} is not equal to the length of their weights {len(weights)}." - ) - - # Set None values to default of 1.0 - # e.g. [{...}, 7] -> [{...}, 7] ; [None, None] -> [1.0, 1.0] - weights = [w if w is not None else 1.0 for w in weights] - - # e.g. [{...}, 7] -> [{expanded dict...}, 7] - scale_expansion_fn = _SET_ADAPTER_SCALE_FN_MAPPING[self.__class__.__name__] - weights = scale_expansion_fn(self, weights) - - set_weights_and_activate_adapters(self, adapter_names, weights) - - def add_adapter(self, adapter_config, adapter_name: str = "default") -> None: - r""" - Adds a new adapter to the current model for training. If no adapter name is passed, a default name is assigned - to the adapter to follow the convention of the PEFT library. - - If you are not familiar with adapters and PEFT methods, we invite you to read more about them in the PEFT - [documentation](https://huggingface.co/docs/peft). - - Args: - adapter_config (`[~peft.PeftConfig]`): - The configuration of the adapter to add; supported adapters are non-prefix tuning and adaption prompt - methods. - adapter_name (`str`, *optional*, defaults to `"default"`): - The name of the adapter to add. If no name is passed, a default name is assigned to the adapter. - """ - check_peft_version(min_version=MIN_PEFT_VERSION) - - if not is_peft_available(): - raise ImportError("PEFT is not available. Please install PEFT to use this function: `pip install peft`.") - - from peft import PeftConfig, inject_adapter_in_model - - if not self._hf_peft_config_loaded: - self._hf_peft_config_loaded = True - elif adapter_name in self.peft_config: - raise ValueError(f"Adapter with name {adapter_name} already exists. Please use a different name.") - - if not isinstance(adapter_config, PeftConfig): - raise ValueError( - f"adapter_config should be an instance of PeftConfig. Got {type(adapter_config)} instead." - ) - - # Unlike transformers, here we don't need to retrieve the name_or_path of the unet as the loading logic is - # handled by the `load_lora_layers` or `StableDiffusionLoraLoaderMixin`. Therefore we set it to `None` here. - adapter_config.base_model_name_or_path = None - inject_adapter_in_model(adapter_config, self, adapter_name) - self.set_adapter(adapter_name) - - def set_adapter(self, adapter_name: str | list[str]) -> None: - """ - Sets a specific adapter by forcing the model to only use that adapter and disables the other adapters. - - If you are not familiar with adapters and PEFT methods, we invite you to read more about them on the PEFT - [documentation](https://huggingface.co/docs/peft). - - Args: - adapter_name (str | list[str])): - The list of adapters to set or the adapter name in the case of a single adapter. - """ - check_peft_version(min_version=MIN_PEFT_VERSION) - - if not self._hf_peft_config_loaded: - raise ValueError("No adapter loaded. Please load an adapter first.") - - if isinstance(adapter_name, str): - adapter_name = [adapter_name] - - missing = set(adapter_name) - set(self.peft_config) - if len(missing) > 0: - raise ValueError( - f"Following adapter(s) could not be found: {', '.join(missing)}. Make sure you are passing the correct adapter name(s)." - f" current loaded adapters are: {list(self.peft_config.keys())}" - ) - - from peft.tuners.tuners_utils import BaseTunerLayer - - _adapters_has_been_set = False - - for _, module in self.named_modules(): - if isinstance(module, BaseTunerLayer): - if hasattr(module, "set_adapter"): - module.set_adapter(adapter_name) - # Previous versions of PEFT does not support multi-adapter inference - elif not hasattr(module, "set_adapter") and len(adapter_name) != 1: - raise ValueError( - "You are trying to set multiple adapters and you have a PEFT version that does not support multi-adapter inference. Please upgrade to the latest version of PEFT." - " `pip install -U peft` or `pip install -U git+https://github.com/huggingface/peft.git`" - ) - else: - module.active_adapter = adapter_name - _adapters_has_been_set = True - - if not _adapters_has_been_set: - raise ValueError( - "Did not succeeded in setting the adapter. Please make sure you are using a model that supports adapters." - ) - - def disable_adapters(self) -> None: - r""" - Disable all adapters attached to the model and fallback to inference with the base model only. - - If you are not familiar with adapters and PEFT methods, we invite you to read more about them on the PEFT - [documentation](https://huggingface.co/docs/peft). - """ - check_peft_version(min_version=MIN_PEFT_VERSION) - - if not self._hf_peft_config_loaded: - raise ValueError("No adapter loaded. Please load an adapter first.") - - from peft.tuners.tuners_utils import BaseTunerLayer - - for _, module in self.named_modules(): - if isinstance(module, BaseTunerLayer): - if hasattr(module, "enable_adapters"): - module.enable_adapters(enabled=False) - else: - # support for older PEFT versions - module.disable_adapters = True - - def enable_adapters(self) -> None: - """ - Enable adapters that are attached to the model. The model uses `self.active_adapters()` to retrieve the list of - adapters to enable. - - If you are not familiar with adapters and PEFT methods, we invite you to read more about them on the PEFT - [documentation](https://huggingface.co/docs/peft). - """ - check_peft_version(min_version=MIN_PEFT_VERSION) - - if not self._hf_peft_config_loaded: - raise ValueError("No adapter loaded. Please load an adapter first.") - - from peft.tuners.tuners_utils import BaseTunerLayer - - for _, module in self.named_modules(): - if isinstance(module, BaseTunerLayer): - if hasattr(module, "enable_adapters"): - module.enable_adapters(enabled=True) - else: - # support for older PEFT versions - module.disable_adapters = False - - def active_adapters(self) -> list[str]: - """ - Gets the current list of active adapters of the model. - - If you are not familiar with adapters and PEFT methods, we invite you to read more about them on the PEFT - [documentation](https://huggingface.co/docs/peft). - """ - check_peft_version(min_version=MIN_PEFT_VERSION) - - if not is_peft_available(): - raise ImportError("PEFT is not available. Please install PEFT to use this function: `pip install peft`.") - - if not self._hf_peft_config_loaded: - raise ValueError("No adapter loaded. Please load an adapter first.") - - from peft.tuners.tuners_utils import BaseTunerLayer - - for _, module in self.named_modules(): - if isinstance(module, BaseTunerLayer): - return module.active_adapter - - def fuse_lora(self, lora_scale=1.0, safe_fusing=False, adapter_names=None): - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for `fuse_lora()`.") - - self.lora_scale = lora_scale - self._safe_fusing = safe_fusing - self.apply(partial(self._fuse_lora_apply, adapter_names=adapter_names)) - - def _fuse_lora_apply(self, module, adapter_names=None): - from peft.tuners.tuners_utils import BaseTunerLayer - - merge_kwargs = {"safe_merge": self._safe_fusing} - - if isinstance(module, BaseTunerLayer): - if self.lora_scale != 1.0: - module.scale_layer(self.lora_scale) - - # For BC with previous PEFT versions, we need to check the signature - # of the `merge` method to see if it supports the `adapter_names` argument. - supported_merge_kwargs = list(inspect.signature(module.merge).parameters) - if "adapter_names" in supported_merge_kwargs: - merge_kwargs["adapter_names"] = adapter_names - elif "adapter_names" not in supported_merge_kwargs and adapter_names is not None: - raise ValueError( - "The `adapter_names` argument is not supported with your PEFT version. Please upgrade" - " to the latest version of PEFT. `pip install -U peft`" - ) - - module.merge(**merge_kwargs) - - def unfuse_lora(self): - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for `unfuse_lora()`.") - self.apply(self._unfuse_lora_apply) - - def _unfuse_lora_apply(self, module): - from peft.tuners.tuners_utils import BaseTunerLayer - - if isinstance(module, BaseTunerLayer): - module.unmerge() - - def unload_lora(self): - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for `unload_lora()`.") - - from ..hooks.group_offloading import _maybe_remove_and_reapply_group_offloading - from ..utils import recurse_remove_peft_layers - - recurse_remove_peft_layers(self) - if hasattr(self, "peft_config"): - del self.peft_config - if hasattr(self, "_hf_peft_config_loaded"): - self._hf_peft_config_loaded = None - - _maybe_remove_and_reapply_group_offloading(self) - - def disable_lora(self): - """ - Disables the active LoRA layers of the underlying model. - - Example: - - ```py - from diffusers import AutoPipelineForText2Image - import torch - - pipeline = AutoPipelineForText2Image.from_pretrained( - "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16 - ).to("cuda") - pipeline.load_lora_weights( - "jbilcke-hf/sdxl-cinematic-1", weight_name="pytorch_lora_weights.safetensors", adapter_name="cinematic" - ) - pipeline.unet.disable_lora() - ``` - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - set_adapter_layers(self, enabled=False) - - def enable_lora(self): - """ - Enables the active LoRA layers of the underlying model. - - Example: - - ```py - from diffusers import AutoPipelineForText2Image - import torch - - pipeline = AutoPipelineForText2Image.from_pretrained( - "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16 - ).to("cuda") - pipeline.load_lora_weights( - "jbilcke-hf/sdxl-cinematic-1", weight_name="pytorch_lora_weights.safetensors", adapter_name="cinematic" - ) - pipeline.unet.enable_lora() - ``` - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - set_adapter_layers(self, enabled=True) - - def delete_adapters(self, adapter_names: list[str] | str): - """ - Delete an adapter's LoRA layers from the underlying model. - - Args: - adapter_names (`list[str, str]`): - The names (single string or list of strings) of the adapter to delete. - - Example: - - ```py - from diffusers import AutoPipelineForText2Image - import torch - - pipeline = AutoPipelineForText2Image.from_pretrained( - "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16 - ).to("cuda") - pipeline.load_lora_weights( - "jbilcke-hf/sdxl-cinematic-1", weight_name="pytorch_lora_weights.safetensors", adapter_names="cinematic" - ) - pipeline.unet.delete_adapters("cinematic") - ``` - """ - if not USE_PEFT_BACKEND: - raise ValueError("PEFT backend is required for this method.") - - if isinstance(adapter_names, str): - adapter_names = [adapter_names] - - for adapter_name in adapter_names: - delete_adapter_layers(self, adapter_name) - - # Pop also the corresponding adapter from the config - if hasattr(self, "peft_config"): - self.peft_config.pop(adapter_name, None) - - _maybe_remove_and_reapply_group_offloading(self) - - def enable_lora_hotswap( - self, target_rank: int = 128, check_compiled: Literal["error", "warn", "ignore"] = "error" - ) -> None: - """Enables the possibility to hotswap LoRA adapters. - - Calling this method is only required when hotswapping adapters and if the model is compiled or if the ranks of - the loaded adapters differ. - - Args: - target_rank (`int`, *optional*, defaults to `128`): - The highest rank among all the adapters that will be loaded. - - check_compiled (`str`, *optional*, defaults to `"error"`): - How to handle the case when the model is already compiled, which should generally be avoided. The - options are: - - "error" (default): raise an error - - "warn": issue a warning - - "ignore": do nothing - """ - if getattr(self, "peft_config", {}): - if check_compiled == "error": - raise RuntimeError("Call `enable_lora_hotswap` before loading the first adapter.") - elif check_compiled == "warn": - logger.warning( - "It is recommended to call `enable_lora_hotswap` before loading the first adapter to avoid recompilation." - ) - elif check_compiled != "ignore": - raise ValueError( - f"check_compiles should be one of 'error', 'warn', or 'ignore', got '{check_compiled}' instead." - ) - - self._prepare_lora_hotswap_kwargs = {"target_rank": target_rank, "check_compiled": check_compiled} diff --git a/diffusers/loaders/single_file.py b/diffusers/loaders/single_file.py deleted file mode 100644 index 881ff9b96a4c0377243c8b1ea9528a1d821b5642..0000000000000000000000000000000000000000 --- a/diffusers/loaders/single_file.py +++ /dev/null @@ -1,567 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import importlib -import inspect -import os - -import torch -from huggingface_hub import snapshot_download -from huggingface_hub.utils import LocalEntryNotFoundError, validate_hf_hub_args -from packaging import version -from typing_extensions import Self - -from ..utils import deprecate, is_transformers_available, logging -from .single_file_utils import ( - SingleFileComponentError, - _is_legacy_scheduler_kwargs, - _is_model_weights_in_cached_folder, - _legacy_load_clip_tokenizer, - _legacy_load_safety_checker, - _legacy_load_scheduler, - create_diffusers_clip_model_from_ldm, - create_diffusers_t5_model_from_checkpoint, - fetch_diffusers_config, - fetch_original_config, - is_clip_model_in_single_file, - is_t5_in_single_file, - load_single_file_checkpoint, -) - - -logger = logging.get_logger(__name__) - -# Legacy behaviour. `from_single_file` does not load the safety checker unless explicitly provided -SINGLE_FILE_OPTIONAL_COMPONENTS = ["safety_checker"] - -if is_transformers_available(): - import transformers - from transformers import PreTrainedModel, PreTrainedTokenizer - - -def load_single_file_sub_model( - library_name, - class_name, - name, - checkpoint, - pipelines, - is_pipeline_module, - cached_model_config_path, - original_config=None, - local_files_only=False, - torch_dtype=None, - is_legacy_loading=False, - disable_mmap=False, - **kwargs, -): - if is_pipeline_module: - pipeline_module = getattr(pipelines, library_name) - class_obj = getattr(pipeline_module, class_name) - else: - # else we just import it from the library. - library = importlib.import_module(library_name) - class_obj = getattr(library, class_name) - - if is_transformers_available(): - transformers_version = version.parse(version.parse(transformers.__version__).base_version) - else: - transformers_version = "N/A" - - is_transformers_model = ( - is_transformers_available() - and issubclass(class_obj, PreTrainedModel) - and transformers_version >= version.parse("4.20.0") - ) - is_tokenizer = ( - is_transformers_available() - and issubclass(class_obj, PreTrainedTokenizer) - and transformers_version >= version.parse("4.20.0") - ) - - diffusers_module = importlib.import_module(__name__.split(".")[0]) - is_diffusers_single_file_model = issubclass(class_obj, diffusers_module.FromOriginalModelMixin) - is_diffusers_model = issubclass(class_obj, diffusers_module.ModelMixin) - is_diffusers_scheduler = issubclass(class_obj, diffusers_module.SchedulerMixin) - - if is_diffusers_single_file_model: - load_method = getattr(class_obj, "from_single_file") - - # We cannot provide two different config options to the `from_single_file` method - # Here we have to ignore loading the config from `cached_model_config_path` if `original_config` is provided - if original_config: - cached_model_config_path = None - - loaded_sub_model = load_method( - pretrained_model_link_or_path_or_dict=checkpoint, - original_config=original_config, - config=cached_model_config_path, - subfolder=name, - torch_dtype=torch_dtype, - local_files_only=local_files_only, - disable_mmap=disable_mmap, - **kwargs, - ) - - elif is_transformers_model and is_clip_model_in_single_file(class_obj, checkpoint): - loaded_sub_model = create_diffusers_clip_model_from_ldm( - class_obj, - checkpoint=checkpoint, - config=cached_model_config_path, - subfolder=name, - torch_dtype=torch_dtype, - local_files_only=local_files_only, - is_legacy_loading=is_legacy_loading, - ) - - elif is_transformers_model and is_t5_in_single_file(checkpoint): - loaded_sub_model = create_diffusers_t5_model_from_checkpoint( - class_obj, - checkpoint=checkpoint, - config=cached_model_config_path, - subfolder=name, - torch_dtype=torch_dtype, - local_files_only=local_files_only, - ) - - elif is_tokenizer and is_legacy_loading: - loaded_sub_model = _legacy_load_clip_tokenizer( - class_obj, checkpoint=checkpoint, config=cached_model_config_path, local_files_only=local_files_only - ) - - elif is_diffusers_scheduler and (is_legacy_loading or _is_legacy_scheduler_kwargs(kwargs)): - loaded_sub_model = _legacy_load_scheduler( - class_obj, checkpoint=checkpoint, component_name=name, original_config=original_config, **kwargs - ) - - else: - if not hasattr(class_obj, "from_pretrained"): - raise ValueError( - ( - f"The component {class_obj.__name__} cannot be loaded as it does not seem to have" - " a supported loading method." - ) - ) - - loading_kwargs = {} - loading_kwargs.update( - { - "pretrained_model_name_or_path": cached_model_config_path, - "subfolder": name, - "local_files_only": local_files_only, - } - ) - - # Schedulers and Tokenizers don't make use of torch_dtype - # Skip passing it to those objects - if issubclass(class_obj, torch.nn.Module): - loading_kwargs.update({"torch_dtype": torch_dtype}) - - if is_diffusers_model or is_transformers_model: - if not _is_model_weights_in_cached_folder(cached_model_config_path, name): - raise SingleFileComponentError( - f"Failed to load {class_name}. Weights for this component appear to be missing in the checkpoint." - ) - - load_method = getattr(class_obj, "from_pretrained") - loaded_sub_model = load_method(**loading_kwargs) - - return loaded_sub_model - - -def _map_component_types_to_config_dict(component_types): - diffusers_module = importlib.import_module(__name__.split(".")[0]) - config_dict = {} - component_types.pop("self", None) - - if is_transformers_available(): - transformers_version = version.parse(version.parse(transformers.__version__).base_version) - else: - transformers_version = "N/A" - - for component_name, component_value in component_types.items(): - is_diffusers_model = issubclass(component_value[0], diffusers_module.ModelMixin) - is_scheduler_enum = component_value[0].__name__ == "KarrasDiffusionSchedulers" - is_scheduler = issubclass(component_value[0], diffusers_module.SchedulerMixin) - - is_transformers_model = ( - is_transformers_available() - and issubclass(component_value[0], PreTrainedModel) - and transformers_version >= version.parse("4.20.0") - ) - is_transformers_tokenizer = ( - is_transformers_available() - and issubclass(component_value[0], PreTrainedTokenizer) - and transformers_version >= version.parse("4.20.0") - ) - - if is_diffusers_model and component_name not in SINGLE_FILE_OPTIONAL_COMPONENTS: - config_dict[component_name] = ["diffusers", component_value[0].__name__] - - elif is_scheduler_enum or is_scheduler: - if is_scheduler_enum: - # Since we cannot fetch a scheduler config from the hub, we default to DDIMScheduler - # if the type hint is a KarrassDiffusionSchedulers enum - config_dict[component_name] = ["diffusers", "DDIMScheduler"] - - elif is_scheduler: - config_dict[component_name] = ["diffusers", component_value[0].__name__] - - elif ( - is_transformers_model or is_transformers_tokenizer - ) and component_name not in SINGLE_FILE_OPTIONAL_COMPONENTS: - config_dict[component_name] = ["transformers", component_value[0].__name__] - - else: - config_dict[component_name] = [None, None] - - return config_dict - - -def _infer_pipeline_config_dict(pipeline_class): - parameters = inspect.signature(pipeline_class.__init__).parameters - required_parameters = {k: v for k, v in parameters.items() if v.default == inspect._empty} - component_types = pipeline_class._get_signature_types() - - # Ignore parameters that are not required for the pipeline - component_types = {k: v for k, v in component_types.items() if k in required_parameters} - config_dict = _map_component_types_to_config_dict(component_types) - - return config_dict - - -def _download_diffusers_model_config_from_hub( - pretrained_model_name_or_path, - cache_dir, - revision, - proxies, - force_download=None, - local_files_only=None, - token=None, -): - allow_patterns = ["**/*.json", "*.json", "*.txt", "**/*.txt", "**/*.model"] - cached_model_path = snapshot_download( - pretrained_model_name_or_path, - cache_dir=cache_dir, - revision=revision, - proxies=proxies, - force_download=force_download, - local_files_only=local_files_only, - token=token, - allow_patterns=allow_patterns, - ) - - return cached_model_path - - -class FromSingleFileMixin: - """ - Load model weights saved in the `.ckpt` format into a [`DiffusionPipeline`]. - """ - - @classmethod - @validate_hf_hub_args - def from_single_file(cls, pretrained_model_link_or_path, **kwargs) -> Self: - r""" - Instantiate a [`DiffusionPipeline`] from pretrained pipeline weights saved in the `.ckpt` or `.safetensors` - format. The pipeline is set in evaluation mode (`model.eval()`) by default. - - Parameters: - pretrained_model_link_or_path (`str` or `os.PathLike`, *optional*): - Can be either: - - A link to the `.ckpt` file (for example - `"https://huggingface.co//blob/main/.ckpt"`) on the Hub. - - A path to a *file* containing all pipeline weights. - dtype (`str` or `torch.dtype`, *optional*): - Override the default `torch.dtype` and load the model with another dtype. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - original_config_file (`str`, *optional*): - The path to the original config file that was used to train the model. If not provided, the config file - will be inferred from the checkpoint file. - config (`str`, *optional*): - Can be either: - - A string, the *repo id* (for example `CompVis/ldm-text2im-large-256`) of a pretrained pipeline - hosted on the Hub. - - A path to a *directory* (for example `./my_pipeline_directory/`) containing the pipeline - component configs in Diffusers format. - disable_mmap ('bool', *optional*, defaults to 'False'): - Whether to disable mmap when loading a Safetensors model. This option can perform better when the model - is on a network mount or hard drive. - kwargs (remaining dictionary of keyword arguments, *optional*): - Can be used to overwrite load and saveable variables (the pipeline components of the specific pipeline - class). The overwritten components are passed directly to the pipelines `__init__` method. See example - below for more information. - - Examples: - - ```py - >>> from diffusers import StableDiffusionPipeline - - >>> # Download pipeline from huggingface.co and cache. - >>> pipeline = StableDiffusionPipeline.from_single_file( - ... "https://huggingface.co/WarriorMama777/OrangeMixs/blob/main/Models/AbyssOrangeMix/AbyssOrangeMix.safetensors" - ... ) - - >>> # Download pipeline from local file - >>> # file is downloaded under ./v1-5-pruned-emaonly.ckpt - >>> pipeline = StableDiffusionPipeline.from_single_file("./v1-5-pruned-emaonly.ckpt") - - >>> # Enable float16 and move to GPU - >>> pipeline = StableDiffusionPipeline.from_single_file( - ... "https://huggingface.co/stable-diffusion-v1-5/stable-diffusion-v1-5/blob/main/v1-5-pruned-emaonly.ckpt", - ... torch_dtype=torch.float16, - ... ) - >>> pipeline.to("cuda") - ``` - - """ - original_config_file = kwargs.pop("original_config_file", None) - config = kwargs.pop("config", None) - original_config = kwargs.pop("original_config", None) - - if original_config_file is not None: - deprecation_message = ( - "`original_config_file` argument is deprecated and will be removed in future versions." - "please use the `original_config` argument instead." - ) - deprecate("original_config_file", "1.0.0", deprecation_message) - original_config = original_config_file - - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - token = kwargs.pop("token", None) - cache_dir = kwargs.pop("cache_dir", None) - local_files_only = kwargs.pop("local_files_only", False) - revision = kwargs.pop("revision", None) - torch_dtype = kwargs.pop("torch_dtype", None) - dtype = kwargs.pop("dtype", None) - torch_dtype = dtype if dtype is not None else torch_dtype - disable_mmap = kwargs.pop("disable_mmap", False) - - is_legacy_loading = False - - if torch_dtype is not None and not isinstance(torch_dtype, torch.dtype): - torch_dtype = torch.float32 - logger.warning( - f"Passed `torch_dtype` {torch_dtype} is not a `torch.dtype`. Defaulting to `torch.float32`." - ) - - # We shouldn't allow configuring individual models components through a Pipeline creation method - # These model kwargs should be deprecated - scaling_factor = kwargs.get("scaling_factor", None) - if scaling_factor is not None: - deprecation_message = ( - "Passing the `scaling_factor` argument to `from_single_file is deprecated " - "and will be ignored in future versions." - ) - deprecate("scaling_factor", "1.0.0", deprecation_message) - - if original_config is not None: - original_config = fetch_original_config(original_config, local_files_only=local_files_only) - - from ..pipelines.pipeline_utils import _get_pipeline_class - - pipeline_class = _get_pipeline_class(cls, config=None) - - checkpoint = load_single_file_checkpoint( - pretrained_model_link_or_path, - force_download=force_download, - proxies=proxies, - token=token, - cache_dir=cache_dir, - local_files_only=local_files_only, - revision=revision, - disable_mmap=disable_mmap, - ) - - if config is None: - config = fetch_diffusers_config(checkpoint) - default_pretrained_model_config_name = config["pretrained_model_name_or_path"] - else: - default_pretrained_model_config_name = config - - if not os.path.isdir(default_pretrained_model_config_name): - # Provided config is a repo_id - if default_pretrained_model_config_name.count("/") > 1: - raise ValueError( - f'The provided config "{config}"' - " is neither a valid local path nor a valid repo id. Please check the parameter." - ) - try: - # Attempt to download the config files for the pipeline - cached_model_config_path = _download_diffusers_model_config_from_hub( - default_pretrained_model_config_name, - cache_dir=cache_dir, - revision=revision, - proxies=proxies, - force_download=force_download, - local_files_only=local_files_only, - token=token, - ) - config_dict = pipeline_class.load_config(cached_model_config_path) - - except LocalEntryNotFoundError: - # `local_files_only=True` but a local diffusers format model config is not available in the cache - # If `original_config` is not provided, we need override `local_files_only` to False - # to fetch the config files from the hub so that we have a way - # to configure the pipeline components. - - if original_config is None: - logger.warning( - "`local_files_only` is True but no local configs were found for this checkpoint.\n" - "Attempting to download the necessary config files for this pipeline.\n" - ) - cached_model_config_path = _download_diffusers_model_config_from_hub( - default_pretrained_model_config_name, - cache_dir=cache_dir, - revision=revision, - proxies=proxies, - force_download=force_download, - local_files_only=False, - token=token, - ) - config_dict = pipeline_class.load_config(cached_model_config_path) - - else: - # For backwards compatibility - # If `original_config` is provided, then we need to assume we are using legacy loading for pipeline components - logger.warning( - "Detected legacy `from_single_file` loading behavior. Attempting to create the pipeline based on inferred components.\n" - "This may lead to errors if the model components are not correctly inferred. \n" - "To avoid this warning, please explicitly pass the `config` argument to `from_single_file` with a path to a local diffusers model repo \n" - "e.g. `from_single_file(, config=) \n" - "or run `from_single_file` with `local_files_only=False` first to update the local cache directory with " - "the necessary config files.\n" - ) - is_legacy_loading = True - cached_model_config_path = None - - config_dict = _infer_pipeline_config_dict(pipeline_class) - config_dict["_class_name"] = pipeline_class.__name__ - - else: - # Provided config is a path to a local directory attempt to load directly. - cached_model_config_path = default_pretrained_model_config_name - config_dict = pipeline_class.load_config(cached_model_config_path) - - # pop out "_ignore_files" as it is only needed for download - config_dict.pop("_ignore_files", None) - - expected_modules, optional_kwargs = pipeline_class._get_signature_keys(cls) - passed_class_obj = {k: kwargs.pop(k) for k in expected_modules if k in kwargs} - passed_pipe_kwargs = {k: kwargs.pop(k) for k in optional_kwargs if k in kwargs} - - init_dict, unused_kwargs, _ = pipeline_class.extract_init_dict(config_dict, **kwargs) - init_kwargs = {k: init_dict.pop(k) for k in optional_kwargs if k in init_dict} - init_kwargs = {**init_kwargs, **passed_pipe_kwargs} - - from diffusers import pipelines - - # remove `null` components - def load_module(name, value): - if value[0] is None: - return False - if name in passed_class_obj and passed_class_obj[name] is None: - return False - if name in SINGLE_FILE_OPTIONAL_COMPONENTS: - return False - - return True - - init_dict = {k: v for k, v in init_dict.items() if load_module(k, v)} - - for name, (library_name, class_name) in logging.tqdm( - sorted(init_dict.items()), desc="Loading pipeline components..." - ): - loaded_sub_model = None - is_pipeline_module = hasattr(pipelines, library_name) - - if name in passed_class_obj: - loaded_sub_model = passed_class_obj[name] - - else: - try: - loaded_sub_model = load_single_file_sub_model( - library_name=library_name, - class_name=class_name, - name=name, - checkpoint=checkpoint, - is_pipeline_module=is_pipeline_module, - cached_model_config_path=cached_model_config_path, - pipelines=pipelines, - torch_dtype=torch_dtype, - original_config=original_config, - local_files_only=local_files_only, - is_legacy_loading=is_legacy_loading, - disable_mmap=disable_mmap, - **kwargs, - ) - except SingleFileComponentError as e: - raise SingleFileComponentError( - ( - f"{e.message}\n" - f"Please load the component before passing it in as an argument to `from_single_file`.\n" - f"\n" - f"{name} = {class_name}.from_pretrained('...')\n" - f"pipe = {pipeline_class.__name__}.from_single_file(, {name}={name})\n" - f"\n" - ) - ) - - init_kwargs[name] = loaded_sub_model - - missing_modules = set(expected_modules) - set(init_kwargs.keys()) - passed_modules = list(passed_class_obj.keys()) - optional_modules = pipeline_class._optional_components - - if len(missing_modules) > 0 and missing_modules <= set(passed_modules + optional_modules): - for module in missing_modules: - init_kwargs[module] = passed_class_obj.get(module, None) - elif len(missing_modules) > 0: - passed_modules = set(list(init_kwargs.keys()) + list(passed_class_obj.keys())) - optional_kwargs - raise ValueError( - f"Pipeline {pipeline_class} expected {expected_modules}, but only {passed_modules} were passed." - ) - - # deprecated kwargs - load_safety_checker = kwargs.pop("load_safety_checker", None) - if load_safety_checker is not None: - deprecation_message = ( - "Please pass instances of `StableDiffusionSafetyChecker` and `AutoImageProcessor`" - "using the `safety_checker` and `feature_extractor` arguments in `from_single_file`" - ) - deprecate("load_safety_checker", "1.0.0", deprecation_message) - - safety_checker_components = _legacy_load_safety_checker(local_files_only, torch_dtype) - init_kwargs.update(safety_checker_components) - - pipe = pipeline_class(**init_kwargs) - - return pipe diff --git a/diffusers/loaders/single_file_model.py b/diffusers/loaders/single_file_model.py deleted file mode 100644 index 56770fd9b6c3df6514f2f582620116ec62e57f4c..0000000000000000000000000000000000000000 --- a/diffusers/loaders/single_file_model.py +++ /dev/null @@ -1,560 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import importlib -import inspect -import re -from contextlib import nullcontext - -import torch -from huggingface_hub.utils import validate_hf_hub_args -from typing_extensions import Self - -from .. import __version__ -from ..models.model_loading_utils import ( - _caching_allocator_warmup, - _determine_device_map, - _expand_device_map, -) -from ..quantizers import DiffusersAutoQuantizer -from ..utils import deprecate, is_accelerate_available, is_torch_version, logging -from ..utils.torch_utils import empty_device_cache -from .single_file_utils import ( - SingleFileComponentError, - convert_animatediff_checkpoint_to_diffusers, - convert_auraflow_transformer_checkpoint_to_diffusers, - convert_autoencoder_dc_checkpoint_to_diffusers, - convert_chroma_transformer_checkpoint_to_diffusers, - convert_controlnet_checkpoint, - convert_cosmos_transformer_checkpoint_to_diffusers, - convert_ernie_image_transformer_checkpoint_to_diffusers, - convert_flux2_transformer_checkpoint_to_diffusers, - convert_flux_transformer_checkpoint_to_diffusers, - convert_hidream_transformer_to_diffusers, - convert_hunyuan_video_transformer_to_diffusers, - convert_ldm_unet_checkpoint, - convert_ldm_vae_checkpoint, - convert_ltx2_audio_vae_to_diffusers, - convert_ltx2_transformer_to_diffusers, - convert_ltx2_vae_to_diffusers, - convert_ltx_transformer_checkpoint_to_diffusers, - convert_ltx_vae_checkpoint_to_diffusers, - convert_lumina2_to_diffusers, - convert_mochi_transformer_checkpoint_to_diffusers, - convert_sana_transformer_to_diffusers, - convert_sd3_transformer_checkpoint_to_diffusers, - convert_stable_cascade_unet_single_file_to_diffusers, - convert_wan_transformer_to_diffusers, - convert_wan_vae_to_diffusers, - convert_z_image_controlnet_checkpoint_to_diffusers, - convert_z_image_transformer_checkpoint_to_diffusers, - create_controlnet_diffusers_config_from_ldm, - create_unet_diffusers_config_from_ldm, - create_vae_diffusers_config_from_ldm, - fetch_diffusers_config, - fetch_original_config, - load_single_file_checkpoint, -) - - -logger = logging.get_logger(__name__) - - -if is_accelerate_available(): - from accelerate import dispatch_model, init_empty_weights - - from ..models.model_loading_utils import load_model_dict_into_meta - -if is_torch_version(">=", "1.9.0") and is_accelerate_available(): - _LOW_CPU_MEM_USAGE_DEFAULT = True -else: - _LOW_CPU_MEM_USAGE_DEFAULT = False - -SINGLE_FILE_LOADABLE_CLASSES = { - "StableCascadeUNet": { - "checkpoint_mapping_fn": convert_stable_cascade_unet_single_file_to_diffusers, - }, - "UNet2DConditionModel": { - "checkpoint_mapping_fn": convert_ldm_unet_checkpoint, - "config_mapping_fn": create_unet_diffusers_config_from_ldm, - "default_subfolder": "unet", - "legacy_kwargs": { - "num_in_channels": "in_channels", # Legacy kwargs supported by `from_single_file` mapped to new args - }, - }, - "AutoencoderKL": { - "checkpoint_mapping_fn": convert_ldm_vae_checkpoint, - "config_mapping_fn": create_vae_diffusers_config_from_ldm, - "default_subfolder": "vae", - }, - "ControlNetModel": { - "checkpoint_mapping_fn": convert_controlnet_checkpoint, - "config_mapping_fn": create_controlnet_diffusers_config_from_ldm, - }, - "SD3Transformer2DModel": { - "checkpoint_mapping_fn": convert_sd3_transformer_checkpoint_to_diffusers, - "default_subfolder": "transformer", - }, - "MotionAdapter": { - "checkpoint_mapping_fn": convert_animatediff_checkpoint_to_diffusers, - }, - "SparseControlNetModel": { - "checkpoint_mapping_fn": convert_animatediff_checkpoint_to_diffusers, - }, - "FluxTransformer2DModel": { - "checkpoint_mapping_fn": convert_flux_transformer_checkpoint_to_diffusers, - "default_subfolder": "transformer", - }, - "ChromaTransformer2DModel": { - "checkpoint_mapping_fn": convert_chroma_transformer_checkpoint_to_diffusers, - "default_subfolder": "transformer", - }, - "ErnieImageTransformer2DModel": { - "checkpoint_mapping_fn": convert_ernie_image_transformer_checkpoint_to_diffusers, - "default_subfolder": "transformer", - }, - "LTXVideoTransformer3DModel": { - "checkpoint_mapping_fn": convert_ltx_transformer_checkpoint_to_diffusers, - "default_subfolder": "transformer", - }, - "AutoencoderKLLTXVideo": { - "checkpoint_mapping_fn": convert_ltx_vae_checkpoint_to_diffusers, - "default_subfolder": "vae", - }, - "AutoencoderDC": {"checkpoint_mapping_fn": convert_autoencoder_dc_checkpoint_to_diffusers}, - "MochiTransformer3DModel": { - "checkpoint_mapping_fn": convert_mochi_transformer_checkpoint_to_diffusers, - "default_subfolder": "transformer", - }, - "HunyuanVideoTransformer3DModel": { - "checkpoint_mapping_fn": convert_hunyuan_video_transformer_to_diffusers, - "default_subfolder": "transformer", - }, - "AuraFlowTransformer2DModel": { - "checkpoint_mapping_fn": convert_auraflow_transformer_checkpoint_to_diffusers, - "default_subfolder": "transformer", - }, - "Lumina2Transformer2DModel": { - "checkpoint_mapping_fn": convert_lumina2_to_diffusers, - "default_subfolder": "transformer", - }, - "SanaTransformer2DModel": { - "checkpoint_mapping_fn": convert_sana_transformer_to_diffusers, - "default_subfolder": "transformer", - }, - "SkyReelsV2Transformer3DModel": { - "checkpoint_mapping_fn": convert_wan_transformer_to_diffusers, - "default_subfolder": "transformer", - }, - "ChronoEditTransformer3DModel": { - "checkpoint_mapping_fn": convert_wan_transformer_to_diffusers, - "default_subfolder": "transformer", - }, - "WanTransformer3DModel": { - "checkpoint_mapping_fn": convert_wan_transformer_to_diffusers, - "default_subfolder": "transformer", - }, - "WanVACETransformer3DModel": { - "checkpoint_mapping_fn": convert_wan_transformer_to_diffusers, - "default_subfolder": "transformer", - }, - "WanAnimateTransformer3DModel": { - "checkpoint_mapping_fn": convert_wan_transformer_to_diffusers, - "default_subfolder": "transformer", - }, - "AutoencoderKLWan": { - "checkpoint_mapping_fn": convert_wan_vae_to_diffusers, - "default_subfolder": "vae", - }, - "HiDreamImageTransformer2DModel": { - "checkpoint_mapping_fn": convert_hidream_transformer_to_diffusers, - "default_subfolder": "transformer", - }, - "CosmosTransformer3DModel": { - "checkpoint_mapping_fn": convert_cosmos_transformer_checkpoint_to_diffusers, - "default_subfolder": "transformer", - }, - "QwenImageTransformer2DModel": { - "checkpoint_mapping_fn": lambda checkpoint, **kwargs: checkpoint, - "default_subfolder": "transformer", - }, - "Flux2Transformer2DModel": { - "checkpoint_mapping_fn": convert_flux2_transformer_checkpoint_to_diffusers, - "default_subfolder": "transformer", - }, - "ZImageTransformer2DModel": { - "checkpoint_mapping_fn": convert_z_image_transformer_checkpoint_to_diffusers, - "default_subfolder": "transformer", - }, - "ZImageControlNetModel": { - "checkpoint_mapping_fn": convert_z_image_controlnet_checkpoint_to_diffusers, - }, - "LTX2VideoTransformer3DModel": { - "checkpoint_mapping_fn": convert_ltx2_transformer_to_diffusers, - "default_subfolder": "transformer", - }, - "AutoencoderKLLTX2Video": { - "checkpoint_mapping_fn": convert_ltx2_vae_to_diffusers, - "default_subfolder": "vae", - }, - "AutoencoderKLLTX2Audio": { - "checkpoint_mapping_fn": convert_ltx2_audio_vae_to_diffusers, - "default_subfolder": "audio_vae", - }, - "MotifVideoTransformer3DModel": { - "checkpoint_mapping_fn": lambda checkpoint, **kwargs: checkpoint, - "default_subfolder": "transformer", - }, -} - - -def _should_convert_state_dict_to_diffusers(model_state_dict, checkpoint_state_dict): - model_state_dict_keys = set(model_state_dict.keys()) - checkpoint_state_dict_keys = set(checkpoint_state_dict.keys()) - is_subset = model_state_dict_keys.issubset(checkpoint_state_dict_keys) - is_match = model_state_dict_keys == checkpoint_state_dict_keys - return not (is_subset and is_match) - - -def _get_single_file_loadable_mapping_class(cls): - diffusers_module = importlib.import_module(__name__.split(".")[0]) - for loadable_class_str in SINGLE_FILE_LOADABLE_CLASSES: - loadable_class = getattr(diffusers_module, loadable_class_str) - - if issubclass(cls, loadable_class): - return loadable_class_str - - return None - - -def _get_mapping_function_kwargs(mapping_fn, **kwargs): - parameters = inspect.signature(mapping_fn).parameters - - mapping_kwargs = {} - for parameter in parameters: - if parameter in kwargs: - mapping_kwargs[parameter] = kwargs[parameter] - - return mapping_kwargs - - -class FromOriginalModelMixin: - """ - Load pretrained weights saved in the `.ckpt` or `.safetensors` format into a model. - """ - - @classmethod - @validate_hf_hub_args - def from_single_file(cls, pretrained_model_link_or_path_or_dict: str | None = None, **kwargs) -> Self: - r""" - Instantiate a model from pretrained weights saved in the original `.ckpt` or `.safetensors` format. The model - is set in evaluation mode (`model.eval()`) by default. - - Parameters: - pretrained_model_link_or_path_or_dict (`str`, *optional*): - Can be either: - - A link to the `.safetensors` or `.ckpt` file (for example - `"https://huggingface.co//blob/main/.safetensors"`) on the Hub. - - A path to a local *file* containing the weights of the component model. - - A state dict containing the component model weights. - config (`str`, *optional*): - - A string, the *repo id* (for example `CompVis/ldm-text2im-large-256`) of a pretrained pipeline hosted - on the Hub. - - A path to a *directory* (for example `./my_pipeline_directory/`) containing the pipeline component - configs in Diffusers format. - subfolder (`str`, *optional*, defaults to `""`): - The subfolder location of a model file within a larger model repository on the Hub or locally. - original_config (`str`, *optional*): - Dict or path to a yaml file containing the configuration for the model in its original format. - If a dict is provided, it will be used to initialize the model configuration. - dtype (`torch.dtype`, *optional*): - Override the default `torch.dtype` and load the model with another dtype. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to True, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 and - is_accelerate_available() else `False`): Speed up model loading only loading the pretrained weights and - not initializing the weights. This also tries to not use more than 1x model size in CPU memory - (including peak memory) while loading the model. Only supported for PyTorch >= 1.9.0. If you are using - an older version of PyTorch, setting this argument to `True` will raise an error. - disable_mmap ('bool', *optional*, defaults to 'False'): - Whether to disable mmap when loading a Safetensors model. This option can perform better when the model - is on a network mount or hard drive, which may not handle the seeky-ness of mmap very well. - kwargs (remaining dictionary of keyword arguments, *optional*): - Can be used to overwrite load and saveable variables (for example the pipeline components of the - specific pipeline class). The overwritten components are directly passed to the pipelines `__init__` - method. See example below for more information. - - ```py - >>> from diffusers import StableCascadeUNet - - >>> ckpt_path = "https://huggingface.co/stabilityai/stable-cascade/blob/main/stage_b_lite.safetensors" - >>> model = StableCascadeUNet.from_single_file(ckpt_path) - ``` - """ - - mapping_class_name = _get_single_file_loadable_mapping_class(cls) - # if class_name not in SINGLE_FILE_LOADABLE_CLASSES: - if mapping_class_name is None: - raise ValueError( - f"FromOriginalModelMixin is currently only compatible with {', '.join(SINGLE_FILE_LOADABLE_CLASSES.keys())}" - ) - - pretrained_model_link_or_path = kwargs.get("pretrained_model_link_or_path", None) - if pretrained_model_link_or_path is not None: - deprecation_message = ( - "Please use `pretrained_model_link_or_path_or_dict` argument instead for model classes" - ) - deprecate("pretrained_model_link_or_path", "1.0.0", deprecation_message) - pretrained_model_link_or_path_or_dict = pretrained_model_link_or_path - - config = kwargs.pop("config", None) - original_config = kwargs.pop("original_config", None) - - if config is not None and original_config is not None: - raise ValueError( - "`from_single_file` cannot accept both `config` and `original_config` arguments. Please provide only one of these arguments" - ) - - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - token = kwargs.pop("token", None) - cache_dir = kwargs.pop("cache_dir", None) - local_files_only = kwargs.pop("local_files_only", None) - subfolder = kwargs.pop("subfolder", None) - revision = kwargs.pop("revision", None) - config_revision = kwargs.pop("config_revision", None) - torch_dtype = kwargs.pop("torch_dtype", None) - dtype = kwargs.pop("dtype", None) - torch_dtype = dtype if dtype is not None else torch_dtype - quantization_config = kwargs.pop("quantization_config", None) - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT) - device = kwargs.pop("device", None) - disable_mmap = kwargs.pop("disable_mmap", False) - device_map = kwargs.pop("device_map", None) - - user_agent = { - "diffusers": __version__, - "file_type": "single_file", - "framework": "pytorch", - } - # In order to ensure popular quantization methods are supported. Can be disable with `disable_telemetry` - if quantization_config is not None: - user_agent["quant"] = quantization_config.quant_method.value - - if torch_dtype is not None and not isinstance(torch_dtype, torch.dtype): - torch_dtype = torch.float32 - logger.warning( - f"Passed `torch_dtype` {torch_dtype} is not a `torch.dtype`. Defaulting to `torch.float32`." - ) - - if isinstance(pretrained_model_link_or_path_or_dict, dict): - checkpoint = pretrained_model_link_or_path_or_dict - else: - checkpoint = load_single_file_checkpoint( - pretrained_model_link_or_path_or_dict, - force_download=force_download, - proxies=proxies, - token=token, - cache_dir=cache_dir, - local_files_only=local_files_only, - revision=revision, - disable_mmap=disable_mmap, - user_agent=user_agent, - ) - if quantization_config is not None: - hf_quantizer = DiffusersAutoQuantizer.from_config(quantization_config) - hf_quantizer.validate_environment() - torch_dtype = hf_quantizer.update_torch_dtype(torch_dtype) - - else: - hf_quantizer = None - - mapping_functions = SINGLE_FILE_LOADABLE_CLASSES[mapping_class_name] - - checkpoint_mapping_fn = mapping_functions["checkpoint_mapping_fn"] - if original_config is not None: - if "config_mapping_fn" in mapping_functions: - config_mapping_fn = mapping_functions["config_mapping_fn"] - else: - config_mapping_fn = None - - if config_mapping_fn is None: - raise ValueError( - ( - f"`original_config` has been provided for {mapping_class_name} but no mapping function" - "was found to convert the original config to a Diffusers config in" - "`diffusers.loaders.single_file_utils`" - ) - ) - - if isinstance(original_config, str): - # If original_config is a URL or filepath fetch the original_config dict - original_config = fetch_original_config(original_config, local_files_only=local_files_only) - - config_mapping_kwargs = _get_mapping_function_kwargs(config_mapping_fn, **kwargs) - diffusers_model_config = config_mapping_fn( - original_config=original_config, - checkpoint=checkpoint, - **config_mapping_kwargs, - ) - else: - if config is not None: - if isinstance(config, str): - default_pretrained_model_config_name = config - else: - raise ValueError( - ( - "Invalid `config` argument. Please provide a string representing a repo id" - "or path to a local Diffusers model repo." - ) - ) - - else: - config = fetch_diffusers_config(checkpoint) - default_pretrained_model_config_name = config["pretrained_model_name_or_path"] - - if "default_subfolder" in mapping_functions: - subfolder = mapping_functions["default_subfolder"] - - subfolder = subfolder or config.pop( - "subfolder", None - ) # some configs contain a subfolder key, e.g. StableCascadeUNet - - diffusers_model_config = cls.load_config( - pretrained_model_name_or_path=default_pretrained_model_config_name, - subfolder=subfolder, - local_files_only=local_files_only, - token=token, - revision=config_revision, - ) - expected_kwargs, optional_kwargs = cls._get_signature_keys(cls) - - # Map legacy kwargs to new kwargs - if "legacy_kwargs" in mapping_functions: - legacy_kwargs = mapping_functions["legacy_kwargs"] - for legacy_key, new_key in legacy_kwargs.items(): - if legacy_key in kwargs: - kwargs[new_key] = kwargs.pop(legacy_key) - - model_kwargs = {k: kwargs.get(k) for k in kwargs if k in expected_kwargs or k in optional_kwargs} - diffusers_model_config.update(model_kwargs) - - ctx = init_empty_weights if low_cpu_mem_usage else nullcontext - with ctx(): - model = cls.from_config(diffusers_model_config) - - model_state_dict = model.state_dict() - - # Check if `_keep_in_fp32_modules` is not None - use_keep_in_fp32_modules = (cls._keep_in_fp32_modules is not None) and ( - (torch_dtype == torch.float16) or hasattr(hf_quantizer, "use_keep_in_fp32_modules") - ) - if use_keep_in_fp32_modules: - keep_in_fp32_modules = cls._keep_in_fp32_modules - if not isinstance(keep_in_fp32_modules, list): - keep_in_fp32_modules = [keep_in_fp32_modules] - - else: - keep_in_fp32_modules = [] - - # Now that the model is loaded, we can determine the `device_map` - device_map = _determine_device_map(model, device_map, None, torch_dtype, keep_in_fp32_modules, hf_quantizer) - if device_map is not None: - expanded_device_map = _expand_device_map(device_map, model_state_dict.keys()) - _caching_allocator_warmup(model, expanded_device_map, torch_dtype, hf_quantizer) - - checkpoint_mapping_kwargs = _get_mapping_function_kwargs(checkpoint_mapping_fn, **kwargs) - - if _should_convert_state_dict_to_diffusers(model_state_dict, checkpoint): - diffusers_format_checkpoint = checkpoint_mapping_fn( - config=diffusers_model_config, - checkpoint=checkpoint, - **checkpoint_mapping_kwargs, - ) - else: - diffusers_format_checkpoint = checkpoint - - if not diffusers_format_checkpoint: - raise SingleFileComponentError( - f"Failed to load {mapping_class_name}. Weights for this component appear to be missing in the checkpoint." - ) - - if hf_quantizer is not None: - hf_quantizer.preprocess_model( - model=model, - device_map=None, - state_dict=diffusers_format_checkpoint, - keep_in_fp32_modules=keep_in_fp32_modules, - ) - - device_map = None - if low_cpu_mem_usage: - param_device = torch.device(device) if device else torch.device("cpu") - empty_state_dict = model.state_dict() - unexpected_keys = [ - param_name for param_name in diffusers_format_checkpoint if param_name not in empty_state_dict - ] - device_map = {"": param_device} - load_model_dict_into_meta( - model, - diffusers_format_checkpoint, - dtype=torch_dtype, - device_map=device_map, - hf_quantizer=hf_quantizer, - keep_in_fp32_modules=keep_in_fp32_modules, - unexpected_keys=unexpected_keys, - ) - empty_device_cache() - else: - _, unexpected_keys = model.load_state_dict(diffusers_format_checkpoint, strict=False) - - if model._keys_to_ignore_on_load_unexpected is not None: - for pat in model._keys_to_ignore_on_load_unexpected: - unexpected_keys = [k for k in unexpected_keys if re.search(pat, k) is None] - - if len(unexpected_keys) > 0: - logger.warning( - f"Some weights of the model checkpoint were not used when initializing {cls.__name__}: \n {[', '.join(unexpected_keys)]}" - ) - - if hf_quantizer is not None: - hf_quantizer.postprocess_model(model) - model.hf_quantizer = hf_quantizer - - if torch_dtype is not None and hf_quantizer is None: - model.to(torch_dtype) - - model.eval() - - if device_map is not None: - device_map_kwargs = {"device_map": device_map} - dispatch_model(model, **device_map_kwargs) - - return model diff --git a/diffusers/loaders/single_file_utils.py b/diffusers/loaders/single_file_utils.py deleted file mode 100644 index 296f32f891f0d362753e8380c089f964bc2ef89a..0000000000000000000000000000000000000000 --- a/diffusers/loaders/single_file_utils.py +++ /dev/null @@ -1,4182 +0,0 @@ -# coding=utf-8 -# Copyright 2025 The HuggingFace Inc. team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -"""Conversion script for the Stable Diffusion checkpoints.""" - -import copy -import os -import re -from contextlib import nullcontext -from io import BytesIO -from urllib.parse import urlparse - -import requests -import torch -import yaml - -from ..models.modeling_utils import load_state_dict -from ..schedulers import ( - DDIMScheduler, - DPMSolverMultistepScheduler, - EDMDPMSolverMultistepScheduler, - EulerAncestralDiscreteScheduler, - EulerDiscreteScheduler, - HeunDiscreteScheduler, - LMSDiscreteScheduler, - PNDMScheduler, -) -from ..utils import ( - SAFETENSORS_WEIGHTS_NAME, - WEIGHTS_NAME, - deprecate, - is_accelerate_available, - is_transformers_available, - logging, -) -from ..utils.constants import DIFFUSERS_REQUEST_TIMEOUT -from ..utils.hub_utils import _get_model_file -from ..utils.torch_utils import empty_device_cache - - -if is_transformers_available(): - from transformers import AutoImageProcessor - -if is_accelerate_available(): - from accelerate import init_empty_weights - - from ..models.model_loading_utils import load_model_dict_into_meta - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - -CHECKPOINT_KEY_NAMES = { - "v1": "model.diffusion_model.output_blocks.11.0.skip_connection.weight", - "v2": "model.diffusion_model.input_blocks.2.1.transformer_blocks.0.attn2.to_k.weight", - "xl_base": "conditioner.embedders.1.model.transformer.resblocks.9.mlp.c_proj.bias", - "xl_refiner": "conditioner.embedders.0.model.transformer.resblocks.9.mlp.c_proj.bias", - "upscale": "model.diffusion_model.input_blocks.10.0.skip_connection.bias", - "controlnet": [ - "control_model.time_embed.0.weight", - "controlnet_cond_embedding.conv_in.weight", - ], - # TODO: find non-Diffusers keys for controlnet_xl - "controlnet_xl": "add_embedding.linear_1.weight", - "controlnet_xl_large": "down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_k.weight", - "controlnet_xl_mid": "down_blocks.1.attentions.0.norm.weight", - "playground-v2-5": "edm_mean", - "inpainting": "model.diffusion_model.input_blocks.0.0.weight", - "clip": "cond_stage_model.transformer.text_model.embeddings.position_embedding.weight", - "clip_sdxl": "conditioner.embedders.0.transformer.text_model.embeddings.position_embedding.weight", - "clip_sd3": "text_encoders.clip_l.transformer.text_model.embeddings.position_embedding.weight", - "open_clip": "cond_stage_model.model.token_embedding.weight", - "open_clip_sdxl": "conditioner.embedders.1.model.positional_embedding", - "open_clip_sdxl_refiner": "conditioner.embedders.0.model.text_projection", - "open_clip_sd3": "text_encoders.clip_g.transformer.text_model.embeddings.position_embedding.weight", - "stable_cascade_stage_b": "down_blocks.1.0.channelwise.0.weight", - "stable_cascade_stage_c": "clip_txt_mapper.weight", - "sd3": [ - "joint_blocks.0.context_block.adaLN_modulation.1.bias", - "model.diffusion_model.joint_blocks.0.context_block.adaLN_modulation.1.bias", - ], - "sd35_large": [ - "joint_blocks.37.x_block.mlp.fc1.weight", - "model.diffusion_model.joint_blocks.37.x_block.mlp.fc1.weight", - ], - "animatediff": "down_blocks.0.motion_modules.0.temporal_transformer.transformer_blocks.0.attention_blocks.0.pos_encoder.pe", - "animatediff_v2": "mid_block.motion_modules.0.temporal_transformer.norm.bias", - "animatediff_sdxl_beta": "up_blocks.2.motion_modules.0.temporal_transformer.norm.weight", - "animatediff_scribble": "controlnet_cond_embedding.conv_in.weight", - "animatediff_rgb": "controlnet_cond_embedding.weight", - "auraflow": [ - "double_layers.0.attn.w2q.weight", - "double_layers.0.attn.w1q.weight", - "cond_seq_linear.weight", - "t_embedder.mlp.0.weight", - ], - "flux": [ - "double_blocks.0.img_attn.norm.key_norm.scale", - "model.diffusion_model.double_blocks.0.img_attn.norm.key_norm.scale", - ], - "ltx-video": [ - "model.diffusion_model.patchify_proj.weight", - "model.diffusion_model.transformer_blocks.27.scale_shift_table", - "patchify_proj.weight", - "transformer_blocks.27.scale_shift_table", - "vae.decoder.last_scale_shift_table", # 0.9.1, 0.9.5, 0.9.7, 0.9.8 - "vae.decoder.up_blocks.9.res_blocks.0.conv1.conv.weight", # 0.9.0 - ], - "autoencoder-dc": "decoder.stages.1.op_list.0.main.conv.conv.bias", - "autoencoder-dc-sana": "encoder.project_in.conv.bias", - "mochi-1-preview": ["model.diffusion_model.blocks.0.attn.qkv_x.weight", "blocks.0.attn.qkv_x.weight"], - "hunyuan-video": "txt_in.individual_token_refiner.blocks.0.adaLN_modulation.1.bias", - "instruct-pix2pix": "model.diffusion_model.input_blocks.0.0.weight", - "lumina2": ["model.diffusion_model.cap_embedder.0.weight", "cap_embedder.0.weight"], - "z-image-turbo": [ - "model.diffusion_model.layers.0.adaLN_modulation.0.weight", - "layers.0.adaLN_modulation.0.weight", - ], - "z-image-turbo-controlnet": "control_all_x_embedder.2-1.weight", - "z-image-turbo-controlnet-2.x": "control_layers.14.adaLN_modulation.0.weight", - "sana": [ - "blocks.0.cross_attn.q_linear.weight", - "blocks.0.cross_attn.q_linear.bias", - "blocks.0.cross_attn.kv_linear.weight", - "blocks.0.cross_attn.kv_linear.bias", - ], - "wan": ["model.diffusion_model.head.modulation", "head.modulation"], - "wan_vae": "decoder.middle.0.residual.0.gamma", - "wan_vace": "vace_blocks.0.after_proj.bias", - "wan_animate": "motion_encoder.dec.direction.weight", - "hidream": "double_stream_blocks.0.block.adaLN_modulation.1.bias", - "cosmos-1.0": [ - "net.x_embedder.proj.1.weight", - "net.blocks.block1.blocks.0.block.attn.to_q.0.weight", - "net.extra_pos_embedder.pos_emb_h", - ], - "cosmos-2.0": [ - "net.x_embedder.proj.1.weight", - "net.blocks.0.self_attn.q_proj.weight", - "net.pos_embedder.dim_spatial_range", - ], - "flux2": ["model.diffusion_model.single_stream_modulation.lin.weight", "single_stream_modulation.lin.weight"], - "ltx2": [ - "model.diffusion_model.av_ca_a2v_gate_adaln_single.emb.timestep_embedder.linear_1.weight", - "vae.per_channel_statistics.mean-of-means", - "audio_vae.per_channel_statistics.mean-of-means", - ], -} - -DIFFUSERS_DEFAULT_PIPELINE_PATHS = { - "xl_base": {"pretrained_model_name_or_path": "stabilityai/stable-diffusion-xl-base-1.0"}, - "xl_refiner": {"pretrained_model_name_or_path": "stabilityai/stable-diffusion-xl-refiner-1.0"}, - "xl_inpaint": {"pretrained_model_name_or_path": "diffusers/stable-diffusion-xl-1.0-inpainting-0.1"}, - "playground-v2-5": {"pretrained_model_name_or_path": "playgroundai/playground-v2.5-1024px-aesthetic"}, - "upscale": {"pretrained_model_name_or_path": "stabilityai/stable-diffusion-x4-upscaler"}, - "inpainting": {"pretrained_model_name_or_path": "stable-diffusion-v1-5/stable-diffusion-inpainting"}, - "inpainting_v2": {"pretrained_model_name_or_path": "stabilityai/stable-diffusion-2-inpainting"}, - "controlnet": {"pretrained_model_name_or_path": "lllyasviel/control_v11p_sd15_canny"}, - "controlnet_xl_large": {"pretrained_model_name_or_path": "diffusers/controlnet-canny-sdxl-1.0"}, - "controlnet_xl_mid": {"pretrained_model_name_or_path": "diffusers/controlnet-canny-sdxl-1.0-mid"}, - "controlnet_xl_small": {"pretrained_model_name_or_path": "diffusers/controlnet-canny-sdxl-1.0-small"}, - "v2": {"pretrained_model_name_or_path": "stabilityai/stable-diffusion-2-1"}, - "v1": {"pretrained_model_name_or_path": "stable-diffusion-v1-5/stable-diffusion-v1-5"}, - "stable_cascade_stage_b": {"pretrained_model_name_or_path": "stabilityai/stable-cascade", "subfolder": "decoder"}, - "stable_cascade_stage_b_lite": { - "pretrained_model_name_or_path": "stabilityai/stable-cascade", - "subfolder": "decoder_lite", - }, - "stable_cascade_stage_c": { - "pretrained_model_name_or_path": "stabilityai/stable-cascade-prior", - "subfolder": "prior", - }, - "stable_cascade_stage_c_lite": { - "pretrained_model_name_or_path": "stabilityai/stable-cascade-prior", - "subfolder": "prior_lite", - }, - "sd3": { - "pretrained_model_name_or_path": "stabilityai/stable-diffusion-3-medium-diffusers", - }, - "sd35_large": { - "pretrained_model_name_or_path": "stabilityai/stable-diffusion-3.5-large", - }, - "sd35_medium": { - "pretrained_model_name_or_path": "stabilityai/stable-diffusion-3.5-medium", - }, - "animatediff_v1": {"pretrained_model_name_or_path": "guoyww/animatediff-motion-adapter-v1-5"}, - "animatediff_v2": {"pretrained_model_name_or_path": "guoyww/animatediff-motion-adapter-v1-5-2"}, - "animatediff_v3": {"pretrained_model_name_or_path": "guoyww/animatediff-motion-adapter-v1-5-3"}, - "animatediff_sdxl_beta": {"pretrained_model_name_or_path": "guoyww/animatediff-motion-adapter-sdxl-beta"}, - "animatediff_scribble": {"pretrained_model_name_or_path": "guoyww/animatediff-sparsectrl-scribble"}, - "animatediff_rgb": {"pretrained_model_name_or_path": "guoyww/animatediff-sparsectrl-rgb"}, - "auraflow": {"pretrained_model_name_or_path": "fal/AuraFlow-v0.3"}, - "flux-dev": {"pretrained_model_name_or_path": "black-forest-labs/FLUX.1-dev"}, - "flux-fill": {"pretrained_model_name_or_path": "black-forest-labs/FLUX.1-Fill-dev"}, - "flux-depth": {"pretrained_model_name_or_path": "black-forest-labs/FLUX.1-Depth-dev"}, - "flux-schnell": {"pretrained_model_name_or_path": "black-forest-labs/FLUX.1-schnell"}, - "flux-2-dev": {"pretrained_model_name_or_path": "black-forest-labs/FLUX.2-dev"}, - "ltx-video": {"pretrained_model_name_or_path": "diffusers/LTX-Video-0.9.0"}, - "ltx-video-0.9.1": {"pretrained_model_name_or_path": "diffusers/LTX-Video-0.9.1"}, - "ltx-video-0.9.5": {"pretrained_model_name_or_path": "Lightricks/LTX-Video-0.9.5"}, - "ltx-video-0.9.7": {"pretrained_model_name_or_path": "Lightricks/LTX-Video-0.9.7-dev"}, - "autoencoder-dc-f128c512": {"pretrained_model_name_or_path": "mit-han-lab/dc-ae-f128c512-mix-1.0-diffusers"}, - "autoencoder-dc-f64c128": {"pretrained_model_name_or_path": "mit-han-lab/dc-ae-f64c128-mix-1.0-diffusers"}, - "autoencoder-dc-f32c32": {"pretrained_model_name_or_path": "mit-han-lab/dc-ae-f32c32-mix-1.0-diffusers"}, - "autoencoder-dc-f32c32-sana": {"pretrained_model_name_or_path": "mit-han-lab/dc-ae-f32c32-sana-1.0-diffusers"}, - "mochi-1-preview": {"pretrained_model_name_or_path": "genmo/mochi-1-preview"}, - "hunyuan-video": {"pretrained_model_name_or_path": "hunyuanvideo-community/HunyuanVideo"}, - "instruct-pix2pix": {"pretrained_model_name_or_path": "timbrooks/instruct-pix2pix"}, - "lumina2": {"pretrained_model_name_or_path": "Alpha-VLLM/Lumina-Image-2.0"}, - "sana": {"pretrained_model_name_or_path": "Efficient-Large-Model/Sana_1600M_1024px_diffusers"}, - "wan-t2v-1.3B": {"pretrained_model_name_or_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"}, - "wan-t2v-14B": {"pretrained_model_name_or_path": "Wan-AI/Wan2.1-T2V-14B-Diffusers"}, - "wan-i2v-14B": {"pretrained_model_name_or_path": "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"}, - "wan-animate-14B": {"pretrained_model_name_or_path": "Wan-AI/Wan2.2-Animate-14B-Diffusers"}, - "wan-vace-1.3B": {"pretrained_model_name_or_path": "Wan-AI/Wan2.1-VACE-1.3B-diffusers"}, - "wan-vace-14B": {"pretrained_model_name_or_path": "Wan-AI/Wan2.1-VACE-14B-diffusers"}, - "hidream": {"pretrained_model_name_or_path": "HiDream-ai/HiDream-I1-Dev"}, - "cosmos-1.0-t2w-7B": {"pretrained_model_name_or_path": "nvidia/Cosmos-1.0-Diffusion-7B-Text2World"}, - "cosmos-1.0-t2w-14B": {"pretrained_model_name_or_path": "nvidia/Cosmos-1.0-Diffusion-14B-Text2World"}, - "cosmos-1.0-v2w-7B": {"pretrained_model_name_or_path": "nvidia/Cosmos-1.0-Diffusion-7B-Video2World"}, - "cosmos-1.0-v2w-14B": {"pretrained_model_name_or_path": "nvidia/Cosmos-1.0-Diffusion-14B-Video2World"}, - "cosmos-2.0-t2i-2B": {"pretrained_model_name_or_path": "nvidia/Cosmos-Predict2-2B-Text2Image"}, - "cosmos-2.0-t2i-14B": {"pretrained_model_name_or_path": "nvidia/Cosmos-Predict2-14B-Text2Image"}, - "cosmos-2.0-v2w-2B": {"pretrained_model_name_or_path": "nvidia/Cosmos-Predict2-2B-Video2World"}, - "cosmos-2.0-v2w-14B": {"pretrained_model_name_or_path": "nvidia/Cosmos-Predict2-14B-Video2World"}, - "z-image-turbo": {"pretrained_model_name_or_path": "Tongyi-MAI/Z-Image-Turbo"}, - "z-image-turbo-controlnet": {"pretrained_model_name_or_path": "hlky/Z-Image-Turbo-Fun-Controlnet-Union"}, - "z-image-turbo-controlnet-2.0": {"pretrained_model_name_or_path": "hlky/Z-Image-Turbo-Fun-Controlnet-Union-2.0"}, - "z-image-turbo-controlnet-2.1": {"pretrained_model_name_or_path": "hlky/Z-Image-Turbo-Fun-Controlnet-Union-2.1"}, - "ltx2-dev": {"pretrained_model_name_or_path": "Lightricks/LTX-2"}, -} - -# Use to configure model sample size when original config is provided -DIFFUSERS_TO_LDM_DEFAULT_IMAGE_SIZE_MAP = { - "xl_base": 1024, - "xl_refiner": 1024, - "xl_inpaint": 1024, - "playground-v2-5": 1024, - "upscale": 512, - "inpainting": 512, - "inpainting_v2": 512, - "controlnet": 512, - "instruct-pix2pix": 512, - "v2": 768, - "v1": 512, -} - - -DIFFUSERS_TO_LDM_MAPPING = { - "unet": { - "layers": { - "time_embedding.linear_1.weight": "time_embed.0.weight", - "time_embedding.linear_1.bias": "time_embed.0.bias", - "time_embedding.linear_2.weight": "time_embed.2.weight", - "time_embedding.linear_2.bias": "time_embed.2.bias", - "conv_in.weight": "input_blocks.0.0.weight", - "conv_in.bias": "input_blocks.0.0.bias", - "conv_norm_out.weight": "out.0.weight", - "conv_norm_out.bias": "out.0.bias", - "conv_out.weight": "out.2.weight", - "conv_out.bias": "out.2.bias", - }, - "class_embed_type": { - "class_embedding.linear_1.weight": "label_emb.0.0.weight", - "class_embedding.linear_1.bias": "label_emb.0.0.bias", - "class_embedding.linear_2.weight": "label_emb.0.2.weight", - "class_embedding.linear_2.bias": "label_emb.0.2.bias", - }, - "addition_embed_type": { - "add_embedding.linear_1.weight": "label_emb.0.0.weight", - "add_embedding.linear_1.bias": "label_emb.0.0.bias", - "add_embedding.linear_2.weight": "label_emb.0.2.weight", - "add_embedding.linear_2.bias": "label_emb.0.2.bias", - }, - }, - "controlnet": { - "layers": { - "time_embedding.linear_1.weight": "time_embed.0.weight", - "time_embedding.linear_1.bias": "time_embed.0.bias", - "time_embedding.linear_2.weight": "time_embed.2.weight", - "time_embedding.linear_2.bias": "time_embed.2.bias", - "conv_in.weight": "input_blocks.0.0.weight", - "conv_in.bias": "input_blocks.0.0.bias", - "controlnet_cond_embedding.conv_in.weight": "input_hint_block.0.weight", - "controlnet_cond_embedding.conv_in.bias": "input_hint_block.0.bias", - "controlnet_cond_embedding.conv_out.weight": "input_hint_block.14.weight", - "controlnet_cond_embedding.conv_out.bias": "input_hint_block.14.bias", - }, - "class_embed_type": { - "class_embedding.linear_1.weight": "label_emb.0.0.weight", - "class_embedding.linear_1.bias": "label_emb.0.0.bias", - "class_embedding.linear_2.weight": "label_emb.0.2.weight", - "class_embedding.linear_2.bias": "label_emb.0.2.bias", - }, - "addition_embed_type": { - "add_embedding.linear_1.weight": "label_emb.0.0.weight", - "add_embedding.linear_1.bias": "label_emb.0.0.bias", - "add_embedding.linear_2.weight": "label_emb.0.2.weight", - "add_embedding.linear_2.bias": "label_emb.0.2.bias", - }, - }, - "vae": { - "encoder.conv_in.weight": "encoder.conv_in.weight", - "encoder.conv_in.bias": "encoder.conv_in.bias", - "encoder.conv_out.weight": "encoder.conv_out.weight", - "encoder.conv_out.bias": "encoder.conv_out.bias", - "encoder.conv_norm_out.weight": "encoder.norm_out.weight", - "encoder.conv_norm_out.bias": "encoder.norm_out.bias", - "decoder.conv_in.weight": "decoder.conv_in.weight", - "decoder.conv_in.bias": "decoder.conv_in.bias", - "decoder.conv_out.weight": "decoder.conv_out.weight", - "decoder.conv_out.bias": "decoder.conv_out.bias", - "decoder.conv_norm_out.weight": "decoder.norm_out.weight", - "decoder.conv_norm_out.bias": "decoder.norm_out.bias", - "quant_conv.weight": "quant_conv.weight", - "quant_conv.bias": "quant_conv.bias", - "post_quant_conv.weight": "post_quant_conv.weight", - "post_quant_conv.bias": "post_quant_conv.bias", - }, - "openclip": { - "layers": { - "text_model.embeddings.position_embedding.weight": "positional_embedding", - "text_model.embeddings.token_embedding.weight": "token_embedding.weight", - "text_model.final_layer_norm.weight": "ln_final.weight", - "text_model.final_layer_norm.bias": "ln_final.bias", - "text_projection.weight": "text_projection", - }, - "transformer": { - "text_model.encoder.layers.": "resblocks.", - "layer_norm1": "ln_1", - "layer_norm2": "ln_2", - ".fc1.": ".c_fc.", - ".fc2.": ".c_proj.", - ".self_attn": ".attn", - "transformer.text_model.final_layer_norm.": "ln_final.", - "transformer.text_model.embeddings.token_embedding.weight": "token_embedding.weight", - "transformer.text_model.embeddings.position_embedding.weight": "positional_embedding", - }, - }, -} - -SD_2_TEXT_ENCODER_KEYS_TO_IGNORE = [ - "cond_stage_model.model.transformer.resblocks.23.attn.in_proj_bias", - "cond_stage_model.model.transformer.resblocks.23.attn.in_proj_weight", - "cond_stage_model.model.transformer.resblocks.23.attn.out_proj.bias", - "cond_stage_model.model.transformer.resblocks.23.attn.out_proj.weight", - "cond_stage_model.model.transformer.resblocks.23.ln_1.bias", - "cond_stage_model.model.transformer.resblocks.23.ln_1.weight", - "cond_stage_model.model.transformer.resblocks.23.ln_2.bias", - "cond_stage_model.model.transformer.resblocks.23.ln_2.weight", - "cond_stage_model.model.transformer.resblocks.23.mlp.c_fc.bias", - "cond_stage_model.model.transformer.resblocks.23.mlp.c_fc.weight", - "cond_stage_model.model.transformer.resblocks.23.mlp.c_proj.bias", - "cond_stage_model.model.transformer.resblocks.23.mlp.c_proj.weight", - "cond_stage_model.model.text_projection", -] - -# To support legacy scheduler_type argument -SCHEDULER_DEFAULT_CONFIG = { - "beta_schedule": "scaled_linear", - "beta_start": 0.00085, - "beta_end": 0.012, - "interpolation_type": "linear", - "num_train_timesteps": 1000, - "prediction_type": "epsilon", - "sample_max_value": 1.0, - "set_alpha_to_one": False, - "skip_prk_steps": True, - "steps_offset": 1, - "timestep_spacing": "leading", -} - -LDM_VAE_KEYS = ["first_stage_model.", "vae."] -LDM_VAE_DEFAULT_SCALING_FACTOR = 0.18215 -PLAYGROUND_VAE_SCALING_FACTOR = 0.5 -LDM_UNET_KEY = "model.diffusion_model." -LDM_CONTROLNET_KEY = "control_model." -LDM_CLIP_PREFIX_TO_REMOVE = [ - "cond_stage_model.transformer.", - "conditioner.embedders.0.transformer.", -] -LDM_OPEN_CLIP_TEXT_PROJECTION_DIM = 1024 -SCHEDULER_LEGACY_KWARGS = ["prediction_type", "scheduler_type"] - -VALID_URL_PREFIXES = ["https://huggingface.co/", "huggingface.co/", "hf.co/", "https://hf.co/"] - - -class SingleFileComponentError(Exception): - def __init__(self, message=None): - self.message = message - super().__init__(self.message) - - -def is_valid_url(url): - result = urlparse(url) - if result.scheme and result.netloc: - return True - - return False - - -def _is_single_file_path_or_url(pretrained_model_name_or_path): - if os.path.isfile(pretrained_model_name_or_path): - return True - - if not is_valid_url(pretrained_model_name_or_path): - return False - - repo_id, weight_name = _extract_repo_id_and_weights_name(pretrained_model_name_or_path) - return bool(repo_id and weight_name) - - -def _extract_repo_id_and_weights_name(pretrained_model_name_or_path): - if not is_valid_url(pretrained_model_name_or_path): - raise ValueError("Invalid `pretrained_model_name_or_path` provided. Please set it to a valid URL.") - - pattern = r"([^/]+)/([^/]+)/(?:blob/main/)?(.+)" - weights_name = None - repo_id = (None,) - for prefix in VALID_URL_PREFIXES: - pretrained_model_name_or_path = pretrained_model_name_or_path.replace(prefix, "") - match = re.match(pattern, pretrained_model_name_or_path) - if not match: - return repo_id, weights_name - - repo_id = f"{match.group(1)}/{match.group(2)}" - weights_name = match.group(3) - - return repo_id, weights_name - - -def _is_model_weights_in_cached_folder(cached_folder, name): - pretrained_model_name_or_path = os.path.join(cached_folder, name) - weights_exist = False - - for weights_name in [WEIGHTS_NAME, SAFETENSORS_WEIGHTS_NAME]: - if os.path.isfile(os.path.join(pretrained_model_name_or_path, weights_name)): - weights_exist = True - - return weights_exist - - -def _is_legacy_scheduler_kwargs(kwargs): - return any(k in SCHEDULER_LEGACY_KWARGS for k in kwargs.keys()) - - -def load_single_file_checkpoint( - pretrained_model_link_or_path, - force_download=False, - proxies=None, - token=None, - cache_dir=None, - local_files_only=None, - revision=None, - disable_mmap=False, - user_agent=None, -): - if user_agent is None: - user_agent = {"file_type": "single_file", "framework": "pytorch"} - - if os.path.isfile(pretrained_model_link_or_path): - pretrained_model_link_or_path = pretrained_model_link_or_path - - else: - repo_id, weights_name = _extract_repo_id_and_weights_name(pretrained_model_link_or_path) - pretrained_model_link_or_path = _get_model_file( - repo_id, - weights_name=weights_name, - force_download=force_download, - cache_dir=cache_dir, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - user_agent=user_agent, - ) - - checkpoint = load_state_dict(pretrained_model_link_or_path, disable_mmap=disable_mmap) - - # some checkpoints contain the model state dict under a "state_dict" key - while "state_dict" in checkpoint: - checkpoint = checkpoint["state_dict"] - - return checkpoint - - -def fetch_original_config(original_config_file, local_files_only=False): - if os.path.isfile(original_config_file): - with open(original_config_file, "r") as fp: - original_config_file = fp.read() - - elif is_valid_url(original_config_file): - if local_files_only: - raise ValueError( - "`local_files_only` is set to True, but a URL was provided as `original_config_file`. " - "Please provide a valid local file path." - ) - - original_config_file = BytesIO(requests.get(original_config_file, timeout=DIFFUSERS_REQUEST_TIMEOUT).content) - - else: - raise ValueError("Invalid `original_config_file` provided. Please set it to a valid file path or URL.") - - original_config = yaml.safe_load(original_config_file) - - return original_config - - -def is_clip_model(checkpoint): - if CHECKPOINT_KEY_NAMES["clip"] in checkpoint: - return True - - return False - - -def is_clip_sdxl_model(checkpoint): - if CHECKPOINT_KEY_NAMES["clip_sdxl"] in checkpoint: - return True - - return False - - -def is_clip_sd3_model(checkpoint): - if CHECKPOINT_KEY_NAMES["clip_sd3"] in checkpoint: - return True - - return False - - -def is_open_clip_model(checkpoint): - if CHECKPOINT_KEY_NAMES["open_clip"] in checkpoint: - return True - - return False - - -def is_open_clip_sdxl_model(checkpoint): - if CHECKPOINT_KEY_NAMES["open_clip_sdxl"] in checkpoint: - return True - - return False - - -def is_open_clip_sd3_model(checkpoint): - if CHECKPOINT_KEY_NAMES["open_clip_sd3"] in checkpoint: - return True - - return False - - -def is_open_clip_sdxl_refiner_model(checkpoint): - if CHECKPOINT_KEY_NAMES["open_clip_sdxl_refiner"] in checkpoint: - return True - - return False - - -def is_clip_model_in_single_file(class_obj, checkpoint): - is_clip_in_checkpoint = any( - [ - is_clip_model(checkpoint), - is_clip_sd3_model(checkpoint), - is_open_clip_model(checkpoint), - is_open_clip_sdxl_model(checkpoint), - is_open_clip_sdxl_refiner_model(checkpoint), - is_open_clip_sd3_model(checkpoint), - ] - ) - if ( - class_obj.__name__ == "CLIPTextModel" or class_obj.__name__ == "CLIPTextModelWithProjection" - ) and is_clip_in_checkpoint: - return True - - return False - - -def infer_diffusers_model_type(checkpoint): - if ( - CHECKPOINT_KEY_NAMES["inpainting"] in checkpoint - and checkpoint[CHECKPOINT_KEY_NAMES["inpainting"]].shape[1] == 9 - ): - if CHECKPOINT_KEY_NAMES["v2"] in checkpoint and checkpoint[CHECKPOINT_KEY_NAMES["v2"]].shape[-1] == 1024: - model_type = "inpainting_v2" - elif CHECKPOINT_KEY_NAMES["xl_base"] in checkpoint: - model_type = "xl_inpaint" - else: - model_type = "inpainting" - - elif CHECKPOINT_KEY_NAMES["v2"] in checkpoint and checkpoint[CHECKPOINT_KEY_NAMES["v2"]].shape[-1] == 1024: - model_type = "v2" - - elif CHECKPOINT_KEY_NAMES["playground-v2-5"] in checkpoint: - model_type = "playground-v2-5" - - elif CHECKPOINT_KEY_NAMES["xl_base"] in checkpoint: - model_type = "xl_base" - - elif CHECKPOINT_KEY_NAMES["xl_refiner"] in checkpoint: - model_type = "xl_refiner" - - elif CHECKPOINT_KEY_NAMES["upscale"] in checkpoint: - model_type = "upscale" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["controlnet"]): - if CHECKPOINT_KEY_NAMES["controlnet_xl"] in checkpoint: - if CHECKPOINT_KEY_NAMES["controlnet_xl_large"] in checkpoint: - model_type = "controlnet_xl_large" - elif CHECKPOINT_KEY_NAMES["controlnet_xl_mid"] in checkpoint: - model_type = "controlnet_xl_mid" - else: - model_type = "controlnet_xl_small" - else: - model_type = "controlnet" - - elif ( - CHECKPOINT_KEY_NAMES["stable_cascade_stage_c"] in checkpoint - and checkpoint[CHECKPOINT_KEY_NAMES["stable_cascade_stage_c"]].shape[0] == 1536 - ): - model_type = "stable_cascade_stage_c_lite" - - elif ( - CHECKPOINT_KEY_NAMES["stable_cascade_stage_c"] in checkpoint - and checkpoint[CHECKPOINT_KEY_NAMES["stable_cascade_stage_c"]].shape[0] == 2048 - ): - model_type = "stable_cascade_stage_c" - - elif ( - CHECKPOINT_KEY_NAMES["stable_cascade_stage_b"] in checkpoint - and checkpoint[CHECKPOINT_KEY_NAMES["stable_cascade_stage_b"]].shape[-1] == 576 - ): - model_type = "stable_cascade_stage_b_lite" - - elif ( - CHECKPOINT_KEY_NAMES["stable_cascade_stage_b"] in checkpoint - and checkpoint[CHECKPOINT_KEY_NAMES["stable_cascade_stage_b"]].shape[-1] == 640 - ): - model_type = "stable_cascade_stage_b" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["sd3"]) and any( - checkpoint[key].shape[-1] == 9216 if key in checkpoint else False for key in CHECKPOINT_KEY_NAMES["sd3"] - ): - if "model.diffusion_model.pos_embed" in checkpoint: - key = "model.diffusion_model.pos_embed" - else: - key = "pos_embed" - - if checkpoint[key].shape[1] == 36864: - model_type = "sd3" - elif checkpoint[key].shape[1] == 147456: - model_type = "sd35_medium" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["sd35_large"]): - model_type = "sd35_large" - - elif CHECKPOINT_KEY_NAMES["animatediff"] in checkpoint: - if CHECKPOINT_KEY_NAMES["animatediff_scribble"] in checkpoint: - model_type = "animatediff_scribble" - - elif CHECKPOINT_KEY_NAMES["animatediff_rgb"] in checkpoint: - model_type = "animatediff_rgb" - - elif CHECKPOINT_KEY_NAMES["animatediff_v2"] in checkpoint: - model_type = "animatediff_v2" - - elif checkpoint[CHECKPOINT_KEY_NAMES["animatediff_sdxl_beta"]].shape[-1] == 320: - model_type = "animatediff_sdxl_beta" - - elif checkpoint[CHECKPOINT_KEY_NAMES["animatediff"]].shape[1] == 24: - model_type = "animatediff_v1" - - else: - model_type = "animatediff_v3" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["flux2"]): - model_type = "flux-2-dev" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["flux"]): - if any( - g in checkpoint for g in ["guidance_in.in_layer.bias", "model.diffusion_model.guidance_in.in_layer.bias"] - ): - if "model.diffusion_model.img_in.weight" in checkpoint: - key = "model.diffusion_model.img_in.weight" - else: - key = "img_in.weight" - - if checkpoint[key].shape[1] == 384: - model_type = "flux-fill" - elif checkpoint[key].shape[1] == 128: - model_type = "flux-depth" - else: - model_type = "flux-dev" - else: - model_type = "flux-schnell" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["ltx-video"]): - has_vae = "vae.encoder.conv_in.conv.bias" in checkpoint - if any(key.endswith("transformer_blocks.47.scale_shift_table") for key in checkpoint): - model_type = "ltx-video-0.9.7" - elif has_vae and checkpoint["vae.encoder.conv_out.conv.weight"].shape[1] == 2048: - model_type = "ltx-video-0.9.5" - elif "vae.decoder.last_time_embedder.timestep_embedder.linear_1.weight" in checkpoint: - model_type = "ltx-video-0.9.1" - else: - model_type = "ltx-video" - - elif CHECKPOINT_KEY_NAMES["autoencoder-dc"] in checkpoint: - encoder_key = "encoder.project_in.conv.conv.bias" - decoder_key = "decoder.project_in.main.conv.weight" - - if CHECKPOINT_KEY_NAMES["autoencoder-dc-sana"] in checkpoint: - model_type = "autoencoder-dc-f32c32-sana" - - elif checkpoint[encoder_key].shape[-1] == 64 and checkpoint[decoder_key].shape[1] == 32: - model_type = "autoencoder-dc-f32c32" - - elif checkpoint[encoder_key].shape[-1] == 64 and checkpoint[decoder_key].shape[1] == 128: - model_type = "autoencoder-dc-f64c128" - - else: - model_type = "autoencoder-dc-f128c512" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["mochi-1-preview"]): - model_type = "mochi-1-preview" - - elif CHECKPOINT_KEY_NAMES["hunyuan-video"] in checkpoint: - model_type = "hunyuan-video" - - elif all(key in checkpoint for key in CHECKPOINT_KEY_NAMES["auraflow"]): - model_type = "auraflow" - - elif ( - CHECKPOINT_KEY_NAMES["instruct-pix2pix"] in checkpoint - and checkpoint[CHECKPOINT_KEY_NAMES["instruct-pix2pix"]].shape[1] == 8 - ): - model_type = "instruct-pix2pix" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["z-image-turbo"]): - model_type = "z-image-turbo" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["lumina2"]): - model_type = "lumina2" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["sana"]): - model_type = "sana" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["wan"]): - if "model.diffusion_model.patch_embedding.weight" in checkpoint: - target_key = "model.diffusion_model.patch_embedding.weight" - else: - target_key = "patch_embedding.weight" - - if CHECKPOINT_KEY_NAMES["wan_vace"] in checkpoint: - if checkpoint[target_key].shape[0] == 1536: - model_type = "wan-vace-1.3B" - elif checkpoint[target_key].shape[0] == 5120: - model_type = "wan-vace-14B" - - if CHECKPOINT_KEY_NAMES["wan_animate"] in checkpoint: - model_type = "wan-animate-14B" - - elif checkpoint[target_key].shape[0] == 1536: - model_type = "wan-t2v-1.3B" - elif checkpoint[target_key].shape[0] == 5120 and checkpoint[target_key].shape[1] == 16: - model_type = "wan-t2v-14B" - else: - model_type = "wan-i2v-14B" - - elif CHECKPOINT_KEY_NAMES["wan_vae"] in checkpoint: - # All Wan models use the same VAE so we can use the same default model repo to fetch the config - model_type = "wan-t2v-14B" - - elif CHECKPOINT_KEY_NAMES["hidream"] in checkpoint: - model_type = "hidream" - - elif all(key in checkpoint for key in CHECKPOINT_KEY_NAMES["cosmos-1.0"]): - x_embedder_shape = checkpoint[CHECKPOINT_KEY_NAMES["cosmos-1.0"][0]].shape - if x_embedder_shape[1] == 68: - model_type = "cosmos-1.0-t2w-7B" if x_embedder_shape[0] == 4096 else "cosmos-1.0-t2w-14B" - elif x_embedder_shape[1] == 72: - model_type = "cosmos-1.0-v2w-7B" if x_embedder_shape[0] == 4096 else "cosmos-1.0-v2w-14B" - else: - raise ValueError(f"Unexpected x_embedder shape: {x_embedder_shape} when loading Cosmos 1.0 model.") - - elif all(key in checkpoint for key in CHECKPOINT_KEY_NAMES["cosmos-2.0"]): - x_embedder_shape = checkpoint[CHECKPOINT_KEY_NAMES["cosmos-2.0"][0]].shape - if x_embedder_shape[1] == 68: - model_type = "cosmos-2.0-t2i-2B" if x_embedder_shape[0] == 2048 else "cosmos-2.0-t2i-14B" - elif x_embedder_shape[1] == 72: - model_type = "cosmos-2.0-v2w-2B" if x_embedder_shape[0] == 2048 else "cosmos-2.0-v2w-14B" - else: - raise ValueError(f"Unexpected x_embedder shape: {x_embedder_shape} when loading Cosmos 2.0 model.") - - elif CHECKPOINT_KEY_NAMES["z-image-turbo-controlnet-2.x"] in checkpoint: - before_proj_weight = checkpoint.get("control_noise_refiner.0.before_proj.weight", None) - if before_proj_weight is None: - model_type = "z-image-turbo-controlnet-2.0" - elif before_proj_weight is not None and torch.all(before_proj_weight == 0.0): - model_type = "z-image-turbo-controlnet-2.0" - else: - model_type = "z-image-turbo-controlnet-2.1" - - elif CHECKPOINT_KEY_NAMES["z-image-turbo-controlnet"] in checkpoint: - model_type = "z-image-turbo-controlnet" - - elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["ltx2"]): - model_type = "ltx2-dev" - - else: - model_type = "v1" - - return model_type - - -def fetch_diffusers_config(checkpoint): - model_type = infer_diffusers_model_type(checkpoint) - model_path = DIFFUSERS_DEFAULT_PIPELINE_PATHS[model_type] - model_path = copy.deepcopy(model_path) - - return model_path - - -def set_image_size(checkpoint, image_size=None): - if image_size: - return image_size - - model_type = infer_diffusers_model_type(checkpoint) - image_size = DIFFUSERS_TO_LDM_DEFAULT_IMAGE_SIZE_MAP[model_type] - - return image_size - - -# Copied from diffusers.pipelines.stable_diffusion.convert_from_ckpt.conv_attn_to_linear -def conv_attn_to_linear(checkpoint): - keys = list(checkpoint.keys()) - attn_keys = ["query.weight", "key.weight", "value.weight"] - for key in keys: - if ".".join(key.split(".")[-2:]) in attn_keys: - if checkpoint[key].ndim > 2: - checkpoint[key] = checkpoint[key][:, :, 0, 0] - elif "proj_attn.weight" in key: - if checkpoint[key].ndim > 2: - checkpoint[key] = checkpoint[key][:, :, 0] - - -def create_unet_diffusers_config_from_ldm( - original_config, checkpoint, image_size=None, upcast_attention=None, num_in_channels=None -): - """ - Creates a config for the diffusers based on the config of the LDM model. - """ - if image_size is not None: - deprecation_message = ( - "Configuring UNet2DConditionModel with the `image_size` argument to `from_single_file`" - "is deprecated and will be ignored in future versions." - ) - deprecate("image_size", "1.0.0", deprecation_message) - - image_size = set_image_size(checkpoint, image_size=image_size) - - if ( - "unet_config" in original_config["model"]["params"] - and original_config["model"]["params"]["unet_config"] is not None - ): - unet_params = original_config["model"]["params"]["unet_config"]["params"] - else: - unet_params = original_config["model"]["params"]["network_config"]["params"] - - if num_in_channels is not None: - deprecation_message = ( - "Configuring UNet2DConditionModel with the `num_in_channels` argument to `from_single_file`" - "is deprecated and will be ignored in future versions." - ) - deprecate("image_size", "1.0.0", deprecation_message) - in_channels = num_in_channels - else: - in_channels = unet_params["in_channels"] - - vae_params = original_config["model"]["params"]["first_stage_config"]["params"]["ddconfig"] - block_out_channels = [unet_params["model_channels"] * mult for mult in unet_params["channel_mult"]] - - down_block_types = [] - resolution = 1 - for i in range(len(block_out_channels)): - block_type = "CrossAttnDownBlock2D" if resolution in unet_params["attention_resolutions"] else "DownBlock2D" - down_block_types.append(block_type) - if i != len(block_out_channels) - 1: - resolution *= 2 - - up_block_types = [] - for i in range(len(block_out_channels)): - block_type = "CrossAttnUpBlock2D" if resolution in unet_params["attention_resolutions"] else "UpBlock2D" - up_block_types.append(block_type) - resolution //= 2 - - if unet_params["transformer_depth"] is not None: - transformer_layers_per_block = ( - unet_params["transformer_depth"] - if isinstance(unet_params["transformer_depth"], int) - else list(unet_params["transformer_depth"]) - ) - else: - transformer_layers_per_block = 1 - - vae_scale_factor = 2 ** (len(vae_params["ch_mult"]) - 1) - - head_dim = unet_params["num_heads"] if "num_heads" in unet_params else None - use_linear_projection = ( - unet_params["use_linear_in_transformer"] if "use_linear_in_transformer" in unet_params else False - ) - if use_linear_projection: - # stable diffusion 2-base-512 and 2-768 - if head_dim is None: - head_dim_mult = unet_params["model_channels"] // unet_params["num_head_channels"] - head_dim = [head_dim_mult * c for c in list(unet_params["channel_mult"])] - - class_embed_type = None - addition_embed_type = None - addition_time_embed_dim = None - projection_class_embeddings_input_dim = None - context_dim = None - - if unet_params["context_dim"] is not None: - context_dim = ( - unet_params["context_dim"] - if isinstance(unet_params["context_dim"], int) - else unet_params["context_dim"][0] - ) - - if "num_classes" in unet_params: - if unet_params["num_classes"] == "sequential": - if context_dim in [2048, 1280]: - # SDXL - addition_embed_type = "text_time" - addition_time_embed_dim = 256 - else: - class_embed_type = "projection" - assert "adm_in_channels" in unet_params - projection_class_embeddings_input_dim = unet_params["adm_in_channels"] - - config = { - "sample_size": image_size // vae_scale_factor, - "in_channels": in_channels, - "down_block_types": down_block_types, - "block_out_channels": block_out_channels, - "layers_per_block": unet_params["num_res_blocks"], - "cross_attention_dim": context_dim, - "attention_head_dim": head_dim, - "use_linear_projection": use_linear_projection, - "class_embed_type": class_embed_type, - "addition_embed_type": addition_embed_type, - "addition_time_embed_dim": addition_time_embed_dim, - "projection_class_embeddings_input_dim": projection_class_embeddings_input_dim, - "transformer_layers_per_block": transformer_layers_per_block, - } - - if upcast_attention is not None: - deprecation_message = ( - "Configuring UNet2DConditionModel with the `upcast_attention` argument to `from_single_file`" - "is deprecated and will be ignored in future versions." - ) - deprecate("image_size", "1.0.0", deprecation_message) - config["upcast_attention"] = upcast_attention - - if "disable_self_attentions" in unet_params: - config["only_cross_attention"] = unet_params["disable_self_attentions"] - - if "num_classes" in unet_params and isinstance(unet_params["num_classes"], int): - config["num_class_embeds"] = unet_params["num_classes"] - - config["out_channels"] = unet_params["out_channels"] - config["up_block_types"] = up_block_types - - return config - - -def create_controlnet_diffusers_config_from_ldm(original_config, checkpoint, image_size=None, **kwargs): - if image_size is not None: - deprecation_message = ( - "Configuring ControlNetModel with the `image_size` argument" - "is deprecated and will be ignored in future versions." - ) - deprecate("image_size", "1.0.0", deprecation_message) - - image_size = set_image_size(checkpoint, image_size=image_size) - - unet_params = original_config["model"]["params"]["control_stage_config"]["params"] - diffusers_unet_config = create_unet_diffusers_config_from_ldm(original_config, image_size=image_size) - - controlnet_config = { - "conditioning_channels": unet_params["hint_channels"], - "in_channels": diffusers_unet_config["in_channels"], - "down_block_types": diffusers_unet_config["down_block_types"], - "block_out_channels": diffusers_unet_config["block_out_channels"], - "layers_per_block": diffusers_unet_config["layers_per_block"], - "cross_attention_dim": diffusers_unet_config["cross_attention_dim"], - "attention_head_dim": diffusers_unet_config["attention_head_dim"], - "use_linear_projection": diffusers_unet_config["use_linear_projection"], - "class_embed_type": diffusers_unet_config["class_embed_type"], - "addition_embed_type": diffusers_unet_config["addition_embed_type"], - "addition_time_embed_dim": diffusers_unet_config["addition_time_embed_dim"], - "projection_class_embeddings_input_dim": diffusers_unet_config["projection_class_embeddings_input_dim"], - "transformer_layers_per_block": diffusers_unet_config["transformer_layers_per_block"], - } - - return controlnet_config - - -def create_vae_diffusers_config_from_ldm(original_config, checkpoint, image_size=None, scaling_factor=None): - """ - Creates a config for the diffusers based on the config of the LDM model. - """ - if image_size is not None: - deprecation_message = ( - "Configuring AutoencoderKL with the `image_size` argument" - "is deprecated and will be ignored in future versions." - ) - deprecate("image_size", "1.0.0", deprecation_message) - - image_size = set_image_size(checkpoint, image_size=image_size) - - if "edm_mean" in checkpoint and "edm_std" in checkpoint: - latents_mean = checkpoint["edm_mean"] - latents_std = checkpoint["edm_std"] - else: - latents_mean = None - latents_std = None - - vae_params = original_config["model"]["params"]["first_stage_config"]["params"]["ddconfig"] - if (scaling_factor is None) and (latents_mean is not None) and (latents_std is not None): - scaling_factor = PLAYGROUND_VAE_SCALING_FACTOR - - elif (scaling_factor is None) and ("scale_factor" in original_config["model"]["params"]): - scaling_factor = original_config["model"]["params"]["scale_factor"] - - elif scaling_factor is None: - scaling_factor = LDM_VAE_DEFAULT_SCALING_FACTOR - - block_out_channels = [vae_params["ch"] * mult for mult in vae_params["ch_mult"]] - down_block_types = ["DownEncoderBlock2D"] * len(block_out_channels) - up_block_types = ["UpDecoderBlock2D"] * len(block_out_channels) - - config = { - "sample_size": image_size, - "in_channels": vae_params["in_channels"], - "out_channels": vae_params["out_ch"], - "down_block_types": down_block_types, - "up_block_types": up_block_types, - "block_out_channels": block_out_channels, - "latent_channels": vae_params["z_channels"], - "layers_per_block": vae_params["num_res_blocks"], - "scaling_factor": scaling_factor, - } - if latents_mean is not None and latents_std is not None: - config.update({"latents_mean": latents_mean, "latents_std": latents_std}) - - return config - - -def update_unet_resnet_ldm_to_diffusers(ldm_keys, new_checkpoint, checkpoint, mapping=None): - for ldm_key in ldm_keys: - diffusers_key = ( - ldm_key.replace("in_layers.0", "norm1") - .replace("in_layers.2", "conv1") - .replace("out_layers.0", "norm2") - .replace("out_layers.3", "conv2") - .replace("emb_layers.1", "time_emb_proj") - .replace("skip_connection", "conv_shortcut") - ) - if mapping: - diffusers_key = diffusers_key.replace(mapping["old"], mapping["new"]) - new_checkpoint[diffusers_key] = checkpoint.get(ldm_key) - - -def update_unet_attention_ldm_to_diffusers(ldm_keys, new_checkpoint, checkpoint, mapping): - for ldm_key in ldm_keys: - diffusers_key = ldm_key.replace(mapping["old"], mapping["new"]) - new_checkpoint[diffusers_key] = checkpoint.get(ldm_key) - - -def update_vae_resnet_ldm_to_diffusers(keys, new_checkpoint, checkpoint, mapping): - for ldm_key in keys: - diffusers_key = ldm_key.replace(mapping["old"], mapping["new"]).replace("nin_shortcut", "conv_shortcut") - new_checkpoint[diffusers_key] = checkpoint.get(ldm_key) - - -def update_vae_attentions_ldm_to_diffusers(keys, new_checkpoint, checkpoint, mapping): - for ldm_key in keys: - diffusers_key = ( - ldm_key.replace(mapping["old"], mapping["new"]) - .replace("norm.weight", "group_norm.weight") - .replace("norm.bias", "group_norm.bias") - .replace("q.weight", "to_q.weight") - .replace("q.bias", "to_q.bias") - .replace("k.weight", "to_k.weight") - .replace("k.bias", "to_k.bias") - .replace("v.weight", "to_v.weight") - .replace("v.bias", "to_v.bias") - .replace("proj_out.weight", "to_out.0.weight") - .replace("proj_out.bias", "to_out.0.bias") - ) - new_checkpoint[diffusers_key] = checkpoint.get(ldm_key) - - # proj_attn.weight has to be converted from conv 1D to linear - shape = new_checkpoint[diffusers_key].shape - - if len(shape) == 3: - new_checkpoint[diffusers_key] = new_checkpoint[diffusers_key][:, :, 0] - elif len(shape) == 4: - new_checkpoint[diffusers_key] = new_checkpoint[diffusers_key][:, :, 0, 0] - - -def convert_stable_cascade_unet_single_file_to_diffusers(checkpoint, **kwargs): - is_stage_c = "clip_txt_mapper.weight" in checkpoint - - if is_stage_c: - state_dict = {} - for key in checkpoint.keys(): - if key.endswith("in_proj_weight"): - weights = checkpoint[key].chunk(3, 0) - state_dict[key.replace("attn.in_proj_weight", "to_q.weight")] = weights[0] - state_dict[key.replace("attn.in_proj_weight", "to_k.weight")] = weights[1] - state_dict[key.replace("attn.in_proj_weight", "to_v.weight")] = weights[2] - elif key.endswith("in_proj_bias"): - weights = checkpoint[key].chunk(3, 0) - state_dict[key.replace("attn.in_proj_bias", "to_q.bias")] = weights[0] - state_dict[key.replace("attn.in_proj_bias", "to_k.bias")] = weights[1] - state_dict[key.replace("attn.in_proj_bias", "to_v.bias")] = weights[2] - elif key.endswith("out_proj.weight"): - weights = checkpoint[key] - state_dict[key.replace("attn.out_proj.weight", "to_out.0.weight")] = weights - elif key.endswith("out_proj.bias"): - weights = checkpoint[key] - state_dict[key.replace("attn.out_proj.bias", "to_out.0.bias")] = weights - else: - state_dict[key] = checkpoint[key] - else: - state_dict = {} - for key in checkpoint.keys(): - if key.endswith("in_proj_weight"): - weights = checkpoint[key].chunk(3, 0) - state_dict[key.replace("attn.in_proj_weight", "to_q.weight")] = weights[0] - state_dict[key.replace("attn.in_proj_weight", "to_k.weight")] = weights[1] - state_dict[key.replace("attn.in_proj_weight", "to_v.weight")] = weights[2] - elif key.endswith("in_proj_bias"): - weights = checkpoint[key].chunk(3, 0) - state_dict[key.replace("attn.in_proj_bias", "to_q.bias")] = weights[0] - state_dict[key.replace("attn.in_proj_bias", "to_k.bias")] = weights[1] - state_dict[key.replace("attn.in_proj_bias", "to_v.bias")] = weights[2] - elif key.endswith("out_proj.weight"): - weights = checkpoint[key] - state_dict[key.replace("attn.out_proj.weight", "to_out.0.weight")] = weights - elif key.endswith("out_proj.bias"): - weights = checkpoint[key] - state_dict[key.replace("attn.out_proj.bias", "to_out.0.bias")] = weights - # rename clip_mapper to clip_txt_pooled_mapper - elif key.endswith("clip_mapper.weight"): - weights = checkpoint[key] - state_dict[key.replace("clip_mapper.weight", "clip_txt_pooled_mapper.weight")] = weights - elif key.endswith("clip_mapper.bias"): - weights = checkpoint[key] - state_dict[key.replace("clip_mapper.bias", "clip_txt_pooled_mapper.bias")] = weights - else: - state_dict[key] = checkpoint[key] - - return state_dict - - -def convert_ldm_unet_checkpoint(checkpoint, config, extract_ema=False, **kwargs): - """ - Takes a state dict and a config, and returns a converted checkpoint. - """ - # extract state_dict for UNet - unet_state_dict = {} - keys = list(checkpoint.keys()) - unet_key = LDM_UNET_KEY - - # at least a 100 parameters have to start with `model_ema` in order for the checkpoint to be EMA - if sum(k.startswith("model_ema") for k in keys) > 100 and extract_ema: - logger.warning("Checkpoint has both EMA and non-EMA weights.") - logger.warning( - "In this conversion only the EMA weights are extracted. If you want to instead extract the non-EMA" - " weights (useful to continue fine-tuning), please make sure to remove the `--extract_ema` flag." - ) - for key in keys: - if key.startswith("model.diffusion_model"): - flat_ema_key = "model_ema." + "".join(key.split(".")[1:]) - unet_state_dict[key.replace(unet_key, "")] = checkpoint.get(flat_ema_key) - else: - if sum(k.startswith("model_ema") for k in keys) > 100: - logger.warning( - "In this conversion only the non-EMA weights are extracted. If you want to instead extract the EMA" - " weights (usually better for inference), please make sure to add the `--extract_ema` flag." - ) - for key in keys: - if key.startswith(unet_key): - unet_state_dict[key.replace(unet_key, "")] = checkpoint.get(key) - - new_checkpoint = {} - ldm_unet_keys = DIFFUSERS_TO_LDM_MAPPING["unet"]["layers"] - for diffusers_key, ldm_key in ldm_unet_keys.items(): - if ldm_key not in unet_state_dict: - continue - new_checkpoint[diffusers_key] = unet_state_dict[ldm_key] - - if ("class_embed_type" in config) and (config["class_embed_type"] in ["timestep", "projection"]): - class_embed_keys = DIFFUSERS_TO_LDM_MAPPING["unet"]["class_embed_type"] - for diffusers_key, ldm_key in class_embed_keys.items(): - new_checkpoint[diffusers_key] = unet_state_dict[ldm_key] - - if ("addition_embed_type" in config) and (config["addition_embed_type"] == "text_time"): - addition_embed_keys = DIFFUSERS_TO_LDM_MAPPING["unet"]["addition_embed_type"] - for diffusers_key, ldm_key in addition_embed_keys.items(): - new_checkpoint[diffusers_key] = unet_state_dict[ldm_key] - - # Relevant to StableDiffusionUpscalePipeline - if "num_class_embeds" in config: - if (config["num_class_embeds"] is not None) and ("label_emb.weight" in unet_state_dict): - new_checkpoint["class_embedding.weight"] = unet_state_dict["label_emb.weight"] - - # Retrieves the keys for the input blocks only - num_input_blocks = len({".".join(layer.split(".")[:2]) for layer in unet_state_dict if "input_blocks" in layer}) - input_blocks = { - layer_id: [key for key in unet_state_dict if f"input_blocks.{layer_id}" in key] - for layer_id in range(num_input_blocks) - } - - # Retrieves the keys for the middle blocks only - num_middle_blocks = len({".".join(layer.split(".")[:2]) for layer in unet_state_dict if "middle_block" in layer}) - middle_blocks = { - layer_id: [key for key in unet_state_dict if f"middle_block.{layer_id}" in key] - for layer_id in range(num_middle_blocks) - } - - # Retrieves the keys for the output blocks only - num_output_blocks = len({".".join(layer.split(".")[:2]) for layer in unet_state_dict if "output_blocks" in layer}) - output_blocks = { - layer_id: [key for key in unet_state_dict if f"output_blocks.{layer_id}" in key] - for layer_id in range(num_output_blocks) - } - - # Down blocks - for i in range(1, num_input_blocks): - block_id = (i - 1) // (config["layers_per_block"] + 1) - layer_in_block_id = (i - 1) % (config["layers_per_block"] + 1) - - resnets = [ - key for key in input_blocks[i] if f"input_blocks.{i}.0" in key and f"input_blocks.{i}.0.op" not in key - ] - update_unet_resnet_ldm_to_diffusers( - resnets, - new_checkpoint, - unet_state_dict, - {"old": f"input_blocks.{i}.0", "new": f"down_blocks.{block_id}.resnets.{layer_in_block_id}"}, - ) - - if f"input_blocks.{i}.0.op.weight" in unet_state_dict: - new_checkpoint[f"down_blocks.{block_id}.downsamplers.0.conv.weight"] = unet_state_dict.get( - f"input_blocks.{i}.0.op.weight" - ) - new_checkpoint[f"down_blocks.{block_id}.downsamplers.0.conv.bias"] = unet_state_dict.get( - f"input_blocks.{i}.0.op.bias" - ) - - attentions = [key for key in input_blocks[i] if f"input_blocks.{i}.1" in key] - if attentions: - update_unet_attention_ldm_to_diffusers( - attentions, - new_checkpoint, - unet_state_dict, - {"old": f"input_blocks.{i}.1", "new": f"down_blocks.{block_id}.attentions.{layer_in_block_id}"}, - ) - - # Mid blocks - for key in middle_blocks.keys(): - diffusers_key = max(key - 1, 0) - if key % 2 == 0: - update_unet_resnet_ldm_to_diffusers( - middle_blocks[key], - new_checkpoint, - unet_state_dict, - mapping={"old": f"middle_block.{key}", "new": f"mid_block.resnets.{diffusers_key}"}, - ) - else: - update_unet_attention_ldm_to_diffusers( - middle_blocks[key], - new_checkpoint, - unet_state_dict, - mapping={"old": f"middle_block.{key}", "new": f"mid_block.attentions.{diffusers_key}"}, - ) - - # Up Blocks - for i in range(num_output_blocks): - block_id = i // (config["layers_per_block"] + 1) - layer_in_block_id = i % (config["layers_per_block"] + 1) - - resnets = [ - key for key in output_blocks[i] if f"output_blocks.{i}.0" in key and f"output_blocks.{i}.0.op" not in key - ] - update_unet_resnet_ldm_to_diffusers( - resnets, - new_checkpoint, - unet_state_dict, - {"old": f"output_blocks.{i}.0", "new": f"up_blocks.{block_id}.resnets.{layer_in_block_id}"}, - ) - - attentions = [ - key for key in output_blocks[i] if f"output_blocks.{i}.1" in key and f"output_blocks.{i}.1.conv" not in key - ] - if attentions: - update_unet_attention_ldm_to_diffusers( - attentions, - new_checkpoint, - unet_state_dict, - {"old": f"output_blocks.{i}.1", "new": f"up_blocks.{block_id}.attentions.{layer_in_block_id}"}, - ) - - if f"output_blocks.{i}.1.conv.weight" in unet_state_dict: - new_checkpoint[f"up_blocks.{block_id}.upsamplers.0.conv.weight"] = unet_state_dict[ - f"output_blocks.{i}.1.conv.weight" - ] - new_checkpoint[f"up_blocks.{block_id}.upsamplers.0.conv.bias"] = unet_state_dict[ - f"output_blocks.{i}.1.conv.bias" - ] - if f"output_blocks.{i}.2.conv.weight" in unet_state_dict: - new_checkpoint[f"up_blocks.{block_id}.upsamplers.0.conv.weight"] = unet_state_dict[ - f"output_blocks.{i}.2.conv.weight" - ] - new_checkpoint[f"up_blocks.{block_id}.upsamplers.0.conv.bias"] = unet_state_dict[ - f"output_blocks.{i}.2.conv.bias" - ] - - return new_checkpoint - - -def convert_controlnet_checkpoint( - checkpoint, - config, - **kwargs, -): - # Return checkpoint if it's already been converted - if "time_embedding.linear_1.weight" in checkpoint: - return checkpoint - # Some controlnet ckpt files are distributed independently from the rest of the - # model components i.e. https://huggingface.co/thibaud/controlnet-sd21/ - if "time_embed.0.weight" in checkpoint: - controlnet_state_dict = checkpoint - - else: - controlnet_state_dict = {} - keys = list(checkpoint.keys()) - controlnet_key = LDM_CONTROLNET_KEY - for key in keys: - if key.startswith(controlnet_key): - controlnet_state_dict[key.replace(controlnet_key, "")] = checkpoint.get(key) - - new_checkpoint = {} - ldm_controlnet_keys = DIFFUSERS_TO_LDM_MAPPING["controlnet"]["layers"] - for diffusers_key, ldm_key in ldm_controlnet_keys.items(): - if ldm_key not in controlnet_state_dict: - continue - new_checkpoint[diffusers_key] = controlnet_state_dict[ldm_key] - - # Retrieves the keys for the input blocks only - num_input_blocks = len( - {".".join(layer.split(".")[:2]) for layer in controlnet_state_dict if "input_blocks" in layer} - ) - input_blocks = { - layer_id: [key for key in controlnet_state_dict if f"input_blocks.{layer_id}" in key] - for layer_id in range(num_input_blocks) - } - - # Down blocks - for i in range(1, num_input_blocks): - block_id = (i - 1) // (config["layers_per_block"] + 1) - layer_in_block_id = (i - 1) % (config["layers_per_block"] + 1) - - resnets = [ - key for key in input_blocks[i] if f"input_blocks.{i}.0" in key and f"input_blocks.{i}.0.op" not in key - ] - update_unet_resnet_ldm_to_diffusers( - resnets, - new_checkpoint, - controlnet_state_dict, - {"old": f"input_blocks.{i}.0", "new": f"down_blocks.{block_id}.resnets.{layer_in_block_id}"}, - ) - - if f"input_blocks.{i}.0.op.weight" in controlnet_state_dict: - new_checkpoint[f"down_blocks.{block_id}.downsamplers.0.conv.weight"] = controlnet_state_dict.get( - f"input_blocks.{i}.0.op.weight" - ) - new_checkpoint[f"down_blocks.{block_id}.downsamplers.0.conv.bias"] = controlnet_state_dict.get( - f"input_blocks.{i}.0.op.bias" - ) - - attentions = [key for key in input_blocks[i] if f"input_blocks.{i}.1" in key] - if attentions: - update_unet_attention_ldm_to_diffusers( - attentions, - new_checkpoint, - controlnet_state_dict, - {"old": f"input_blocks.{i}.1", "new": f"down_blocks.{block_id}.attentions.{layer_in_block_id}"}, - ) - - # controlnet down blocks - for i in range(num_input_blocks): - new_checkpoint[f"controlnet_down_blocks.{i}.weight"] = controlnet_state_dict.get(f"zero_convs.{i}.0.weight") - new_checkpoint[f"controlnet_down_blocks.{i}.bias"] = controlnet_state_dict.get(f"zero_convs.{i}.0.bias") - - # Retrieves the keys for the middle blocks only - num_middle_blocks = len( - {".".join(layer.split(".")[:2]) for layer in controlnet_state_dict if "middle_block" in layer} - ) - middle_blocks = { - layer_id: [key for key in controlnet_state_dict if f"middle_block.{layer_id}" in key] - for layer_id in range(num_middle_blocks) - } - - # Mid blocks - for key in middle_blocks.keys(): - diffusers_key = max(key - 1, 0) - if key % 2 == 0: - update_unet_resnet_ldm_to_diffusers( - middle_blocks[key], - new_checkpoint, - controlnet_state_dict, - mapping={"old": f"middle_block.{key}", "new": f"mid_block.resnets.{diffusers_key}"}, - ) - else: - update_unet_attention_ldm_to_diffusers( - middle_blocks[key], - new_checkpoint, - controlnet_state_dict, - mapping={"old": f"middle_block.{key}", "new": f"mid_block.attentions.{diffusers_key}"}, - ) - - # mid block - new_checkpoint["controlnet_mid_block.weight"] = controlnet_state_dict.get("middle_block_out.0.weight") - new_checkpoint["controlnet_mid_block.bias"] = controlnet_state_dict.get("middle_block_out.0.bias") - - # controlnet cond embedding blocks - cond_embedding_blocks = { - ".".join(layer.split(".")[:2]) - for layer in controlnet_state_dict - if "input_hint_block" in layer and ("input_hint_block.0" not in layer) and ("input_hint_block.14" not in layer) - } - num_cond_embedding_blocks = len(cond_embedding_blocks) - - for idx in range(1, num_cond_embedding_blocks + 1): - diffusers_idx = idx - 1 - cond_block_id = 2 * idx - - new_checkpoint[f"controlnet_cond_embedding.blocks.{diffusers_idx}.weight"] = controlnet_state_dict.get( - f"input_hint_block.{cond_block_id}.weight" - ) - new_checkpoint[f"controlnet_cond_embedding.blocks.{diffusers_idx}.bias"] = controlnet_state_dict.get( - f"input_hint_block.{cond_block_id}.bias" - ) - - return new_checkpoint - - -def convert_ldm_vae_checkpoint(checkpoint, config): - # extract state dict for VAE - # remove the LDM_VAE_KEY prefix from the ldm checkpoint keys so that it is easier to map them to diffusers keys - vae_state_dict = {} - keys = list(checkpoint.keys()) - vae_key = "" - for ldm_vae_key in LDM_VAE_KEYS: - if any(k.startswith(ldm_vae_key) for k in keys): - vae_key = ldm_vae_key - - for key in keys: - if key.startswith(vae_key): - vae_state_dict[key.replace(vae_key, "")] = checkpoint.get(key) - - new_checkpoint = {} - vae_diffusers_ldm_map = DIFFUSERS_TO_LDM_MAPPING["vae"] - for diffusers_key, ldm_key in vae_diffusers_ldm_map.items(): - if ldm_key not in vae_state_dict: - continue - new_checkpoint[diffusers_key] = vae_state_dict[ldm_key] - - # Retrieves the keys for the encoder down blocks only - num_down_blocks = len(config["down_block_types"]) - down_blocks = { - layer_id: [key for key in vae_state_dict if f"down.{layer_id}" in key] for layer_id in range(num_down_blocks) - } - - for i in range(num_down_blocks): - resnets = [key for key in down_blocks[i] if f"down.{i}" in key and f"down.{i}.downsample" not in key] - update_vae_resnet_ldm_to_diffusers( - resnets, - new_checkpoint, - vae_state_dict, - mapping={"old": f"down.{i}.block", "new": f"down_blocks.{i}.resnets"}, - ) - if f"encoder.down.{i}.downsample.conv.weight" in vae_state_dict: - new_checkpoint[f"encoder.down_blocks.{i}.downsamplers.0.conv.weight"] = vae_state_dict.get( - f"encoder.down.{i}.downsample.conv.weight" - ) - new_checkpoint[f"encoder.down_blocks.{i}.downsamplers.0.conv.bias"] = vae_state_dict.get( - f"encoder.down.{i}.downsample.conv.bias" - ) - - mid_resnets = [key for key in vae_state_dict if "encoder.mid.block" in key] - num_mid_res_blocks = 2 - for i in range(1, num_mid_res_blocks + 1): - resnets = [key for key in mid_resnets if f"encoder.mid.block_{i}" in key] - update_vae_resnet_ldm_to_diffusers( - resnets, - new_checkpoint, - vae_state_dict, - mapping={"old": f"mid.block_{i}", "new": f"mid_block.resnets.{i - 1}"}, - ) - - mid_attentions = [key for key in vae_state_dict if "encoder.mid.attn" in key] - update_vae_attentions_ldm_to_diffusers( - mid_attentions, new_checkpoint, vae_state_dict, mapping={"old": "mid.attn_1", "new": "mid_block.attentions.0"} - ) - - # Retrieves the keys for the decoder up blocks only - num_up_blocks = len(config["up_block_types"]) - up_blocks = { - layer_id: [key for key in vae_state_dict if f"up.{layer_id}" in key] for layer_id in range(num_up_blocks) - } - - for i in range(num_up_blocks): - block_id = num_up_blocks - 1 - i - resnets = [ - key for key in up_blocks[block_id] if f"up.{block_id}" in key and f"up.{block_id}.upsample" not in key - ] - update_vae_resnet_ldm_to_diffusers( - resnets, - new_checkpoint, - vae_state_dict, - mapping={"old": f"up.{block_id}.block", "new": f"up_blocks.{i}.resnets"}, - ) - if f"decoder.up.{block_id}.upsample.conv.weight" in vae_state_dict: - new_checkpoint[f"decoder.up_blocks.{i}.upsamplers.0.conv.weight"] = vae_state_dict[ - f"decoder.up.{block_id}.upsample.conv.weight" - ] - new_checkpoint[f"decoder.up_blocks.{i}.upsamplers.0.conv.bias"] = vae_state_dict[ - f"decoder.up.{block_id}.upsample.conv.bias" - ] - - mid_resnets = [key for key in vae_state_dict if "decoder.mid.block" in key] - num_mid_res_blocks = 2 - for i in range(1, num_mid_res_blocks + 1): - resnets = [key for key in mid_resnets if f"decoder.mid.block_{i}" in key] - update_vae_resnet_ldm_to_diffusers( - resnets, - new_checkpoint, - vae_state_dict, - mapping={"old": f"mid.block_{i}", "new": f"mid_block.resnets.{i - 1}"}, - ) - - mid_attentions = [key for key in vae_state_dict if "decoder.mid.attn" in key] - update_vae_attentions_ldm_to_diffusers( - mid_attentions, new_checkpoint, vae_state_dict, mapping={"old": "mid.attn_1", "new": "mid_block.attentions.0"} - ) - conv_attn_to_linear(new_checkpoint) - - return new_checkpoint - - -def convert_ldm_clip_checkpoint(checkpoint, remove_prefix=None): - keys = list(checkpoint.keys()) - text_model_dict = {} - - remove_prefixes = [] - remove_prefixes.extend(LDM_CLIP_PREFIX_TO_REMOVE) - if remove_prefix: - remove_prefixes.append(remove_prefix) - - for key in keys: - for prefix in remove_prefixes: - if key.startswith(prefix): - diffusers_key = key.replace(prefix, "") - text_model_dict[diffusers_key] = checkpoint.get(key) - - return text_model_dict - - -def convert_open_clip_checkpoint( - text_model, - checkpoint, - prefix="cond_stage_model.model.", -): - text_model_dict = {} - text_proj_key = prefix + "text_projection" - - if text_proj_key in checkpoint: - text_proj_dim = int(checkpoint[text_proj_key].shape[0]) - elif hasattr(text_model.config, "hidden_size"): - text_proj_dim = text_model.config.hidden_size - else: - text_proj_dim = LDM_OPEN_CLIP_TEXT_PROJECTION_DIM - - keys = list(checkpoint.keys()) - keys_to_ignore = SD_2_TEXT_ENCODER_KEYS_TO_IGNORE - - openclip_diffusers_ldm_map = DIFFUSERS_TO_LDM_MAPPING["openclip"]["layers"] - for diffusers_key, ldm_key in openclip_diffusers_ldm_map.items(): - ldm_key = prefix + ldm_key - if ldm_key not in checkpoint: - continue - if ldm_key in keys_to_ignore: - continue - if ldm_key.endswith("text_projection"): - text_model_dict[diffusers_key] = checkpoint[ldm_key].T.contiguous() - else: - text_model_dict[diffusers_key] = checkpoint[ldm_key] - - for key in keys: - if key in keys_to_ignore: - continue - - if not key.startswith(prefix + "transformer."): - continue - - diffusers_key = key.replace(prefix + "transformer.", "") - transformer_diffusers_to_ldm_map = DIFFUSERS_TO_LDM_MAPPING["openclip"]["transformer"] - for new_key, old_key in transformer_diffusers_to_ldm_map.items(): - diffusers_key = ( - diffusers_key.replace(old_key, new_key).replace(".in_proj_weight", "").replace(".in_proj_bias", "") - ) - - if key.endswith(".in_proj_weight"): - weight_value = checkpoint.get(key) - - text_model_dict[diffusers_key + ".q_proj.weight"] = weight_value[:text_proj_dim, :].clone().detach() - text_model_dict[diffusers_key + ".k_proj.weight"] = ( - weight_value[text_proj_dim : text_proj_dim * 2, :].clone().detach() - ) - text_model_dict[diffusers_key + ".v_proj.weight"] = weight_value[text_proj_dim * 2 :, :].clone().detach() - - elif key.endswith(".in_proj_bias"): - weight_value = checkpoint.get(key) - text_model_dict[diffusers_key + ".q_proj.bias"] = weight_value[:text_proj_dim].clone().detach() - text_model_dict[diffusers_key + ".k_proj.bias"] = ( - weight_value[text_proj_dim : text_proj_dim * 2].clone().detach() - ) - text_model_dict[diffusers_key + ".v_proj.bias"] = weight_value[text_proj_dim * 2 :].clone().detach() - else: - text_model_dict[diffusers_key] = checkpoint.get(key) - - return text_model_dict - - -def create_diffusers_clip_model_from_ldm( - cls, - checkpoint, - subfolder="", - config=None, - torch_dtype=None, - local_files_only=None, - is_legacy_loading=False, -): - if config: - config = {"pretrained_model_name_or_path": config} - else: - config = fetch_diffusers_config(checkpoint) - - # For backwards compatibility - # Older versions of `from_single_file` expected CLIP configs to be placed in their original transformers model repo - # in the cache_dir, rather than in a subfolder of the Diffusers model - if is_legacy_loading: - logger.warning( - ( - "Detected legacy CLIP loading behavior. Please run `from_single_file` with `local_files_only=False once to update " - "the local cache directory with the necessary CLIP model config files. " - "Attempting to load CLIP model from legacy cache directory." - ) - ) - - if is_clip_model(checkpoint) or is_clip_sdxl_model(checkpoint): - clip_config = "openai/clip-vit-large-patch14" - config["pretrained_model_name_or_path"] = clip_config - subfolder = "" - - elif is_open_clip_model(checkpoint): - clip_config = "stabilityai/stable-diffusion-2" - config["pretrained_model_name_or_path"] = clip_config - subfolder = "text_encoder" - - else: - clip_config = "laion/CLIP-ViT-bigG-14-laion2B-39B-b160k" - config["pretrained_model_name_or_path"] = clip_config - subfolder = "" - - model_config = cls.config_class.from_pretrained(**config, subfolder=subfolder, local_files_only=local_files_only) - ctx = init_empty_weights if is_accelerate_available() else nullcontext - with ctx(): - model = cls(model_config) - - # `CLIPTextModel` was flattened in transformers >=5.6; `CLIPTextModelWithProjection` still wraps via `text_model`. - has_text_model_wrapper = hasattr(model, "text_model") - text_model = model.text_model if has_text_model_wrapper else model - position_embedding_dim = text_model.embeddings.position_embedding.weight.shape[-1] - - if is_clip_model(checkpoint): - diffusers_format_checkpoint = convert_ldm_clip_checkpoint(checkpoint) - - elif ( - is_clip_sdxl_model(checkpoint) - and checkpoint[CHECKPOINT_KEY_NAMES["clip_sdxl"]].shape[-1] == position_embedding_dim - ): - diffusers_format_checkpoint = convert_ldm_clip_checkpoint(checkpoint) - - elif ( - is_clip_sd3_model(checkpoint) - and checkpoint[CHECKPOINT_KEY_NAMES["clip_sd3"]].shape[-1] == position_embedding_dim - ): - diffusers_format_checkpoint = convert_ldm_clip_checkpoint(checkpoint, "text_encoders.clip_l.transformer.") - diffusers_format_checkpoint["text_projection.weight"] = torch.eye(position_embedding_dim) - - elif is_open_clip_model(checkpoint): - prefix = "cond_stage_model.model." - diffusers_format_checkpoint = convert_open_clip_checkpoint(model, checkpoint, prefix=prefix) - - elif ( - is_open_clip_sdxl_model(checkpoint) - and checkpoint[CHECKPOINT_KEY_NAMES["open_clip_sdxl"]].shape[-1] == position_embedding_dim - ): - prefix = "conditioner.embedders.1.model." - diffusers_format_checkpoint = convert_open_clip_checkpoint(model, checkpoint, prefix=prefix) - - elif is_open_clip_sdxl_refiner_model(checkpoint): - prefix = "conditioner.embedders.0.model." - diffusers_format_checkpoint = convert_open_clip_checkpoint(model, checkpoint, prefix=prefix) - - elif ( - is_open_clip_sd3_model(checkpoint) - and checkpoint[CHECKPOINT_KEY_NAMES["open_clip_sd3"]].shape[-1] == position_embedding_dim - ): - diffusers_format_checkpoint = convert_ldm_clip_checkpoint(checkpoint, "text_encoders.clip_g.transformer.") - - else: - raise ValueError("The provided checkpoint does not seem to contain a valid CLIP model.") - - if not has_text_model_wrapper: - diffusers_format_checkpoint = { - k.removeprefix("text_model."): v for k, v in diffusers_format_checkpoint.items() - } - - if is_accelerate_available(): - load_model_dict_into_meta(model, diffusers_format_checkpoint, dtype=torch_dtype) - empty_device_cache() - else: - model.load_state_dict(diffusers_format_checkpoint, strict=False) - - if torch_dtype is not None: - model.to(torch_dtype) - - model.eval() - - return model - - -def _legacy_load_scheduler( - cls, - checkpoint, - component_name, - original_config=None, - **kwargs, -): - scheduler_type = kwargs.get("scheduler_type", None) - prediction_type = kwargs.get("prediction_type", None) - - if scheduler_type is not None: - deprecation_message = ( - "Please pass an instance of a Scheduler object directly to the `scheduler` argument in `from_single_file`\n\n" - "Example:\n\n" - "from diffusers import StableDiffusionPipeline, DDIMScheduler\n\n" - "scheduler = DDIMScheduler()\n" - "pipe = StableDiffusionPipeline.from_single_file(, scheduler=scheduler)\n" - ) - deprecate("scheduler_type", "1.0.0", deprecation_message) - - if prediction_type is not None: - deprecation_message = ( - "Please configure an instance of a Scheduler with the appropriate `prediction_type` and " - "pass the object directly to the `scheduler` argument in `from_single_file`.\n\n" - "Example:\n\n" - "from diffusers import StableDiffusionPipeline, DDIMScheduler\n\n" - 'scheduler = DDIMScheduler(prediction_type="v_prediction")\n' - "pipe = StableDiffusionPipeline.from_single_file(, scheduler=scheduler)\n" - ) - deprecate("prediction_type", "1.0.0", deprecation_message) - - scheduler_config = SCHEDULER_DEFAULT_CONFIG - model_type = infer_diffusers_model_type(checkpoint=checkpoint) - - global_step = checkpoint["global_step"] if "global_step" in checkpoint else None - - if original_config: - num_train_timesteps = getattr(original_config["model"]["params"], "timesteps", 1000) - else: - num_train_timesteps = 1000 - - scheduler_config["num_train_timesteps"] = num_train_timesteps - - if model_type == "v2": - if prediction_type is None: - # NOTE: For stable diffusion 2 base it is recommended to pass `prediction_type=="epsilon"` # as it relies on a brittle global step parameter here - prediction_type = "epsilon" if global_step == 875000 else "v_prediction" - - else: - prediction_type = prediction_type or "epsilon" - - scheduler_config["prediction_type"] = prediction_type - - if model_type in ["xl_base", "xl_refiner"]: - scheduler_type = "euler" - elif model_type == "playground": - scheduler_type = "edm_dpm_solver_multistep" - else: - if original_config: - beta_start = original_config["model"]["params"].get("linear_start") - beta_end = original_config["model"]["params"].get("linear_end") - - else: - beta_start = 0.02 - beta_end = 0.085 - - scheduler_config["beta_start"] = beta_start - scheduler_config["beta_end"] = beta_end - scheduler_config["beta_schedule"] = "scaled_linear" - scheduler_config["clip_sample"] = False - scheduler_config["set_alpha_to_one"] = False - - # to deal with an edge case StableDiffusionUpscale pipeline has two schedulers - if component_name == "low_res_scheduler": - return cls.from_config( - { - "beta_end": 0.02, - "beta_schedule": "scaled_linear", - "beta_start": 0.0001, - "clip_sample": True, - "num_train_timesteps": 1000, - "prediction_type": "epsilon", - "trained_betas": None, - "variance_type": "fixed_small", - } - ) - - if scheduler_type is None: - return cls.from_config(scheduler_config) - - elif scheduler_type == "pndm": - scheduler_config["skip_prk_steps"] = True - scheduler = PNDMScheduler.from_config(scheduler_config) - - elif scheduler_type == "lms": - scheduler = LMSDiscreteScheduler.from_config(scheduler_config) - - elif scheduler_type == "heun": - scheduler = HeunDiscreteScheduler.from_config(scheduler_config) - - elif scheduler_type == "euler": - scheduler = EulerDiscreteScheduler.from_config(scheduler_config) - - elif scheduler_type == "euler-ancestral": - scheduler = EulerAncestralDiscreteScheduler.from_config(scheduler_config) - - elif scheduler_type == "dpm": - scheduler = DPMSolverMultistepScheduler.from_config(scheduler_config) - - elif scheduler_type == "ddim": - scheduler = DDIMScheduler.from_config(scheduler_config) - - elif scheduler_type == "edm_dpm_solver_multistep": - scheduler_config = { - "algorithm_type": "dpmsolver++", - "dynamic_thresholding_ratio": 0.995, - "euler_at_final": False, - "final_sigmas_type": "zero", - "lower_order_final": True, - "num_train_timesteps": 1000, - "prediction_type": "epsilon", - "rho": 7.0, - "sample_max_value": 1.0, - "sigma_data": 0.5, - "sigma_max": 80.0, - "sigma_min": 0.002, - "solver_order": 2, - "solver_type": "midpoint", - "thresholding": False, - } - scheduler = EDMDPMSolverMultistepScheduler(**scheduler_config) - - else: - raise ValueError(f"Scheduler of type {scheduler_type} doesn't exist!") - - return scheduler - - -def _legacy_load_clip_tokenizer(cls, checkpoint, config=None, local_files_only=False): - if config: - config = {"pretrained_model_name_or_path": config} - else: - config = fetch_diffusers_config(checkpoint) - - if is_clip_model(checkpoint) or is_clip_sdxl_model(checkpoint): - clip_config = "openai/clip-vit-large-patch14" - config["pretrained_model_name_or_path"] = clip_config - subfolder = "" - - elif is_open_clip_model(checkpoint): - clip_config = "stabilityai/stable-diffusion-2" - config["pretrained_model_name_or_path"] = clip_config - subfolder = "tokenizer" - - else: - clip_config = "laion/CLIP-ViT-bigG-14-laion2B-39B-b160k" - config["pretrained_model_name_or_path"] = clip_config - subfolder = "" - - tokenizer = cls.from_pretrained(**config, subfolder=subfolder, local_files_only=local_files_only) - - return tokenizer - - -def _legacy_load_safety_checker(local_files_only, torch_dtype): - # Support for loading safety checker components using the deprecated - # `load_safety_checker` argument. - - from ..pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker - - feature_extractor = AutoImageProcessor.from_pretrained( - "CompVis/stable-diffusion-safety-checker", local_files_only=local_files_only, torch_dtype=torch_dtype - ) - safety_checker = StableDiffusionSafetyChecker.from_pretrained( - "CompVis/stable-diffusion-safety-checker", local_files_only=local_files_only, torch_dtype=torch_dtype - ) - - return {"safety_checker": safety_checker, "feature_extractor": feature_extractor} - - -# in SD3 original implementation of AdaLayerNormContinuous, it split linear projection output into shift, scale; -# while in diffusers it split into scale, shift. Here we swap the linear projection weights in order to be able to use diffusers implementation -def swap_scale_shift(weight, dim): - shift, scale = weight.chunk(2, dim=0) - new_weight = torch.cat([scale, shift], dim=0) - return new_weight - - -def swap_proj_gate(weight): - proj, gate = weight.chunk(2, dim=0) - new_weight = torch.cat([gate, proj], dim=0) - return new_weight - - -def get_attn2_layers(state_dict): - attn2_layers = [] - for key in state_dict.keys(): - if "attn2." in key: - # Extract the layer number from the key - layer_num = int(key.split(".")[1]) - attn2_layers.append(layer_num) - - return tuple(sorted(set(attn2_layers))) - - -def get_caption_projection_dim(state_dict): - caption_projection_dim = state_dict["context_embedder.weight"].shape[0] - return caption_projection_dim - - -def convert_sd3_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {} - keys = list(checkpoint.keys()) - for k in keys: - if "model.diffusion_model." in k: - checkpoint[k.replace("model.diffusion_model.", "")] = checkpoint.pop(k) - - num_layers = list(set(int(k.split(".", 2)[1]) for k in checkpoint if "joint_blocks" in k))[-1] + 1 # noqa: C401 - dual_attention_layers = get_attn2_layers(checkpoint) - - caption_projection_dim = get_caption_projection_dim(checkpoint) - has_qk_norm = any("ln_q" in key for key in checkpoint.keys()) - - # Positional and patch embeddings. - converted_state_dict["pos_embed.pos_embed"] = checkpoint.pop("pos_embed") - converted_state_dict["pos_embed.proj.weight"] = checkpoint.pop("x_embedder.proj.weight") - converted_state_dict["pos_embed.proj.bias"] = checkpoint.pop("x_embedder.proj.bias") - - # Timestep embeddings. - converted_state_dict["time_text_embed.timestep_embedder.linear_1.weight"] = checkpoint.pop( - "t_embedder.mlp.0.weight" - ) - converted_state_dict["time_text_embed.timestep_embedder.linear_1.bias"] = checkpoint.pop("t_embedder.mlp.0.bias") - converted_state_dict["time_text_embed.timestep_embedder.linear_2.weight"] = checkpoint.pop( - "t_embedder.mlp.2.weight" - ) - converted_state_dict["time_text_embed.timestep_embedder.linear_2.bias"] = checkpoint.pop("t_embedder.mlp.2.bias") - - # Context projections. - converted_state_dict["context_embedder.weight"] = checkpoint.pop("context_embedder.weight") - converted_state_dict["context_embedder.bias"] = checkpoint.pop("context_embedder.bias") - - # Pooled context projection. - converted_state_dict["time_text_embed.text_embedder.linear_1.weight"] = checkpoint.pop("y_embedder.mlp.0.weight") - converted_state_dict["time_text_embed.text_embedder.linear_1.bias"] = checkpoint.pop("y_embedder.mlp.0.bias") - converted_state_dict["time_text_embed.text_embedder.linear_2.weight"] = checkpoint.pop("y_embedder.mlp.2.weight") - converted_state_dict["time_text_embed.text_embedder.linear_2.bias"] = checkpoint.pop("y_embedder.mlp.2.bias") - - # Transformer blocks 🎸. - for i in range(num_layers): - # Q, K, V - sample_q, sample_k, sample_v = torch.chunk( - checkpoint.pop(f"joint_blocks.{i}.x_block.attn.qkv.weight"), 3, dim=0 - ) - context_q, context_k, context_v = torch.chunk( - checkpoint.pop(f"joint_blocks.{i}.context_block.attn.qkv.weight"), 3, dim=0 - ) - sample_q_bias, sample_k_bias, sample_v_bias = torch.chunk( - checkpoint.pop(f"joint_blocks.{i}.x_block.attn.qkv.bias"), 3, dim=0 - ) - context_q_bias, context_k_bias, context_v_bias = torch.chunk( - checkpoint.pop(f"joint_blocks.{i}.context_block.attn.qkv.bias"), 3, dim=0 - ) - - converted_state_dict[f"transformer_blocks.{i}.attn.to_q.weight"] = torch.cat([sample_q]) - converted_state_dict[f"transformer_blocks.{i}.attn.to_q.bias"] = torch.cat([sample_q_bias]) - converted_state_dict[f"transformer_blocks.{i}.attn.to_k.weight"] = torch.cat([sample_k]) - converted_state_dict[f"transformer_blocks.{i}.attn.to_k.bias"] = torch.cat([sample_k_bias]) - converted_state_dict[f"transformer_blocks.{i}.attn.to_v.weight"] = torch.cat([sample_v]) - converted_state_dict[f"transformer_blocks.{i}.attn.to_v.bias"] = torch.cat([sample_v_bias]) - - converted_state_dict[f"transformer_blocks.{i}.attn.add_q_proj.weight"] = torch.cat([context_q]) - converted_state_dict[f"transformer_blocks.{i}.attn.add_q_proj.bias"] = torch.cat([context_q_bias]) - converted_state_dict[f"transformer_blocks.{i}.attn.add_k_proj.weight"] = torch.cat([context_k]) - converted_state_dict[f"transformer_blocks.{i}.attn.add_k_proj.bias"] = torch.cat([context_k_bias]) - converted_state_dict[f"transformer_blocks.{i}.attn.add_v_proj.weight"] = torch.cat([context_v]) - converted_state_dict[f"transformer_blocks.{i}.attn.add_v_proj.bias"] = torch.cat([context_v_bias]) - - # qk norm - if has_qk_norm: - converted_state_dict[f"transformer_blocks.{i}.attn.norm_q.weight"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.attn.ln_q.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.attn.norm_k.weight"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.attn.ln_k.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.attn.norm_added_q.weight"] = checkpoint.pop( - f"joint_blocks.{i}.context_block.attn.ln_q.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.attn.norm_added_k.weight"] = checkpoint.pop( - f"joint_blocks.{i}.context_block.attn.ln_k.weight" - ) - - # output projections. - converted_state_dict[f"transformer_blocks.{i}.attn.to_out.0.weight"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.attn.proj.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.attn.to_out.0.bias"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.attn.proj.bias" - ) - if not (i == num_layers - 1): - converted_state_dict[f"transformer_blocks.{i}.attn.to_add_out.weight"] = checkpoint.pop( - f"joint_blocks.{i}.context_block.attn.proj.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.attn.to_add_out.bias"] = checkpoint.pop( - f"joint_blocks.{i}.context_block.attn.proj.bias" - ) - - if i in dual_attention_layers: - # Q, K, V - sample_q2, sample_k2, sample_v2 = torch.chunk( - checkpoint.pop(f"joint_blocks.{i}.x_block.attn2.qkv.weight"), 3, dim=0 - ) - sample_q2_bias, sample_k2_bias, sample_v2_bias = torch.chunk( - checkpoint.pop(f"joint_blocks.{i}.x_block.attn2.qkv.bias"), 3, dim=0 - ) - converted_state_dict[f"transformer_blocks.{i}.attn2.to_q.weight"] = torch.cat([sample_q2]) - converted_state_dict[f"transformer_blocks.{i}.attn2.to_q.bias"] = torch.cat([sample_q2_bias]) - converted_state_dict[f"transformer_blocks.{i}.attn2.to_k.weight"] = torch.cat([sample_k2]) - converted_state_dict[f"transformer_blocks.{i}.attn2.to_k.bias"] = torch.cat([sample_k2_bias]) - converted_state_dict[f"transformer_blocks.{i}.attn2.to_v.weight"] = torch.cat([sample_v2]) - converted_state_dict[f"transformer_blocks.{i}.attn2.to_v.bias"] = torch.cat([sample_v2_bias]) - - # qk norm - if has_qk_norm: - converted_state_dict[f"transformer_blocks.{i}.attn2.norm_q.weight"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.attn2.ln_q.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.attn2.norm_k.weight"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.attn2.ln_k.weight" - ) - - # output projections. - converted_state_dict[f"transformer_blocks.{i}.attn2.to_out.0.weight"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.attn2.proj.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.attn2.to_out.0.bias"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.attn2.proj.bias" - ) - - # norms. - converted_state_dict[f"transformer_blocks.{i}.norm1.linear.weight"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.adaLN_modulation.1.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.norm1.linear.bias"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.adaLN_modulation.1.bias" - ) - if not (i == num_layers - 1): - converted_state_dict[f"transformer_blocks.{i}.norm1_context.linear.weight"] = checkpoint.pop( - f"joint_blocks.{i}.context_block.adaLN_modulation.1.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.norm1_context.linear.bias"] = checkpoint.pop( - f"joint_blocks.{i}.context_block.adaLN_modulation.1.bias" - ) - else: - converted_state_dict[f"transformer_blocks.{i}.norm1_context.linear.weight"] = swap_scale_shift( - checkpoint.pop(f"joint_blocks.{i}.context_block.adaLN_modulation.1.weight"), - dim=caption_projection_dim, - ) - converted_state_dict[f"transformer_blocks.{i}.norm1_context.linear.bias"] = swap_scale_shift( - checkpoint.pop(f"joint_blocks.{i}.context_block.adaLN_modulation.1.bias"), - dim=caption_projection_dim, - ) - - # ffs. - converted_state_dict[f"transformer_blocks.{i}.ff.net.0.proj.weight"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.mlp.fc1.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.ff.net.0.proj.bias"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.mlp.fc1.bias" - ) - converted_state_dict[f"transformer_blocks.{i}.ff.net.2.weight"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.mlp.fc2.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.ff.net.2.bias"] = checkpoint.pop( - f"joint_blocks.{i}.x_block.mlp.fc2.bias" - ) - if not (i == num_layers - 1): - converted_state_dict[f"transformer_blocks.{i}.ff_context.net.0.proj.weight"] = checkpoint.pop( - f"joint_blocks.{i}.context_block.mlp.fc1.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.ff_context.net.0.proj.bias"] = checkpoint.pop( - f"joint_blocks.{i}.context_block.mlp.fc1.bias" - ) - converted_state_dict[f"transformer_blocks.{i}.ff_context.net.2.weight"] = checkpoint.pop( - f"joint_blocks.{i}.context_block.mlp.fc2.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.ff_context.net.2.bias"] = checkpoint.pop( - f"joint_blocks.{i}.context_block.mlp.fc2.bias" - ) - - # Final blocks. - converted_state_dict["proj_out.weight"] = checkpoint.pop("final_layer.linear.weight") - converted_state_dict["proj_out.bias"] = checkpoint.pop("final_layer.linear.bias") - converted_state_dict["norm_out.linear.weight"] = swap_scale_shift( - checkpoint.pop("final_layer.adaLN_modulation.1.weight"), dim=caption_projection_dim - ) - converted_state_dict["norm_out.linear.bias"] = swap_scale_shift( - checkpoint.pop("final_layer.adaLN_modulation.1.bias"), dim=caption_projection_dim - ) - - return converted_state_dict - - -def is_t5_in_single_file(checkpoint): - if "text_encoders.t5xxl.transformer.shared.weight" in checkpoint: - return True - - return False - - -def convert_sd3_t5_checkpoint_to_diffusers(checkpoint): - keys = list(checkpoint.keys()) - text_model_dict = {} - - remove_prefixes = ["text_encoders.t5xxl.transformer."] - - for key in keys: - for prefix in remove_prefixes: - if key.startswith(prefix): - diffusers_key = key.replace(prefix, "") - text_model_dict[diffusers_key] = checkpoint.get(key) - - return text_model_dict - - -def create_diffusers_t5_model_from_checkpoint( - cls, - checkpoint, - subfolder="", - config=None, - torch_dtype=None, - local_files_only=None, -): - if config: - config = {"pretrained_model_name_or_path": config} - else: - config = fetch_diffusers_config(checkpoint) - - model_config = cls.config_class.from_pretrained(**config, subfolder=subfolder, local_files_only=local_files_only) - ctx = init_empty_weights if is_accelerate_available() else nullcontext - with ctx(): - model = cls(model_config) - - diffusers_format_checkpoint = convert_sd3_t5_checkpoint_to_diffusers(checkpoint) - - if is_accelerate_available(): - load_model_dict_into_meta(model, diffusers_format_checkpoint, dtype=torch_dtype) - empty_device_cache() - else: - model.load_state_dict(diffusers_format_checkpoint) - - use_keep_in_fp32_modules = (cls._keep_in_fp32_modules is not None) and (torch_dtype == torch.float16) - if use_keep_in_fp32_modules: - keep_in_fp32_modules = model._keep_in_fp32_modules - else: - keep_in_fp32_modules = [] - - if keep_in_fp32_modules is not None: - for name, param in model.named_parameters(): - if any(module_to_keep_in_fp32 in name.split(".") for module_to_keep_in_fp32 in keep_in_fp32_modules): - # param = param.to(torch.float32) does not work here as only in the local scope. - param.data = param.data.to(torch.float32) - - return model - - -def convert_animatediff_checkpoint_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {} - for k, v in checkpoint.items(): - if "pos_encoder" in k: - continue - - else: - converted_state_dict[ - k.replace(".norms.0", ".norm1") - .replace(".norms.1", ".norm2") - .replace(".ff_norm", ".norm3") - .replace(".attention_blocks.0", ".attn1") - .replace(".attention_blocks.1", ".attn2") - .replace(".temporal_transformer", "") - ] = v - - return converted_state_dict - - -def convert_flux_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {} - keys = list(checkpoint.keys()) - - for k in keys: - if "model.diffusion_model." in k: - checkpoint[k.replace("model.diffusion_model.", "")] = checkpoint.pop(k) - - num_layers = list(set(int(k.split(".", 2)[1]) for k in checkpoint if "double_blocks." in k))[-1] + 1 # noqa: C401 - num_single_layers = list(set(int(k.split(".", 2)[1]) for k in checkpoint if "single_blocks." in k))[-1] + 1 # noqa: C401 - mlp_ratio = 4.0 - inner_dim = 3072 - - # in SD3 original implementation of AdaLayerNormContinuous, it split linear projection output into shift, scale; - # while in diffusers it split into scale, shift. Here we swap the linear projection weights in order to be able to use diffusers implementation - def swap_scale_shift(weight): - shift, scale = weight.chunk(2, dim=0) - new_weight = torch.cat([scale, shift], dim=0) - return new_weight - - ## time_text_embed.timestep_embedder <- time_in - converted_state_dict["time_text_embed.timestep_embedder.linear_1.weight"] = checkpoint.pop( - "time_in.in_layer.weight" - ) - converted_state_dict["time_text_embed.timestep_embedder.linear_1.bias"] = checkpoint.pop("time_in.in_layer.bias") - converted_state_dict["time_text_embed.timestep_embedder.linear_2.weight"] = checkpoint.pop( - "time_in.out_layer.weight" - ) - converted_state_dict["time_text_embed.timestep_embedder.linear_2.bias"] = checkpoint.pop("time_in.out_layer.bias") - - ## time_text_embed.text_embedder <- vector_in - converted_state_dict["time_text_embed.text_embedder.linear_1.weight"] = checkpoint.pop("vector_in.in_layer.weight") - converted_state_dict["time_text_embed.text_embedder.linear_1.bias"] = checkpoint.pop("vector_in.in_layer.bias") - converted_state_dict["time_text_embed.text_embedder.linear_2.weight"] = checkpoint.pop( - "vector_in.out_layer.weight" - ) - converted_state_dict["time_text_embed.text_embedder.linear_2.bias"] = checkpoint.pop("vector_in.out_layer.bias") - - # guidance - has_guidance = any("guidance" in k for k in checkpoint) - if has_guidance: - converted_state_dict["time_text_embed.guidance_embedder.linear_1.weight"] = checkpoint.pop( - "guidance_in.in_layer.weight" - ) - converted_state_dict["time_text_embed.guidance_embedder.linear_1.bias"] = checkpoint.pop( - "guidance_in.in_layer.bias" - ) - converted_state_dict["time_text_embed.guidance_embedder.linear_2.weight"] = checkpoint.pop( - "guidance_in.out_layer.weight" - ) - converted_state_dict["time_text_embed.guidance_embedder.linear_2.bias"] = checkpoint.pop( - "guidance_in.out_layer.bias" - ) - - # context_embedder - converted_state_dict["context_embedder.weight"] = checkpoint.pop("txt_in.weight") - converted_state_dict["context_embedder.bias"] = checkpoint.pop("txt_in.bias") - - # x_embedder - converted_state_dict["x_embedder.weight"] = checkpoint.pop("img_in.weight") - converted_state_dict["x_embedder.bias"] = checkpoint.pop("img_in.bias") - - # double transformer blocks - for i in range(num_layers): - block_prefix = f"transformer_blocks.{i}." - # norms. - ## norm1 - converted_state_dict[f"{block_prefix}norm1.linear.weight"] = checkpoint.pop( - f"double_blocks.{i}.img_mod.lin.weight" - ) - converted_state_dict[f"{block_prefix}norm1.linear.bias"] = checkpoint.pop( - f"double_blocks.{i}.img_mod.lin.bias" - ) - ## norm1_context - converted_state_dict[f"{block_prefix}norm1_context.linear.weight"] = checkpoint.pop( - f"double_blocks.{i}.txt_mod.lin.weight" - ) - converted_state_dict[f"{block_prefix}norm1_context.linear.bias"] = checkpoint.pop( - f"double_blocks.{i}.txt_mod.lin.bias" - ) - # Q, K, V - sample_q, sample_k, sample_v = torch.chunk(checkpoint.pop(f"double_blocks.{i}.img_attn.qkv.weight"), 3, dim=0) - context_q, context_k, context_v = torch.chunk( - checkpoint.pop(f"double_blocks.{i}.txt_attn.qkv.weight"), 3, dim=0 - ) - sample_q_bias, sample_k_bias, sample_v_bias = torch.chunk( - checkpoint.pop(f"double_blocks.{i}.img_attn.qkv.bias"), 3, dim=0 - ) - context_q_bias, context_k_bias, context_v_bias = torch.chunk( - checkpoint.pop(f"double_blocks.{i}.txt_attn.qkv.bias"), 3, dim=0 - ) - converted_state_dict[f"{block_prefix}attn.to_q.weight"] = torch.cat([sample_q]) - converted_state_dict[f"{block_prefix}attn.to_q.bias"] = torch.cat([sample_q_bias]) - converted_state_dict[f"{block_prefix}attn.to_k.weight"] = torch.cat([sample_k]) - converted_state_dict[f"{block_prefix}attn.to_k.bias"] = torch.cat([sample_k_bias]) - converted_state_dict[f"{block_prefix}attn.to_v.weight"] = torch.cat([sample_v]) - converted_state_dict[f"{block_prefix}attn.to_v.bias"] = torch.cat([sample_v_bias]) - converted_state_dict[f"{block_prefix}attn.add_q_proj.weight"] = torch.cat([context_q]) - converted_state_dict[f"{block_prefix}attn.add_q_proj.bias"] = torch.cat([context_q_bias]) - converted_state_dict[f"{block_prefix}attn.add_k_proj.weight"] = torch.cat([context_k]) - converted_state_dict[f"{block_prefix}attn.add_k_proj.bias"] = torch.cat([context_k_bias]) - converted_state_dict[f"{block_prefix}attn.add_v_proj.weight"] = torch.cat([context_v]) - converted_state_dict[f"{block_prefix}attn.add_v_proj.bias"] = torch.cat([context_v_bias]) - # qk_norm - converted_state_dict[f"{block_prefix}attn.norm_q.weight"] = checkpoint.pop( - f"double_blocks.{i}.img_attn.norm.query_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_k.weight"] = checkpoint.pop( - f"double_blocks.{i}.img_attn.norm.key_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_added_q.weight"] = checkpoint.pop( - f"double_blocks.{i}.txt_attn.norm.query_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_added_k.weight"] = checkpoint.pop( - f"double_blocks.{i}.txt_attn.norm.key_norm.scale" - ) - # ff img_mlp - converted_state_dict[f"{block_prefix}ff.net.0.proj.weight"] = checkpoint.pop( - f"double_blocks.{i}.img_mlp.0.weight" - ) - converted_state_dict[f"{block_prefix}ff.net.0.proj.bias"] = checkpoint.pop(f"double_blocks.{i}.img_mlp.0.bias") - converted_state_dict[f"{block_prefix}ff.net.2.weight"] = checkpoint.pop(f"double_blocks.{i}.img_mlp.2.weight") - converted_state_dict[f"{block_prefix}ff.net.2.bias"] = checkpoint.pop(f"double_blocks.{i}.img_mlp.2.bias") - converted_state_dict[f"{block_prefix}ff_context.net.0.proj.weight"] = checkpoint.pop( - f"double_blocks.{i}.txt_mlp.0.weight" - ) - converted_state_dict[f"{block_prefix}ff_context.net.0.proj.bias"] = checkpoint.pop( - f"double_blocks.{i}.txt_mlp.0.bias" - ) - converted_state_dict[f"{block_prefix}ff_context.net.2.weight"] = checkpoint.pop( - f"double_blocks.{i}.txt_mlp.2.weight" - ) - converted_state_dict[f"{block_prefix}ff_context.net.2.bias"] = checkpoint.pop( - f"double_blocks.{i}.txt_mlp.2.bias" - ) - # output projections. - converted_state_dict[f"{block_prefix}attn.to_out.0.weight"] = checkpoint.pop( - f"double_blocks.{i}.img_attn.proj.weight" - ) - converted_state_dict[f"{block_prefix}attn.to_out.0.bias"] = checkpoint.pop( - f"double_blocks.{i}.img_attn.proj.bias" - ) - converted_state_dict[f"{block_prefix}attn.to_add_out.weight"] = checkpoint.pop( - f"double_blocks.{i}.txt_attn.proj.weight" - ) - converted_state_dict[f"{block_prefix}attn.to_add_out.bias"] = checkpoint.pop( - f"double_blocks.{i}.txt_attn.proj.bias" - ) - - # single transformer blocks - for i in range(num_single_layers): - block_prefix = f"single_transformer_blocks.{i}." - # norm.linear <- single_blocks.0.modulation.lin - converted_state_dict[f"{block_prefix}norm.linear.weight"] = checkpoint.pop( - f"single_blocks.{i}.modulation.lin.weight" - ) - converted_state_dict[f"{block_prefix}norm.linear.bias"] = checkpoint.pop( - f"single_blocks.{i}.modulation.lin.bias" - ) - # Q, K, V, mlp - mlp_hidden_dim = int(inner_dim * mlp_ratio) - split_size = (inner_dim, inner_dim, inner_dim, mlp_hidden_dim) - q, k, v, mlp = torch.split(checkpoint.pop(f"single_blocks.{i}.linear1.weight"), split_size, dim=0) - q_bias, k_bias, v_bias, mlp_bias = torch.split( - checkpoint.pop(f"single_blocks.{i}.linear1.bias"), split_size, dim=0 - ) - converted_state_dict[f"{block_prefix}attn.to_q.weight"] = torch.cat([q]) - converted_state_dict[f"{block_prefix}attn.to_q.bias"] = torch.cat([q_bias]) - converted_state_dict[f"{block_prefix}attn.to_k.weight"] = torch.cat([k]) - converted_state_dict[f"{block_prefix}attn.to_k.bias"] = torch.cat([k_bias]) - converted_state_dict[f"{block_prefix}attn.to_v.weight"] = torch.cat([v]) - converted_state_dict[f"{block_prefix}attn.to_v.bias"] = torch.cat([v_bias]) - converted_state_dict[f"{block_prefix}proj_mlp.weight"] = torch.cat([mlp]) - converted_state_dict[f"{block_prefix}proj_mlp.bias"] = torch.cat([mlp_bias]) - # qk norm - converted_state_dict[f"{block_prefix}attn.norm_q.weight"] = checkpoint.pop( - f"single_blocks.{i}.norm.query_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_k.weight"] = checkpoint.pop( - f"single_blocks.{i}.norm.key_norm.scale" - ) - # output projections. - converted_state_dict[f"{block_prefix}proj_out.weight"] = checkpoint.pop(f"single_blocks.{i}.linear2.weight") - converted_state_dict[f"{block_prefix}proj_out.bias"] = checkpoint.pop(f"single_blocks.{i}.linear2.bias") - - converted_state_dict["proj_out.weight"] = checkpoint.pop("final_layer.linear.weight") - converted_state_dict["proj_out.bias"] = checkpoint.pop("final_layer.linear.bias") - converted_state_dict["norm_out.linear.weight"] = swap_scale_shift( - checkpoint.pop("final_layer.adaLN_modulation.1.weight") - ) - converted_state_dict["norm_out.linear.bias"] = swap_scale_shift( - checkpoint.pop("final_layer.adaLN_modulation.1.bias") - ) - - return converted_state_dict - - -def convert_ltx_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {key: checkpoint.pop(key) for key in list(checkpoint.keys()) if "vae" not in key} - - TRANSFORMER_KEYS_RENAME_DICT = { - "model.diffusion_model.": "", - "patchify_proj": "proj_in", - "adaln_single": "time_embed", - "q_norm": "norm_q", - "k_norm": "norm_k", - } - - TRANSFORMER_SPECIAL_KEYS_REMAP = {} - - for key in list(converted_state_dict.keys()): - new_key = key - for replace_key, rename_key in TRANSFORMER_KEYS_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - converted_state_dict[new_key] = converted_state_dict.pop(key) - - for key in list(converted_state_dict.keys()): - for special_key, handler_fn_inplace in TRANSFORMER_SPECIAL_KEYS_REMAP.items(): - if special_key not in key: - continue - handler_fn_inplace(key, converted_state_dict) - - return converted_state_dict - - -def convert_ltx_vae_checkpoint_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {key: checkpoint.pop(key) for key in list(checkpoint.keys()) if "vae." in key} - - def remove_keys_(key: str, state_dict): - state_dict.pop(key) - - VAE_KEYS_RENAME_DICT = { - # common - "vae.": "", - # decoder - "up_blocks.0": "mid_block", - "up_blocks.1": "up_blocks.0", - "up_blocks.2": "up_blocks.1.upsamplers.0", - "up_blocks.3": "up_blocks.1", - "up_blocks.4": "up_blocks.2.conv_in", - "up_blocks.5": "up_blocks.2.upsamplers.0", - "up_blocks.6": "up_blocks.2", - "up_blocks.7": "up_blocks.3.conv_in", - "up_blocks.8": "up_blocks.3.upsamplers.0", - "up_blocks.9": "up_blocks.3", - # encoder - "down_blocks.0": "down_blocks.0", - "down_blocks.1": "down_blocks.0.downsamplers.0", - "down_blocks.2": "down_blocks.0.conv_out", - "down_blocks.3": "down_blocks.1", - "down_blocks.4": "down_blocks.1.downsamplers.0", - "down_blocks.5": "down_blocks.1.conv_out", - "down_blocks.6": "down_blocks.2", - "down_blocks.7": "down_blocks.2.downsamplers.0", - "down_blocks.8": "down_blocks.3", - "down_blocks.9": "mid_block", - # common - "conv_shortcut": "conv_shortcut.conv", - "res_blocks": "resnets", - "norm3.norm": "norm3", - "per_channel_statistics.mean-of-means": "latents_mean", - "per_channel_statistics.std-of-means": "latents_std", - } - - VAE_091_RENAME_DICT = { - # decoder - "up_blocks.0": "mid_block", - "up_blocks.1": "up_blocks.0.upsamplers.0", - "up_blocks.2": "up_blocks.0", - "up_blocks.3": "up_blocks.1.upsamplers.0", - "up_blocks.4": "up_blocks.1", - "up_blocks.5": "up_blocks.2.upsamplers.0", - "up_blocks.6": "up_blocks.2", - "up_blocks.7": "up_blocks.3.upsamplers.0", - "up_blocks.8": "up_blocks.3", - # common - "last_time_embedder": "time_embedder", - "last_scale_shift_table": "scale_shift_table", - } - - VAE_095_RENAME_DICT = { - # decoder - "up_blocks.0": "mid_block", - "up_blocks.1": "up_blocks.0.upsamplers.0", - "up_blocks.2": "up_blocks.0", - "up_blocks.3": "up_blocks.1.upsamplers.0", - "up_blocks.4": "up_blocks.1", - "up_blocks.5": "up_blocks.2.upsamplers.0", - "up_blocks.6": "up_blocks.2", - "up_blocks.7": "up_blocks.3.upsamplers.0", - "up_blocks.8": "up_blocks.3", - # encoder - "down_blocks.0": "down_blocks.0", - "down_blocks.1": "down_blocks.0.downsamplers.0", - "down_blocks.2": "down_blocks.1", - "down_blocks.3": "down_blocks.1.downsamplers.0", - "down_blocks.4": "down_blocks.2", - "down_blocks.5": "down_blocks.2.downsamplers.0", - "down_blocks.6": "down_blocks.3", - "down_blocks.7": "down_blocks.3.downsamplers.0", - "down_blocks.8": "mid_block", - # common - "last_time_embedder": "time_embedder", - "last_scale_shift_table": "scale_shift_table", - } - - VAE_SPECIAL_KEYS_REMAP = { - "per_channel_statistics.channel": remove_keys_, - "per_channel_statistics.mean-of-means": remove_keys_, - "per_channel_statistics.mean-of-stds": remove_keys_, - } - - if converted_state_dict["vae.encoder.conv_out.conv.weight"].shape[1] == 2048: - VAE_KEYS_RENAME_DICT.update(VAE_095_RENAME_DICT) - elif "vae.decoder.last_time_embedder.timestep_embedder.linear_1.weight" in converted_state_dict: - VAE_KEYS_RENAME_DICT.update(VAE_091_RENAME_DICT) - - for key in list(converted_state_dict.keys()): - new_key = key - for replace_key, rename_key in VAE_KEYS_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - converted_state_dict[new_key] = converted_state_dict.pop(key) - - for key in list(converted_state_dict.keys()): - for special_key, handler_fn_inplace in VAE_SPECIAL_KEYS_REMAP.items(): - if special_key not in key: - continue - handler_fn_inplace(key, converted_state_dict) - - return converted_state_dict - - -def convert_autoencoder_dc_checkpoint_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {key: checkpoint.pop(key) for key in list(checkpoint.keys())} - - def remap_qkv_(key: str, state_dict): - qkv = state_dict.pop(key) - q, k, v = torch.chunk(qkv, 3, dim=0) - parent_module, _, _ = key.rpartition(".qkv.conv.weight") - state_dict[f"{parent_module}.to_q.weight"] = q.squeeze() - state_dict[f"{parent_module}.to_k.weight"] = k.squeeze() - state_dict[f"{parent_module}.to_v.weight"] = v.squeeze() - - def remap_proj_conv_(key: str, state_dict): - parent_module, _, _ = key.rpartition(".proj.conv.weight") - state_dict[f"{parent_module}.to_out.weight"] = state_dict.pop(key).squeeze() - - AE_KEYS_RENAME_DICT = { - # common - "main.": "", - "op_list.": "", - "context_module": "attn", - "local_module": "conv_out", - # NOTE: The below two lines work because scales in the available configs only have a tuple length of 1 - # If there were more scales, there would be more layers, so a loop would be better to handle this - "aggreg.0.0": "to_qkv_multiscale.0.proj_in", - "aggreg.0.1": "to_qkv_multiscale.0.proj_out", - "depth_conv.conv": "conv_depth", - "inverted_conv.conv": "conv_inverted", - "point_conv.conv": "conv_point", - "point_conv.norm": "norm", - "conv.conv.": "conv.", - "conv1.conv": "conv1", - "conv2.conv": "conv2", - "conv2.norm": "norm", - "proj.norm": "norm_out", - # encoder - "encoder.project_in.conv": "encoder.conv_in", - "encoder.project_out.0.conv": "encoder.conv_out", - "encoder.stages": "encoder.down_blocks", - # decoder - "decoder.project_in.conv": "decoder.conv_in", - "decoder.project_out.0": "decoder.norm_out", - "decoder.project_out.2.conv": "decoder.conv_out", - "decoder.stages": "decoder.up_blocks", - } - - AE_F32C32_F64C128_F128C512_KEYS = { - "encoder.project_in.conv": "encoder.conv_in.conv", - "decoder.project_out.2.conv": "decoder.conv_out.conv", - } - - AE_SPECIAL_KEYS_REMAP = { - "qkv.conv.weight": remap_qkv_, - "proj.conv.weight": remap_proj_conv_, - } - if "encoder.project_in.conv.bias" not in converted_state_dict: - AE_KEYS_RENAME_DICT.update(AE_F32C32_F64C128_F128C512_KEYS) - - for key in list(converted_state_dict.keys()): - new_key = key[:] - for replace_key, rename_key in AE_KEYS_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - converted_state_dict[new_key] = converted_state_dict.pop(key) - - for key in list(converted_state_dict.keys()): - for special_key, handler_fn_inplace in AE_SPECIAL_KEYS_REMAP.items(): - if special_key not in key: - continue - handler_fn_inplace(key, converted_state_dict) - - return converted_state_dict - - -def convert_mochi_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {} - - # Comfy checkpoints add this prefix - keys = list(checkpoint.keys()) - for k in keys: - if "model.diffusion_model." in k: - checkpoint[k.replace("model.diffusion_model.", "")] = checkpoint.pop(k) - - # Convert patch_embed - converted_state_dict["patch_embed.proj.weight"] = checkpoint.pop("x_embedder.proj.weight") - converted_state_dict["patch_embed.proj.bias"] = checkpoint.pop("x_embedder.proj.bias") - - # Convert time_embed - converted_state_dict["time_embed.timestep_embedder.linear_1.weight"] = checkpoint.pop("t_embedder.mlp.0.weight") - converted_state_dict["time_embed.timestep_embedder.linear_1.bias"] = checkpoint.pop("t_embedder.mlp.0.bias") - converted_state_dict["time_embed.timestep_embedder.linear_2.weight"] = checkpoint.pop("t_embedder.mlp.2.weight") - converted_state_dict["time_embed.timestep_embedder.linear_2.bias"] = checkpoint.pop("t_embedder.mlp.2.bias") - converted_state_dict["time_embed.pooler.to_kv.weight"] = checkpoint.pop("t5_y_embedder.to_kv.weight") - converted_state_dict["time_embed.pooler.to_kv.bias"] = checkpoint.pop("t5_y_embedder.to_kv.bias") - converted_state_dict["time_embed.pooler.to_q.weight"] = checkpoint.pop("t5_y_embedder.to_q.weight") - converted_state_dict["time_embed.pooler.to_q.bias"] = checkpoint.pop("t5_y_embedder.to_q.bias") - converted_state_dict["time_embed.pooler.to_out.weight"] = checkpoint.pop("t5_y_embedder.to_out.weight") - converted_state_dict["time_embed.pooler.to_out.bias"] = checkpoint.pop("t5_y_embedder.to_out.bias") - converted_state_dict["time_embed.caption_proj.weight"] = checkpoint.pop("t5_yproj.weight") - converted_state_dict["time_embed.caption_proj.bias"] = checkpoint.pop("t5_yproj.bias") - - # Convert transformer blocks - num_layers = 48 - for i in range(num_layers): - block_prefix = f"transformer_blocks.{i}." - old_prefix = f"blocks.{i}." - - # norm1 - converted_state_dict[block_prefix + "norm1.linear.weight"] = checkpoint.pop(old_prefix + "mod_x.weight") - converted_state_dict[block_prefix + "norm1.linear.bias"] = checkpoint.pop(old_prefix + "mod_x.bias") - if i < num_layers - 1: - converted_state_dict[block_prefix + "norm1_context.linear.weight"] = checkpoint.pop( - old_prefix + "mod_y.weight" - ) - converted_state_dict[block_prefix + "norm1_context.linear.bias"] = checkpoint.pop( - old_prefix + "mod_y.bias" - ) - else: - converted_state_dict[block_prefix + "norm1_context.linear_1.weight"] = checkpoint.pop( - old_prefix + "mod_y.weight" - ) - converted_state_dict[block_prefix + "norm1_context.linear_1.bias"] = checkpoint.pop( - old_prefix + "mod_y.bias" - ) - - # Visual attention - qkv_weight = checkpoint.pop(old_prefix + "attn.qkv_x.weight") - q, k, v = qkv_weight.chunk(3, dim=0) - - converted_state_dict[block_prefix + "attn1.to_q.weight"] = q - converted_state_dict[block_prefix + "attn1.to_k.weight"] = k - converted_state_dict[block_prefix + "attn1.to_v.weight"] = v - converted_state_dict[block_prefix + "attn1.norm_q.weight"] = checkpoint.pop( - old_prefix + "attn.q_norm_x.weight" - ) - converted_state_dict[block_prefix + "attn1.norm_k.weight"] = checkpoint.pop( - old_prefix + "attn.k_norm_x.weight" - ) - converted_state_dict[block_prefix + "attn1.to_out.0.weight"] = checkpoint.pop( - old_prefix + "attn.proj_x.weight" - ) - converted_state_dict[block_prefix + "attn1.to_out.0.bias"] = checkpoint.pop(old_prefix + "attn.proj_x.bias") - - # Context attention - qkv_weight = checkpoint.pop(old_prefix + "attn.qkv_y.weight") - q, k, v = qkv_weight.chunk(3, dim=0) - - converted_state_dict[block_prefix + "attn1.add_q_proj.weight"] = q - converted_state_dict[block_prefix + "attn1.add_k_proj.weight"] = k - converted_state_dict[block_prefix + "attn1.add_v_proj.weight"] = v - converted_state_dict[block_prefix + "attn1.norm_added_q.weight"] = checkpoint.pop( - old_prefix + "attn.q_norm_y.weight" - ) - converted_state_dict[block_prefix + "attn1.norm_added_k.weight"] = checkpoint.pop( - old_prefix + "attn.k_norm_y.weight" - ) - if i < num_layers - 1: - converted_state_dict[block_prefix + "attn1.to_add_out.weight"] = checkpoint.pop( - old_prefix + "attn.proj_y.weight" - ) - converted_state_dict[block_prefix + "attn1.to_add_out.bias"] = checkpoint.pop( - old_prefix + "attn.proj_y.bias" - ) - - # MLP - converted_state_dict[block_prefix + "ff.net.0.proj.weight"] = swap_proj_gate( - checkpoint.pop(old_prefix + "mlp_x.w1.weight") - ) - converted_state_dict[block_prefix + "ff.net.2.weight"] = checkpoint.pop(old_prefix + "mlp_x.w2.weight") - if i < num_layers - 1: - converted_state_dict[block_prefix + "ff_context.net.0.proj.weight"] = swap_proj_gate( - checkpoint.pop(old_prefix + "mlp_y.w1.weight") - ) - converted_state_dict[block_prefix + "ff_context.net.2.weight"] = checkpoint.pop( - old_prefix + "mlp_y.w2.weight" - ) - - # Output layers - converted_state_dict["norm_out.linear.weight"] = swap_scale_shift(checkpoint.pop("final_layer.mod.weight"), dim=0) - converted_state_dict["norm_out.linear.bias"] = swap_scale_shift(checkpoint.pop("final_layer.mod.bias"), dim=0) - converted_state_dict["proj_out.weight"] = checkpoint.pop("final_layer.linear.weight") - converted_state_dict["proj_out.bias"] = checkpoint.pop("final_layer.linear.bias") - - converted_state_dict["pos_frequencies"] = checkpoint.pop("pos_frequencies") - - return converted_state_dict - - -def convert_hunyuan_video_transformer_to_diffusers(checkpoint, **kwargs): - def remap_norm_scale_shift_(key, state_dict): - weight = state_dict.pop(key) - shift, scale = weight.chunk(2, dim=0) - new_weight = torch.cat([scale, shift], dim=0) - state_dict[key.replace("final_layer.adaLN_modulation.1", "norm_out.linear")] = new_weight - - def remap_txt_in_(key, state_dict): - def rename_key(key): - new_key = key.replace("individual_token_refiner.blocks", "token_refiner.refiner_blocks") - new_key = new_key.replace("adaLN_modulation.1", "norm_out.linear") - new_key = new_key.replace("txt_in", "context_embedder") - new_key = new_key.replace("t_embedder.mlp.0", "time_text_embed.timestep_embedder.linear_1") - new_key = new_key.replace("t_embedder.mlp.2", "time_text_embed.timestep_embedder.linear_2") - new_key = new_key.replace("c_embedder", "time_text_embed.text_embedder") - new_key = new_key.replace("mlp", "ff") - return new_key - - if "self_attn_qkv" in key: - weight = state_dict.pop(key) - to_q, to_k, to_v = weight.chunk(3, dim=0) - state_dict[rename_key(key.replace("self_attn_qkv", "attn.to_q"))] = to_q - state_dict[rename_key(key.replace("self_attn_qkv", "attn.to_k"))] = to_k - state_dict[rename_key(key.replace("self_attn_qkv", "attn.to_v"))] = to_v - else: - state_dict[rename_key(key)] = state_dict.pop(key) - - def remap_img_attn_qkv_(key, state_dict): - weight = state_dict.pop(key) - to_q, to_k, to_v = weight.chunk(3, dim=0) - state_dict[key.replace("img_attn_qkv", "attn.to_q")] = to_q - state_dict[key.replace("img_attn_qkv", "attn.to_k")] = to_k - state_dict[key.replace("img_attn_qkv", "attn.to_v")] = to_v - - def remap_txt_attn_qkv_(key, state_dict): - weight = state_dict.pop(key) - to_q, to_k, to_v = weight.chunk(3, dim=0) - state_dict[key.replace("txt_attn_qkv", "attn.add_q_proj")] = to_q - state_dict[key.replace("txt_attn_qkv", "attn.add_k_proj")] = to_k - state_dict[key.replace("txt_attn_qkv", "attn.add_v_proj")] = to_v - - def remap_single_transformer_blocks_(key, state_dict): - hidden_size = 3072 - - if "linear1.weight" in key: - linear1_weight = state_dict.pop(key) - split_size = (hidden_size, hidden_size, hidden_size, linear1_weight.size(0) - 3 * hidden_size) - q, k, v, mlp = torch.split(linear1_weight, split_size, dim=0) - new_key = key.replace("single_blocks", "single_transformer_blocks").removesuffix(".linear1.weight") - state_dict[f"{new_key}.attn.to_q.weight"] = q - state_dict[f"{new_key}.attn.to_k.weight"] = k - state_dict[f"{new_key}.attn.to_v.weight"] = v - state_dict[f"{new_key}.proj_mlp.weight"] = mlp - - elif "linear1.bias" in key: - linear1_bias = state_dict.pop(key) - split_size = (hidden_size, hidden_size, hidden_size, linear1_bias.size(0) - 3 * hidden_size) - q_bias, k_bias, v_bias, mlp_bias = torch.split(linear1_bias, split_size, dim=0) - new_key = key.replace("single_blocks", "single_transformer_blocks").removesuffix(".linear1.bias") - state_dict[f"{new_key}.attn.to_q.bias"] = q_bias - state_dict[f"{new_key}.attn.to_k.bias"] = k_bias - state_dict[f"{new_key}.attn.to_v.bias"] = v_bias - state_dict[f"{new_key}.proj_mlp.bias"] = mlp_bias - - else: - new_key = key.replace("single_blocks", "single_transformer_blocks") - new_key = new_key.replace("linear2", "proj_out") - new_key = new_key.replace("q_norm", "attn.norm_q") - new_key = new_key.replace("k_norm", "attn.norm_k") - state_dict[new_key] = state_dict.pop(key) - - TRANSFORMER_KEYS_RENAME_DICT = { - "img_in": "x_embedder", - "time_in.mlp.0": "time_text_embed.timestep_embedder.linear_1", - "time_in.mlp.2": "time_text_embed.timestep_embedder.linear_2", - "guidance_in.mlp.0": "time_text_embed.guidance_embedder.linear_1", - "guidance_in.mlp.2": "time_text_embed.guidance_embedder.linear_2", - "vector_in.in_layer": "time_text_embed.text_embedder.linear_1", - "vector_in.out_layer": "time_text_embed.text_embedder.linear_2", - "double_blocks": "transformer_blocks", - "img_attn_q_norm": "attn.norm_q", - "img_attn_k_norm": "attn.norm_k", - "img_attn_proj": "attn.to_out.0", - "txt_attn_q_norm": "attn.norm_added_q", - "txt_attn_k_norm": "attn.norm_added_k", - "txt_attn_proj": "attn.to_add_out", - "img_mod.linear": "norm1.linear", - "img_norm1": "norm1.norm", - "img_norm2": "norm2", - "img_mlp": "ff", - "txt_mod.linear": "norm1_context.linear", - "txt_norm1": "norm1.norm", - "txt_norm2": "norm2_context", - "txt_mlp": "ff_context", - "self_attn_proj": "attn.to_out.0", - "modulation.linear": "norm.linear", - "pre_norm": "norm.norm", - "final_layer.norm_final": "norm_out.norm", - "final_layer.linear": "proj_out", - "fc1": "net.0.proj", - "fc2": "net.2", - "input_embedder": "proj_in", - } - - TRANSFORMER_SPECIAL_KEYS_REMAP = { - "txt_in": remap_txt_in_, - "img_attn_qkv": remap_img_attn_qkv_, - "txt_attn_qkv": remap_txt_attn_qkv_, - "single_blocks": remap_single_transformer_blocks_, - "final_layer.adaLN_modulation.1": remap_norm_scale_shift_, - } - - def update_state_dict_(state_dict, old_key, new_key): - state_dict[new_key] = state_dict.pop(old_key) - - for key in list(checkpoint.keys()): - new_key = key[:] - for replace_key, rename_key in TRANSFORMER_KEYS_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - update_state_dict_(checkpoint, key, new_key) - - for key in list(checkpoint.keys()): - for special_key, handler_fn_inplace in TRANSFORMER_SPECIAL_KEYS_REMAP.items(): - if special_key not in key: - continue - handler_fn_inplace(key, checkpoint) - - return checkpoint - - -def convert_auraflow_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {} - state_dict_keys = list(checkpoint.keys()) - - # Handle register tokens and positional embeddings - converted_state_dict["register_tokens"] = checkpoint.pop("register_tokens", None) - - # Handle time step projection - converted_state_dict["time_step_proj.linear_1.weight"] = checkpoint.pop("t_embedder.mlp.0.weight", None) - converted_state_dict["time_step_proj.linear_1.bias"] = checkpoint.pop("t_embedder.mlp.0.bias", None) - converted_state_dict["time_step_proj.linear_2.weight"] = checkpoint.pop("t_embedder.mlp.2.weight", None) - converted_state_dict["time_step_proj.linear_2.bias"] = checkpoint.pop("t_embedder.mlp.2.bias", None) - - # Handle context embedder - converted_state_dict["context_embedder.weight"] = checkpoint.pop("cond_seq_linear.weight", None) - - # Calculate the number of layers - def calculate_layers(keys, key_prefix): - layers = set() - for k in keys: - if key_prefix in k: - layer_num = int(k.split(".")[1]) # get the layer number - layers.add(layer_num) - return len(layers) - - mmdit_layers = calculate_layers(state_dict_keys, key_prefix="double_layers") - single_dit_layers = calculate_layers(state_dict_keys, key_prefix="single_layers") - - # MMDiT blocks - for i in range(mmdit_layers): - # Feed-forward - path_mapping = {"mlpX": "ff", "mlpC": "ff_context"} - weight_mapping = {"c_fc1": "linear_1", "c_fc2": "linear_2", "c_proj": "out_projection"} - for orig_k, diffuser_k in path_mapping.items(): - for k, v in weight_mapping.items(): - converted_state_dict[f"joint_transformer_blocks.{i}.{diffuser_k}.{v}.weight"] = checkpoint.pop( - f"double_layers.{i}.{orig_k}.{k}.weight", None - ) - - # Norms - path_mapping = {"modX": "norm1", "modC": "norm1_context"} - for orig_k, diffuser_k in path_mapping.items(): - converted_state_dict[f"joint_transformer_blocks.{i}.{diffuser_k}.linear.weight"] = checkpoint.pop( - f"double_layers.{i}.{orig_k}.1.weight", None - ) - - # Attentions - x_attn_mapping = {"w2q": "to_q", "w2k": "to_k", "w2v": "to_v", "w2o": "to_out.0"} - context_attn_mapping = {"w1q": "add_q_proj", "w1k": "add_k_proj", "w1v": "add_v_proj", "w1o": "to_add_out"} - for attn_mapping in [x_attn_mapping, context_attn_mapping]: - for k, v in attn_mapping.items(): - converted_state_dict[f"joint_transformer_blocks.{i}.attn.{v}.weight"] = checkpoint.pop( - f"double_layers.{i}.attn.{k}.weight", None - ) - - # Single-DiT blocks - for i in range(single_dit_layers): - # Feed-forward - mapping = {"c_fc1": "linear_1", "c_fc2": "linear_2", "c_proj": "out_projection"} - for k, v in mapping.items(): - converted_state_dict[f"single_transformer_blocks.{i}.ff.{v}.weight"] = checkpoint.pop( - f"single_layers.{i}.mlp.{k}.weight", None - ) - - # Norms - converted_state_dict[f"single_transformer_blocks.{i}.norm1.linear.weight"] = checkpoint.pop( - f"single_layers.{i}.modCX.1.weight", None - ) - - # Attentions - x_attn_mapping = {"w1q": "to_q", "w1k": "to_k", "w1v": "to_v", "w1o": "to_out.0"} - for k, v in x_attn_mapping.items(): - converted_state_dict[f"single_transformer_blocks.{i}.attn.{v}.weight"] = checkpoint.pop( - f"single_layers.{i}.attn.{k}.weight", None - ) - # Final blocks - converted_state_dict["proj_out.weight"] = checkpoint.pop("final_linear.weight", None) - - # Handle the final norm layer - norm_weight = checkpoint.pop("modF.1.weight", None) - if norm_weight is not None: - converted_state_dict["norm_out.linear.weight"] = swap_scale_shift(norm_weight, dim=None) - else: - converted_state_dict["norm_out.linear.weight"] = None - - converted_state_dict["pos_embed.pos_embed"] = checkpoint.pop("positional_encoding") - converted_state_dict["pos_embed.proj.weight"] = checkpoint.pop("init_x_linear.weight") - converted_state_dict["pos_embed.proj.bias"] = checkpoint.pop("init_x_linear.bias") - - return converted_state_dict - - -def convert_lumina2_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {} - - # Original Lumina-Image-2 has an extra norm parameter that is unused - # We just remove it here - checkpoint.pop("norm_final.weight", None) - - # Comfy checkpoints add this prefix - keys = list(checkpoint.keys()) - for k in keys: - if "model.diffusion_model." in k: - checkpoint[k.replace("model.diffusion_model.", "")] = checkpoint.pop(k) - - LUMINA_KEY_MAP = { - "cap_embedder": "time_caption_embed.caption_embedder", - "t_embedder.mlp.0": "time_caption_embed.timestep_embedder.linear_1", - "t_embedder.mlp.2": "time_caption_embed.timestep_embedder.linear_2", - "attention": "attn", - ".out.": ".to_out.0.", - "k_norm": "norm_k", - "q_norm": "norm_q", - "w1": "linear_1", - "w2": "linear_2", - "w3": "linear_3", - "adaLN_modulation.1": "norm1.linear", - } - ATTENTION_NORM_MAP = { - "attention_norm1": "norm1.norm", - "attention_norm2": "norm2", - } - CONTEXT_REFINER_MAP = { - "context_refiner.0.attention_norm1": "context_refiner.0.norm1", - "context_refiner.0.attention_norm2": "context_refiner.0.norm2", - "context_refiner.1.attention_norm1": "context_refiner.1.norm1", - "context_refiner.1.attention_norm2": "context_refiner.1.norm2", - } - FINAL_LAYER_MAP = { - "final_layer.adaLN_modulation.1": "norm_out.linear_1", - "final_layer.linear": "norm_out.linear_2", - } - - def convert_lumina_attn_to_diffusers(tensor, diffusers_key): - q_dim = 2304 - k_dim = v_dim = 768 - - to_q, to_k, to_v = torch.split(tensor, [q_dim, k_dim, v_dim], dim=0) - - return { - diffusers_key.replace("qkv", "to_q"): to_q, - diffusers_key.replace("qkv", "to_k"): to_k, - diffusers_key.replace("qkv", "to_v"): to_v, - } - - for key in keys: - diffusers_key = key - for k, v in CONTEXT_REFINER_MAP.items(): - diffusers_key = diffusers_key.replace(k, v) - for k, v in FINAL_LAYER_MAP.items(): - diffusers_key = diffusers_key.replace(k, v) - for k, v in ATTENTION_NORM_MAP.items(): - diffusers_key = diffusers_key.replace(k, v) - for k, v in LUMINA_KEY_MAP.items(): - diffusers_key = diffusers_key.replace(k, v) - - if "qkv" in diffusers_key: - converted_state_dict.update(convert_lumina_attn_to_diffusers(checkpoint.pop(key), diffusers_key)) - else: - converted_state_dict[diffusers_key] = checkpoint.pop(key) - - return converted_state_dict - - -def convert_sana_transformer_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {} - keys = list(checkpoint.keys()) - for k in keys: - if "model.diffusion_model." in k: - checkpoint[k.replace("model.diffusion_model.", "")] = checkpoint.pop(k) - - num_layers = list(set(int(k.split(".", 2)[1]) for k in checkpoint if "blocks" in k))[-1] + 1 # noqa: C401 - - # Positional and patch embeddings. - checkpoint.pop("pos_embed") - converted_state_dict["patch_embed.proj.weight"] = checkpoint.pop("x_embedder.proj.weight") - converted_state_dict["patch_embed.proj.bias"] = checkpoint.pop("x_embedder.proj.bias") - - # Timestep embeddings. - converted_state_dict["time_embed.emb.timestep_embedder.linear_1.weight"] = checkpoint.pop( - "t_embedder.mlp.0.weight" - ) - converted_state_dict["time_embed.emb.timestep_embedder.linear_1.bias"] = checkpoint.pop("t_embedder.mlp.0.bias") - converted_state_dict["time_embed.emb.timestep_embedder.linear_2.weight"] = checkpoint.pop( - "t_embedder.mlp.2.weight" - ) - converted_state_dict["time_embed.emb.timestep_embedder.linear_2.bias"] = checkpoint.pop("t_embedder.mlp.2.bias") - converted_state_dict["time_embed.linear.weight"] = checkpoint.pop("t_block.1.weight") - converted_state_dict["time_embed.linear.bias"] = checkpoint.pop("t_block.1.bias") - - # Caption Projection. - checkpoint.pop("y_embedder.y_embedding") - converted_state_dict["caption_projection.linear_1.weight"] = checkpoint.pop("y_embedder.y_proj.fc1.weight") - converted_state_dict["caption_projection.linear_1.bias"] = checkpoint.pop("y_embedder.y_proj.fc1.bias") - converted_state_dict["caption_projection.linear_2.weight"] = checkpoint.pop("y_embedder.y_proj.fc2.weight") - converted_state_dict["caption_projection.linear_2.bias"] = checkpoint.pop("y_embedder.y_proj.fc2.bias") - converted_state_dict["caption_norm.weight"] = checkpoint.pop("attention_y_norm.weight") - - for i in range(num_layers): - converted_state_dict[f"transformer_blocks.{i}.scale_shift_table"] = checkpoint.pop( - f"blocks.{i}.scale_shift_table" - ) - - # Self-Attention - sample_q, sample_k, sample_v = torch.chunk(checkpoint.pop(f"blocks.{i}.attn.qkv.weight"), 3, dim=0) - converted_state_dict[f"transformer_blocks.{i}.attn1.to_q.weight"] = torch.cat([sample_q]) - converted_state_dict[f"transformer_blocks.{i}.attn1.to_k.weight"] = torch.cat([sample_k]) - converted_state_dict[f"transformer_blocks.{i}.attn1.to_v.weight"] = torch.cat([sample_v]) - - # Output Projections - converted_state_dict[f"transformer_blocks.{i}.attn1.to_out.0.weight"] = checkpoint.pop( - f"blocks.{i}.attn.proj.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.attn1.to_out.0.bias"] = checkpoint.pop( - f"blocks.{i}.attn.proj.bias" - ) - - # Cross-Attention - converted_state_dict[f"transformer_blocks.{i}.attn2.to_q.weight"] = checkpoint.pop( - f"blocks.{i}.cross_attn.q_linear.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.attn2.to_q.bias"] = checkpoint.pop( - f"blocks.{i}.cross_attn.q_linear.bias" - ) - - linear_sample_k, linear_sample_v = torch.chunk( - checkpoint.pop(f"blocks.{i}.cross_attn.kv_linear.weight"), 2, dim=0 - ) - linear_sample_k_bias, linear_sample_v_bias = torch.chunk( - checkpoint.pop(f"blocks.{i}.cross_attn.kv_linear.bias"), 2, dim=0 - ) - converted_state_dict[f"transformer_blocks.{i}.attn2.to_k.weight"] = linear_sample_k - converted_state_dict[f"transformer_blocks.{i}.attn2.to_v.weight"] = linear_sample_v - converted_state_dict[f"transformer_blocks.{i}.attn2.to_k.bias"] = linear_sample_k_bias - converted_state_dict[f"transformer_blocks.{i}.attn2.to_v.bias"] = linear_sample_v_bias - - # Output Projections - converted_state_dict[f"transformer_blocks.{i}.attn2.to_out.0.weight"] = checkpoint.pop( - f"blocks.{i}.cross_attn.proj.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.attn2.to_out.0.bias"] = checkpoint.pop( - f"blocks.{i}.cross_attn.proj.bias" - ) - - # MLP - converted_state_dict[f"transformer_blocks.{i}.ff.conv_inverted.weight"] = checkpoint.pop( - f"blocks.{i}.mlp.inverted_conv.conv.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.ff.conv_inverted.bias"] = checkpoint.pop( - f"blocks.{i}.mlp.inverted_conv.conv.bias" - ) - converted_state_dict[f"transformer_blocks.{i}.ff.conv_depth.weight"] = checkpoint.pop( - f"blocks.{i}.mlp.depth_conv.conv.weight" - ) - converted_state_dict[f"transformer_blocks.{i}.ff.conv_depth.bias"] = checkpoint.pop( - f"blocks.{i}.mlp.depth_conv.conv.bias" - ) - converted_state_dict[f"transformer_blocks.{i}.ff.conv_point.weight"] = checkpoint.pop( - f"blocks.{i}.mlp.point_conv.conv.weight" - ) - - # Final layer - converted_state_dict["proj_out.weight"] = checkpoint.pop("final_layer.linear.weight") - converted_state_dict["proj_out.bias"] = checkpoint.pop("final_layer.linear.bias") - converted_state_dict["scale_shift_table"] = checkpoint.pop("final_layer.scale_shift_table") - - return converted_state_dict - - -def convert_wan_transformer_to_diffusers(checkpoint, **kwargs): - def generate_motion_encoder_mappings(): - mappings = { - "motion_encoder.dec.direction.weight": "motion_encoder.motion_synthesis_weight", - "motion_encoder.enc.net_app.convs.0.0.weight": "motion_encoder.conv_in.weight", - "motion_encoder.enc.net_app.convs.0.1.bias": "motion_encoder.conv_in.act_fn.bias", - "motion_encoder.enc.net_app.convs.8.weight": "motion_encoder.conv_out.weight", - "motion_encoder.enc.fc": "motion_encoder.motion_network", - } - - for i in range(7): - conv_idx = i + 1 - mappings.update( - { - f"motion_encoder.enc.net_app.convs.{conv_idx}.conv1.0.weight": f"motion_encoder.res_blocks.{i}.conv1.weight", - f"motion_encoder.enc.net_app.convs.{conv_idx}.conv1.1.bias": f"motion_encoder.res_blocks.{i}.conv1.act_fn.bias", - f"motion_encoder.enc.net_app.convs.{conv_idx}.conv2.1.weight": f"motion_encoder.res_blocks.{i}.conv2.weight", - f"motion_encoder.enc.net_app.convs.{conv_idx}.conv2.2.bias": f"motion_encoder.res_blocks.{i}.conv2.act_fn.bias", - f"motion_encoder.enc.net_app.convs.{conv_idx}.skip.1.weight": f"motion_encoder.res_blocks.{i}.conv_skip.weight", - } - ) - - return mappings - - def generate_face_adapter_mappings(): - return { - "face_adapter.fuser_blocks": "face_adapter", - ".k_norm.": ".norm_k.", - ".q_norm.": ".norm_q.", - ".linear1_q.": ".to_q.", - ".linear2.": ".to_out.", - "conv1_local.conv": "conv1_local", - "conv2.conv": "conv2", - "conv3.conv": "conv3", - } - - def split_tensor_handler(key, state_dict, split_pattern, target_keys): - tensor = state_dict.pop(key) - split_idx = tensor.shape[0] // 2 - - new_key_1 = key.replace(split_pattern, target_keys[0]) - new_key_2 = key.replace(split_pattern, target_keys[1]) - - state_dict[new_key_1] = tensor[:split_idx] - state_dict[new_key_2] = tensor[split_idx:] - - def reshape_bias_handler(key, state_dict): - if "motion_encoder.enc.net_app.convs." in key and ".bias" in key: - state_dict[key] = state_dict[key][0, :, 0, 0] - - converted_state_dict = {} - - # Strip model.diffusion_model prefix - keys = list(checkpoint.keys()) - for k in keys: - if "model.diffusion_model." in k: - checkpoint[k.replace("model.diffusion_model.", "")] = checkpoint.pop(k) - - # Base transformer mappings - TRANSFORMER_KEYS_RENAME_DICT = { - "time_embedding.0": "condition_embedder.time_embedder.linear_1", - "time_embedding.2": "condition_embedder.time_embedder.linear_2", - "text_embedding.0": "condition_embedder.text_embedder.linear_1", - "text_embedding.2": "condition_embedder.text_embedder.linear_2", - "time_projection.1": "condition_embedder.time_proj", - "cross_attn": "attn2", - "self_attn": "attn1", - ".o.": ".to_out.0.", - ".q.": ".to_q.", - ".k.": ".to_k.", - ".v.": ".to_v.", - ".k_img.": ".add_k_proj.", - ".v_img.": ".add_v_proj.", - ".norm_k_img.": ".norm_added_k.", - "head.modulation": "scale_shift_table", - "head.head": "proj_out", - "modulation": "scale_shift_table", - "ffn.0": "ffn.net.0.proj", - "ffn.2": "ffn.net.2", - # Hack to swap the layer names - "norm2": "norm__placeholder", - "norm3": "norm2", - "norm__placeholder": "norm3", - # I2V model - "img_emb.proj.0": "condition_embedder.image_embedder.norm1", - "img_emb.proj.1": "condition_embedder.image_embedder.ff.net.0.proj", - "img_emb.proj.3": "condition_embedder.image_embedder.ff.net.2", - "img_emb.proj.4": "condition_embedder.image_embedder.norm2", - # VACE model - "before_proj": "proj_in", - "after_proj": "proj_out", - } - - SPECIAL_KEYS_HANDLERS = {} - if any("face_adapter" in k for k in checkpoint.keys()): - TRANSFORMER_KEYS_RENAME_DICT.update(generate_face_adapter_mappings()) - SPECIAL_KEYS_HANDLERS[".linear1_kv."] = (split_tensor_handler, [".to_k.", ".to_v."]) - - if any("motion_encoder" in k for k in checkpoint.keys()): - TRANSFORMER_KEYS_RENAME_DICT.update(generate_motion_encoder_mappings()) - - for key in list(checkpoint.keys()): - reshape_bias_handler(key, checkpoint) - - for key in list(checkpoint.keys()): - new_key = key - for replace_key, rename_key in TRANSFORMER_KEYS_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - converted_state_dict[new_key] = checkpoint.pop(key) - - for key in list(converted_state_dict.keys()): - for pattern, (handler_fn, target_keys) in SPECIAL_KEYS_HANDLERS.items(): - if pattern not in key: - continue - handler_fn(key, converted_state_dict, pattern, target_keys) - break - - return converted_state_dict - - -def convert_wan_vae_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {} - - # Create mappings for specific components - middle_key_mapping = { - # Encoder middle block - "encoder.middle.0.residual.0.gamma": "encoder.mid_block.resnets.0.norm1.gamma", - "encoder.middle.0.residual.2.bias": "encoder.mid_block.resnets.0.conv1.bias", - "encoder.middle.0.residual.2.weight": "encoder.mid_block.resnets.0.conv1.weight", - "encoder.middle.0.residual.3.gamma": "encoder.mid_block.resnets.0.norm2.gamma", - "encoder.middle.0.residual.6.bias": "encoder.mid_block.resnets.0.conv2.bias", - "encoder.middle.0.residual.6.weight": "encoder.mid_block.resnets.0.conv2.weight", - "encoder.middle.2.residual.0.gamma": "encoder.mid_block.resnets.1.norm1.gamma", - "encoder.middle.2.residual.2.bias": "encoder.mid_block.resnets.1.conv1.bias", - "encoder.middle.2.residual.2.weight": "encoder.mid_block.resnets.1.conv1.weight", - "encoder.middle.2.residual.3.gamma": "encoder.mid_block.resnets.1.norm2.gamma", - "encoder.middle.2.residual.6.bias": "encoder.mid_block.resnets.1.conv2.bias", - "encoder.middle.2.residual.6.weight": "encoder.mid_block.resnets.1.conv2.weight", - # Decoder middle block - "decoder.middle.0.residual.0.gamma": "decoder.mid_block.resnets.0.norm1.gamma", - "decoder.middle.0.residual.2.bias": "decoder.mid_block.resnets.0.conv1.bias", - "decoder.middle.0.residual.2.weight": "decoder.mid_block.resnets.0.conv1.weight", - "decoder.middle.0.residual.3.gamma": "decoder.mid_block.resnets.0.norm2.gamma", - "decoder.middle.0.residual.6.bias": "decoder.mid_block.resnets.0.conv2.bias", - "decoder.middle.0.residual.6.weight": "decoder.mid_block.resnets.0.conv2.weight", - "decoder.middle.2.residual.0.gamma": "decoder.mid_block.resnets.1.norm1.gamma", - "decoder.middle.2.residual.2.bias": "decoder.mid_block.resnets.1.conv1.bias", - "decoder.middle.2.residual.2.weight": "decoder.mid_block.resnets.1.conv1.weight", - "decoder.middle.2.residual.3.gamma": "decoder.mid_block.resnets.1.norm2.gamma", - "decoder.middle.2.residual.6.bias": "decoder.mid_block.resnets.1.conv2.bias", - "decoder.middle.2.residual.6.weight": "decoder.mid_block.resnets.1.conv2.weight", - } - - # Create a mapping for attention blocks - attention_mapping = { - # Encoder middle attention - "encoder.middle.1.norm.gamma": "encoder.mid_block.attentions.0.norm.gamma", - "encoder.middle.1.to_qkv.weight": "encoder.mid_block.attentions.0.to_qkv.weight", - "encoder.middle.1.to_qkv.bias": "encoder.mid_block.attentions.0.to_qkv.bias", - "encoder.middle.1.proj.weight": "encoder.mid_block.attentions.0.proj.weight", - "encoder.middle.1.proj.bias": "encoder.mid_block.attentions.0.proj.bias", - # Decoder middle attention - "decoder.middle.1.norm.gamma": "decoder.mid_block.attentions.0.norm.gamma", - "decoder.middle.1.to_qkv.weight": "decoder.mid_block.attentions.0.to_qkv.weight", - "decoder.middle.1.to_qkv.bias": "decoder.mid_block.attentions.0.to_qkv.bias", - "decoder.middle.1.proj.weight": "decoder.mid_block.attentions.0.proj.weight", - "decoder.middle.1.proj.bias": "decoder.mid_block.attentions.0.proj.bias", - } - - # Create a mapping for the head components - head_mapping = { - # Encoder head - "encoder.head.0.gamma": "encoder.norm_out.gamma", - "encoder.head.2.bias": "encoder.conv_out.bias", - "encoder.head.2.weight": "encoder.conv_out.weight", - # Decoder head - "decoder.head.0.gamma": "decoder.norm_out.gamma", - "decoder.head.2.bias": "decoder.conv_out.bias", - "decoder.head.2.weight": "decoder.conv_out.weight", - } - - # Create a mapping for the quant components - quant_mapping = { - "conv1.weight": "quant_conv.weight", - "conv1.bias": "quant_conv.bias", - "conv2.weight": "post_quant_conv.weight", - "conv2.bias": "post_quant_conv.bias", - } - - # Process each key in the state dict - for key, value in checkpoint.items(): - # Handle middle block keys using the mapping - if key in middle_key_mapping: - new_key = middle_key_mapping[key] - converted_state_dict[new_key] = value - # Handle attention blocks using the mapping - elif key in attention_mapping: - new_key = attention_mapping[key] - converted_state_dict[new_key] = value - # Handle head keys using the mapping - elif key in head_mapping: - new_key = head_mapping[key] - converted_state_dict[new_key] = value - # Handle quant keys using the mapping - elif key in quant_mapping: - new_key = quant_mapping[key] - converted_state_dict[new_key] = value - # Handle encoder conv1 - elif key == "encoder.conv1.weight": - converted_state_dict["encoder.conv_in.weight"] = value - elif key == "encoder.conv1.bias": - converted_state_dict["encoder.conv_in.bias"] = value - # Handle decoder conv1 - elif key == "decoder.conv1.weight": - converted_state_dict["decoder.conv_in.weight"] = value - elif key == "decoder.conv1.bias": - converted_state_dict["decoder.conv_in.bias"] = value - # Handle encoder downsamples - elif key.startswith("encoder.downsamples."): - # Convert to down_blocks - new_key = key.replace("encoder.downsamples.", "encoder.down_blocks.") - - # Convert residual block naming but keep the original structure - if ".residual.0.gamma" in new_key: - new_key = new_key.replace(".residual.0.gamma", ".norm1.gamma") - elif ".residual.2.bias" in new_key: - new_key = new_key.replace(".residual.2.bias", ".conv1.bias") - elif ".residual.2.weight" in new_key: - new_key = new_key.replace(".residual.2.weight", ".conv1.weight") - elif ".residual.3.gamma" in new_key: - new_key = new_key.replace(".residual.3.gamma", ".norm2.gamma") - elif ".residual.6.bias" in new_key: - new_key = new_key.replace(".residual.6.bias", ".conv2.bias") - elif ".residual.6.weight" in new_key: - new_key = new_key.replace(".residual.6.weight", ".conv2.weight") - elif ".shortcut.bias" in new_key: - new_key = new_key.replace(".shortcut.bias", ".conv_shortcut.bias") - elif ".shortcut.weight" in new_key: - new_key = new_key.replace(".shortcut.weight", ".conv_shortcut.weight") - - converted_state_dict[new_key] = value - - # Handle decoder upsamples - elif key.startswith("decoder.upsamples."): - # Convert to up_blocks - parts = key.split(".") - block_idx = int(parts[2]) - - # Group residual blocks - if "residual" in key: - if block_idx in [0, 1, 2]: - new_block_idx = 0 - resnet_idx = block_idx - elif block_idx in [4, 5, 6]: - new_block_idx = 1 - resnet_idx = block_idx - 4 - elif block_idx in [8, 9, 10]: - new_block_idx = 2 - resnet_idx = block_idx - 8 - elif block_idx in [12, 13, 14]: - new_block_idx = 3 - resnet_idx = block_idx - 12 - else: - # Keep as is for other blocks - converted_state_dict[key] = value - continue - - # Convert residual block naming - if ".residual.0.gamma" in key: - new_key = f"decoder.up_blocks.{new_block_idx}.resnets.{resnet_idx}.norm1.gamma" - elif ".residual.2.bias" in key: - new_key = f"decoder.up_blocks.{new_block_idx}.resnets.{resnet_idx}.conv1.bias" - elif ".residual.2.weight" in key: - new_key = f"decoder.up_blocks.{new_block_idx}.resnets.{resnet_idx}.conv1.weight" - elif ".residual.3.gamma" in key: - new_key = f"decoder.up_blocks.{new_block_idx}.resnets.{resnet_idx}.norm2.gamma" - elif ".residual.6.bias" in key: - new_key = f"decoder.up_blocks.{new_block_idx}.resnets.{resnet_idx}.conv2.bias" - elif ".residual.6.weight" in key: - new_key = f"decoder.up_blocks.{new_block_idx}.resnets.{resnet_idx}.conv2.weight" - else: - new_key = key - - converted_state_dict[new_key] = value - - # Handle shortcut connections - elif ".shortcut." in key: - if block_idx == 4: - new_key = key.replace(".shortcut.", ".resnets.0.conv_shortcut.") - new_key = new_key.replace("decoder.upsamples.4", "decoder.up_blocks.1") - else: - new_key = key.replace("decoder.upsamples.", "decoder.up_blocks.") - new_key = new_key.replace(".shortcut.", ".conv_shortcut.") - - converted_state_dict[new_key] = value - - # Handle upsamplers - elif ".resample." in key or ".time_conv." in key: - if block_idx == 3: - new_key = key.replace(f"decoder.upsamples.{block_idx}", "decoder.up_blocks.0.upsamplers.0") - elif block_idx == 7: - new_key = key.replace(f"decoder.upsamples.{block_idx}", "decoder.up_blocks.1.upsamplers.0") - elif block_idx == 11: - new_key = key.replace(f"decoder.upsamples.{block_idx}", "decoder.up_blocks.2.upsamplers.0") - else: - new_key = key.replace("decoder.upsamples.", "decoder.up_blocks.") - - converted_state_dict[new_key] = value - else: - new_key = key.replace("decoder.upsamples.", "decoder.up_blocks.") - converted_state_dict[new_key] = value - else: - # Keep other keys unchanged - converted_state_dict[key] = value - - return converted_state_dict - - -def convert_hidream_transformer_to_diffusers(checkpoint, **kwargs): - keys = list(checkpoint.keys()) - for k in keys: - if "model.diffusion_model." in k: - checkpoint[k.replace("model.diffusion_model.", "")] = checkpoint.pop(k) - - return checkpoint - - -def convert_chroma_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {} - keys = list(checkpoint.keys()) - - for k in keys: - if "model.diffusion_model." in k: - checkpoint[k.replace("model.diffusion_model.", "")] = checkpoint.pop(k) - - num_layers = list(set(int(k.split(".", 2)[1]) for k in checkpoint if "double_blocks." in k))[-1] + 1 # noqa: C401 - num_single_layers = list(set(int(k.split(".", 2)[1]) for k in checkpoint if "single_blocks." in k))[-1] + 1 # noqa: C401 - num_guidance_layers = ( - list(set(int(k.split(".", 3)[2]) for k in checkpoint if "distilled_guidance_layer.layers." in k))[-1] + 1 # noqa: C401 - ) - mlp_ratio = 4.0 - inner_dim = 3072 - - # in SD3 original implementation of AdaLayerNormContinuous, it split linear projection output into shift, scale; - # while in diffusers it split into scale, shift. Here we swap the linear projection weights in order to be able to use diffusers implementation - def swap_scale_shift(weight): - shift, scale = weight.chunk(2, dim=0) - new_weight = torch.cat([scale, shift], dim=0) - return new_weight - - # guidance - converted_state_dict["distilled_guidance_layer.in_proj.bias"] = checkpoint.pop( - "distilled_guidance_layer.in_proj.bias" - ) - converted_state_dict["distilled_guidance_layer.in_proj.weight"] = checkpoint.pop( - "distilled_guidance_layer.in_proj.weight" - ) - converted_state_dict["distilled_guidance_layer.out_proj.bias"] = checkpoint.pop( - "distilled_guidance_layer.out_proj.bias" - ) - converted_state_dict["distilled_guidance_layer.out_proj.weight"] = checkpoint.pop( - "distilled_guidance_layer.out_proj.weight" - ) - for i in range(num_guidance_layers): - block_prefix = f"distilled_guidance_layer.layers.{i}." - converted_state_dict[f"{block_prefix}linear_1.bias"] = checkpoint.pop( - f"distilled_guidance_layer.layers.{i}.in_layer.bias" - ) - converted_state_dict[f"{block_prefix}linear_1.weight"] = checkpoint.pop( - f"distilled_guidance_layer.layers.{i}.in_layer.weight" - ) - converted_state_dict[f"{block_prefix}linear_2.bias"] = checkpoint.pop( - f"distilled_guidance_layer.layers.{i}.out_layer.bias" - ) - converted_state_dict[f"{block_prefix}linear_2.weight"] = checkpoint.pop( - f"distilled_guidance_layer.layers.{i}.out_layer.weight" - ) - converted_state_dict[f"distilled_guidance_layer.norms.{i}.weight"] = checkpoint.pop( - f"distilled_guidance_layer.norms.{i}.scale" - ) - - # context_embedder - converted_state_dict["context_embedder.weight"] = checkpoint.pop("txt_in.weight") - converted_state_dict["context_embedder.bias"] = checkpoint.pop("txt_in.bias") - - # x_embedder - converted_state_dict["x_embedder.weight"] = checkpoint.pop("img_in.weight") - converted_state_dict["x_embedder.bias"] = checkpoint.pop("img_in.bias") - - # double transformer blocks - for i in range(num_layers): - block_prefix = f"transformer_blocks.{i}." - # Q, K, V - sample_q, sample_k, sample_v = torch.chunk(checkpoint.pop(f"double_blocks.{i}.img_attn.qkv.weight"), 3, dim=0) - context_q, context_k, context_v = torch.chunk( - checkpoint.pop(f"double_blocks.{i}.txt_attn.qkv.weight"), 3, dim=0 - ) - sample_q_bias, sample_k_bias, sample_v_bias = torch.chunk( - checkpoint.pop(f"double_blocks.{i}.img_attn.qkv.bias"), 3, dim=0 - ) - context_q_bias, context_k_bias, context_v_bias = torch.chunk( - checkpoint.pop(f"double_blocks.{i}.txt_attn.qkv.bias"), 3, dim=0 - ) - converted_state_dict[f"{block_prefix}attn.to_q.weight"] = torch.cat([sample_q]) - converted_state_dict[f"{block_prefix}attn.to_q.bias"] = torch.cat([sample_q_bias]) - converted_state_dict[f"{block_prefix}attn.to_k.weight"] = torch.cat([sample_k]) - converted_state_dict[f"{block_prefix}attn.to_k.bias"] = torch.cat([sample_k_bias]) - converted_state_dict[f"{block_prefix}attn.to_v.weight"] = torch.cat([sample_v]) - converted_state_dict[f"{block_prefix}attn.to_v.bias"] = torch.cat([sample_v_bias]) - converted_state_dict[f"{block_prefix}attn.add_q_proj.weight"] = torch.cat([context_q]) - converted_state_dict[f"{block_prefix}attn.add_q_proj.bias"] = torch.cat([context_q_bias]) - converted_state_dict[f"{block_prefix}attn.add_k_proj.weight"] = torch.cat([context_k]) - converted_state_dict[f"{block_prefix}attn.add_k_proj.bias"] = torch.cat([context_k_bias]) - converted_state_dict[f"{block_prefix}attn.add_v_proj.weight"] = torch.cat([context_v]) - converted_state_dict[f"{block_prefix}attn.add_v_proj.bias"] = torch.cat([context_v_bias]) - # qk_norm - converted_state_dict[f"{block_prefix}attn.norm_q.weight"] = checkpoint.pop( - f"double_blocks.{i}.img_attn.norm.query_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_k.weight"] = checkpoint.pop( - f"double_blocks.{i}.img_attn.norm.key_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_added_q.weight"] = checkpoint.pop( - f"double_blocks.{i}.txt_attn.norm.query_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_added_k.weight"] = checkpoint.pop( - f"double_blocks.{i}.txt_attn.norm.key_norm.scale" - ) - # ff img_mlp - converted_state_dict[f"{block_prefix}ff.net.0.proj.weight"] = checkpoint.pop( - f"double_blocks.{i}.img_mlp.0.weight" - ) - converted_state_dict[f"{block_prefix}ff.net.0.proj.bias"] = checkpoint.pop(f"double_blocks.{i}.img_mlp.0.bias") - converted_state_dict[f"{block_prefix}ff.net.2.weight"] = checkpoint.pop(f"double_blocks.{i}.img_mlp.2.weight") - converted_state_dict[f"{block_prefix}ff.net.2.bias"] = checkpoint.pop(f"double_blocks.{i}.img_mlp.2.bias") - converted_state_dict[f"{block_prefix}ff_context.net.0.proj.weight"] = checkpoint.pop( - f"double_blocks.{i}.txt_mlp.0.weight" - ) - converted_state_dict[f"{block_prefix}ff_context.net.0.proj.bias"] = checkpoint.pop( - f"double_blocks.{i}.txt_mlp.0.bias" - ) - converted_state_dict[f"{block_prefix}ff_context.net.2.weight"] = checkpoint.pop( - f"double_blocks.{i}.txt_mlp.2.weight" - ) - converted_state_dict[f"{block_prefix}ff_context.net.2.bias"] = checkpoint.pop( - f"double_blocks.{i}.txt_mlp.2.bias" - ) - # output projections. - converted_state_dict[f"{block_prefix}attn.to_out.0.weight"] = checkpoint.pop( - f"double_blocks.{i}.img_attn.proj.weight" - ) - converted_state_dict[f"{block_prefix}attn.to_out.0.bias"] = checkpoint.pop( - f"double_blocks.{i}.img_attn.proj.bias" - ) - converted_state_dict[f"{block_prefix}attn.to_add_out.weight"] = checkpoint.pop( - f"double_blocks.{i}.txt_attn.proj.weight" - ) - converted_state_dict[f"{block_prefix}attn.to_add_out.bias"] = checkpoint.pop( - f"double_blocks.{i}.txt_attn.proj.bias" - ) - - # single transformer blocks - for i in range(num_single_layers): - block_prefix = f"single_transformer_blocks.{i}." - # Q, K, V, mlp - mlp_hidden_dim = int(inner_dim * mlp_ratio) - split_size = (inner_dim, inner_dim, inner_dim, mlp_hidden_dim) - q, k, v, mlp = torch.split(checkpoint.pop(f"single_blocks.{i}.linear1.weight"), split_size, dim=0) - q_bias, k_bias, v_bias, mlp_bias = torch.split( - checkpoint.pop(f"single_blocks.{i}.linear1.bias"), split_size, dim=0 - ) - converted_state_dict[f"{block_prefix}attn.to_q.weight"] = torch.cat([q]) - converted_state_dict[f"{block_prefix}attn.to_q.bias"] = torch.cat([q_bias]) - converted_state_dict[f"{block_prefix}attn.to_k.weight"] = torch.cat([k]) - converted_state_dict[f"{block_prefix}attn.to_k.bias"] = torch.cat([k_bias]) - converted_state_dict[f"{block_prefix}attn.to_v.weight"] = torch.cat([v]) - converted_state_dict[f"{block_prefix}attn.to_v.bias"] = torch.cat([v_bias]) - converted_state_dict[f"{block_prefix}proj_mlp.weight"] = torch.cat([mlp]) - converted_state_dict[f"{block_prefix}proj_mlp.bias"] = torch.cat([mlp_bias]) - # qk norm - converted_state_dict[f"{block_prefix}attn.norm_q.weight"] = checkpoint.pop( - f"single_blocks.{i}.norm.query_norm.scale" - ) - converted_state_dict[f"{block_prefix}attn.norm_k.weight"] = checkpoint.pop( - f"single_blocks.{i}.norm.key_norm.scale" - ) - # output projections. - converted_state_dict[f"{block_prefix}proj_out.weight"] = checkpoint.pop(f"single_blocks.{i}.linear2.weight") - converted_state_dict[f"{block_prefix}proj_out.bias"] = checkpoint.pop(f"single_blocks.{i}.linear2.bias") - - converted_state_dict["proj_out.weight"] = checkpoint.pop("final_layer.linear.weight") - converted_state_dict["proj_out.bias"] = checkpoint.pop("final_layer.linear.bias") - - return converted_state_dict - - -def convert_cosmos_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): - converted_state_dict = {key: checkpoint.pop(key) for key in list(checkpoint.keys())} - - def remove_keys_(key: str, state_dict): - state_dict.pop(key) - - def rename_transformer_blocks_(key: str, state_dict): - block_index = int(key.split(".")[1].removeprefix("block")) - new_key = key - old_prefix = f"blocks.block{block_index}" - new_prefix = f"transformer_blocks.{block_index}" - new_key = new_prefix + new_key.removeprefix(old_prefix) - state_dict[new_key] = state_dict.pop(key) - - TRANSFORMER_KEYS_RENAME_DICT_COSMOS_1_0 = { - "t_embedder.1": "time_embed.t_embedder", - "affline_norm": "time_embed.norm", - ".blocks.0.block.attn": ".attn1", - ".blocks.1.block.attn": ".attn2", - ".blocks.2.block": ".ff", - ".blocks.0.adaLN_modulation.1": ".norm1.linear_1", - ".blocks.0.adaLN_modulation.2": ".norm1.linear_2", - ".blocks.1.adaLN_modulation.1": ".norm2.linear_1", - ".blocks.1.adaLN_modulation.2": ".norm2.linear_2", - ".blocks.2.adaLN_modulation.1": ".norm3.linear_1", - ".blocks.2.adaLN_modulation.2": ".norm3.linear_2", - "to_q.0": "to_q", - "to_q.1": "norm_q", - "to_k.0": "to_k", - "to_k.1": "norm_k", - "to_v.0": "to_v", - "layer1": "net.0.proj", - "layer2": "net.2", - "proj.1": "proj", - "x_embedder": "patch_embed", - "extra_pos_embedder": "learnable_pos_embed", - "final_layer.adaLN_modulation.1": "norm_out.linear_1", - "final_layer.adaLN_modulation.2": "norm_out.linear_2", - "final_layer.linear": "proj_out", - } - - TRANSFORMER_SPECIAL_KEYS_REMAP_COSMOS_1_0 = { - "blocks.block": rename_transformer_blocks_, - "logvar.0.freqs": remove_keys_, - "logvar.0.phases": remove_keys_, - "logvar.1.weight": remove_keys_, - "pos_embedder.seq": remove_keys_, - } - - TRANSFORMER_KEYS_RENAME_DICT_COSMOS_2_0 = { - "t_embedder.1": "time_embed.t_embedder", - "t_embedding_norm": "time_embed.norm", - "blocks": "transformer_blocks", - "adaln_modulation_self_attn.1": "norm1.linear_1", - "adaln_modulation_self_attn.2": "norm1.linear_2", - "adaln_modulation_cross_attn.1": "norm2.linear_1", - "adaln_modulation_cross_attn.2": "norm2.linear_2", - "adaln_modulation_mlp.1": "norm3.linear_1", - "adaln_modulation_mlp.2": "norm3.linear_2", - "self_attn": "attn1", - "cross_attn": "attn2", - "q_proj": "to_q", - "k_proj": "to_k", - "v_proj": "to_v", - "output_proj": "to_out.0", - "q_norm": "norm_q", - "k_norm": "norm_k", - "mlp.layer1": "ff.net.0.proj", - "mlp.layer2": "ff.net.2", - "x_embedder.proj.1": "patch_embed.proj", - "final_layer.adaln_modulation.1": "norm_out.linear_1", - "final_layer.adaln_modulation.2": "norm_out.linear_2", - "final_layer.linear": "proj_out", - } - - TRANSFORMER_SPECIAL_KEYS_REMAP_COSMOS_2_0 = { - "accum_video_sample_counter": remove_keys_, - "accum_image_sample_counter": remove_keys_, - "accum_iteration": remove_keys_, - "accum_train_in_hours": remove_keys_, - "pos_embedder.seq": remove_keys_, - "pos_embedder.dim_spatial_range": remove_keys_, - "pos_embedder.dim_temporal_range": remove_keys_, - "_extra_state": remove_keys_, - } - - PREFIX_KEY = "net." - if "net.blocks.block1.blocks.0.block.attn.to_q.0.weight" in checkpoint: - TRANSFORMER_KEYS_RENAME_DICT = TRANSFORMER_KEYS_RENAME_DICT_COSMOS_1_0 - TRANSFORMER_SPECIAL_KEYS_REMAP = TRANSFORMER_SPECIAL_KEYS_REMAP_COSMOS_1_0 - else: - TRANSFORMER_KEYS_RENAME_DICT = TRANSFORMER_KEYS_RENAME_DICT_COSMOS_2_0 - TRANSFORMER_SPECIAL_KEYS_REMAP = TRANSFORMER_SPECIAL_KEYS_REMAP_COSMOS_2_0 - - state_dict_keys = list(converted_state_dict.keys()) - for key in state_dict_keys: - new_key = key[:] - if new_key.startswith(PREFIX_KEY): - new_key = new_key.removeprefix(PREFIX_KEY) - for replace_key, rename_key in TRANSFORMER_KEYS_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - converted_state_dict[new_key] = converted_state_dict.pop(key) - - state_dict_keys = list(converted_state_dict.keys()) - for key in state_dict_keys: - for special_key, handler_fn_inplace in TRANSFORMER_SPECIAL_KEYS_REMAP.items(): - if special_key not in key: - continue - handler_fn_inplace(key, converted_state_dict) - - return converted_state_dict - - -def convert_flux2_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): - FLUX2_TRANSFORMER_KEYS_RENAME_DICT = { - # Image and text input projections - "img_in": "x_embedder", - "txt_in": "context_embedder", - # Timestep and guidance embeddings - "time_in.in_layer": "time_guidance_embed.timestep_embedder.linear_1", - "time_in.out_layer": "time_guidance_embed.timestep_embedder.linear_2", - "guidance_in.in_layer": "time_guidance_embed.guidance_embedder.linear_1", - "guidance_in.out_layer": "time_guidance_embed.guidance_embedder.linear_2", - # Modulation parameters - "double_stream_modulation_img.lin": "double_stream_modulation_img.linear", - "double_stream_modulation_txt.lin": "double_stream_modulation_txt.linear", - "single_stream_modulation.lin": "single_stream_modulation.linear", - # Final output layer - # "final_layer.adaLN_modulation.1": "norm_out.linear", # Handle separately since we need to swap mod params - "final_layer.linear": "proj_out", - } - - FLUX2_TRANSFORMER_ADA_LAYER_NORM_KEY_MAP = { - "final_layer.adaLN_modulation.1": "norm_out.linear", - } - - FLUX2_TRANSFORMER_DOUBLE_BLOCK_KEY_MAP = { - # Handle fused QKV projections separately as we need to break into Q, K, V projections - "img_attn.norm.query_norm": "attn.norm_q", - "img_attn.norm.key_norm": "attn.norm_k", - "img_attn.proj": "attn.to_out.0", - "img_mlp.0": "ff.linear_in", - "img_mlp.2": "ff.linear_out", - "txt_attn.norm.query_norm": "attn.norm_added_q", - "txt_attn.norm.key_norm": "attn.norm_added_k", - "txt_attn.proj": "attn.to_add_out", - "txt_mlp.0": "ff_context.linear_in", - "txt_mlp.2": "ff_context.linear_out", - } - - FLUX2_TRANSFORMER_SINGLE_BLOCK_KEY_MAP = { - "linear1": "attn.to_qkv_mlp_proj", - "linear2": "attn.to_out", - "norm.query_norm": "attn.norm_q", - "norm.key_norm": "attn.norm_k", - } - - def convert_flux2_single_stream_blocks(key: str, state_dict: dict[str, object]) -> None: - # Skip if not a weight, bias, or scale - if ".weight" not in key and ".bias" not in key and ".scale" not in key: - return - - # Mapping: - # - single_blocks.{N}.linear1 --> single_transformer_blocks.{N}.attn.to_qkv_mlp_proj - # - single_blocks.{N}.linear2 --> single_transformer_blocks.{N}.attn.to_out - # - single_blocks.{N}.norm.query_norm.scale --> single_transformer_blocks.{N}.attn.norm_q.weight - # - single_blocks.{N}.norm.key_norm.scale --> single_transformer_blocks.{N}.attn.norm_k.weight - new_prefix = "single_transformer_blocks" - if "single_blocks." in key: - parts = key.split(".") - block_idx = parts[1] - within_block_name = ".".join(parts[2:-1]) - param_type = parts[-1] - - if param_type == "scale": - param_type = "weight" - - new_within_block_name = FLUX2_TRANSFORMER_SINGLE_BLOCK_KEY_MAP[within_block_name] - new_key = ".".join([new_prefix, block_idx, new_within_block_name, param_type]) - - param = state_dict.pop(key) - state_dict[new_key] = param - - return - - def convert_ada_layer_norm_weights(key: str, state_dict: dict[str, object]) -> None: - # Skip if not a weight - if ".weight" not in key: - return - - # If adaLN_modulation is in the key, swap scale and shift parameters - # Original implementation is (shift, scale); diffusers implementation is (scale, shift) - if "adaLN_modulation" in key: - key_without_param_type, param_type = key.rsplit(".", maxsplit=1) - # Assume all such keys are in the AdaLayerNorm key map - new_key_without_param_type = FLUX2_TRANSFORMER_ADA_LAYER_NORM_KEY_MAP[key_without_param_type] - new_key = ".".join([new_key_without_param_type, param_type]) - - swapped_weight = swap_scale_shift(state_dict.pop(key), 0) - state_dict[new_key] = swapped_weight - - return - - def convert_flux2_double_stream_blocks(key: str, state_dict: dict[str, object]) -> None: - # Skip if not a weight, bias, or scale - if ".weight" not in key and ".bias" not in key and ".scale" not in key: - return - - new_prefix = "transformer_blocks" - if "double_blocks." in key: - parts = key.split(".") - block_idx = parts[1] - modality_block_name = parts[2] # img_attn, img_mlp, txt_attn, txt_mlp - within_block_name = ".".join(parts[2:-1]) - param_type = parts[-1] - - if param_type == "scale": - param_type = "weight" - - if "qkv" in within_block_name: - fused_qkv_weight = state_dict.pop(key) - to_q_weight, to_k_weight, to_v_weight = torch.chunk(fused_qkv_weight, 3, dim=0) - if "img" in modality_block_name: - # double_blocks.{N}.img_attn.qkv --> transformer_blocks.{N}.attn.{to_q|to_k|to_v} - to_q_weight, to_k_weight, to_v_weight = torch.chunk(fused_qkv_weight, 3, dim=0) - new_q_name = "attn.to_q" - new_k_name = "attn.to_k" - new_v_name = "attn.to_v" - elif "txt" in modality_block_name: - # double_blocks.{N}.txt_attn.qkv --> transformer_blocks.{N}.attn.{add_q_proj|add_k_proj|add_v_proj} - to_q_weight, to_k_weight, to_v_weight = torch.chunk(fused_qkv_weight, 3, dim=0) - new_q_name = "attn.add_q_proj" - new_k_name = "attn.add_k_proj" - new_v_name = "attn.add_v_proj" - new_q_key = ".".join([new_prefix, block_idx, new_q_name, param_type]) - new_k_key = ".".join([new_prefix, block_idx, new_k_name, param_type]) - new_v_key = ".".join([new_prefix, block_idx, new_v_name, param_type]) - state_dict[new_q_key] = to_q_weight - state_dict[new_k_key] = to_k_weight - state_dict[new_v_key] = to_v_weight - else: - new_within_block_name = FLUX2_TRANSFORMER_DOUBLE_BLOCK_KEY_MAP[within_block_name] - new_key = ".".join([new_prefix, block_idx, new_within_block_name, param_type]) - - param = state_dict.pop(key) - state_dict[new_key] = param - return - - def update_state_dict(state_dict: dict[str, object], old_key: str, new_key: str) -> None: - state_dict[new_key] = state_dict.pop(old_key) - - TRANSFORMER_SPECIAL_KEYS_REMAP = { - "adaLN_modulation": convert_ada_layer_norm_weights, - "double_blocks": convert_flux2_double_stream_blocks, - "single_blocks": convert_flux2_single_stream_blocks, - } - - converted_state_dict = {key: checkpoint.pop(key) for key in list(checkpoint.keys())} - - # Handle official code --> diffusers key remapping via the remap dict - for key in list(converted_state_dict.keys()): - new_key = key[:] - for replace_key, rename_key in FLUX2_TRANSFORMER_KEYS_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - - update_state_dict(converted_state_dict, key, new_key) - - # Handle any special logic which can't be expressed by a simple 1:1 remapping with the handlers in - # special_keys_remap - for key in list(converted_state_dict.keys()): - for special_key, handler_fn_inplace in TRANSFORMER_SPECIAL_KEYS_REMAP.items(): - if special_key not in key: - continue - handler_fn_inplace(key, converted_state_dict) - - return converted_state_dict - - -def convert_z_image_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): - Z_IMAGE_KEYS_RENAME_DICT = { - "final_layer.": "all_final_layer.2-1.", - "x_embedder.": "all_x_embedder.2-1.", - ".attention.out.bias": ".attention.to_out.0.bias", - ".attention.k_norm.weight": ".attention.norm_k.weight", - ".attention.q_norm.weight": ".attention.norm_q.weight", - ".attention.out.weight": ".attention.to_out.0.weight", - "model.diffusion_model.": "", - } - - def convert_z_image_fused_attention(key: str, state_dict: dict[str, object]) -> None: - if ".attention.qkv.weight" not in key: - return - - fused_qkv_weight = state_dict.pop(key) - to_q_weight, to_k_weight, to_v_weight = torch.chunk(fused_qkv_weight, 3, dim=0) - new_q_name = key.replace(".attention.qkv.weight", ".attention.to_q.weight") - new_k_name = key.replace(".attention.qkv.weight", ".attention.to_k.weight") - new_v_name = key.replace(".attention.qkv.weight", ".attention.to_v.weight") - - state_dict[new_q_name] = to_q_weight - state_dict[new_k_name] = to_k_weight - state_dict[new_v_name] = to_v_weight - return - - TRANSFORMER_SPECIAL_KEYS_REMAP = { - ".attention.qkv.weight": convert_z_image_fused_attention, - } - - def update_state_dict(state_dict: dict[str, object], old_key: str, new_key: str) -> None: - state_dict[new_key] = state_dict.pop(old_key) - - converted_state_dict = {key: checkpoint.pop(key) for key in list(checkpoint.keys())} - - # Handle single file --> diffusers key remapping via the remap dict - for key in list(converted_state_dict.keys()): - new_key = key[:] - for replace_key, rename_key in Z_IMAGE_KEYS_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - - update_state_dict(converted_state_dict, key, new_key) - - if "norm_final.weight" in converted_state_dict.keys(): - _ = converted_state_dict.pop("norm_final.weight") - - # Handle any special logic which can't be expressed by a simple 1:1 remapping with the handlers in - # special_keys_remap - for key in list(converted_state_dict.keys()): - for special_key, handler_fn_inplace in TRANSFORMER_SPECIAL_KEYS_REMAP.items(): - if special_key not in key: - continue - handler_fn_inplace(key, converted_state_dict) - - return converted_state_dict - - -def convert_z_image_controlnet_checkpoint_to_diffusers(checkpoint, config, **kwargs): - if config["add_control_noise_refiner"] is None: - return checkpoint - elif config["add_control_noise_refiner"] == "control_noise_refiner": - return checkpoint - elif config["add_control_noise_refiner"] == "control_layers": - converted_state_dict = { - key: checkpoint.pop(key) for key in list(checkpoint.keys()) if not key.startswith("control_noise_refiner.") - } - return converted_state_dict - else: - raise ValueError("Unknown Z-Image Turbo ControlNet type.") - - -def convert_ltx2_transformer_to_diffusers(checkpoint, **kwargs): - LTX_2_0_TRANSFORMER_KEYS_RENAME_DICT = { - # Transformer prefix - "model.diffusion_model.": "", - # Input Patchify Projections - "patchify_proj": "proj_in", - "audio_patchify_proj": "audio_proj_in", - # Modulation Parameters - # Handle adaln_single --> time_embed, audioln_single --> audio_time_embed separately as the original keys are - # substrings of the other modulation parameters below - "av_ca_video_scale_shift_adaln_single": "av_cross_attn_video_scale_shift", - "av_ca_a2v_gate_adaln_single": "av_cross_attn_video_a2v_gate", - "av_ca_audio_scale_shift_adaln_single": "av_cross_attn_audio_scale_shift", - "av_ca_v2a_gate_adaln_single": "av_cross_attn_audio_v2a_gate", - # Transformer Blocks - # Per-Block Cross Attention Modulation Parameters - "scale_shift_table_a2v_ca_video": "video_a2v_cross_attn_scale_shift_table", - "scale_shift_table_a2v_ca_audio": "audio_a2v_cross_attn_scale_shift_table", - # Attention QK Norms - "q_norm": "norm_q", - "k_norm": "norm_k", - } - - def update_state_dict_inplace(state_dict, old_key: str, new_key: str) -> None: - state_dict[new_key] = state_dict.pop(old_key) - - def remove_keys_inplace(key: str, state_dict) -> None: - state_dict.pop(key) - - def convert_ltx2_transformer_adaln_single(key: str, state_dict) -> None: - # Skip if not a weight, bias - if ".weight" not in key and ".bias" not in key: - return - - if key.startswith("adaln_single."): - new_key = key.replace("adaln_single.", "time_embed.") - param = state_dict.pop(key) - state_dict[new_key] = param - - if key.startswith("audio_adaln_single."): - new_key = key.replace("audio_adaln_single.", "audio_time_embed.") - param = state_dict.pop(key) - state_dict[new_key] = param - - return - - LTX_2_0_TRANSFORMER_SPECIAL_KEYS_REMAP = { - "video_embeddings_connector": remove_keys_inplace, - "audio_embeddings_connector": remove_keys_inplace, - "adaln_single": convert_ltx2_transformer_adaln_single, - } - - converted_state_dict = {key: checkpoint.pop(key) for key in list(checkpoint.keys())} - - # Handle official code --> diffusers key remapping via the remap dict - for key in list(converted_state_dict.keys()): - new_key = key[:] - for replace_key, rename_key in LTX_2_0_TRANSFORMER_KEYS_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - - update_state_dict_inplace(converted_state_dict, key, new_key) - - # Handle any special logic which can't be expressed by a simple 1:1 remapping with the handlers in - # special_keys_remap - for key in list(converted_state_dict.keys()): - for special_key, handler_fn_inplace in LTX_2_0_TRANSFORMER_SPECIAL_KEYS_REMAP.items(): - if special_key not in key: - continue - handler_fn_inplace(key, converted_state_dict) - - return converted_state_dict - - -def convert_ltx2_vae_to_diffusers(checkpoint, **kwargs): - LTX_2_0_VIDEO_VAE_RENAME_DICT = { - # Video VAE prefix - "vae.": "", - # Encoder - "down_blocks.0": "down_blocks.0", - "down_blocks.1": "down_blocks.0.downsamplers.0", - "down_blocks.2": "down_blocks.1", - "down_blocks.3": "down_blocks.1.downsamplers.0", - "down_blocks.4": "down_blocks.2", - "down_blocks.5": "down_blocks.2.downsamplers.0", - "down_blocks.6": "down_blocks.3", - "down_blocks.7": "down_blocks.3.downsamplers.0", - "down_blocks.8": "mid_block", - # Decoder - "up_blocks.0": "mid_block", - "up_blocks.1": "up_blocks.0.upsamplers.0", - "up_blocks.2": "up_blocks.0", - "up_blocks.3": "up_blocks.1.upsamplers.0", - "up_blocks.4": "up_blocks.1", - "up_blocks.5": "up_blocks.2.upsamplers.0", - "up_blocks.6": "up_blocks.2", - # Common - # For all 3D ResNets - "res_blocks": "resnets", - "per_channel_statistics.mean-of-means": "latents_mean", - "per_channel_statistics.std-of-means": "latents_std", - } - - def update_state_dict_inplace(state_dict, old_key: str, new_key: str) -> None: - state_dict[new_key] = state_dict.pop(old_key) - - def remove_keys_inplace(key: str, state_dict) -> None: - state_dict.pop(key) - - LTX_2_0_VAE_SPECIAL_KEYS_REMAP = { - "per_channel_statistics.channel": remove_keys_inplace, - "per_channel_statistics.mean-of-stds": remove_keys_inplace, - } - - converted_state_dict = {key: checkpoint.pop(key) for key in list(checkpoint.keys())} - - # Handle official code --> diffusers key remapping via the remap dict - for key in list(converted_state_dict.keys()): - new_key = key[:] - for replace_key, rename_key in LTX_2_0_VIDEO_VAE_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - - update_state_dict_inplace(converted_state_dict, key, new_key) - - # Handle any special logic which can't be expressed by a simple 1:1 remapping with the handlers in - # special_keys_remap - for key in list(converted_state_dict.keys()): - for special_key, handler_fn_inplace in LTX_2_0_VAE_SPECIAL_KEYS_REMAP.items(): - if special_key not in key: - continue - handler_fn_inplace(key, converted_state_dict) - - return converted_state_dict - - -def convert_ltx2_audio_vae_to_diffusers(checkpoint, **kwargs): - LTX_2_0_AUDIO_VAE_RENAME_DICT = { - # Audio VAE prefix - "audio_vae.": "", - "per_channel_statistics.mean-of-means": "latents_mean", - "per_channel_statistics.std-of-means": "latents_std", - } - - def update_state_dict_inplace(state_dict, old_key: str, new_key: str) -> None: - state_dict[new_key] = state_dict.pop(old_key) - - converted_state_dict = {key: checkpoint.pop(key) for key in list(checkpoint.keys())} - - # Handle official code --> diffusers key remapping via the remap dict - for key in list(converted_state_dict.keys()): - new_key = key[:] - for replace_key, rename_key in LTX_2_0_AUDIO_VAE_RENAME_DICT.items(): - new_key = new_key.replace(replace_key, rename_key) - - update_state_dict_inplace(converted_state_dict, key, new_key) - - return converted_state_dict - - -def convert_ernie_image_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): - keys = list(checkpoint.keys()) - - for k in keys: - if "model.diffusion_model." in k: - checkpoint[k.replace("model.diffusion_model.", "")] = checkpoint.pop(k) - - return checkpoint diff --git a/diffusers/loaders/textual_inversion.py b/diffusers/loaders/textual_inversion.py deleted file mode 100644 index 72ae0c169b1adafa906e73f92dc6559b114bfaa1..0000000000000000000000000000000000000000 --- a/diffusers/loaders/textual_inversion.py +++ /dev/null @@ -1,605 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from __future__ import annotations - -import json - -import safetensors -import torch -from huggingface_hub.utils import validate_hf_hub_args -from tokenizers import Tokenizer as TokenizerFast -from torch import nn - -from ..models.modeling_utils import load_state_dict -from ..utils import ( - _get_model_file, - is_accelerate_available, - is_transformers_available, - logging, -) - - -if is_transformers_available(): - from transformers import PreTrainedModel, PreTrainedTokenizer - -if is_accelerate_available(): - from accelerate.hooks import AlignDevicesHook, CpuOffload, remove_hook_from_module - -logger = logging.get_logger(__name__) - -TEXT_INVERSION_NAME = "learned_embeds.bin" -TEXT_INVERSION_NAME_SAFE = "learned_embeds.safetensors" - - -@validate_hf_hub_args -def load_textual_inversion_state_dicts(pretrained_model_name_or_paths, **kwargs): - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - hf_token = kwargs.pop("hf_token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = { - "file_type": "text_inversion", - "framework": "pytorch", - } - state_dicts = [] - for pretrained_model_name_or_path in pretrained_model_name_or_paths: - if not isinstance(pretrained_model_name_or_path, (dict, torch.Tensor)): - # 3.1. Load textual inversion file - model_file = None - - # Let's first try to load .safetensors weights - if (use_safetensors and weight_name is None) or ( - weight_name is not None and weight_name.endswith(".safetensors") - ): - try: - model_file = _get_model_file( - pretrained_model_name_or_path, - weights_name=weight_name or TEXT_INVERSION_NAME_SAFE, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=hf_token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - ) - state_dict = safetensors.torch.load_file(model_file, device="cpu") - except Exception as e: - if not allow_pickle: - raise e - - model_file = None - - if model_file is None: - model_file = _get_model_file( - pretrained_model_name_or_path, - weights_name=weight_name or TEXT_INVERSION_NAME, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=hf_token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - ) - state_dict = load_state_dict(model_file) - else: - state_dict = pretrained_model_name_or_path - - state_dicts.append(state_dict) - - return state_dicts - - -class TextualInversionLoaderMixin: - r""" - Load Textual Inversion tokens and embeddings to the tokenizer and text encoder. - """ - - def maybe_convert_prompt(self, prompt: str | list[str], tokenizer: "PreTrainedTokenizer"): # noqa: F821 - r""" - Processes prompts that include a special token corresponding to a multi-vector textual inversion embedding to - be replaced with multiple special tokens each corresponding to one of the vectors. If the prompt has no textual - inversion token or if the textual inversion token is a single vector, the input prompt is returned. - - Parameters: - prompt (`str` or list of `str`): - The prompt or prompts to guide the image generation. - tokenizer (`PreTrainedTokenizer`): - The tokenizer responsible for encoding the prompt into input tokens. - - Returns: - `str` or list of `str`: The converted prompt - """ - if not isinstance(prompt, list): - prompts = [prompt] - else: - prompts = prompt - - prompts = [self._maybe_convert_prompt(p, tokenizer) for p in prompts] - - if not isinstance(prompt, list): - return prompts[0] - - return prompts - - def _maybe_convert_prompt(self, prompt: str, tokenizer: "PreTrainedTokenizer"): # noqa: F821 - r""" - Maybe convert a prompt into a "multi vector"-compatible prompt. If the prompt includes a token that corresponds - to a multi-vector textual inversion embedding, this function will process the prompt so that the special token - is replaced with multiple special tokens each corresponding to one of the vectors. If the prompt has no textual - inversion token or a textual inversion token that is a single vector, the input prompt is simply returned. - - Parameters: - prompt (`str`): - The prompt to guide the image generation. - tokenizer (`PreTrainedTokenizer`): - The tokenizer responsible for encoding the prompt into input tokens. - - Returns: - `str`: The converted prompt - """ - tokens = tokenizer.tokenize(prompt) - unique_tokens = set(tokens) - for token in unique_tokens: - if token in tokenizer.added_tokens_encoder: - replacement = token - i = 1 - while f"{token}_{i}" in tokenizer.added_tokens_encoder: - replacement += f" {token}_{i}" - i += 1 - - prompt = prompt.replace(token, replacement) - - return prompt - - def _check_text_inv_inputs(self, tokenizer, text_encoder, pretrained_model_name_or_paths, tokens): - if tokenizer is None: - raise ValueError( - f"{self.__class__.__name__} requires `self.tokenizer` or passing a `tokenizer` of type `PreTrainedTokenizer` for calling" - f" `{self.load_textual_inversion.__name__}`" - ) - - if text_encoder is None: - raise ValueError( - f"{self.__class__.__name__} requires `self.text_encoder` or passing a `text_encoder` of type `PreTrainedModel` for calling" - f" `{self.load_textual_inversion.__name__}`" - ) - - if len(pretrained_model_name_or_paths) > 1 and len(pretrained_model_name_or_paths) != len(tokens): - raise ValueError( - f"You have passed a list of models of length {len(pretrained_model_name_or_paths)}, and list of tokens of length {len(tokens)} " - f"Make sure both lists have the same length." - ) - - valid_tokens = [t for t in tokens if t is not None] - if len(set(valid_tokens)) < len(valid_tokens): - raise ValueError(f"You have passed a list of tokens that contains duplicates: {tokens}") - - @staticmethod - def _retrieve_tokens_and_embeddings(tokens, state_dicts, tokenizer): - all_tokens = [] - all_embeddings = [] - for state_dict, token in zip(state_dicts, tokens): - if isinstance(state_dict, torch.Tensor): - if token is None: - raise ValueError( - "You are trying to load a textual inversion embedding that has been saved as a PyTorch tensor. Make sure to pass the name of the corresponding token in this case: `token=...`." - ) - loaded_token = token - embedding = state_dict - elif len(state_dict) == 1: - # diffusers - loaded_token, embedding = next(iter(state_dict.items())) - elif "string_to_param" in state_dict: - # A1111 - loaded_token = state_dict["name"] - embedding = state_dict["string_to_param"]["*"] - else: - raise ValueError( - f"Loaded state dictionary is incorrect: {state_dict}. \n\n" - "Please verify that the loaded state dictionary of the textual embedding either only has a single key or includes the `string_to_param`" - " input key." - ) - - if token is not None and loaded_token != token: - logger.info(f"The loaded token: {loaded_token} is overwritten by the passed token {token}.") - else: - token = loaded_token - - if token in tokenizer.get_vocab(): - raise ValueError( - f"Token {token} already in tokenizer vocabulary. Please choose a different token name or remove {token} and embedding from the tokenizer and text encoder." - ) - - all_tokens.append(token) - all_embeddings.append(embedding) - - return all_tokens, all_embeddings - - @staticmethod - def _extend_tokens_and_embeddings(tokens, embeddings, tokenizer): - all_tokens = [] - all_embeddings = [] - - for embedding, token in zip(embeddings, tokens): - if f"{token}_1" in tokenizer.get_vocab(): - multi_vector_tokens = [token] - i = 1 - while f"{token}_{i}" in tokenizer.added_tokens_encoder: - multi_vector_tokens.append(f"{token}_{i}") - i += 1 - - raise ValueError( - f"Multi-vector Token {multi_vector_tokens} already in tokenizer vocabulary. Please choose a different token name or remove the {multi_vector_tokens} and embedding from the tokenizer and text encoder." - ) - - is_multi_vector = len(embedding.shape) > 1 and embedding.shape[0] > 1 - if is_multi_vector: - all_tokens += [token] + [f"{token}_{i}" for i in range(1, embedding.shape[0])] - all_embeddings += [e for e in embedding] # noqa: C416 - else: - all_tokens += [token] - all_embeddings += [embedding[0]] if len(embedding.shape) > 1 else [embedding] - - return all_tokens, all_embeddings - - @validate_hf_hub_args - def load_textual_inversion( - self, - pretrained_model_name_or_path: str | list[str] | dict[str, torch.Tensor] | list[dict[str, torch.Tensor]], - token: str | list[str] | None = None, - tokenizer: "PreTrainedTokenizer" | None = None, # noqa: F821 - text_encoder: "PreTrainedModel" | None = None, # noqa: F821 - **kwargs, - ): - r""" - Load Textual Inversion embeddings into the text encoder of [`StableDiffusionPipeline`] (both 🤗 Diffusers and - Automatic1111 formats are supported). - - Parameters: - pretrained_model_name_or_path (`str` or `os.PathLike` or `list[str or os.PathLike]` or `Dict` or `list[Dict]`): - Can be either one of the following or a list of them: - - - A string, the *model id* (for example `sd-concepts-library/low-poly-hd-logos-icons`) of a - pretrained model hosted on the Hub. - - A path to a *directory* (for example `./my_text_inversion_directory/`) containing the textual - inversion weights. - - A path to a *file* (for example `./my_text_inversions.pt`) containing textual inversion weights. - - A [torch state - dict](https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict). - - token (`str` or `list[str]`, *optional*): - Override the token to use for the textual inversion weights. If `pretrained_model_name_or_path` is a - list, then `token` must also be a list of equal length. - text_encoder ([`~transformers.CLIPTextModel`], *optional*): - Frozen text-encoder ([clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14)). - If not specified, function will take self.tokenizer. - tokenizer ([`~transformers.CLIPTokenizer`], *optional*): - A `CLIPTokenizer` to tokenize text. If not specified, function will take self.tokenizer. - weight_name (`str`, *optional*): - Name of a custom weight file. This should be used when: - - - The saved textual inversion file is in 🤗 Diffusers format, but was saved under a specific weight - name such as `text_inv.bin`. - - The saved textual inversion file is in the Automatic1111 format. - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - hf_token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - subfolder (`str`, *optional*, defaults to `""`): - The subfolder location of a model file within a larger model repository on the Hub or locally. - mirror (`str`, *optional*): - Mirror source to resolve accessibility issues if you're downloading a model in China. We do not - guarantee the timeliness or safety of the source, and you should refer to the mirror site for more - information. - - Example: - - To load a Textual Inversion embedding vector in 🤗 Diffusers format: - - ```py - from diffusers import StableDiffusionPipeline - import torch - - model_id = "stable-diffusion-v1-5/stable-diffusion-v1-5" - pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=torch.float16).to("cuda") - - pipe.load_textual_inversion("sd-concepts-library/cat-toy") - - prompt = "A backpack" - - image = pipe(prompt, num_inference_steps=50).images[0] - image.save("cat-backpack.png") - ``` - - To load a Textual Inversion embedding vector in Automatic1111 format, make sure to download the vector first - (for example from [civitAI](https://civitai.com/models/3036?modelVersionId=9857)) and then load the vector - locally: - - ```py - from diffusers import StableDiffusionPipeline - import torch - - model_id = "stable-diffusion-v1-5/stable-diffusion-v1-5" - pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=torch.float16).to("cuda") - - pipe.load_textual_inversion("./charturnerv2.pt", token="charturnerv2") - - prompt = "charturnerv2, multiple views of the same character in the same outfit, a character turnaround of a woman wearing a black jacket and red shirt, best quality, intricate details." - - image = pipe(prompt, num_inference_steps=50).images[0] - image.save("character.png") - ``` - - """ - # 1. Set correct tokenizer and text encoder - tokenizer = tokenizer or getattr(self, "tokenizer", None) - text_encoder = text_encoder or getattr(self, "text_encoder", None) - - # 2. Normalize inputs - pretrained_model_name_or_paths = ( - [pretrained_model_name_or_path] - if not isinstance(pretrained_model_name_or_path, list) - else pretrained_model_name_or_path - ) - tokens = [token] if not isinstance(token, list) else token - if tokens[0] is None: - tokens = tokens * len(pretrained_model_name_or_paths) - - # 3. Check inputs - self._check_text_inv_inputs(tokenizer, text_encoder, pretrained_model_name_or_paths, tokens) - - # 4. Load state dicts of textual embeddings - state_dicts = load_textual_inversion_state_dicts(pretrained_model_name_or_paths, **kwargs) - - # 4.1 Handle the special case when state_dict is a tensor that contains n embeddings for n tokens - if len(tokens) > 1 and len(state_dicts) == 1: - if isinstance(state_dicts[0], torch.Tensor): - state_dicts = list(state_dicts[0]) - if len(tokens) != len(state_dicts): - raise ValueError( - f"You have passed a state_dict contains {len(state_dicts)} embeddings, and list of tokens of length {len(tokens)} " - f"Make sure both have the same length." - ) - - # 4. Retrieve tokens and embeddings - tokens, embeddings = self._retrieve_tokens_and_embeddings(tokens, state_dicts, tokenizer) - - # 5. Extend tokens and embeddings for multi vector - tokens, embeddings = self._extend_tokens_and_embeddings(tokens, embeddings, tokenizer) - - # 6. Make sure all embeddings have the correct size - expected_emb_dim = text_encoder.get_input_embeddings().weight.shape[-1] - if any(expected_emb_dim != emb.shape[-1] for emb in embeddings): - raise ValueError( - "Loaded embeddings are of incorrect shape. Expected each textual inversion embedding " - "to be of shape {input_embeddings.shape[-1]}, but are {embeddings.shape[-1]} " - ) - - # 7. Now we can be sure that loading the embedding matrix works - # < Unsafe code: - - # 7.1 Offload all hooks in case the pipeline was cpu offloaded before make sure, we offload and onload again - is_model_cpu_offload = False - is_sequential_cpu_offload = False - if self.hf_device_map is None: - for _, component in self.components.items(): - if isinstance(component, nn.Module): - if hasattr(component, "_hf_hook"): - is_model_cpu_offload = isinstance(getattr(component, "_hf_hook"), CpuOffload) - is_sequential_cpu_offload = ( - isinstance(getattr(component, "_hf_hook"), AlignDevicesHook) - or hasattr(component._hf_hook, "hooks") - and isinstance(component._hf_hook.hooks[0], AlignDevicesHook) - ) - logger.info( - "Accelerate hooks detected. Since you have called `load_textual_inversion()`, the previous hooks will be first removed. Then the textual inversion parameters will be loaded and the hooks will be applied again." - ) - if is_sequential_cpu_offload or is_model_cpu_offload: - remove_hook_from_module(component, recurse=is_sequential_cpu_offload) - - # 7.2 save expected device and dtype - device = text_encoder.device - dtype = text_encoder.dtype - - # 7.3 Increase token embedding matrix - text_encoder.resize_token_embeddings(len(tokenizer) + len(tokens)) - input_embeddings = text_encoder.get_input_embeddings().weight - - # 7.4 Load token and embedding - for token, embedding in zip(tokens, embeddings): - # add tokens and get ids - tokenizer.add_tokens(token) - token_id = tokenizer.convert_tokens_to_ids(token) - input_embeddings.data[token_id] = embedding - logger.info(f"Loaded textual inversion embedding for {token}.") - - input_embeddings.to(dtype=dtype, device=device) - - # 7.5 Offload the model again - if is_model_cpu_offload: - self.enable_model_cpu_offload(device=device) - elif is_sequential_cpu_offload: - self.enable_sequential_cpu_offload(device=device) - - # / Unsafe Code > - - def unload_textual_inversion( - self, - tokens: str | list[str] | None = None, - tokenizer: "PreTrainedTokenizer" | None = None, - text_encoder: "PreTrainedModel" | None = None, - ): - r""" - Unload Textual Inversion embeddings from the text encoder of [`StableDiffusionPipeline`] - - Example: - ```py - from diffusers import AutoPipelineForText2Image - import torch - - pipeline = AutoPipelineForText2Image.from_pretrained("stable-diffusion-v1-5/stable-diffusion-v1-5") - - # Example 1 - pipeline.load_textual_inversion("sd-concepts-library/gta5-artwork") - pipeline.load_textual_inversion("sd-concepts-library/moeb-style") - - # Remove all token embeddings - pipeline.unload_textual_inversion() - - # Example 2 - pipeline.load_textual_inversion("sd-concepts-library/moeb-style") - pipeline.load_textual_inversion("sd-concepts-library/gta5-artwork") - - # Remove just one token - pipeline.unload_textual_inversion("") - - # Example 3: unload from SDXL - pipeline = AutoPipelineForText2Image.from_pretrained("stabilityai/stable-diffusion-xl-base-1.0") - embedding_path = hf_hub_download( - repo_id="linoyts/web_y2k", filename="web_y2k_emb.safetensors", repo_type="model" - ) - - # load embeddings to the text encoders - state_dict = load_file(embedding_path) - - # load embeddings of text_encoder 1 (CLIP ViT-L/14) - pipeline.load_textual_inversion( - state_dict["clip_l"], - tokens=["", ""], - text_encoder=pipeline.text_encoder, - tokenizer=pipeline.tokenizer, - ) - # load embeddings of text_encoder 2 (CLIP ViT-G/14) - pipeline.load_textual_inversion( - state_dict["clip_g"], - tokens=["", ""], - text_encoder=pipeline.text_encoder_2, - tokenizer=pipeline.tokenizer_2, - ) - - # Unload explicitly from both text encoders and tokenizers - pipeline.unload_textual_inversion( - tokens=["", ""], text_encoder=pipeline.text_encoder, tokenizer=pipeline.tokenizer - ) - pipeline.unload_textual_inversion( - tokens=["", ""], text_encoder=pipeline.text_encoder_2, tokenizer=pipeline.tokenizer_2 - ) - ``` - """ - - tokenizer = tokenizer or getattr(self, "tokenizer", None) - text_encoder = text_encoder or getattr(self, "text_encoder", None) - - # Get textual inversion tokens and ids - token_ids = [] - last_special_token_id = None - - if tokens: - if isinstance(tokens, str): - tokens = [tokens] - for added_token_id, added_token in tokenizer.added_tokens_decoder.items(): - if not added_token.special: - if added_token.content in tokens: - token_ids.append(added_token_id) - else: - last_special_token_id = added_token_id - if len(token_ids) == 0: - raise ValueError("No tokens to remove found") - else: - tokens = [] - for added_token_id, added_token in tokenizer.added_tokens_decoder.items(): - if not added_token.special: - token_ids.append(added_token_id) - tokens.append(added_token.content) - else: - last_special_token_id = added_token_id - - # Fast tokenizers (v5+) - if hasattr(tokenizer, "_tokenizer"): - # Fast tokenizers: serialize, filter tokens, reload - tokenizer_json = json.loads(tokenizer._tokenizer.to_str()) - new_id = last_special_token_id + 1 - filtered = [] - for tok in tokenizer_json.get("added_tokens", []): - if tok.get("content") in set(tokens): - continue - if not tok.get("special", False): - tok["id"] = new_id - new_id += 1 - filtered.append(tok) - tokenizer_json["added_tokens"] = filtered - tokenizer._tokenizer = TokenizerFast.from_str(json.dumps(tokenizer_json)) - else: - # Slow tokenizers - for token_id, token_to_remove in zip(token_ids, tokens): - del tokenizer._added_tokens_decoder[token_id] - del tokenizer._added_tokens_encoder[token_to_remove] - - key_id = 1 - for token_id in list(tokenizer.added_tokens_decoder.keys()): - if token_id > last_special_token_id and token_id > last_special_token_id + key_id: - token = tokenizer._added_tokens_decoder[token_id] - tokenizer._added_tokens_decoder[last_special_token_id + key_id] = token - del tokenizer._added_tokens_decoder[token_id] - tokenizer._added_tokens_encoder[token.content] = last_special_token_id + key_id - key_id += 1 - if hasattr(tokenizer, "_update_trie"): - tokenizer._update_trie() - if hasattr(tokenizer, "_update_total_vocab_size"): - tokenizer._update_total_vocab_size() - - # Delete from text encoder - text_embedding_dim = text_encoder.get_input_embeddings().embedding_dim - temp_text_embedding_weights = text_encoder.get_input_embeddings().weight - text_embedding_weights = temp_text_embedding_weights[: last_special_token_id + 1] - to_append = [] - for i in range(last_special_token_id + 1, temp_text_embedding_weights.shape[0]): - if i not in token_ids: - to_append.append(temp_text_embedding_weights[i].unsqueeze(0)) - if len(to_append) > 0: - to_append = torch.cat(to_append, dim=0) - text_embedding_weights = torch.cat([text_embedding_weights, to_append], dim=0) - text_embeddings_filtered = nn.Embedding(text_embedding_weights.shape[0], text_embedding_dim) - text_embeddings_filtered.weight.data = text_embedding_weights - text_encoder.set_input_embeddings(text_embeddings_filtered) diff --git a/diffusers/loaders/transformer_flux.py b/diffusers/loaders/transformer_flux.py deleted file mode 100644 index 632f6601a6f697b616180075f101346a30cdf225..0000000000000000000000000000000000000000 --- a/diffusers/loaders/transformer_flux.py +++ /dev/null @@ -1,179 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from contextlib import nullcontext - -from ..models.embeddings import ( - ImageProjection, - MultiIPAdapterImageProjection, -) -from ..models.model_loading_utils import load_model_dict_into_meta -from ..models.modeling_utils import _LOW_CPU_MEM_USAGE_DEFAULT -from ..utils import is_accelerate_available, is_torch_version, logging -from ..utils.torch_utils import empty_device_cache - - -if is_accelerate_available(): - pass - -logger = logging.get_logger(__name__) - - -class FluxTransformer2DLoadersMixin: - """ - Load layers into a [`FluxTransformer2DModel`]. - """ - - def _convert_ip_adapter_image_proj_to_diffusers(self, state_dict, low_cpu_mem_usage=_LOW_CPU_MEM_USAGE_DEFAULT): - if low_cpu_mem_usage: - if is_accelerate_available(): - from accelerate import init_empty_weights - - else: - low_cpu_mem_usage = False - logger.warning( - "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the" - " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install" - " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip" - " install accelerate\n```\n." - ) - - if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `low_cpu_mem_usage=False`." - ) - - updated_state_dict = {} - image_projection = None - init_context = init_empty_weights if low_cpu_mem_usage else nullcontext - - if "proj.weight" in state_dict: - # IP-Adapter - num_image_text_embeds = 4 - if state_dict["proj.weight"].shape[0] == 65536: - num_image_text_embeds = 16 - clip_embeddings_dim = state_dict["proj.weight"].shape[-1] - cross_attention_dim = state_dict["proj.weight"].shape[0] // num_image_text_embeds - - with init_context(): - image_projection = ImageProjection( - cross_attention_dim=cross_attention_dim, - image_embed_dim=clip_embeddings_dim, - num_image_text_embeds=num_image_text_embeds, - ) - - for key, value in state_dict.items(): - diffusers_name = key.replace("proj", "image_embeds") - updated_state_dict[diffusers_name] = value - - if not low_cpu_mem_usage: - image_projection.load_state_dict(updated_state_dict, strict=True) - else: - device_map = {"": self.device} - load_model_dict_into_meta(image_projection, updated_state_dict, device_map=device_map, dtype=self.dtype) - empty_device_cache() - - return image_projection - - def _convert_ip_adapter_attn_to_diffusers(self, state_dicts, low_cpu_mem_usage=_LOW_CPU_MEM_USAGE_DEFAULT): - from ..models.transformers.transformer_flux import FluxIPAdapterAttnProcessor - - if low_cpu_mem_usage: - if is_accelerate_available(): - from accelerate import init_empty_weights - - else: - low_cpu_mem_usage = False - logger.warning( - "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the" - " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install" - " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip" - " install accelerate\n```\n." - ) - - if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `low_cpu_mem_usage=False`." - ) - - # set ip-adapter cross-attention processors & load state_dict - attn_procs = {} - key_id = 0 - init_context = init_empty_weights if low_cpu_mem_usage else nullcontext - for name in self.attn_processors.keys(): - if name.startswith("single_transformer_blocks"): - attn_processor_class = self.attn_processors[name].__class__ - attn_procs[name] = attn_processor_class() - else: - cross_attention_dim = self.config.joint_attention_dim - hidden_size = self.inner_dim - attn_processor_class = FluxIPAdapterAttnProcessor - num_image_text_embeds = [] - for state_dict in state_dicts: - if "proj.weight" in state_dict["image_proj"]: - num_image_text_embed = 4 - if state_dict["image_proj"]["proj.weight"].shape[0] == 65536: - num_image_text_embed = 16 - # IP-Adapter - num_image_text_embeds += [num_image_text_embed] - - with init_context(): - attn_procs[name] = attn_processor_class( - hidden_size=hidden_size, - cross_attention_dim=cross_attention_dim, - scale=1.0, - num_tokens=num_image_text_embeds, - dtype=self.dtype, - device=self.device, - ) - - value_dict = {} - for i, state_dict in enumerate(state_dicts): - value_dict.update({f"to_k_ip.{i}.weight": state_dict["ip_adapter"][f"{key_id}.to_k_ip.weight"]}) - value_dict.update({f"to_v_ip.{i}.weight": state_dict["ip_adapter"][f"{key_id}.to_v_ip.weight"]}) - value_dict.update({f"to_k_ip.{i}.bias": state_dict["ip_adapter"][f"{key_id}.to_k_ip.bias"]}) - value_dict.update({f"to_v_ip.{i}.bias": state_dict["ip_adapter"][f"{key_id}.to_v_ip.bias"]}) - - if not low_cpu_mem_usage: - attn_procs[name].load_state_dict(value_dict) - else: - device_map = {"": self.device} - dtype = self.dtype - load_model_dict_into_meta(attn_procs[name], value_dict, device_map=device_map, dtype=dtype) - - key_id += 1 - - empty_device_cache() - - return attn_procs - - def _load_ip_adapter_weights(self, state_dicts, low_cpu_mem_usage=_LOW_CPU_MEM_USAGE_DEFAULT): - if not isinstance(state_dicts, list): - state_dicts = [state_dicts] - - self.encoder_hid_proj = None - - attn_procs = self._convert_ip_adapter_attn_to_diffusers(state_dicts, low_cpu_mem_usage=low_cpu_mem_usage) - self.set_attn_processor(attn_procs) - - image_projection_layers = [] - for state_dict in state_dicts: - image_projection_layer = self._convert_ip_adapter_image_proj_to_diffusers( - state_dict["image_proj"], low_cpu_mem_usage=low_cpu_mem_usage - ) - image_projection_layers.append(image_projection_layer) - - self.encoder_hid_proj = MultiIPAdapterImageProjection(image_projection_layers) - self.config.encoder_hid_dim_type = "ip_image_proj" diff --git a/diffusers/loaders/transformer_sd3.py b/diffusers/loaders/transformer_sd3.py deleted file mode 100644 index 7fc90bf7dda42265cc833ee4d5b5cd06ab595f49..0000000000000000000000000000000000000000 --- a/diffusers/loaders/transformer_sd3.py +++ /dev/null @@ -1,174 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from contextlib import nullcontext - -from ..models.attention_processor import SD3IPAdapterJointAttnProcessor2_0 -from ..models.embeddings import IPAdapterTimeImageProjection -from ..models.model_loading_utils import load_model_dict_into_meta -from ..models.modeling_utils import _LOW_CPU_MEM_USAGE_DEFAULT -from ..utils import is_accelerate_available, is_torch_version, logging -from ..utils.torch_utils import empty_device_cache - - -logger = logging.get_logger(__name__) - - -class SD3Transformer2DLoadersMixin: - """Load IP-Adapters and LoRA layers into a `[SD3Transformer2DModel]`.""" - - def _convert_ip_adapter_attn_to_diffusers( - self, state_dict: dict, low_cpu_mem_usage: bool = _LOW_CPU_MEM_USAGE_DEFAULT - ) -> dict: - if low_cpu_mem_usage: - if is_accelerate_available(): - from accelerate import init_empty_weights - - else: - low_cpu_mem_usage = False - logger.warning( - "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the" - " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install" - " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip" - " install accelerate\n```\n." - ) - - if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `low_cpu_mem_usage=False`." - ) - - # IP-Adapter cross attention parameters - hidden_size = self.config.attention_head_dim * self.config.num_attention_heads - ip_hidden_states_dim = self.config.attention_head_dim * self.config.num_attention_heads - timesteps_emb_dim = state_dict["0.norm_ip.linear.weight"].shape[1] - - # Dict where key is transformer layer index, value is attention processor's state dict - # ip_adapter state dict keys example: "0.norm_ip.linear.weight" - layer_state_dict = {idx: {} for idx in range(len(self.attn_processors))} - for key, weights in state_dict.items(): - idx, name = key.split(".", maxsplit=1) - layer_state_dict[int(idx)][name] = weights - - # Create IP-Adapter attention processor & load state_dict - attn_procs = {} - init_context = init_empty_weights if low_cpu_mem_usage else nullcontext - for idx, name in enumerate(self.attn_processors.keys()): - with init_context(): - attn_procs[name] = SD3IPAdapterJointAttnProcessor2_0( - hidden_size=hidden_size, - ip_hidden_states_dim=ip_hidden_states_dim, - head_dim=self.config.attention_head_dim, - timesteps_emb_dim=timesteps_emb_dim, - ) - - if not low_cpu_mem_usage: - attn_procs[name].load_state_dict(layer_state_dict[idx], strict=True) - else: - device_map = {"": self.device} - load_model_dict_into_meta( - attn_procs[name], layer_state_dict[idx], device_map=device_map, dtype=self.dtype - ) - - empty_device_cache() - - return attn_procs - - def _convert_ip_adapter_image_proj_to_diffusers( - self, state_dict: dict, low_cpu_mem_usage: bool = _LOW_CPU_MEM_USAGE_DEFAULT - ) -> IPAdapterTimeImageProjection: - if low_cpu_mem_usage: - if is_accelerate_available(): - from accelerate import init_empty_weights - - else: - low_cpu_mem_usage = False - logger.warning( - "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the" - " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install" - " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip" - " install accelerate\n```\n." - ) - - if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `low_cpu_mem_usage=False`." - ) - - init_context = init_empty_weights if low_cpu_mem_usage else nullcontext - - # Convert to diffusers - updated_state_dict = {} - for key, value in state_dict.items(): - # InstantX/SD3.5-Large-IP-Adapter - if key.startswith("layers."): - idx = key.split(".")[1] - key = key.replace(f"layers.{idx}.0.norm1", f"layers.{idx}.ln0") - key = key.replace(f"layers.{idx}.0.norm2", f"layers.{idx}.ln1") - key = key.replace(f"layers.{idx}.0.to_q", f"layers.{idx}.attn.to_q") - key = key.replace(f"layers.{idx}.0.to_kv", f"layers.{idx}.attn.to_kv") - key = key.replace(f"layers.{idx}.0.to_out", f"layers.{idx}.attn.to_out.0") - key = key.replace(f"layers.{idx}.1.0", f"layers.{idx}.adaln_norm") - key = key.replace(f"layers.{idx}.1.1", f"layers.{idx}.ff.net.0.proj") - key = key.replace(f"layers.{idx}.1.3", f"layers.{idx}.ff.net.2") - key = key.replace(f"layers.{idx}.2.1", f"layers.{idx}.adaln_proj") - updated_state_dict[key] = value - - # Image projection parameters - embed_dim = updated_state_dict["proj_in.weight"].shape[1] - output_dim = updated_state_dict["proj_out.weight"].shape[0] - hidden_dim = updated_state_dict["proj_in.weight"].shape[0] - heads = updated_state_dict["layers.0.attn.to_q.weight"].shape[0] // 64 - num_queries = updated_state_dict["latents"].shape[1] - timestep_in_dim = updated_state_dict["time_embedding.linear_1.weight"].shape[1] - - # Image projection - with init_context(): - image_proj = IPAdapterTimeImageProjection( - embed_dim=embed_dim, - output_dim=output_dim, - hidden_dim=hidden_dim, - heads=heads, - num_queries=num_queries, - timestep_in_dim=timestep_in_dim, - ) - - if not low_cpu_mem_usage: - image_proj.load_state_dict(updated_state_dict, strict=True) - else: - device_map = {"": self.device} - load_model_dict_into_meta(image_proj, updated_state_dict, device_map=device_map, dtype=self.dtype) - empty_device_cache() - - return image_proj - - def _load_ip_adapter_weights(self, state_dict: dict, low_cpu_mem_usage: bool = _LOW_CPU_MEM_USAGE_DEFAULT) -> None: - """Sets IP-Adapter attention processors, image projection, and loads state_dict. - - Args: - state_dict (`Dict`): - State dict with keys "ip_adapter", which contains parameters for attention processors, and - "image_proj", which contains parameters for image projection net. - low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 else `False`): - Speed up model loading only loading the pretrained weights and not initializing the weights. This also - tries to not use more than 1x model size in CPU memory (including peak memory) while loading the model. - Only supported for PyTorch >= 1.9.0. If you are using an older version of PyTorch, setting this - argument to `True` will raise an error. - """ - - attn_procs = self._convert_ip_adapter_attn_to_diffusers(state_dict["ip_adapter"], low_cpu_mem_usage) - self.set_attn_processor(attn_procs) - - self.image_proj = self._convert_ip_adapter_image_proj_to_diffusers(state_dict["image_proj"], low_cpu_mem_usage) diff --git a/diffusers/loaders/unet.py b/diffusers/loaders/unet.py deleted file mode 100644 index 116d7d6646475e941af6e40c6569adedf4488584..0000000000000000000000000000000000000000 --- a/diffusers/loaders/unet.py +++ /dev/null @@ -1,779 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import os -from collections import defaultdict -from contextlib import nullcontext -from pathlib import Path -from typing import Callable - -import safetensors -import torch -import torch.nn.functional as F -from huggingface_hub.utils import validate_hf_hub_args - -from ..models.embeddings import ( - ImageProjection, - IPAdapterFaceIDImageProjection, - IPAdapterFaceIDPlusImageProjection, - IPAdapterFullImageProjection, - IPAdapterPlusImageProjection, - MultiIPAdapterImageProjection, -) -from ..models.model_loading_utils import load_model_dict_into_meta -from ..models.modeling_utils import _LOW_CPU_MEM_USAGE_DEFAULT, load_state_dict -from ..utils import ( - _get_model_file, - is_accelerate_available, - is_torch_version, - logging, -) -from ..utils.torch_utils import empty_device_cache -from .lora_base import _func_optionally_disable_offloading -from .lora_pipeline import LORA_WEIGHT_NAME, LORA_WEIGHT_NAME_SAFE, TEXT_ENCODER_NAME, UNET_NAME -from .utils import AttnProcsLayers - - -logger = logging.get_logger(__name__) - - -CUSTOM_DIFFUSION_WEIGHT_NAME = "pytorch_custom_diffusion_weights.bin" -CUSTOM_DIFFUSION_WEIGHT_NAME_SAFE = "pytorch_custom_diffusion_weights.safetensors" - - -class UNet2DConditionLoadersMixin: - """ - Load LoRA layers into a [`UNet2DCondtionModel`]. - """ - - text_encoder_name = TEXT_ENCODER_NAME - unet_name = UNET_NAME - - @validate_hf_hub_args - def load_attn_procs(self, pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], **kwargs): - r""" - Load pretrained Custom Diffusion attention processor layers into [`UNet2DConditionModel`]. Attention processor - layers have to be defined in - [`attention_processor.py`](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py) - and be a `torch.nn.Module` class. To load LoRA layers, use [`~loaders.PeftAdapterMixin.load_lora_adapter`] - instead. - - Parameters: - pretrained_model_name_or_path_or_dict (`str` or `os.PathLike` or `dict`): - Can be either: - - - A string, the model id (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a directory (for example `./my_model_directory`) containing the model weights saved - with [`ModelMixin.save_pretrained`]. - - A [torch state - dict](https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict). - - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - local_files_only (`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - subfolder (`str`, *optional*, defaults to `""`): - The subfolder location of a model file within a larger model repository on the Hub or locally. - weight_name (`str`, *optional*, defaults to None): - Name of the serialized state dict file. - - Example: - - ```py - import torch - from diffusers import DiffusionPipeline - - pipeline = DiffusionPipeline.from_pretrained( - "CompVis/stable-diffusion-v1-4", - torch_dtype=torch.float16, - ).to("cuda") - pipeline.unet.load_attn_procs("path-to-save-model", weight_name="pytorch_custom_diffusion_weights.bin") - ``` - """ - from ..hooks.group_offloading import _maybe_remove_and_reapply_group_offloading - - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - weight_name = kwargs.pop("weight_name", None) - use_safetensors = kwargs.pop("use_safetensors", None) - _pipeline = kwargs.pop("_pipeline", None) - allow_pickle = False - - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - user_agent = {"file_type": "attn_procs_weights", "framework": "pytorch"} - - model_file = None - if not isinstance(pretrained_model_name_or_path_or_dict, dict): - # Let's first try to load .safetensors weights - if (use_safetensors and weight_name is None) or ( - weight_name is not None and weight_name.endswith(".safetensors") - ): - try: - model_file = _get_model_file( - pretrained_model_name_or_path_or_dict, - weights_name=weight_name or LORA_WEIGHT_NAME_SAFE, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - ) - state_dict = safetensors.torch.load_file(model_file, device="cpu") - except IOError as e: - if not allow_pickle: - raise e - # try loading non-safetensors weights - pass - if model_file is None: - model_file = _get_model_file( - pretrained_model_name_or_path_or_dict, - weights_name=weight_name or LORA_WEIGHT_NAME, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - ) - state_dict = load_state_dict(model_file) - else: - state_dict = pretrained_model_name_or_path_or_dict - - is_custom_diffusion = any("custom_diffusion" in k for k in state_dict.keys()) - if not is_custom_diffusion: - raise ValueError( - f"{model_file} does not seem to be in the correct format expected by Custom Diffusion training." - ) - - attn_processors = self._process_custom_diffusion(state_dict=state_dict) - - # - - def _process_custom_diffusion(self, state_dict): - from ..models.attention_processor import CustomDiffusionAttnProcessor - - attn_processors = {} - custom_diffusion_grouped_dict = defaultdict(dict) - for key, value in state_dict.items(): - if len(value) == 0: - custom_diffusion_grouped_dict[key] = {} - else: - if "to_out" in key: - attn_processor_key, sub_key = ".".join(key.split(".")[:-3]), ".".join(key.split(".")[-3:]) - else: - attn_processor_key, sub_key = ".".join(key.split(".")[:-2]), ".".join(key.split(".")[-2:]) - custom_diffusion_grouped_dict[attn_processor_key][sub_key] = value - - for key, value_dict in custom_diffusion_grouped_dict.items(): - if len(value_dict) == 0: - attn_processors[key] = CustomDiffusionAttnProcessor( - train_kv=False, train_q_out=False, hidden_size=None, cross_attention_dim=None - ) - else: - cross_attention_dim = value_dict["to_k_custom_diffusion.weight"].shape[1] - hidden_size = value_dict["to_k_custom_diffusion.weight"].shape[0] - train_q_out = True if "to_q_custom_diffusion.weight" in value_dict else False - attn_processors[key] = CustomDiffusionAttnProcessor( - train_kv=True, - train_q_out=train_q_out, - hidden_size=hidden_size, - cross_attention_dim=cross_attention_dim, - ) - attn_processors[key].load_state_dict(value_dict) - - return attn_processors - - @classmethod - # Copied from diffusers.loaders.lora_base.LoraBaseMixin._optionally_disable_offloading - def _optionally_disable_offloading(cls, _pipeline): - return _func_optionally_disable_offloading(_pipeline=_pipeline) - - def save_attn_procs( - self, - save_directory: str | os.PathLike, - is_main_process: bool = True, - weight_name: str = None, - save_function: Callable = None, - safe_serialization: bool = True, - **kwargs, - ): - r""" - Save Custom Diffusion attention processor layers to a directory so that it can be reloaded with the - [`~loaders.UNet2DConditionLoadersMixin.load_attn_procs`] method. To save LoRA layers, use - [`~loaders.PeftAdapterMixin.save_lora_adapter`] instead. - - Arguments: - save_directory (`str` or `os.PathLike`): - Directory to save an attention processor to (will be created if it doesn't exist). - is_main_process (`bool`, *optional*, defaults to `True`): - Whether the process calling this is the main process or not. Useful during distributed training and you - need to call this function on all processes. In this case, set `is_main_process=True` only on the main - process to avoid race conditions. - save_function (`Callable`): - The function to use to save the state dictionary. Useful during distributed training when you need to - replace `torch.save` with another method. Can be configured with the environment variable - `DIFFUSERS_SAVE_MODE`. - safe_serialization (`bool`, *optional*, defaults to `True`): - Whether to save the model using `safetensors` or with `pickle`. - - Example: - - ```py - import torch - from diffusers import DiffusionPipeline - - pipeline = DiffusionPipeline.from_pretrained( - "CompVis/stable-diffusion-v1-4", - torch_dtype=torch.float16, - ).to("cuda") - pipeline.unet.load_attn_procs("path-to-save-model", weight_name="pytorch_custom_diffusion_weights.bin") - pipeline.unet.save_attn_procs("path-to-save-model", weight_name="pytorch_custom_diffusion_weights.bin") - ``` - """ - from ..models.attention_processor import ( - CustomDiffusionAttnProcessor, - CustomDiffusionAttnProcessor2_0, - CustomDiffusionXFormersAttnProcessor, - ) - - if os.path.isfile(save_directory): - logger.error(f"Provided path ({save_directory}) should be a directory, not a file") - return - - is_custom_diffusion = any( - isinstance( - x, - (CustomDiffusionAttnProcessor, CustomDiffusionAttnProcessor2_0, CustomDiffusionXFormersAttnProcessor), - ) - for (_, x) in self.attn_processors.items() - ) - if not is_custom_diffusion: - raise ValueError( - "`save_attn_procs()` only supports saving Custom Diffusion attention processors. Please use " - "`save_lora_adapter()` to save LoRA layers." - ) - - state_dict = self._get_custom_diffusion_state_dict() - if save_function is None and safe_serialization: - # safetensors does not support saving dicts with non-tensor values - empty_state_dict = {k: v for k, v in state_dict.items() if not isinstance(v, torch.Tensor)} - if len(empty_state_dict) > 0: - logger.warning( - f"Safetensors does not support saving dicts with non-tensor values. " - f"The following keys will be ignored: {empty_state_dict.keys()}" - ) - state_dict = {k: v for k, v in state_dict.items() if isinstance(v, torch.Tensor)} - - if save_function is None: - if safe_serialization: - - def save_function(weights, filename): - return safetensors.torch.save_file(weights, filename, metadata={"format": "pt"}) - - else: - save_function = torch.save - - os.makedirs(save_directory, exist_ok=True) - - if weight_name is None: - if safe_serialization: - weight_name = CUSTOM_DIFFUSION_WEIGHT_NAME_SAFE - else: - weight_name = CUSTOM_DIFFUSION_WEIGHT_NAME - - # Save the model - save_path = Path(save_directory, weight_name).as_posix() - save_function(state_dict, save_path) - logger.info(f"Model weights saved in {save_path}") - - def _get_custom_diffusion_state_dict(self): - from ..models.attention_processor import ( - CustomDiffusionAttnProcessor, - CustomDiffusionAttnProcessor2_0, - CustomDiffusionXFormersAttnProcessor, - ) - - model_to_save = AttnProcsLayers( - { - y: x - for (y, x) in self.attn_processors.items() - if isinstance( - x, - ( - CustomDiffusionAttnProcessor, - CustomDiffusionAttnProcessor2_0, - CustomDiffusionXFormersAttnProcessor, - ), - ) - } - ) - state_dict = model_to_save.state_dict() - for name, attn in self.attn_processors.items(): - if len(attn.state_dict()) == 0: - state_dict[name] = {} - - return state_dict - - def _convert_ip_adapter_image_proj_to_diffusers(self, state_dict, low_cpu_mem_usage=_LOW_CPU_MEM_USAGE_DEFAULT): - if low_cpu_mem_usage: - if is_accelerate_available(): - from accelerate import init_empty_weights - - else: - low_cpu_mem_usage = False - logger.warning( - "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the" - " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install" - " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip" - " install accelerate\n```\n." - ) - - if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `low_cpu_mem_usage=False`." - ) - - updated_state_dict = {} - image_projection = None - init_context = init_empty_weights if low_cpu_mem_usage else nullcontext - - if "proj.weight" in state_dict: - # IP-Adapter - num_image_text_embeds = 4 - clip_embeddings_dim = state_dict["proj.weight"].shape[-1] - cross_attention_dim = state_dict["proj.weight"].shape[0] // 4 - - with init_context(): - image_projection = ImageProjection( - cross_attention_dim=cross_attention_dim, - image_embed_dim=clip_embeddings_dim, - num_image_text_embeds=num_image_text_embeds, - ) - - for key, value in state_dict.items(): - diffusers_name = key.replace("proj", "image_embeds") - updated_state_dict[diffusers_name] = value - - elif "proj.3.weight" in state_dict: - # IP-Adapter Full - clip_embeddings_dim = state_dict["proj.0.weight"].shape[0] - cross_attention_dim = state_dict["proj.3.weight"].shape[0] - - with init_context(): - image_projection = IPAdapterFullImageProjection( - cross_attention_dim=cross_attention_dim, image_embed_dim=clip_embeddings_dim - ) - - for key, value in state_dict.items(): - diffusers_name = key.replace("proj.0", "ff.net.0.proj") - diffusers_name = diffusers_name.replace("proj.2", "ff.net.2") - diffusers_name = diffusers_name.replace("proj.3", "norm") - updated_state_dict[diffusers_name] = value - - elif "perceiver_resampler.proj_in.weight" in state_dict: - # IP-Adapter Face ID Plus - id_embeddings_dim = state_dict["proj.0.weight"].shape[1] - embed_dims = state_dict["perceiver_resampler.proj_in.weight"].shape[0] - hidden_dims = state_dict["perceiver_resampler.proj_in.weight"].shape[1] - output_dims = state_dict["perceiver_resampler.proj_out.weight"].shape[0] - heads = state_dict["perceiver_resampler.layers.0.0.to_q.weight"].shape[0] // 64 - - with init_context(): - image_projection = IPAdapterFaceIDPlusImageProjection( - embed_dims=embed_dims, - output_dims=output_dims, - hidden_dims=hidden_dims, - heads=heads, - id_embeddings_dim=id_embeddings_dim, - ) - - for key, value in state_dict.items(): - diffusers_name = key.replace("perceiver_resampler.", "") - diffusers_name = diffusers_name.replace("0.to", "attn.to") - diffusers_name = diffusers_name.replace("0.1.0.", "0.ff.0.") - diffusers_name = diffusers_name.replace("0.1.1.weight", "0.ff.1.net.0.proj.weight") - diffusers_name = diffusers_name.replace("0.1.3.weight", "0.ff.1.net.2.weight") - diffusers_name = diffusers_name.replace("1.1.0.", "1.ff.0.") - diffusers_name = diffusers_name.replace("1.1.1.weight", "1.ff.1.net.0.proj.weight") - diffusers_name = diffusers_name.replace("1.1.3.weight", "1.ff.1.net.2.weight") - diffusers_name = diffusers_name.replace("2.1.0.", "2.ff.0.") - diffusers_name = diffusers_name.replace("2.1.1.weight", "2.ff.1.net.0.proj.weight") - diffusers_name = diffusers_name.replace("2.1.3.weight", "2.ff.1.net.2.weight") - diffusers_name = diffusers_name.replace("3.1.0.", "3.ff.0.") - diffusers_name = diffusers_name.replace("3.1.1.weight", "3.ff.1.net.0.proj.weight") - diffusers_name = diffusers_name.replace("3.1.3.weight", "3.ff.1.net.2.weight") - diffusers_name = diffusers_name.replace("layers.0.0", "layers.0.ln0") - diffusers_name = diffusers_name.replace("layers.0.1", "layers.0.ln1") - diffusers_name = diffusers_name.replace("layers.1.0", "layers.1.ln0") - diffusers_name = diffusers_name.replace("layers.1.1", "layers.1.ln1") - diffusers_name = diffusers_name.replace("layers.2.0", "layers.2.ln0") - diffusers_name = diffusers_name.replace("layers.2.1", "layers.2.ln1") - diffusers_name = diffusers_name.replace("layers.3.0", "layers.3.ln0") - diffusers_name = diffusers_name.replace("layers.3.1", "layers.3.ln1") - - if "norm1" in diffusers_name: - updated_state_dict[diffusers_name.replace("0.norm1", "0")] = value - elif "norm2" in diffusers_name: - updated_state_dict[diffusers_name.replace("0.norm2", "1")] = value - elif "to_kv" in diffusers_name: - v_chunk = value.chunk(2, dim=0) - updated_state_dict[diffusers_name.replace("to_kv", "to_k")] = v_chunk[0] - updated_state_dict[diffusers_name.replace("to_kv", "to_v")] = v_chunk[1] - elif "to_out" in diffusers_name: - updated_state_dict[diffusers_name.replace("to_out", "to_out.0")] = value - elif "proj.0.weight" == diffusers_name: - updated_state_dict["proj.net.0.proj.weight"] = value - elif "proj.0.bias" == diffusers_name: - updated_state_dict["proj.net.0.proj.bias"] = value - elif "proj.2.weight" == diffusers_name: - updated_state_dict["proj.net.2.weight"] = value - elif "proj.2.bias" == diffusers_name: - updated_state_dict["proj.net.2.bias"] = value - else: - updated_state_dict[diffusers_name] = value - - elif "norm.weight" in state_dict: - # IP-Adapter Face ID - id_embeddings_dim_in = state_dict["proj.0.weight"].shape[1] - id_embeddings_dim_out = state_dict["proj.0.weight"].shape[0] - multiplier = id_embeddings_dim_out // id_embeddings_dim_in - norm_layer = "norm.weight" - cross_attention_dim = state_dict[norm_layer].shape[0] - num_tokens = state_dict["proj.2.weight"].shape[0] // cross_attention_dim - - with init_context(): - image_projection = IPAdapterFaceIDImageProjection( - cross_attention_dim=cross_attention_dim, - image_embed_dim=id_embeddings_dim_in, - mult=multiplier, - num_tokens=num_tokens, - ) - - for key, value in state_dict.items(): - diffusers_name = key.replace("proj.0", "ff.net.0.proj") - diffusers_name = diffusers_name.replace("proj.2", "ff.net.2") - updated_state_dict[diffusers_name] = value - - else: - # IP-Adapter Plus - num_image_text_embeds = state_dict["latents"].shape[1] - embed_dims = state_dict["proj_in.weight"].shape[1] - output_dims = state_dict["proj_out.weight"].shape[0] - hidden_dims = state_dict["latents"].shape[2] - attn_key_present = any("attn" in k for k in state_dict) - heads = ( - state_dict["layers.0.attn.to_q.weight"].shape[0] // 64 - if attn_key_present - else state_dict["layers.0.0.to_q.weight"].shape[0] // 64 - ) - - with init_context(): - image_projection = IPAdapterPlusImageProjection( - embed_dims=embed_dims, - output_dims=output_dims, - hidden_dims=hidden_dims, - heads=heads, - num_queries=num_image_text_embeds, - ) - - for key, value in state_dict.items(): - diffusers_name = key.replace("0.to", "2.to") - - diffusers_name = diffusers_name.replace("0.0.norm1", "0.ln0") - diffusers_name = diffusers_name.replace("0.0.norm2", "0.ln1") - diffusers_name = diffusers_name.replace("1.0.norm1", "1.ln0") - diffusers_name = diffusers_name.replace("1.0.norm2", "1.ln1") - diffusers_name = diffusers_name.replace("2.0.norm1", "2.ln0") - diffusers_name = diffusers_name.replace("2.0.norm2", "2.ln1") - diffusers_name = diffusers_name.replace("3.0.norm1", "3.ln0") - diffusers_name = diffusers_name.replace("3.0.norm2", "3.ln1") - - if "to_kv" in diffusers_name: - parts = diffusers_name.split(".") - parts[2] = "attn" - diffusers_name = ".".join(parts) - v_chunk = value.chunk(2, dim=0) - updated_state_dict[diffusers_name.replace("to_kv", "to_k")] = v_chunk[0] - updated_state_dict[diffusers_name.replace("to_kv", "to_v")] = v_chunk[1] - elif "to_q" in diffusers_name: - parts = diffusers_name.split(".") - parts[2] = "attn" - diffusers_name = ".".join(parts) - updated_state_dict[diffusers_name] = value - elif "to_out" in diffusers_name: - parts = diffusers_name.split(".") - parts[2] = "attn" - diffusers_name = ".".join(parts) - updated_state_dict[diffusers_name.replace("to_out", "to_out.0")] = value - else: - diffusers_name = diffusers_name.replace("0.1.0", "0.ff.0") - diffusers_name = diffusers_name.replace("0.1.1", "0.ff.1.net.0.proj") - diffusers_name = diffusers_name.replace("0.1.3", "0.ff.1.net.2") - - diffusers_name = diffusers_name.replace("1.1.0", "1.ff.0") - diffusers_name = diffusers_name.replace("1.1.1", "1.ff.1.net.0.proj") - diffusers_name = diffusers_name.replace("1.1.3", "1.ff.1.net.2") - - diffusers_name = diffusers_name.replace("2.1.0", "2.ff.0") - diffusers_name = diffusers_name.replace("2.1.1", "2.ff.1.net.0.proj") - diffusers_name = diffusers_name.replace("2.1.3", "2.ff.1.net.2") - - diffusers_name = diffusers_name.replace("3.1.0", "3.ff.0") - diffusers_name = diffusers_name.replace("3.1.1", "3.ff.1.net.0.proj") - diffusers_name = diffusers_name.replace("3.1.3", "3.ff.1.net.2") - updated_state_dict[diffusers_name] = value - - if not low_cpu_mem_usage: - image_projection.load_state_dict(updated_state_dict, strict=True) - else: - device_map = {"": self.device} - load_model_dict_into_meta(image_projection, updated_state_dict, device_map=device_map, dtype=self.dtype) - empty_device_cache() - - return image_projection - - def _convert_ip_adapter_attn_to_diffusers(self, state_dicts, low_cpu_mem_usage=_LOW_CPU_MEM_USAGE_DEFAULT): - from ..models.attention_processor import ( - IPAdapterAttnProcessor, - IPAdapterAttnProcessor2_0, - IPAdapterXFormersAttnProcessor, - ) - - if low_cpu_mem_usage: - if is_accelerate_available(): - from accelerate import init_empty_weights - - else: - low_cpu_mem_usage = False - logger.warning( - "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the" - " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install" - " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip" - " install accelerate\n```\n." - ) - - if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `low_cpu_mem_usage=False`." - ) - - # set ip-adapter cross-attention processors & load state_dict - attn_procs = {} - key_id = 1 - init_context = init_empty_weights if low_cpu_mem_usage else nullcontext - for name in self.attn_processors.keys(): - cross_attention_dim = None if name.endswith("attn1.processor") else self.config.cross_attention_dim - if name.startswith("mid_block"): - hidden_size = self.config.block_out_channels[-1] - elif name.startswith("up_blocks"): - block_id = int(name[len("up_blocks.")]) - hidden_size = list(reversed(self.config.block_out_channels))[block_id] - elif name.startswith("down_blocks"): - block_id = int(name[len("down_blocks.")]) - hidden_size = self.config.block_out_channels[block_id] - - if cross_attention_dim is None or "motion_modules" in name: - attn_processor_class = self.attn_processors[name].__class__ - attn_procs[name] = attn_processor_class() - else: - if "XFormers" in str(self.attn_processors[name].__class__): - attn_processor_class = IPAdapterXFormersAttnProcessor - else: - attn_processor_class = ( - IPAdapterAttnProcessor2_0 - if hasattr(F, "scaled_dot_product_attention") - else IPAdapterAttnProcessor - ) - num_image_text_embeds = [] - for state_dict in state_dicts: - if "proj.weight" in state_dict["image_proj"]: - # IP-Adapter - num_image_text_embeds += [4] - elif "proj.3.weight" in state_dict["image_proj"]: - # IP-Adapter Full Face - num_image_text_embeds += [257] # 256 CLIP tokens + 1 CLS token - elif "perceiver_resampler.proj_in.weight" in state_dict["image_proj"]: - # IP-Adapter Face ID Plus - num_image_text_embeds += [4] - elif "norm.weight" in state_dict["image_proj"]: - # IP-Adapter Face ID - num_image_text_embeds += [4] - else: - # IP-Adapter Plus - num_image_text_embeds += [state_dict["image_proj"]["latents"].shape[1]] - - with init_context(): - attn_procs[name] = attn_processor_class( - hidden_size=hidden_size, - cross_attention_dim=cross_attention_dim, - scale=1.0, - num_tokens=num_image_text_embeds, - ) - - value_dict = {} - for i, state_dict in enumerate(state_dicts): - value_dict.update({f"to_k_ip.{i}.weight": state_dict["ip_adapter"][f"{key_id}.to_k_ip.weight"]}) - value_dict.update({f"to_v_ip.{i}.weight": state_dict["ip_adapter"][f"{key_id}.to_v_ip.weight"]}) - - if not low_cpu_mem_usage: - attn_procs[name].load_state_dict(value_dict) - else: - device = next(iter(value_dict.values())).device - dtype = next(iter(value_dict.values())).dtype - device_map = {"": device} - load_model_dict_into_meta(attn_procs[name], value_dict, device_map=device_map, dtype=dtype) - - key_id += 2 - - empty_device_cache() - - return attn_procs - - def _load_ip_adapter_weights(self, state_dicts, low_cpu_mem_usage=_LOW_CPU_MEM_USAGE_DEFAULT): - if not isinstance(state_dicts, list): - state_dicts = [state_dicts] - - # Kolors Unet already has a `encoder_hid_proj` - if ( - self.encoder_hid_proj is not None - and self.config.encoder_hid_dim_type == "text_proj" - and not hasattr(self, "text_encoder_hid_proj") - ): - self.text_encoder_hid_proj = self.encoder_hid_proj - - # Set encoder_hid_proj after loading ip_adapter weights, - # because `IPAdapterPlusImageProjection` also has `attn_processors`. - self.encoder_hid_proj = None - - attn_procs = self._convert_ip_adapter_attn_to_diffusers(state_dicts, low_cpu_mem_usage=low_cpu_mem_usage) - self.set_attn_processor(attn_procs) - - # convert IP-Adapter Image Projection layers to diffusers - image_projection_layers = [] - for state_dict in state_dicts: - image_projection_layer = self._convert_ip_adapter_image_proj_to_diffusers( - state_dict["image_proj"], low_cpu_mem_usage=low_cpu_mem_usage - ) - image_projection_layers.append(image_projection_layer) - - self.encoder_hid_proj = MultiIPAdapterImageProjection(image_projection_layers) - self.config.encoder_hid_dim_type = "ip_image_proj" - - self.to(dtype=self.dtype, device=self.device) - - def _load_ip_adapter_loras(self, state_dicts): - lora_dicts = {} - for key_id, name in enumerate(self.attn_processors.keys()): - for i, state_dict in enumerate(state_dicts): - if f"{key_id}.to_k_lora.down.weight" in state_dict["ip_adapter"]: - if i not in lora_dicts: - lora_dicts[i] = {} - lora_dicts[i].update( - { - f"unet.{name}.to_k_lora.down.weight": state_dict["ip_adapter"][ - f"{key_id}.to_k_lora.down.weight" - ] - } - ) - lora_dicts[i].update( - { - f"unet.{name}.to_q_lora.down.weight": state_dict["ip_adapter"][ - f"{key_id}.to_q_lora.down.weight" - ] - } - ) - lora_dicts[i].update( - { - f"unet.{name}.to_v_lora.down.weight": state_dict["ip_adapter"][ - f"{key_id}.to_v_lora.down.weight" - ] - } - ) - lora_dicts[i].update( - { - f"unet.{name}.to_out_lora.down.weight": state_dict["ip_adapter"][ - f"{key_id}.to_out_lora.down.weight" - ] - } - ) - lora_dicts[i].update( - {f"unet.{name}.to_k_lora.up.weight": state_dict["ip_adapter"][f"{key_id}.to_k_lora.up.weight"]} - ) - lora_dicts[i].update( - {f"unet.{name}.to_q_lora.up.weight": state_dict["ip_adapter"][f"{key_id}.to_q_lora.up.weight"]} - ) - lora_dicts[i].update( - {f"unet.{name}.to_v_lora.up.weight": state_dict["ip_adapter"][f"{key_id}.to_v_lora.up.weight"]} - ) - lora_dicts[i].update( - { - f"unet.{name}.to_out_lora.up.weight": state_dict["ip_adapter"][ - f"{key_id}.to_out_lora.up.weight" - ] - } - ) - return lora_dicts diff --git a/diffusers/loaders/unet_loader_utils.py b/diffusers/loaders/unet_loader_utils.py deleted file mode 100644 index 15ccab6f45a115f7f686ad2f66907bfb733e6ac4..0000000000000000000000000000000000000000 --- a/diffusers/loaders/unet_loader_utils.py +++ /dev/null @@ -1,164 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import copy -from typing import TYPE_CHECKING - -from torch import nn - -from ..utils import logging - - -if TYPE_CHECKING: - # import here to avoid circular imports - from ..models import UNet2DConditionModel - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _translate_into_actual_layer_name(name): - """Translate user-friendly name (e.g. 'mid') into actual layer name (e.g. 'mid_block.attentions.0')""" - if name == "mid": - return "mid_block.attentions.0" - - updown, block, attn = name.split(".") - - updown = updown.replace("down", "down_blocks").replace("up", "up_blocks") - block = block.replace("block_", "") - attn = "attentions." + attn - - return ".".join((updown, block, attn)) - - -def _maybe_expand_lora_scales(unet: "UNet2DConditionModel", weight_scales: list[float | dict], default_scale=1.0): - blocks_with_transformer = { - "down": [i for i, block in enumerate(unet.down_blocks) if hasattr(block, "attentions")], - "up": [i for i, block in enumerate(unet.up_blocks) if hasattr(block, "attentions")], - } - transformer_per_block = {"down": unet.config.layers_per_block, "up": unet.config.layers_per_block + 1} - - expanded_weight_scales = [ - _maybe_expand_lora_scales_for_one_adapter( - weight_for_adapter, - blocks_with_transformer, - transformer_per_block, - model=unet, - default_scale=default_scale, - ) - for weight_for_adapter in weight_scales - ] - - return expanded_weight_scales - - -def _maybe_expand_lora_scales_for_one_adapter( - scales: float | dict, - blocks_with_transformer: dict[str, int], - transformer_per_block: dict[str, int], - model: nn.Module, - default_scale: float = 1.0, -): - """ - Expands the inputs into a more granular dictionary. See the example below for more details. - - Parameters: - scales (`float | Dict`): - Scales dict to expand. - blocks_with_transformer (`dict[str, int]`): - Dict with keys 'up' and 'down', showing which blocks have transformer layers - transformer_per_block (`dict[str, int]`): - Dict with keys 'up' and 'down', showing how many transformer layers each block has - - E.g. turns - ```python - scales = {"down": 2, "mid": 3, "up": {"block_0": 4, "block_1": [5, 6, 7]}} - blocks_with_transformer = {"down": [1, 2], "up": [0, 1]} - transformer_per_block = {"down": 2, "up": 3} - ``` - into - ```python - { - "down.block_1.0": 2, - "down.block_1.1": 2, - "down.block_2.0": 2, - "down.block_2.1": 2, - "mid": 3, - "up.block_0.0": 4, - "up.block_0.1": 4, - "up.block_0.2": 4, - "up.block_1.0": 5, - "up.block_1.1": 6, - "up.block_1.2": 7, - } - ``` - """ - if sorted(blocks_with_transformer.keys()) != ["down", "up"]: - raise ValueError("blocks_with_transformer needs to be a dict with keys `'down' and `'up'`") - - if sorted(transformer_per_block.keys()) != ["down", "up"]: - raise ValueError("transformer_per_block needs to be a dict with keys `'down' and `'up'`") - - if not isinstance(scales, dict): - # don't expand if scales is a single number - return scales - - scales = copy.deepcopy(scales) - - if "mid" not in scales: - scales["mid"] = default_scale - elif isinstance(scales["mid"], list): - if len(scales["mid"]) == 1: - scales["mid"] = scales["mid"][0] - else: - raise ValueError(f"Expected 1 scales for mid, got {len(scales['mid'])}.") - - for updown in ["up", "down"]: - if updown not in scales: - scales[updown] = default_scale - - # eg {"down": 1} to {"down": {"block_1": 1, "block_2": 1}}} - if not isinstance(scales[updown], dict): - scales[updown] = {f"block_{i}": copy.deepcopy(scales[updown]) for i in blocks_with_transformer[updown]} - - # eg {"down": {"block_1": 1}} to {"down": {"block_1": [1, 1]}} - for i in blocks_with_transformer[updown]: - block = f"block_{i}" - # set not assigned blocks to default scale - if block not in scales[updown]: - scales[updown][block] = default_scale - if not isinstance(scales[updown][block], list): - scales[updown][block] = [scales[updown][block] for _ in range(transformer_per_block[updown])] - elif len(scales[updown][block]) == 1: - # a list specifying scale to each masked IP input - scales[updown][block] = scales[updown][block] * transformer_per_block[updown] - elif len(scales[updown][block]) != transformer_per_block[updown]: - raise ValueError( - f"Expected {transformer_per_block[updown]} scales for {updown}.{block}, got {len(scales[updown][block])}." - ) - - # eg {"down": "block_1": [1, 1]}} to {"down.block_1.0": 1, "down.block_1.1": 1} - for i in blocks_with_transformer[updown]: - block = f"block_{i}" - for tf_idx, value in enumerate(scales[updown][block]): - scales[f"{updown}.{block}.{tf_idx}"] = value - - del scales[updown] - - state_dict = model.state_dict() - for layer in scales.keys(): - if not any(_translate_into_actual_layer_name(layer) in module for module in state_dict.keys()): - raise ValueError( - f"Can't set lora scale for layer {layer}. It either doesn't exist in this unet or it has no attentions." - ) - - return {_translate_into_actual_layer_name(name): weight for name, weight in scales.items()} diff --git a/diffusers/loaders/utils.py b/diffusers/loaders/utils.py deleted file mode 100644 index 9e484559fa5422f91ac99ba9879975d278c6c240..0000000000000000000000000000000000000000 --- a/diffusers/loaders/utils.py +++ /dev/null @@ -1,58 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch - - -class AttnProcsLayers(torch.nn.Module): - def __init__(self, state_dict: dict[str, torch.Tensor]): - super().__init__() - self.layers = torch.nn.ModuleList(state_dict.values()) - self.mapping = dict(enumerate(state_dict.keys())) - self.rev_mapping = {v: k for k, v in enumerate(state_dict.keys())} - - # .processor for unet, .self_attn for text encoder - self.split_keys = [".processor", ".self_attn"] - - # we add a hook to state_dict() and load_state_dict() so that the - # naming fits with `unet.attn_processors` - def map_to(module, state_dict, *args, **kwargs): - new_state_dict = {} - for key, value in state_dict.items(): - num = int(key.split(".")[1]) # 0 is always "layers" - new_key = key.replace(f"layers.{num}", module.mapping[num]) - new_state_dict[new_key] = value - - return new_state_dict - - def remap_key(key, state_dict): - for k in self.split_keys: - if k in key: - return key.split(k)[0] + k - - raise ValueError( - f"There seems to be a problem with the state_dict: {set(state_dict.keys())}. {key} has to have one of {self.split_keys}." - ) - - def map_from(module, state_dict, *args, **kwargs): - all_keys = list(state_dict.keys()) - for key in all_keys: - replace_key = remap_key(key, state_dict) - new_key = key.replace(replace_key, f"layers.{module.rev_mapping[replace_key]}") - state_dict[new_key] = state_dict[key] - del state_dict[key] - - self._register_state_dict_hook(map_to) - self._register_load_state_dict_pre_hook(map_from, with_module=True) diff --git a/diffusers/models/README.md b/diffusers/models/README.md deleted file mode 100644 index fb91f59411265660e01d8b4bcc0b99e8b8fe9d55..0000000000000000000000000000000000000000 --- a/diffusers/models/README.md +++ /dev/null @@ -1,3 +0,0 @@ -# Models - -For more detail on the models, please refer to the [docs](https://huggingface.co/docs/diffusers/api/models/overview). \ No newline at end of file diff --git a/diffusers/models/__init__.py b/diffusers/models/__init__.py deleted file mode 100644 index a78702ab34facb7425a191920f6b596fb20fb197..0000000000000000000000000000000000000000 --- a/diffusers/models/__init__.py +++ /dev/null @@ -1,309 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import TYPE_CHECKING - -from ..utils import ( - DIFFUSERS_SLOW_IMPORT, - _LazyModule, - is_torch_available, -) - - -_import_structure = {} - -if is_torch_available(): - _import_structure["_modeling_parallel"] = ["ContextParallelConfig", "ParallelConfig"] - _import_structure["adapter"] = ["MultiAdapter", "T2IAdapter"] - _import_structure["attention_dispatch"] = ["AttentionBackendName", "attention_backend"] - _import_structure["auto_model"] = ["AutoModel"] - _import_structure["autoencoders.autoencoder_asym_kl"] = ["AsymmetricAutoencoderKL"] - _import_structure["autoencoders.autoencoder_cosmos3_audio"] = ["Cosmos3AVAEAudioTokenizer"] - _import_structure["autoencoders.autoencoder_dc"] = ["AutoencoderDC"] - _import_structure["autoencoders.autoencoder_kl"] = ["AutoencoderKL"] - _import_structure["autoencoders.autoencoder_kl_allegro"] = ["AutoencoderKLAllegro"] - _import_structure["autoencoders.autoencoder_kl_cogvideox"] = ["AutoencoderKLCogVideoX"] - _import_structure["autoencoders.autoencoder_kl_cosmos"] = ["AutoencoderKLCosmos"] - _import_structure["autoencoders.autoencoder_kl_flux2"] = ["AutoencoderKLFlux2"] - _import_structure["autoencoders.autoencoder_kl_hunyuan_video"] = ["AutoencoderKLHunyuanVideo"] - _import_structure["autoencoders.autoencoder_kl_hunyuanimage"] = ["AutoencoderKLHunyuanImage"] - _import_structure["autoencoders.autoencoder_kl_hunyuanimage_refiner"] = ["AutoencoderKLHunyuanImageRefiner"] - _import_structure["autoencoders.autoencoder_kl_hunyuanvideo15"] = ["AutoencoderKLHunyuanVideo15"] - _import_structure["autoencoders.autoencoder_kl_kvae"] = ["AutoencoderKLKVAE"] - _import_structure["autoencoders.autoencoder_kl_kvae_video"] = ["AutoencoderKLKVAEVideo"] - _import_structure["autoencoders.autoencoder_kl_ltx"] = ["AutoencoderKLLTXVideo"] - _import_structure["autoencoders.autoencoder_kl_ltx2"] = ["AutoencoderKLLTX2Video"] - _import_structure["autoencoders.autoencoder_kl_ltx2_audio"] = ["AutoencoderKLLTX2Audio"] - _import_structure["autoencoders.autoencoder_kl_magvit"] = ["AutoencoderKLMagvit"] - _import_structure["autoencoders.autoencoder_kl_minimax_h3"] = ["AutoencoderKLMiniMaxH3"] - _import_structure["autoencoders.autoencoder_kl_minimax_h3_audio"] = ["AutoencoderKLMiniMaxH3Audio"] - _import_structure["autoencoders.autoencoder_kl_mochi"] = ["AutoencoderKLMochi"] - _import_structure["autoencoders.autoencoder_kl_qwenimage"] = ["AutoencoderKLQwenImage"] - _import_structure["autoencoders.autoencoder_kl_temporal_decoder"] = ["AutoencoderKLTemporalDecoder"] - _import_structure["autoencoders.autoencoder_kl_wan"] = ["AutoencoderKLWan"] - _import_structure["autoencoders.autoencoder_longcat_audio_dit"] = ["LongCatAudioDiTVae"] - _import_structure["autoencoders.autoencoder_oobleck"] = ["AutoencoderOobleck"] - _import_structure["autoencoders.autoencoder_rae"] = ["AutoencoderRAE"] - _import_structure["autoencoders.autoencoder_tiny"] = ["AutoencoderTiny"] - _import_structure["autoencoders.autoencoder_vidtok"] = ["AutoencoderVidTok"] - _import_structure["autoencoders.consistency_decoder_vae"] = ["ConsistencyDecoderVAE"] - _import_structure["autoencoders.vq_model"] = ["VQModel"] - _import_structure["cache_utils"] = ["CacheMixin"] - _import_structure["condition_embedders.condition_embedder_anima"] = ["AnimaTextConditioner"] - _import_structure["controlnets.controlnet"] = ["ControlNetModel"] - _import_structure["controlnets.controlnet_cosmos"] = ["CosmosControlNetModel"] - _import_structure["controlnets.controlnet_flux"] = ["FluxControlNetModel", "FluxMultiControlNetModel"] - _import_structure["controlnets.controlnet_hunyuan"] = [ - "HunyuanDiT2DControlNetModel", - "HunyuanDiT2DMultiControlNetModel", - ] - _import_structure["controlnets.controlnet_qwenimage"] = [ - "QwenImageControlNetModel", - "QwenImageMultiControlNetModel", - ] - _import_structure["controlnets.controlnet_sana"] = ["SanaControlNetModel"] - _import_structure["controlnets.controlnet_sd3"] = ["SD3ControlNetModel", "SD3MultiControlNetModel"] - _import_structure["controlnets.controlnet_sparsectrl"] = ["SparseControlNetModel"] - _import_structure["controlnets.controlnet_union"] = ["ControlNetUnionModel"] - _import_structure["controlnets.controlnet_xs"] = ["ControlNetXSAdapter", "UNetControlNetXSModel"] - _import_structure["controlnets.controlnet_z_image"] = ["ZImageControlNetModel"] - _import_structure["controlnets.multicontrolnet"] = ["MultiControlNetModel"] - _import_structure["controlnets.multicontrolnet_union"] = ["MultiControlNetUnionModel"] - _import_structure["embeddings"] = ["ImageProjection"] - _import_structure["modeling_utils"] = ["ModelMixin"] - _import_structure["transformers.ace_step_transformer"] = ["AceStepTransformer1DModel"] - _import_structure["transformers.auraflow_transformer_2d"] = ["AuraFlowTransformer2DModel"] - _import_structure["transformers.cogvideox_transformer_3d"] = ["CogVideoXTransformer3DModel"] - _import_structure["transformers.consisid_transformer_3d"] = ["ConsisIDTransformer3DModel"] - _import_structure["transformers.dit_transformer_2d"] = ["DiTTransformer2DModel"] - _import_structure["transformers.dual_transformer_2d"] = ["DualTransformer2DModel"] - _import_structure["transformers.hunyuan_transformer_2d"] = ["HunyuanDiT2DModel"] - _import_structure["transformers.latte_transformer_3d"] = ["LatteTransformer3DModel"] - _import_structure["transformers.lumina_nextdit2d"] = ["LuminaNextDiT2DModel"] - _import_structure["transformers.pixart_transformer_2d"] = ["PixArtTransformer2DModel"] - _import_structure["transformers.prior_transformer"] = ["PriorTransformer"] - _import_structure["transformers.sana_transformer"] = ["SanaTransformer2DModel"] - _import_structure["transformers.stable_audio_transformer"] = ["StableAudioDiTModel"] - _import_structure["transformers.t5_film_transformer"] = ["T5FilmDecoder"] - _import_structure["transformers.transformer_2d"] = ["Transformer2DModel"] - _import_structure["transformers.transformer_2d_dreamlite"] = ["DreamLiteTransformer2DModel"] - _import_structure["transformers.transformer_allegro"] = ["AllegroTransformer3DModel"] - _import_structure["transformers.transformer_anyflow"] = ["AnyFlowTransformer3DModel"] - _import_structure["transformers.transformer_anyflow_far"] = ["AnyFlowFARTransformer3DModel"] - _import_structure["transformers.transformer_bria"] = ["BriaTransformer2DModel"] - _import_structure["transformers.transformer_bria_fibo"] = ["BriaFiboTransformer2DModel"] - _import_structure["transformers.transformer_chroma"] = ["ChromaTransformer2DModel"] - _import_structure["transformers.transformer_chronoedit"] = ["ChronoEditTransformer3DModel"] - _import_structure["transformers.transformer_cogview3plus"] = ["CogView3PlusTransformer2DModel"] - _import_structure["transformers.transformer_cogview4"] = ["CogView4Transformer2DModel"] - _import_structure["transformers.transformer_cosmos"] = ["CosmosTransformer3DModel"] - _import_structure["transformers.transformer_cosmos3"] = ["Cosmos3OmniTransformer"] - _import_structure["transformers.transformer_easyanimate"] = ["EasyAnimateTransformer3DModel"] - _import_structure["transformers.transformer_ernie_image"] = ["ErnieImageTransformer2DModel"] - _import_structure["transformers.transformer_flux"] = ["FluxTransformer2DModel"] - _import_structure["transformers.transformer_flux2"] = ["Flux2Transformer2DModel"] - _import_structure["transformers.transformer_glm_image"] = ["GlmImageTransformer2DModel"] - _import_structure["transformers.transformer_helios"] = ["HeliosTransformer3DModel"] - _import_structure["transformers.transformer_hidream_image"] = ["HiDreamImageTransformer2DModel"] - _import_structure["transformers.transformer_hunyuan_video"] = ["HunyuanVideoTransformer3DModel"] - _import_structure["transformers.transformer_hunyuan_video15"] = ["HunyuanVideo15Transformer3DModel"] - _import_structure["transformers.transformer_hunyuan_video_framepack"] = ["HunyuanVideoFramepackTransformer3DModel"] - _import_structure["transformers.transformer_hunyuanimage"] = ["HunyuanImageTransformer2DModel"] - _import_structure["transformers.transformer_ideogram4"] = ["Ideogram4Transformer2DModel"] - _import_structure["transformers.transformer_joyimage"] = ["JoyImageEditTransformer3DModel"] - _import_structure["transformers.transformer_joyimage_edit_plus"] = ["JoyImageEditPlusTransformer3DModel"] - _import_structure["transformers.transformer_kandinsky"] = ["Kandinsky5Transformer3DModel"] - _import_structure["transformers.transformer_krea2"] = ["Krea2Transformer2DModel"] - _import_structure["transformers.transformer_longcat_audio_dit"] = ["LongCatAudioDiTTransformer"] - _import_structure["transformers.transformer_longcat_image"] = ["LongCatImageTransformer2DModel"] - _import_structure["transformers.transformer_ltx"] = ["LTXVideoTransformer3DModel"] - _import_structure["transformers.transformer_ltx2"] = ["LTX2VideoTransformer3DModel"] - _import_structure["transformers.transformer_lumina2"] = ["Lumina2Transformer2DModel"] - _import_structure["transformers.transformer_minimax_h3"] = ["MiniMaxH3Transformer3DModel"] - _import_structure["transformers.transformer_mochi"] = ["MochiTransformer3DModel"] - _import_structure["transformers.transformer_motif_video"] = ["MotifVideoTransformer3DModel"] - _import_structure["transformers.transformer_nucleusmoe_image"] = ["NucleusMoEImageTransformer2DModel"] - _import_structure["transformers.transformer_omnigen"] = ["OmniGenTransformer2DModel"] - _import_structure["transformers.transformer_ovis_image"] = ["OvisImageTransformer2DModel"] - _import_structure["transformers.transformer_prx"] = ["PRXTransformer2DModel"] - _import_structure["transformers.transformer_qwenimage"] = ["QwenImageTransformer2DModel"] - _import_structure["transformers.transformer_sana_video"] = ["SanaVideoTransformer3DModel"] - _import_structure["transformers.transformer_sd3"] = ["SD3Transformer2DModel"] - _import_structure["transformers.transformer_skyreels_v2"] = ["SkyReelsV2Transformer3DModel"] - _import_structure["transformers.transformer_temporal"] = ["TransformerTemporalModel"] - _import_structure["transformers.transformer_wan"] = ["WanTransformer3DModel"] - _import_structure["transformers.transformer_wan_animate"] = ["WanAnimateTransformer3DModel"] - _import_structure["transformers.transformer_wan_vace"] = ["WanVACETransformer3DModel"] - _import_structure["transformers.transformer_z_image"] = ["ZImageTransformer2DModel"] - _import_structure["unets.unet_1d"] = ["UNet1DModel"] - _import_structure["unets.unet_2d"] = ["UNet2DModel"] - _import_structure["unets.unet_2d_condition"] = ["UNet2DConditionModel"] - _import_structure["unets.unet_3d_condition"] = ["UNet3DConditionModel"] - _import_structure["unets.unet_dreamlite"] = ["DreamLiteUNetModel"] - _import_structure["unets.unet_i2vgen_xl"] = ["I2VGenXLUNet"] - _import_structure["unets.unet_kandinsky3"] = ["Kandinsky3UNet"] - _import_structure["unets.unet_motion_model"] = ["MotionAdapter", "UNetMotionModel"] - _import_structure["unets.unet_spatio_temporal_condition"] = ["UNetSpatioTemporalConditionModel"] - _import_structure["unets.unet_stable_cascade"] = ["StableCascadeUNet"] - _import_structure["unets.uvit_2d"] = ["UVit2DModel"] - - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - if is_torch_available(): - from ._modeling_parallel import ContextParallelConfig, ParallelConfig - from .adapter import MultiAdapter, T2IAdapter - from .attention_dispatch import AttentionBackendName, attention_backend - from .auto_model import AutoModel - from .autoencoders import ( - AsymmetricAutoencoderKL, - AutoencoderDC, - AutoencoderKL, - AutoencoderKLAllegro, - AutoencoderKLCogVideoX, - AutoencoderKLCosmos, - AutoencoderKLFlux2, - AutoencoderKLHunyuanImage, - AutoencoderKLHunyuanImageRefiner, - AutoencoderKLHunyuanVideo, - AutoencoderKLHunyuanVideo15, - AutoencoderKLKVAE, - AutoencoderKLKVAEVideo, - AutoencoderKLLTX2Audio, - AutoencoderKLLTX2Video, - AutoencoderKLLTXVideo, - AutoencoderKLMagvit, - AutoencoderKLMiniMaxH3, - AutoencoderKLMiniMaxH3Audio, - AutoencoderKLMochi, - AutoencoderKLQwenImage, - AutoencoderKLTemporalDecoder, - AutoencoderKLWan, - AutoencoderOobleck, - AutoencoderRAE, - AutoencoderTiny, - AutoencoderVidTok, - ConsistencyDecoderVAE, - Cosmos3AVAEAudioTokenizer, - LongCatAudioDiTVae, - VQModel, - ) - from .cache_utils import CacheMixin - from .condition_embedders import AnimaTextConditioner - from .controlnets import ( - ControlNetModel, - ControlNetUnionModel, - ControlNetXSAdapter, - CosmosControlNetModel, - FluxControlNetModel, - FluxMultiControlNetModel, - HunyuanDiT2DControlNetModel, - HunyuanDiT2DMultiControlNetModel, - MultiControlNetModel, - MultiControlNetUnionModel, - QwenImageControlNetModel, - QwenImageMultiControlNetModel, - SanaControlNetModel, - SD3ControlNetModel, - SD3MultiControlNetModel, - SparseControlNetModel, - UNetControlNetXSModel, - ZImageControlNetModel, - ) - from .embeddings import ImageProjection - from .modeling_utils import ModelMixin - from .transformers import ( - AceStepTransformer1DModel, - AllegroTransformer3DModel, - AnyFlowFARTransformer3DModel, - AnyFlowTransformer3DModel, - AuraFlowTransformer2DModel, - BriaFiboTransformer2DModel, - BriaTransformer2DModel, - ChromaTransformer2DModel, - ChronoEditTransformer3DModel, - CogVideoXTransformer3DModel, - CogView3PlusTransformer2DModel, - CogView4Transformer2DModel, - ConsisIDTransformer3DModel, - Cosmos3OmniTransformer, - CosmosTransformer3DModel, - DiTTransformer2DModel, - DreamLiteTransformer2DModel, - DualTransformer2DModel, - EasyAnimateTransformer3DModel, - ErnieImageTransformer2DModel, - Flux2Transformer2DModel, - FluxTransformer2DModel, - GlmImageTransformer2DModel, - HeliosTransformer3DModel, - HiDreamImageTransformer2DModel, - HunyuanDiT2DModel, - HunyuanImageTransformer2DModel, - HunyuanVideo15Transformer3DModel, - HunyuanVideoFramepackTransformer3DModel, - HunyuanVideoTransformer3DModel, - Ideogram4Transformer2DModel, - JoyImageEditPlusTransformer3DModel, - JoyImageEditTransformer3DModel, - Kandinsky5Transformer3DModel, - Krea2Transformer2DModel, - LatteTransformer3DModel, - LongCatAudioDiTTransformer, - LongCatImageTransformer2DModel, - LTX2VideoTransformer3DModel, - LTXVideoTransformer3DModel, - Lumina2Transformer2DModel, - LuminaNextDiT2DModel, - MiniMaxH3Transformer3DModel, - MochiTransformer3DModel, - MotifVideoTransformer3DModel, - NucleusMoEImageTransformer2DModel, - OmniGenTransformer2DModel, - OvisImageTransformer2DModel, - PixArtTransformer2DModel, - PriorTransformer, - PRXTransformer2DModel, - QwenImageTransformer2DModel, - SanaTransformer2DModel, - SanaVideoTransformer3DModel, - SD3Transformer2DModel, - SkyReelsV2Transformer3DModel, - StableAudioDiTModel, - T5FilmDecoder, - Transformer2DModel, - TransformerTemporalModel, - WanAnimateTransformer3DModel, - WanTransformer3DModel, - WanVACETransformer3DModel, - ZImageTransformer2DModel, - ) - from .unets import ( - DreamLiteUNetModel, - I2VGenXLUNet, - Kandinsky3UNet, - MotionAdapter, - StableCascadeUNet, - UNet1DModel, - UNet2DConditionModel, - UNet2DModel, - UNet3DConditionModel, - UNetMotionModel, - UNetSpatioTemporalConditionModel, - UVit2DModel, - ) - -else: - import sys - - sys.modules[__name__] = _LazyModule(__name__, globals()["__file__"], _import_structure, module_spec=__spec__) diff --git a/diffusers/models/_modeling_parallel.py b/diffusers/models/_modeling_parallel.py deleted file mode 100644 index f5693f1033cf2ad0cf3911a98e97623c61781fc5..0000000000000000000000000000000000000000 --- a/diffusers/models/_modeling_parallel.py +++ /dev/null @@ -1,325 +0,0 @@ -# 🚨🚨🚨 Experimental parallelism support for Diffusers 🚨🚨🚨 -# Experimental changes are subject to change and APIs may break without warning. - -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass -from typing import TYPE_CHECKING, Literal - -import torch -import torch.distributed as dist - -from ..utils import get_logger - - -if TYPE_CHECKING: - pass - - -logger = get_logger(__name__) # pylint: disable=invalid-name - - -# TODO(aryan): add support for the following: -# - Unified Attention -# - More dispatcher attention backends -# - CFG/Data Parallel -# - Tensor Parallel - - -@dataclass -class ContextParallelConfig: - """ - Configuration for context parallelism. - - Args: - ring_degree (`int`, *optional*, defaults to `1`): - Number of devices to use for Ring Attention. Sequence is split across devices. Each device computes - attention between its local Q and KV chunks passed sequentially around ring. Lower memory (only holds 1/N - of KV at a time), overlaps compute with communication, but requires N iterations to see all tokens. Best - for long sequences with limited memory/bandwidth. Number of devices to use for ring attention within a - context parallel region. Must be a divisor of the total number of devices in the context parallel mesh. - ulysses_degree (`int`, *optional*, defaults to `1`): - Number of devices to use for Ulysses Attention. Sequence split is across devices. Each device computes - local QKV, then all-gathers all KV chunks to compute full attention in one pass. Higher memory (stores all - KV), requires high-bandwidth all-to-all communication, but lower latency. Best for moderate sequences with - good interconnect bandwidth. - convert_to_fp32 (`bool`, *optional*, defaults to `True`): - Whether to convert output and LSE to float32 for ring attention numerical stability. - rotate_method (`str`, *optional*, defaults to `"allgather"`): - Method to use for rotating key/value states across devices in ring attention. Currently, only `"allgather"` - is supported. - ulysses_anything (`bool`, *optional*, defaults to `False`): - Whether to enable "Ulysses Anything" mode, which supports arbitrary sequence lengths and head counts that - are not evenly divisible by `ulysses_degree`. When enabled, `ulysses_degree` must be greater than 1 and - `ring_degree` must be 1. - ring_anything (`bool`, *optional*, defaults to `False`): - Whether to enable "Ring Anything" mode, which supports arbitrary sequence lengths. When enabled, - `ring_degree` must be greater than 1 and `ulysses_degree` must be 1. - mesh (`torch.distributed.device_mesh.DeviceMesh`, *optional*): - A custom device mesh to use for context parallelism. If provided, this mesh will be used instead of - creating a new one. This is useful when combining context parallelism with other parallelism strategies - (e.g., FSDP, tensor parallelism) that share the same device mesh. The mesh must have both "ring" and - "ulysses" dimensions. Use size 1 for dimensions not being used (e.g., `mesh_shape=(2, 1, 4)` with - `mesh_dim_names=("ring", "ulysses", "fsdp")` for ring attention only with FSDP). - - """ - - ring_degree: int | None = None - ulysses_degree: int | None = None - convert_to_fp32: bool = True - # TODO: support alltoall - rotate_method: Literal["allgather", "alltoall"] = "allgather" - mesh: torch.distributed.device_mesh.DeviceMesh | None = None - # Whether to enable ulysses anything attention to support - # any sequence lengths and any head numbers. - ulysses_anything: bool = False - # Whether to enable ring anything attention to support any sequence lengths. - ring_anything: bool = False - - _rank: int = None - _world_size: int = None - _device: torch.device = None - _mesh: torch.distributed.device_mesh.DeviceMesh = None - _flattened_mesh: torch.distributed.device_mesh.DeviceMesh = None - _ring_mesh: torch.distributed.device_mesh.DeviceMesh = None - _ulysses_mesh: torch.distributed.device_mesh.DeviceMesh = None - _ring_local_rank: int = None - _ulysses_local_rank: int = None - - def __post_init__(self): - if self.ring_degree is None: - self.ring_degree = 1 - if self.ulysses_degree is None: - self.ulysses_degree = 1 - - if self.ring_degree == 1 and self.ulysses_degree == 1: - raise ValueError( - "Either ring_degree or ulysses_degree must be greater than 1 in order to use context parallel inference" - ) - if self.ring_degree < 1 or self.ulysses_degree < 1: - raise ValueError("`ring_degree` and `ulysses_degree` must be greater than or equal to 1.") - if self.rotate_method != "allgather": - raise NotImplementedError( - f"Only rotate_method='allgather' is supported for now, but got {self.rotate_method}." - ) - if self.ulysses_anything: - if self.ulysses_degree == 1: - raise ValueError("ulysses_degree must be greater than 1 for ulysses_anything to be enabled.") - if self.ring_degree > 1: - raise ValueError("ulysses_anything cannot be enabled when ring_degree > 1.") - if self.ring_anything: - if self.ring_degree == 1: - raise ValueError("ring_degree must be greater than 1 for ring_anything to be enabled.") - if self.ulysses_degree > 1: - raise ValueError("ring_anything cannot be enabled when ulysses_degree > 1.") - if self.ulysses_anything and self.ring_anything: - raise ValueError("ulysses_anything and ring_anything cannot both be enabled.") - - @property - def mesh_shape(self) -> tuple[int, int]: - return (self.ring_degree, self.ulysses_degree) - - @property - def mesh_dim_names(self) -> tuple[str, str]: - """Dimension names for the device mesh.""" - return ("ring", "ulysses") - - def setup(self, rank: int, world_size: int, device: torch.device, mesh: torch.distributed.device_mesh.DeviceMesh): - self._rank = rank - self._world_size = world_size - self._device = device - self._mesh = mesh - - if self.ulysses_degree * self.ring_degree > world_size: - raise ValueError( - f"The product of `ring_degree` ({self.ring_degree}) and `ulysses_degree` ({self.ulysses_degree}) must not exceed the world size ({world_size})." - ) - - self._flattened_mesh = self._mesh["ring", "ulysses"]._flatten() - self._ring_mesh = self._mesh["ring"] - self._ulysses_mesh = self._mesh["ulysses"] - self._ring_local_rank = self._ring_mesh.get_local_rank() - self._ulysses_local_rank = self._ulysses_mesh.get_local_rank() - - -@dataclass -class ParallelConfig: - """ - Configuration for applying different parallelisms. - - Args: - context_parallel_config (`ContextParallelConfig`, *optional*): - Configuration for context parallelism. - """ - - context_parallel_config: ContextParallelConfig | None = None - - _rank: int = None - _world_size: int = None - _device: torch.device = None - _mesh: torch.distributed.device_mesh.DeviceMesh = None - - def setup( - self, - rank: int, - world_size: int, - device: torch.device, - *, - mesh: torch.distributed.device_mesh.DeviceMesh | None = None, - ): - self._rank = rank - self._world_size = world_size - self._device = device - self._mesh = mesh - if self.context_parallel_config is not None: - self.context_parallel_config.setup(rank, world_size, device, mesh) - - -@dataclass(frozen=True) -class ContextParallelInput: - """ - Configuration for splitting an input tensor across context parallel region. - - Args: - split_dim (`int`): - The dimension along which to split the tensor. - expected_dims (`int`, *optional*): - The expected number of dimensions of the tensor. If provided, a check will be performed to ensure that the - tensor has the expected number of dimensions before splitting. - split_output (`bool`, *optional*, defaults to `False`): - Whether to split the output tensor of the layer along the given `split_dim` instead of the input tensor. - This is useful for layers whose outputs should be split after it does some preprocessing on the inputs (ex: - RoPE). - """ - - split_dim: int - expected_dims: int | None = None - split_output: bool = False - - def __repr__(self): - return f"ContextParallelInput(split_dim={self.split_dim}, expected_dims={self.expected_dims}, split_output={self.split_output})" - - -@dataclass(frozen=True) -class ContextParallelOutput: - """ - Configuration for gathering an output tensor across context parallel region. - - Args: - gather_dim (`int`): - The dimension along which to gather the tensor. - expected_dims (`int`, *optional*): - The expected number of dimensions of the tensor. If provided, a check will be performed to ensure that the - tensor has the expected number of dimensions before gathering. - """ - - gather_dim: int - expected_dims: int | None = None - - def __repr__(self): - return f"ContextParallelOutput(gather_dim={self.gather_dim}, expected_dims={self.expected_dims})" - - -# A dictionary where keys denote the input to be split across context parallel region, and the -# value denotes the sharding configuration. -# If the key is a string, it denotes the name of the parameter in the forward function. -# If the key is an integer, split_output must be set to True, and it denotes the index of the output -# to be split across context parallel region. -ContextParallelInputType = dict[ - str | int, ContextParallelInput | list[ContextParallelInput] | tuple[ContextParallelInput, ...] -] - -# A dictionary where keys denote the output to be gathered across context parallel region, and the -# value denotes the gathering configuration. -ContextParallelOutputType = ContextParallelOutput | list[ContextParallelOutput] | tuple[ContextParallelOutput, ...] - -# A dictionary where keys denote the module id, and the value denotes how the inputs/outputs of -# the module should be split/gathered across context parallel region. -ContextParallelModelPlan = dict[str, ContextParallelInputType | ContextParallelOutputType] - - -# Example of a ContextParallelModelPlan (QwenImageTransformer2DModel): -# -# Each model should define a _cp_plan attribute that contains information on how to shard/gather -# tensors at different stages of the forward: -# -# ```python -# _cp_plan = { -# "": { -# "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), -# "encoder_hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), -# "encoder_hidden_states_mask": ContextParallelInput(split_dim=1, expected_dims=2, split_output=False), -# }, -# "pos_embed": { -# 0: ContextParallelInput(split_dim=0, expected_dims=2, split_output=True), -# 1: ContextParallelInput(split_dim=0, expected_dims=2, split_output=True), -# }, -# "proj_out": ContextParallelOutput(gather_dim=1, expected_dims=3), -# } -# ``` -# -# The dictionary is a set of module names mapped to their respective CP plan. The inputs/outputs of layers will be -# split/gathered according to this at the respective module level. Here, the following happens: -# - "": -# we specify that we want to split the various inputs across the sequence dim in the pre-forward hook (i.e. before -# the actual forward logic of the QwenImageTransformer2DModel is run, we will splitthe inputs) -# - "pos_embed": -# we specify that we want to split the outputs of the RoPE layer. Since there are two outputs (imag & text freqs), -# we can individually specify how they should be split -# - "proj_out": -# before returning to the user, we gather the entire sequence on each rank in the post-forward hook (after the linear -# layer forward has run). -# -# ContextParallelInput: -# specifies how to split the input tensor in the pre-forward or post-forward hook of the layer it is attached to -# -# ContextParallelOutput: -# specifies how to gather the input tensor in the post-forward hook in the layer it is attached to - - -# Below are utility functions for distributed communication in context parallelism. -def gather_size_by_comm(size: int, group: dist.ProcessGroup) -> list[int]: - r"""Gather the local size from all ranks. - size: int, local size return: list[int], list of size from all ranks - """ - # NOTE(Serving/CP Safety): - # Do NOT cache this collective result. - # - # In "Ulysses Anything" mode, `size` (e.g. per-rank local seq_len / S_LOCAL) - # may legitimately differ across ranks. If we cache based on the *local* `size`, - # different ranks can have different cache hit/miss patterns across time. - # - # That can lead to a catastrophic distributed hang: - # - some ranks hit cache and *skip* dist.all_gather() - # - other ranks miss cache and *enter* dist.all_gather() - # This mismatched collective participation will stall the process group and - # eventually trigger NCCL watchdog timeouts (often surfacing later as ALLTOALL - # timeouts in Ulysses attention). - world_size = dist.get_world_size(group=group) - # HACK: Use Gloo backend for all_gather to avoid H2D and D2H overhead - comm_backends = str(dist.get_backend(group=group)) - # NOTE: e.g., dist.init_process_group(backend="cpu:gloo,cuda:nccl") - gather_device = "cpu" if "cpu" in comm_backends else torch.accelerator.current_accelerator() - gathered_sizes = [torch.empty((1,), device=gather_device, dtype=torch.int64) for _ in range(world_size)] - dist.all_gather( - gathered_sizes, - torch.tensor([size], device=gather_device, dtype=torch.int64), - group=group, - ) - - gathered_sizes = [s[0].item() for s in gathered_sizes] - # NOTE: DON'T use tolist here due to graph break - Explanation: - # Backend compiler `inductor` failed with aten._local_scalar_dense.default - return gathered_sizes diff --git a/diffusers/models/activations.py b/diffusers/models/activations.py deleted file mode 100644 index 2d1fdb5f7d8303ae605afe4c6905af38950e4acb..0000000000000000000000000000000000000000 --- a/diffusers/models/activations.py +++ /dev/null @@ -1,178 +0,0 @@ -# coding=utf-8 -# Copyright 2025 HuggingFace Inc. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -import torch.nn.functional as F -from torch import nn - -from ..utils import deprecate -from ..utils.import_utils import is_torch_npu_available, is_torch_version - - -if is_torch_npu_available(): - import torch_npu - -ACT2CLS = { - "swish": nn.SiLU, - "silu": nn.SiLU, - "mish": nn.Mish, - "gelu": nn.GELU, - "relu": nn.ReLU, -} - - -def get_activation(act_fn: str) -> nn.Module: - """Helper function to get activation function from string. - - Args: - act_fn (str): Name of activation function. - - Returns: - nn.Module: Activation function. - """ - - act_fn = act_fn.lower() - if act_fn in ACT2CLS: - return ACT2CLS[act_fn]() - else: - raise ValueError(f"activation function {act_fn} not found in ACT2FN mapping {list(ACT2CLS.keys())}") - - -class FP32SiLU(nn.Module): - r""" - SiLU activation function with input upcasted to torch.float32. - """ - - def __init__(self): - super().__init__() - - def forward(self, inputs: torch.Tensor) -> torch.Tensor: - return F.silu(inputs.float(), inplace=False).to(inputs.dtype) - - -class GELU(nn.Module): - r""" - GELU activation function with tanh approximation support with `approximate="tanh"`. - - Parameters: - dim_in (`int`): The number of channels in the input. - dim_out (`int`): The number of channels in the output. - approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation. - bias (`bool`, defaults to True): Whether to use a bias in the linear layer. - """ - - def __init__(self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True): - super().__init__() - self.proj = nn.Linear(dim_in, dim_out, bias=bias) - self.approximate = approximate - - def gelu(self, gate: torch.Tensor) -> torch.Tensor: - if gate.device.type == "mps" and is_torch_version("<", "2.0.0"): - # fp16 gelu not supported on mps before torch 2.0 - return F.gelu(gate.to(dtype=torch.float32), approximate=self.approximate).to(dtype=gate.dtype) - return F.gelu(gate, approximate=self.approximate) - - def forward(self, hidden_states): - hidden_states = self.proj(hidden_states) - hidden_states = self.gelu(hidden_states) - return hidden_states - - -class GEGLU(nn.Module): - r""" - A [variant](https://huggingface.co/papers/2002.05202) of the gated linear unit activation function. - - Parameters: - dim_in (`int`): The number of channels in the input. - dim_out (`int`): The number of channels in the output. - bias (`bool`, defaults to True): Whether to use a bias in the linear layer. - """ - - def __init__(self, dim_in: int, dim_out: int, bias: bool = True): - super().__init__() - self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias) - - def gelu(self, gate: torch.Tensor) -> torch.Tensor: - if gate.device.type == "mps" and is_torch_version("<", "2.0.0"): - # fp16 gelu not supported on mps before torch 2.0 - return F.gelu(gate.to(dtype=torch.float32)).to(dtype=gate.dtype) - return F.gelu(gate) - - def forward(self, hidden_states, *args, **kwargs): - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - hidden_states = self.proj(hidden_states) - if is_torch_npu_available(): - # using torch_npu.npu_geglu can run faster and save memory on NPU. - return torch_npu.npu_geglu(hidden_states, dim=-1, approximate=1)[0] - else: - hidden_states, gate = hidden_states.chunk(2, dim=-1) - return hidden_states * self.gelu(gate) - - -class SwiGLU(nn.Module): - r""" - A [variant](https://huggingface.co/papers/2002.05202) of the gated linear unit activation function. It's similar to - `GEGLU` but uses SiLU / Swish instead of GeLU. - - Parameters: - dim_in (`int`): The number of channels in the input. - dim_out (`int`): The number of channels in the output. - bias (`bool`, defaults to True): Whether to use a bias in the linear layer. - """ - - def __init__(self, dim_in: int, dim_out: int, bias: bool = True): - super().__init__() - - self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias) - self.activation = nn.SiLU() - - def forward(self, hidden_states): - hidden_states = self.proj(hidden_states) - hidden_states, gate = hidden_states.chunk(2, dim=-1) - return hidden_states * self.activation(gate) - - -class ApproximateGELU(nn.Module): - r""" - The approximate form of the Gaussian Error Linear Unit (GELU). For more details, see section 2 of this - [paper](https://huggingface.co/papers/1606.08415). - - Parameters: - dim_in (`int`): The number of channels in the input. - dim_out (`int`): The number of channels in the output. - bias (`bool`, defaults to True): Whether to use a bias in the linear layer. - """ - - def __init__(self, dim_in: int, dim_out: int, bias: bool = True): - super().__init__() - self.proj = nn.Linear(dim_in, dim_out, bias=bias) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = self.proj(x) - return x * torch.sigmoid(1.702 * x) - - -class LinearActivation(nn.Module): - def __init__(self, dim_in: int, dim_out: int, bias: bool = True, activation: str = "silu"): - super().__init__() - - self.proj = nn.Linear(dim_in, dim_out, bias=bias) - self.activation = get_activation(activation) - - def forward(self, hidden_states): - hidden_states = self.proj(hidden_states) - return self.activation(hidden_states) diff --git a/diffusers/models/adapter.py b/diffusers/models/adapter.py deleted file mode 100644 index 2072749c65ae8f626f17c62819c975a1d6019176..0000000000000000000000000000000000000000 --- a/diffusers/models/adapter.py +++ /dev/null @@ -1,596 +0,0 @@ -# Copyright 2022 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import os -from typing import Callable - -import torch -import torch.nn as nn - -from ..configuration_utils import ConfigMixin, register_to_config -from ..utils import logging -from .modeling_utils import ModelMixin - - -logger = logging.get_logger(__name__) - - -class MultiAdapter(ModelMixin): - r""" - MultiAdapter is a wrapper model that contains multiple adapter models and merges their outputs according to - user-assigned weighting. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for common methods such as downloading - or saving. - - Args: - adapters (`list[T2IAdapter]`, *optional*, defaults to None): - A list of `T2IAdapter` model instances. - """ - - def __init__(self, adapters: list["T2IAdapter"]): - super(MultiAdapter, self).__init__() - - self.num_adapter = len(adapters) - self.adapters = nn.ModuleList(adapters) - - if len(adapters) == 0: - raise ValueError("Expecting at least one adapter") - - if len(adapters) == 1: - raise ValueError("For a single adapter, please use the `T2IAdapter` class instead of `MultiAdapter`") - - # The outputs from each adapter are added together with a weight. - # This means that the change in dimensions from downsampling must - # be the same for all adapters. Inductively, it also means the - # downscale_factor and total_downscale_factor must be the same for all - # adapters. - first_adapter_total_downscale_factor = adapters[0].total_downscale_factor - first_adapter_downscale_factor = adapters[0].downscale_factor - for idx in range(1, len(adapters)): - if ( - adapters[idx].total_downscale_factor != first_adapter_total_downscale_factor - or adapters[idx].downscale_factor != first_adapter_downscale_factor - ): - raise ValueError( - f"Expecting all adapters to have the same downscaling behavior, but got:\n" - f"adapters[0].total_downscale_factor={first_adapter_total_downscale_factor}\n" - f"adapters[0].downscale_factor={first_adapter_downscale_factor}\n" - f"adapter[`{idx}`].total_downscale_factor={adapters[idx].total_downscale_factor}\n" - f"adapter[`{idx}`].downscale_factor={adapters[idx].downscale_factor}" - ) - - self.total_downscale_factor = first_adapter_total_downscale_factor - self.downscale_factor = first_adapter_downscale_factor - - def forward(self, xs: torch.Tensor, adapter_weights: list[float] | None = None) -> list[torch.Tensor]: - r""" - Args: - xs (`torch.Tensor`): - A tensor of shape (batch, channel, height, width) representing input images for multiple adapter - models, concatenated along dimension 1(channel dimension). The `channel` dimension should be equal to - `num_adapter` * number of channel per image. - - adapter_weights (`list[float]`, *optional*, defaults to None): - A list of floats representing the weights which will be multiplied by each adapter's output before - summing them together. If `None`, equal weights will be used for all adapters. - - Returns: - `list[torch.Tensor]`: - A list of feature tensors, one per scale, obtained by summing the per-scale features of each adapter - weighted by `adapter_weights`. - """ - if adapter_weights is None: - adapter_weights = torch.tensor([1 / self.num_adapter] * self.num_adapter) - else: - adapter_weights = torch.tensor(adapter_weights) - - accume_state = None - for x, w, adapter in zip(xs, adapter_weights, self.adapters): - features = adapter(x) - if accume_state is None: - accume_state = features - for i in range(len(accume_state)): - accume_state[i] = w * accume_state[i] - else: - for i in range(len(features)): - accume_state[i] += w * features[i] - return accume_state - - def save_pretrained( - self, - save_directory: str | os.PathLike, - is_main_process: bool = True, - save_function: Callable = None, - safe_serialization: bool = True, - variant: str | None = None, - ): - """ - Save a model and its configuration file to a specified directory, allowing it to be re-loaded with the - `[`~models.adapter.MultiAdapter.from_pretrained`]` class method. - - Args: - save_directory (`str` or `os.PathLike`): - The directory where the model will be saved. If the directory does not exist, it will be created. - is_main_process (`bool`, optional, defaults=True): - Indicates whether current process is the main process or not. Useful for distributed training (e.g., - TPUs) and need to call this function on all processes. In this case, set `is_main_process=True` only - for the main process to avoid race conditions. - save_function (`Callable`): - Function used to save the state dictionary. Useful for distributed training (e.g., TPUs) to replace - `torch.save` with another method. Can also be configured using`DIFFUSERS_SAVE_MODE` environment - variable. - safe_serialization (`bool`, optional, defaults=True): - If `True`, save the model using `safetensors`. If `False`, save the model with `pickle`. - variant (`str`, *optional*): - If specified, weights are saved in the format `pytorch_model..bin`. - """ - idx = 0 - model_path_to_save = save_directory - for adapter in self.adapters: - adapter.save_pretrained( - model_path_to_save, - is_main_process=is_main_process, - save_function=save_function, - safe_serialization=safe_serialization, - variant=variant, - ) - - idx += 1 - model_path_to_save = model_path_to_save + f"_{idx}" - - @classmethod - def from_pretrained(cls, pretrained_model_path: str | os.PathLike | None, **kwargs): - r""" - Instantiate a pretrained `MultiAdapter` model from multiple pre-trained adapter models. - - The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated). To train - the model, set it back to training mode using `model.train()`. - - Warnings: - *Weights from XXX not initialized from pretrained model* means that the weights of XXX are not pretrained - with the rest of the model. It is up to you to train those weights with a downstream fine-tuning. *Weights - from XXX not used in YYY* means that the layer XXX is not used by YYY, so those weights are discarded. - - Args: - pretrained_model_path (`os.PathLike`): - A path to a *directory* containing model weights saved using - [`~diffusers.models.adapter.MultiAdapter.save_pretrained`], e.g., `./my_model_directory/adapter`. - dtype (`torch.dtype`, *optional*): - Override the default `torch.dtype` and load the model under this dtype. - output_loading_info(`bool`, *optional*, defaults to `False`): - Whether or not to also return a dictionary containing missing keys, unexpected keys and error messages. - device_map (`str` or `dict[str, int | str | torch.device]`, *optional*): - A map that specifies where each submodule should go. It doesn't need to be refined to each - parameter/buffer name, once a given module name is inside, every submodule of it will be sent to the - same device. - - To have Accelerate compute the most optimized `device_map` automatically, set `device_map="auto"`. For - more information about each option see [designing a device - map](https://hf.co/docs/accelerate/main/en/usage_guides/big_modeling#designing-a-device-map). - max_memory (`Dict`, *optional*): - A dictionary mapping device identifiers to their maximum memory. Default to the maximum memory - available for each GPU and the available CPU RAM if unset. - low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 else `False`): - Speed up model loading by not initializing the weights and only loading the pre-trained weights. This - also tries to not use more than 1x model size in CPU memory (including peak memory) while loading the - model. This is only supported when torch version >= 1.9.0. If you are using an older version of torch, - setting this argument to `True` will raise an error. - variant (`str`, *optional*): - If specified, load weights from a `variant` file (*e.g.* pytorch_model..bin). - use_safetensors (`bool`, *optional*, defaults to `None`): - If `None`, the `safetensors` weights will be downloaded if available **and** if`safetensors` library is - installed. If `True`, the model will be forcibly loaded from`safetensors` weights. If `False`, - `safetensors` is not used. - """ - idx = 0 - adapters = [] - - # load adapter and append to list until no adapter directory exists anymore - # first adapter has to be saved under `./mydirectory/adapter` to be compliant with `DiffusionPipeline.from_pretrained` - # second, third, ... adapters have to be saved under `./mydirectory/adapter_1`, `./mydirectory/adapter_2`, ... - model_path_to_load = pretrained_model_path - while os.path.isdir(model_path_to_load): - adapter = T2IAdapter.from_pretrained(model_path_to_load, **kwargs) - adapters.append(adapter) - - idx += 1 - model_path_to_load = pretrained_model_path + f"_{idx}" - - logger.info(f"{len(adapters)} adapters loaded from {pretrained_model_path}.") - - if len(adapters) == 0: - raise ValueError( - f"No T2IAdapters found under {os.path.dirname(pretrained_model_path)}. Expected at least {pretrained_model_path + '_0'}." - ) - - return cls(adapters) - - -class T2IAdapter(ModelMixin, ConfigMixin): - r""" - A simple ResNet-like model that accepts images containing control signals such as keyposes and depth. The model - generates multiple feature maps that are used as additional conditioning in [`UNet2DConditionModel`]. The model's - architecture follows the original implementation of - [Adapter](https://github.com/TencentARC/T2I-Adapter/blob/686de4681515662c0ac2ffa07bf5dda83af1038a/ldm/modules/encoders/adapter.py#L97) - and - [AdapterLight](https://github.com/TencentARC/T2I-Adapter/blob/686de4681515662c0ac2ffa07bf5dda83af1038a/ldm/modules/encoders/adapter.py#L235). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for the common methods, such as - downloading or saving. - - Args: - in_channels (`int`, *optional*, defaults to `3`): - The number of channels in the adapter's input (*control image*). Set it to 1 if you're using a gray scale - image. - channels (`list[int]`, *optional*, defaults to `(320, 640, 1280, 1280)`): - The number of channels in each downsample block's output hidden state. The `len(block_out_channels)` - determines the number of downsample blocks in the adapter. - num_res_blocks (`int`, *optional*, defaults to `2`): - Number of ResNet blocks in each downsample block. - downscale_factor (`int`, *optional*, defaults to `8`): - A factor that determines the total downscale factor of the Adapter. - adapter_type (`str`, *optional*, defaults to `full_adapter`): - Adapter type (`full_adapter` or `full_adapter_xl` or `light_adapter`) to use. - """ - - @register_to_config - def __init__( - self, - in_channels: int = 3, - channels: list[int] = [320, 640, 1280, 1280], - num_res_blocks: int = 2, - downscale_factor: int = 8, - adapter_type: str = "full_adapter", - ): - super().__init__() - - if adapter_type == "full_adapter": - self.adapter = FullAdapter(in_channels, channels, num_res_blocks, downscale_factor) - elif adapter_type == "full_adapter_xl": - self.adapter = FullAdapterXL(in_channels, channels, num_res_blocks, downscale_factor) - elif adapter_type == "light_adapter": - self.adapter = LightAdapter(in_channels, channels, num_res_blocks, downscale_factor) - else: - raise ValueError( - f"Unsupported adapter_type: '{adapter_type}'. Choose either 'full_adapter' or " - "'full_adapter_xl' or 'light_adapter'." - ) - - def forward(self, x: torch.Tensor) -> list[torch.Tensor]: - r""" - This function processes the input tensor `x` through the adapter model and returns a list of feature tensors, - each representing information extracted at a different scale from the input. The length of the list is - determined by the number of downsample blocks in the Adapter, as specified by the `channels` and - `num_res_blocks` parameters during initialization. - - Args: - x (`torch.Tensor`): - The input tensor to process through the adapter model. - - Returns: - `list[torch.Tensor]`: - A list of feature tensors, each representing information extracted at a different scale from the input. - The length of the list equals the number of downsample blocks in the adapter. - """ - return self.adapter(x) - - @property - def total_downscale_factor(self): - return self.adapter.total_downscale_factor - - @property - def downscale_factor(self): - """The downscale factor applied in the T2I-Adapter's initial pixel unshuffle operation. If an input image's dimensions are - not evenly divisible by the downscale_factor then an exception will be raised. - """ - return self.adapter.unshuffle.downscale_factor - - -# full adapter - - -class FullAdapter(nn.Module): - r""" - See [`T2IAdapter`] for more information. - """ - - def __init__( - self, - in_channels: int = 3, - channels: list[int] = [320, 640, 1280, 1280], - num_res_blocks: int = 2, - downscale_factor: int = 8, - ): - super().__init__() - - in_channels = in_channels * downscale_factor**2 - - self.unshuffle = nn.PixelUnshuffle(downscale_factor) - self.conv_in = nn.Conv2d(in_channels, channels[0], kernel_size=3, padding=1) - - self.body = nn.ModuleList( - [ - AdapterBlock(channels[0], channels[0], num_res_blocks), - *[ - AdapterBlock(channels[i - 1], channels[i], num_res_blocks, down=True) - for i in range(1, len(channels)) - ], - ] - ) - - self.total_downscale_factor = downscale_factor * 2 ** (len(channels) - 1) - - def forward(self, x: torch.Tensor) -> list[torch.Tensor]: - r""" - This method processes the input tensor `x` through the FullAdapter model and performs operations including - pixel unshuffling, convolution, and a stack of AdapterBlocks. It returns a list of feature tensors, each - capturing information at a different stage of processing within the FullAdapter model. The number of feature - tensors in the list is determined by the number of downsample blocks specified during initialization. - """ - x = self.unshuffle(x) - x = self.conv_in(x) - - features = [] - - for block in self.body: - x = block(x) - features.append(x) - - return features - - -class FullAdapterXL(nn.Module): - r""" - See [`T2IAdapter`] for more information. - """ - - def __init__( - self, - in_channels: int = 3, - channels: list[int] = [320, 640, 1280, 1280], - num_res_blocks: int = 2, - downscale_factor: int = 16, - ): - super().__init__() - - in_channels = in_channels * downscale_factor**2 - - self.unshuffle = nn.PixelUnshuffle(downscale_factor) - self.conv_in = nn.Conv2d(in_channels, channels[0], kernel_size=3, padding=1) - - self.body = [] - # blocks to extract XL features with dimensions of [320, 64, 64], [640, 64, 64], [1280, 32, 32], [1280, 32, 32] - for i in range(len(channels)): - if i == 1: - self.body.append(AdapterBlock(channels[i - 1], channels[i], num_res_blocks)) - elif i == 2: - self.body.append(AdapterBlock(channels[i - 1], channels[i], num_res_blocks, down=True)) - else: - self.body.append(AdapterBlock(channels[i], channels[i], num_res_blocks)) - - self.body = nn.ModuleList(self.body) - # XL has only one downsampling AdapterBlock. - self.total_downscale_factor = downscale_factor * 2 - - def forward(self, x: torch.Tensor) -> list[torch.Tensor]: - r""" - This method takes the tensor x as input and processes it through FullAdapterXL model. It consists of operations - including unshuffling pixels, applying convolution layer and appending each block into list of feature tensors. - """ - x = self.unshuffle(x) - x = self.conv_in(x) - - features = [] - - for block in self.body: - x = block(x) - features.append(x) - - return features - - -class AdapterBlock(nn.Module): - r""" - An AdapterBlock is a helper model that contains multiple ResNet-like blocks. It is used in the `FullAdapter` and - `FullAdapterXL` models. - - Args: - in_channels (`int`): - Number of channels of AdapterBlock's input. - out_channels (`int`): - Number of channels of AdapterBlock's output. - num_res_blocks (`int`): - Number of ResNet blocks in the AdapterBlock. - down (`bool`, *optional*, defaults to `False`): - If `True`, perform downsampling on AdapterBlock's input. - """ - - def __init__(self, in_channels: int, out_channels: int, num_res_blocks: int, down: bool = False): - super().__init__() - - self.downsample = None - if down: - self.downsample = nn.AvgPool2d(kernel_size=2, stride=2, ceil_mode=True) - - self.in_conv = None - if in_channels != out_channels: - self.in_conv = nn.Conv2d(in_channels, out_channels, kernel_size=1) - - self.resnets = nn.Sequential( - *[AdapterResnetBlock(out_channels) for _ in range(num_res_blocks)], - ) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - r""" - This method takes tensor x as input and performs operations downsampling and convolutional layers if the - self.downsample and self.in_conv properties of AdapterBlock model are specified. Then it applies a series of - residual blocks to the input tensor. - """ - if self.downsample is not None: - x = self.downsample(x) - - if self.in_conv is not None: - x = self.in_conv(x) - - x = self.resnets(x) - - return x - - -class AdapterResnetBlock(nn.Module): - r""" - An `AdapterResnetBlock` is a helper model that implements a ResNet-like block. - - Args: - channels (`int`): - Number of channels of AdapterResnetBlock's input and output. - """ - - def __init__(self, channels: int): - super().__init__() - self.block1 = nn.Conv2d(channels, channels, kernel_size=3, padding=1) - self.act = nn.ReLU() - self.block2 = nn.Conv2d(channels, channels, kernel_size=1) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - r""" - This method takes input tensor x and applies a convolutional layer, ReLU activation, and another convolutional - layer on the input tensor. It returns addition with the input tensor. - """ - - h = self.act(self.block1(x)) - h = self.block2(h) - - return h + x - - -# light adapter - - -class LightAdapter(nn.Module): - r""" - See [`T2IAdapter`] for more information. - """ - - def __init__( - self, - in_channels: int = 3, - channels: list[int] = [320, 640, 1280], - num_res_blocks: int = 4, - downscale_factor: int = 8, - ): - super().__init__() - - in_channels = in_channels * downscale_factor**2 - - self.unshuffle = nn.PixelUnshuffle(downscale_factor) - - self.body = nn.ModuleList( - [ - LightAdapterBlock(in_channels, channels[0], num_res_blocks), - *[ - LightAdapterBlock(channels[i], channels[i + 1], num_res_blocks, down=True) - for i in range(len(channels) - 1) - ], - LightAdapterBlock(channels[-1], channels[-1], num_res_blocks, down=True), - ] - ) - - self.total_downscale_factor = downscale_factor * (2 ** len(channels)) - - def forward(self, x: torch.Tensor) -> list[torch.Tensor]: - r""" - This method takes the input tensor x and performs downscaling and appends it in list of feature tensors. Each - feature tensor corresponds to a different level of processing within the LightAdapter. - """ - x = self.unshuffle(x) - - features = [] - - for block in self.body: - x = block(x) - features.append(x) - - return features - - -class LightAdapterBlock(nn.Module): - r""" - A `LightAdapterBlock` is a helper model that contains multiple `LightAdapterResnetBlocks`. It is used in the - `LightAdapter` model. - - Args: - in_channels (`int`): - Number of channels of LightAdapterBlock's input. - out_channels (`int`): - Number of channels of LightAdapterBlock's output. - num_res_blocks (`int`): - Number of LightAdapterResnetBlocks in the LightAdapterBlock. - down (`bool`, *optional*, defaults to `False`): - If `True`, perform downsampling on LightAdapterBlock's input. - """ - - def __init__(self, in_channels: int, out_channels: int, num_res_blocks: int, down: bool = False): - super().__init__() - mid_channels = out_channels // 4 - - self.downsample = None - if down: - self.downsample = nn.AvgPool2d(kernel_size=2, stride=2, ceil_mode=True) - - self.in_conv = nn.Conv2d(in_channels, mid_channels, kernel_size=1) - self.resnets = nn.Sequential(*[LightAdapterResnetBlock(mid_channels) for _ in range(num_res_blocks)]) - self.out_conv = nn.Conv2d(mid_channels, out_channels, kernel_size=1) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - r""" - This method takes tensor x as input and performs downsampling if required. Then it applies in convolution - layer, a sequence of residual blocks, and out convolutional layer. - """ - if self.downsample is not None: - x = self.downsample(x) - - x = self.in_conv(x) - x = self.resnets(x) - x = self.out_conv(x) - - return x - - -class LightAdapterResnetBlock(nn.Module): - """ - A `LightAdapterResnetBlock` is a helper model that implements a ResNet-like block with a slightly different - architecture than `AdapterResnetBlock`. - - Args: - channels (`int`): - Number of channels of LightAdapterResnetBlock's input and output. - """ - - def __init__(self, channels: int): - super().__init__() - self.block1 = nn.Conv2d(channels, channels, kernel_size=3, padding=1) - self.act = nn.ReLU() - self.block2 = nn.Conv2d(channels, channels, kernel_size=3, padding=1) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - r""" - This function takes input tensor x and processes it through one convolutional layer, ReLU activation, and - another convolutional layer and adds it to input tensor. - """ - - h = self.act(self.block1(x)) - h = self.block2(h) - - return h + x diff --git a/diffusers/models/attention.py b/diffusers/models/attention.py deleted file mode 100644 index 5d949050397428ec447ba050f5716a34e457138a..0000000000000000000000000000000000000000 --- a/diffusers/models/attention.py +++ /dev/null @@ -1,1742 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any, Callable - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ..utils import deprecate, logging -from ..utils.import_utils import is_torch_npu_available, is_torch_xla_available, is_xformers_available -from ..utils.torch_utils import maybe_allow_in_graph -from .activations import GEGLU, GELU, ApproximateGELU, FP32SiLU, LinearActivation, SwiGLU -from .attention_processor import Attention, AttentionProcessor, JointAttnProcessor2_0 -from .embeddings import SinusoidalPositionalEmbedding -from .normalization import AdaLayerNorm, AdaLayerNormContinuous, AdaLayerNormZero, RMSNorm, SD35AdaLayerNormZeroX - - -if is_xformers_available(): - import xformers as xops -else: - xops = None - - -logger = logging.get_logger(__name__) - - -class AttentionMixin: - @property - def attn_processors(self) -> dict[str, AttentionProcessor]: - r""" - Returns: - `dict` of attention processors: A dictionary containing all attention processors used in the model with - indexed by its weight name. - """ - # set recursively - processors = {} - - def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: dict[str, AttentionProcessor]): - if hasattr(module, "get_processor"): - processors[f"{name}.processor"] = module.get_processor() - - for sub_name, child in module.named_children(): - fn_recursive_add_processors(f"{name}.{sub_name}", child, processors) - - return processors - - for name, module in self.named_children(): - fn_recursive_add_processors(name, module, processors) - - return processors - - def set_attn_processor(self, processor: AttentionProcessor | dict[str, AttentionProcessor]): - r""" - Sets the attention processor to use to compute attention. - - Parameters: - processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`): - The instantiated processor class or a dictionary of processor classes that will be set as the processor - for **all** `Attention` layers. - - If `processor` is a dict, the key needs to define the path to the corresponding cross attention - processor. This is strongly recommended when setting trainable attention processors. - - """ - count = len(self.attn_processors.keys()) - - if isinstance(processor, dict) and len(processor) != count: - raise ValueError( - f"A dict of processors was passed, but the number of processors {len(processor)} does not match the" - f" number of attention layers: {count}. Please make sure to pass {count} processor classes." - ) - - def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor): - if hasattr(module, "set_processor"): - if not isinstance(processor, dict): - module.set_processor(processor) - else: - module.set_processor(processor.pop(f"{name}.processor")) - - for sub_name, child in module.named_children(): - fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor) - - for name, module in self.named_children(): - fn_recursive_attn_processor(name, module, processor) - - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - """ - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - for module in self.modules(): - if isinstance(module, AttentionModuleMixin) and module._supports_qkv_fusion: - module.fuse_projections() - - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - """ - for module in self.modules(): - if isinstance(module, AttentionModuleMixin) and module._supports_qkv_fusion: - module.unfuse_projections() - - -class AttentionModuleMixin: - _default_processor_cls = None - _available_processors = [] - _supports_qkv_fusion = True - fused_projections = False - - def set_processor(self, processor: AttentionProcessor) -> None: - """ - Set the attention processor to use. - - Args: - processor (`AttnProcessor`): - The attention processor to use. - """ - # if current processor is in `self._modules` and if passed `processor` is not, we need to - # pop `processor` from `self._modules` - if ( - hasattr(self, "processor") - and isinstance(self.processor, torch.nn.Module) - and not isinstance(processor, torch.nn.Module) - ): - logger.info(f"You are removing possibly trained weights of {self.processor} with {processor}") - self._modules.pop("processor") - - self.processor = processor - - def get_processor(self, return_deprecated_lora: bool = False) -> "AttentionProcessor": - """ - Get the attention processor in use. - - Args: - return_deprecated_lora (`bool`, *optional*, defaults to `False`): - Set to `True` to return the deprecated LoRA attention processor. - - Returns: - "AttentionProcessor": The attention processor in use. - """ - if not return_deprecated_lora: - return self.processor - - def set_attention_backend(self, backend: str): - from .attention_dispatch import AttentionBackendName - - available_backends = {x.value for x in AttentionBackendName.__members__.values()} - if backend not in available_backends: - raise ValueError(f"`{backend=}` must be one of the following: " + ", ".join(available_backends)) - - backend = AttentionBackendName(backend.lower()) - self.processor._attention_backend = backend - - def set_use_npu_flash_attention(self, use_npu_flash_attention: bool) -> None: - """ - Set whether to use NPU flash attention from `torch_npu` or not. - - Args: - use_npu_flash_attention (`bool`): Whether to use NPU flash attention or not. - """ - - if use_npu_flash_attention: - if not is_torch_npu_available(): - raise ImportError("torch_npu is not available") - - self.set_attention_backend("_native_npu") - - def set_use_xla_flash_attention( - self, - use_xla_flash_attention: bool, - partition_spec: tuple[str | None, ...] | None = None, - is_flux=False, - ) -> None: - """ - Set whether to use XLA flash attention from `torch_xla` or not. - - Args: - use_xla_flash_attention (`bool`): - Whether to use pallas flash attention kernel from `torch_xla` or not. - partition_spec (`tuple[]`, *optional*): - Specify the partition specification if using SPMD. Otherwise None. - is_flux (`bool`, *optional*, defaults to `False`): - Whether the model is a Flux model. - """ - if use_xla_flash_attention: - if not is_torch_xla_available(): - raise ImportError("torch_xla is not available") - - self.set_attention_backend("_native_xla") - - def set_use_memory_efficient_attention_xformers( - self, use_memory_efficient_attention_xformers: bool, attention_op: Callable | None = None - ) -> None: - """ - Set whether to use memory efficient attention from `xformers` or not. - - Args: - use_memory_efficient_attention_xformers (`bool`): - Whether to use memory efficient attention from `xformers` or not. - attention_op (`Callable`, *optional*): - The attention operation to use. Defaults to `None` which uses the default attention operation from - `xformers`. - """ - if use_memory_efficient_attention_xformers: - if not is_xformers_available(): - raise ModuleNotFoundError( - "Refer to https://github.com/facebookresearch/xformers for more information on how to install xformers", - name="xformers", - ) - elif not torch.cuda.is_available(): - raise ValueError( - "torch.cuda.is_available() should be True but is False. xformers' memory efficient attention is" - " only available for GPU " - ) - else: - try: - # Make sure we can run the memory efficient attention - if is_xformers_available(): - dtype = None - if attention_op is not None: - op_fw, op_bw = attention_op - dtype, *_ = op_fw.SUPPORTED_DTYPES - q = torch.randn((1, 2, 40), device="cuda", dtype=dtype) - _ = xops.ops.memory_efficient_attention(q, q, q) - except Exception as e: - raise e - - self.set_attention_backend("xformers") - - @torch.no_grad() - def fuse_projections(self): - """ - Fuse the query, key, and value projections into a single projection for efficiency. - """ - # Skip if the AttentionModuleMixin subclass does not support fusion (for example, the QKV projections in Flux2 - # single stream blocks are always fused) - if not self._supports_qkv_fusion: - logger.debug( - f"{self.__class__.__name__} does not support fusing QKV projections, so `fuse_projections` will no-op." - ) - return - - # Skip if already fused - if getattr(self, "fused_projections", False): - return - - device = self.to_q.weight.data.device - dtype = self.to_q.weight.data.dtype - - if hasattr(self, "is_cross_attention") and self.is_cross_attention: - # Fuse cross-attention key-value projections - concatenated_weights = torch.cat([self.to_k.weight.data, self.to_v.weight.data]) - in_features = concatenated_weights.shape[1] - out_features = concatenated_weights.shape[0] - - self.to_kv = nn.Linear(in_features, out_features, bias=self.use_bias, device=device, dtype=dtype) - self.to_kv.weight.copy_(concatenated_weights) - if hasattr(self, "use_bias") and self.use_bias: - concatenated_bias = torch.cat([self.to_k.bias.data, self.to_v.bias.data]) - self.to_kv.bias.copy_(concatenated_bias) - else: - # Fuse self-attention projections - concatenated_weights = torch.cat([self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data]) - in_features = concatenated_weights.shape[1] - out_features = concatenated_weights.shape[0] - - self.to_qkv = nn.Linear(in_features, out_features, bias=self.use_bias, device=device, dtype=dtype) - self.to_qkv.weight.copy_(concatenated_weights) - if hasattr(self, "use_bias") and self.use_bias: - concatenated_bias = torch.cat([self.to_q.bias.data, self.to_k.bias.data, self.to_v.bias.data]) - self.to_qkv.bias.copy_(concatenated_bias) - - # Handle added projections for models like SD3, Flux, etc. - if ( - getattr(self, "add_q_proj", None) is not None - and getattr(self, "add_k_proj", None) is not None - and getattr(self, "add_v_proj", None) is not None - ): - concatenated_weights = torch.cat( - [self.add_q_proj.weight.data, self.add_k_proj.weight.data, self.add_v_proj.weight.data] - ) - in_features = concatenated_weights.shape[1] - out_features = concatenated_weights.shape[0] - - self.to_added_qkv = nn.Linear( - in_features, out_features, bias=self.added_proj_bias, device=device, dtype=dtype - ) - self.to_added_qkv.weight.copy_(concatenated_weights) - if self.added_proj_bias: - concatenated_bias = torch.cat( - [self.add_q_proj.bias.data, self.add_k_proj.bias.data, self.add_v_proj.bias.data] - ) - self.to_added_qkv.bias.copy_(concatenated_bias) - - self.fused_projections = True - - @torch.no_grad() - def unfuse_projections(self): - """ - Unfuse the query, key, and value projections back to separate projections. - """ - # Skip if the AttentionModuleMixin subclass does not support fusion (for example, the QKV projections in Flux2 - # single stream blocks are always fused) - if not self._supports_qkv_fusion: - return - - # Skip if not fused - if not getattr(self, "fused_projections", False): - return - - # Remove fused projection layers - if hasattr(self, "to_qkv"): - delattr(self, "to_qkv") - - if hasattr(self, "to_kv"): - delattr(self, "to_kv") - - if hasattr(self, "to_added_qkv"): - delattr(self, "to_added_qkv") - - self.fused_projections = False - - def set_attention_slice(self, slice_size: int) -> None: - """ - Set the slice size for attention computation. - - Args: - slice_size (`int`): - The slice size for attention computation. - """ - if hasattr(self, "sliceable_head_dim") and slice_size is not None and slice_size > self.sliceable_head_dim: - raise ValueError(f"slice_size {slice_size} has to be smaller or equal to {self.sliceable_head_dim}.") - - processor = None - - # Try to get a compatible processor for sliced attention - if slice_size is not None: - processor = self._get_compatible_processor("sliced") - - # If no processor was found or slice_size is None, use default processor - if processor is None: - processor = self.default_processor_cls() - - self.set_processor(processor) - - def batch_to_head_dim(self, tensor: torch.Tensor) -> torch.Tensor: - """ - Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size // heads, seq_len, dim * heads]`. - - Args: - tensor (`torch.Tensor`): The tensor to reshape. - - Returns: - `torch.Tensor`: The reshaped tensor. - """ - head_size = self.heads - batch_size, seq_len, dim = tensor.shape - tensor = tensor.reshape(batch_size // head_size, head_size, seq_len, dim) - tensor = tensor.permute(0, 2, 1, 3).reshape(batch_size // head_size, seq_len, dim * head_size) - return tensor - - def head_to_batch_dim(self, tensor: torch.Tensor, out_dim: int = 3) -> torch.Tensor: - """ - Reshape the tensor for multi-head attention processing. - - Args: - tensor (`torch.Tensor`): The tensor to reshape. - out_dim (`int`, *optional*, defaults to `3`): The output dimension of the tensor. - - Returns: - `torch.Tensor`: The reshaped tensor. - """ - head_size = self.heads - if tensor.ndim == 3: - batch_size, seq_len, dim = tensor.shape - extra_dim = 1 - else: - batch_size, extra_dim, seq_len, dim = tensor.shape - tensor = tensor.reshape(batch_size, seq_len * extra_dim, head_size, dim // head_size) - tensor = tensor.permute(0, 2, 1, 3) - - if out_dim == 3: - tensor = tensor.reshape(batch_size * head_size, seq_len * extra_dim, dim // head_size) - - return tensor - - def get_attention_scores( - self, query: torch.Tensor, key: torch.Tensor, attention_mask: torch.Tensor | None = None - ) -> torch.Tensor: - """ - Compute the attention scores. - - Args: - query (`torch.Tensor`): The query tensor. - key (`torch.Tensor`): The key tensor. - attention_mask (`torch.Tensor`, *optional*): The attention mask to use. - - Returns: - `torch.Tensor`: The attention probabilities/scores. - """ - dtype = query.dtype - if self.upcast_attention: - query = query.float() - key = key.float() - - if attention_mask is None: - baddbmm_input = torch.empty( - query.shape[0], query.shape[1], key.shape[1], dtype=query.dtype, device=query.device - ) - beta = 0 - else: - baddbmm_input = attention_mask - beta = 1 - - attention_scores = torch.baddbmm( - baddbmm_input, - query, - key.transpose(-1, -2), - beta=beta, - alpha=self.scale, - ) - del baddbmm_input - - if self.upcast_softmax: - attention_scores = attention_scores.float() - - attention_probs = attention_scores.softmax(dim=-1) - del attention_scores - - attention_probs = attention_probs.to(dtype) - - return attention_probs - - def prepare_attention_mask( - self, attention_mask: torch.Tensor, target_length: int, batch_size: int, out_dim: int = 3 - ) -> torch.Tensor: - """ - Prepare the attention mask for the attention computation. - - Args: - attention_mask (`torch.Tensor`): The attention mask to prepare. - target_length (`int`): The target length of the attention mask. - batch_size (`int`): The batch size for repeating the attention mask. - out_dim (`int`, *optional*, defaults to `3`): Output dimension. - - Returns: - `torch.Tensor`: The prepared attention mask. - """ - head_size = self.heads - if attention_mask is None: - return attention_mask - - current_length: int = attention_mask.shape[-1] - if current_length != target_length: - if attention_mask.device.type == "mps": - # HACK: MPS: Does not support padding by greater than dimension of input tensor. - # Instead, we can manually construct the padding tensor. - padding_shape = (attention_mask.shape[0], attention_mask.shape[1], target_length) - padding = torch.zeros(padding_shape, dtype=attention_mask.dtype, device=attention_mask.device) - attention_mask = torch.cat([attention_mask, padding], dim=2) - else: - # TODO: for pipelines such as stable-diffusion, padding cross-attn mask: - # we want to instead pad by (0, remaining_length), where remaining_length is: - # remaining_length: int = target_length - current_length - # TODO: re-enable tests/models/test_models_unet_2d_condition.py#test_model_xattn_padding - attention_mask = F.pad(attention_mask, (0, target_length), value=0.0) - - if out_dim == 3: - if attention_mask.shape[0] < batch_size * head_size: - attention_mask = attention_mask.repeat_interleave(head_size, dim=0) - elif out_dim == 4: - attention_mask = attention_mask.unsqueeze(1) - attention_mask = attention_mask.repeat_interleave(head_size, dim=1) - - return attention_mask - - def norm_encoder_hidden_states(self, encoder_hidden_states: torch.Tensor) -> torch.Tensor: - """ - Normalize the encoder hidden states. - - Args: - encoder_hidden_states (`torch.Tensor`): Hidden states of the encoder. - - Returns: - `torch.Tensor`: The normalized encoder hidden states. - """ - assert self.norm_cross is not None, "self.norm_cross must be defined to call self.norm_encoder_hidden_states" - if isinstance(self.norm_cross, nn.LayerNorm): - encoder_hidden_states = self.norm_cross(encoder_hidden_states) - elif isinstance(self.norm_cross, nn.GroupNorm): - # Group norm norms along the channels dimension and expects - # input to be in the shape of (N, C, *). In this case, we want - # to norm along the hidden dimension, so we need to move - # (batch_size, sequence_length, hidden_size) -> - # (batch_size, hidden_size, sequence_length) - encoder_hidden_states = encoder_hidden_states.transpose(1, 2) - encoder_hidden_states = self.norm_cross(encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states.transpose(1, 2) - else: - assert False - - return encoder_hidden_states - - -def _chunked_feed_forward(ff: nn.Module, hidden_states: torch.Tensor, chunk_dim: int, chunk_size: int): - # "feed_forward_chunk_size" can be used to save memory - if hidden_states.shape[chunk_dim] % chunk_size != 0: - raise ValueError( - f"`hidden_states` dimension to be chunked: {hidden_states.shape[chunk_dim]} has to be divisible by chunk size: {chunk_size}. Make sure to set an appropriate `chunk_size` when calling `unet.enable_forward_chunking`." - ) - - num_chunks = hidden_states.shape[chunk_dim] // chunk_size - ff_output = torch.cat( - [ff(hid_slice) for hid_slice in hidden_states.chunk(num_chunks, dim=chunk_dim)], - dim=chunk_dim, - ) - return ff_output - - -@maybe_allow_in_graph -class GatedSelfAttentionDense(nn.Module): - r""" - A gated self-attention dense layer that combines visual features and object features. - - Parameters: - query_dim (`int`): The number of channels in the query. - context_dim (`int`): The number of channels in the context. - n_heads (`int`): The number of heads to use for attention. - d_head (`int`): The number of channels in each head. - """ - - def __init__(self, query_dim: int, context_dim: int, n_heads: int, d_head: int): - super().__init__() - - # we need a linear projection since we need cat visual feature and obj feature - self.linear = nn.Linear(context_dim, query_dim) - - self.attn = Attention(query_dim=query_dim, heads=n_heads, dim_head=d_head) - self.ff = FeedForward(query_dim, activation_fn="geglu") - - self.norm1 = nn.LayerNorm(query_dim) - self.norm2 = nn.LayerNorm(query_dim) - - self.register_parameter("alpha_attn", nn.Parameter(torch.tensor(0.0))) - self.register_parameter("alpha_dense", nn.Parameter(torch.tensor(0.0))) - - self.enabled = True - - def forward(self, x: torch.Tensor, objs: torch.Tensor) -> torch.Tensor: - if not self.enabled: - return x - - n_visual = x.shape[1] - objs = self.linear(objs) - - x = x + self.alpha_attn.tanh() * self.attn(self.norm1(torch.cat([x, objs], dim=1)))[:, :n_visual, :] - x = x + self.alpha_dense.tanh() * self.ff(self.norm2(x)) - - return x - - -@maybe_allow_in_graph -class JointTransformerBlock(nn.Module): - r""" - A Transformer block following the MMDiT architecture, introduced in Stable Diffusion 3. - - Reference: https://huggingface.co/papers/2403.03206 - - Parameters: - dim (`int`): The number of channels in the input and output. - num_attention_heads (`int`): The number of heads to use for multi-head attention. - attention_head_dim (`int`): The number of channels in each head. - context_pre_only (`bool`): Boolean to determine if we should add some blocks associated with the - processing of `context` conditions. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - context_pre_only: bool = False, - qk_norm: str | None = None, - use_dual_attention: bool = False, - ): - super().__init__() - - self.use_dual_attention = use_dual_attention - self.context_pre_only = context_pre_only - context_norm_type = "ada_norm_continous" if context_pre_only else "ada_norm_zero" - - if use_dual_attention: - self.norm1 = SD35AdaLayerNormZeroX(dim) - else: - self.norm1 = AdaLayerNormZero(dim) - - if context_norm_type == "ada_norm_continous": - self.norm1_context = AdaLayerNormContinuous( - dim, dim, elementwise_affine=False, eps=1e-6, bias=True, norm_type="layer_norm" - ) - elif context_norm_type == "ada_norm_zero": - self.norm1_context = AdaLayerNormZero(dim) - else: - raise ValueError( - f"Unknown context_norm_type: {context_norm_type}, currently only support `ada_norm_continous`, `ada_norm_zero`" - ) - - if hasattr(F, "scaled_dot_product_attention"): - processor = JointAttnProcessor2_0() - else: - raise ValueError( - "The current PyTorch version does not support the `scaled_dot_product_attention` function." - ) - - self.attn = Attention( - query_dim=dim, - cross_attention_dim=None, - added_kv_proj_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - context_pre_only=context_pre_only, - bias=True, - processor=processor, - qk_norm=qk_norm, - eps=1e-6, - ) - - if use_dual_attention: - self.attn2 = Attention( - query_dim=dim, - cross_attention_dim=None, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - bias=True, - processor=processor, - qk_norm=qk_norm, - eps=1e-6, - ) - else: - self.attn2 = None - - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - if not context_pre_only: - self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff_context = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - else: - self.norm2_context = None - self.ff_context = None - - # let chunk size default to None - self._chunk_size = None - self._chunk_dim = 0 - - # Copied from diffusers.models.attention.BasicTransformerBlock.set_chunk_feed_forward - def set_chunk_feed_forward(self, chunk_size: int | None, dim: int = 0): - # Sets chunk feed-forward - self._chunk_size = chunk_size - self._chunk_dim = dim - - def forward( - self, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor, - temb: torch.FloatTensor, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - joint_attention_kwargs = joint_attention_kwargs or {} - if self.use_dual_attention: - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp, norm_hidden_states2, gate_msa2 = self.norm1( - hidden_states, emb=temb - ) - else: - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - - if self.context_pre_only: - norm_encoder_hidden_states = self.norm1_context(encoder_hidden_states, temb) - else: - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - - # Attention. - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - **joint_attention_kwargs, - ) - - # Process attention outputs for the `hidden_states`. - attn_output = gate_msa.unsqueeze(1) * attn_output - hidden_states = hidden_states + attn_output - - if self.use_dual_attention: - attn_output2 = self.attn2(hidden_states=norm_hidden_states2, **joint_attention_kwargs) - attn_output2 = gate_msa2.unsqueeze(1) * attn_output2 - hidden_states = hidden_states + attn_output2 - - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - if self._chunk_size is not None: - # "feed_forward_chunk_size" can be used to save memory - ff_output = _chunked_feed_forward(self.ff, norm_hidden_states, self._chunk_dim, self._chunk_size) - else: - ff_output = self.ff(norm_hidden_states) - ff_output = gate_mlp.unsqueeze(1) * ff_output - - hidden_states = hidden_states + ff_output - - # Process attention outputs for the `encoder_hidden_states`. - if self.context_pre_only: - encoder_hidden_states = None - else: - context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output - encoder_hidden_states = encoder_hidden_states + context_attn_output - - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - if self._chunk_size is not None: - # "feed_forward_chunk_size" can be used to save memory - context_ff_output = _chunked_feed_forward( - self.ff_context, norm_encoder_hidden_states, self._chunk_dim, self._chunk_size - ) - else: - context_ff_output = self.ff_context(norm_encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - - return encoder_hidden_states, hidden_states - - -@maybe_allow_in_graph -class BasicTransformerBlock(nn.Module): - r""" - A basic Transformer block. - - Parameters: - dim (`int`): The number of channels in the input and output. - num_attention_heads (`int`): The number of heads to use for multi-head attention. - attention_head_dim (`int`): The number of channels in each head. - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention. - activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward. - num_embeds_ada_norm (: - obj: `int`, *optional*): The number of diffusion steps used during training. See `Transformer2DModel`. - attention_bias (: - obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter. - only_cross_attention (`bool`, *optional*): - Whether to use only cross-attention layers. In this case two cross attention layers are used. - double_self_attention (`bool`, *optional*): - Whether to use two self-attention layers. In this case no cross attention layers are used. - upcast_attention (`bool`, *optional*): - Whether to upcast the attention computation to float32. This is useful for mixed precision training. - norm_elementwise_affine (`bool`, *optional*, defaults to `True`): - Whether to use learnable elementwise affine parameters for normalization. - norm_type (`str`, *optional*, defaults to `"layer_norm"`): - The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`. - final_dropout (`bool` *optional*, defaults to False): - Whether to apply a final dropout after the last feed-forward layer. - attention_type (`str`, *optional*, defaults to `"default"`): - The type of attention to use. Can be `"default"` or `"gated"` or `"gated-text-image"`. - positional_embeddings (`str`, *optional*, defaults to `None`): - The type of positional embeddings to apply to. - num_positional_embeddings (`int`, *optional*, defaults to `None`): - The maximum number of positional embeddings to apply. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - dropout=0.0, - cross_attention_dim: int | None = None, - activation_fn: str = "geglu", - num_embeds_ada_norm: int | None = None, - attention_bias: bool = False, - only_cross_attention: bool = False, - double_self_attention: bool = False, - upcast_attention: bool = False, - norm_elementwise_affine: bool = True, - norm_type: str = "layer_norm", # 'layer_norm', 'ada_norm', 'ada_norm_zero', 'ada_norm_single', 'ada_norm_continuous', 'layer_norm_i2vgen' - norm_eps: float = 1e-5, - final_dropout: bool = False, - attention_type: str = "default", - positional_embeddings: str | None = None, - num_positional_embeddings: int | None = None, - ada_norm_continous_conditioning_embedding_dim: int | None = None, - ada_norm_bias: int | None = None, - ff_inner_dim: int | None = None, - ff_bias: bool = True, - attention_out_bias: bool = True, - ): - super().__init__() - self.dim = dim - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - self.dropout = dropout - self.cross_attention_dim = cross_attention_dim - self.activation_fn = activation_fn - self.attention_bias = attention_bias - self.double_self_attention = double_self_attention - self.norm_elementwise_affine = norm_elementwise_affine - self.positional_embeddings = positional_embeddings - self.num_positional_embeddings = num_positional_embeddings - self.only_cross_attention = only_cross_attention - - # We keep these boolean flags for backward-compatibility. - self.use_ada_layer_norm_zero = (num_embeds_ada_norm is not None) and norm_type == "ada_norm_zero" - self.use_ada_layer_norm = (num_embeds_ada_norm is not None) and norm_type == "ada_norm" - self.use_ada_layer_norm_single = norm_type == "ada_norm_single" - self.use_layer_norm = norm_type == "layer_norm" - self.use_ada_layer_norm_continuous = norm_type == "ada_norm_continuous" - - if norm_type in ("ada_norm", "ada_norm_zero") and num_embeds_ada_norm is None: - raise ValueError( - f"`norm_type` is set to {norm_type}, but `num_embeds_ada_norm` is not defined. Please make sure to" - f" define `num_embeds_ada_norm` if setting `norm_type` to {norm_type}." - ) - - self.norm_type = norm_type - self.num_embeds_ada_norm = num_embeds_ada_norm - - if positional_embeddings and (num_positional_embeddings is None): - raise ValueError( - "If `positional_embedding` type is defined, `num_positition_embeddings` must also be defined." - ) - - if positional_embeddings == "sinusoidal": - self.pos_embed = SinusoidalPositionalEmbedding(dim, max_seq_length=num_positional_embeddings) - else: - self.pos_embed = None - - # Define 3 blocks. Each block has its own normalization layer. - # 1. Self-Attn - if norm_type == "ada_norm": - self.norm1 = AdaLayerNorm(dim, num_embeds_ada_norm) - elif norm_type == "ada_norm_zero": - self.norm1 = AdaLayerNormZero(dim, num_embeds_ada_norm) - elif norm_type == "ada_norm_continuous": - self.norm1 = AdaLayerNormContinuous( - dim, - ada_norm_continous_conditioning_embedding_dim, - norm_elementwise_affine, - norm_eps, - ada_norm_bias, - "rms_norm", - ) - else: - self.norm1 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps) - - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - cross_attention_dim=cross_attention_dim if only_cross_attention else None, - upcast_attention=upcast_attention, - out_bias=attention_out_bias, - ) - - # 2. Cross-Attn - if cross_attention_dim is not None or double_self_attention: - # We currently only use AdaLayerNormZero for self attention where there will only be one attention block. - # I.e. the number of returned modulation chunks from AdaLayerZero would not make sense if returned during - # the second cross attention block. - if norm_type == "ada_norm": - self.norm2 = AdaLayerNorm(dim, num_embeds_ada_norm) - elif norm_type == "ada_norm_continuous": - self.norm2 = AdaLayerNormContinuous( - dim, - ada_norm_continous_conditioning_embedding_dim, - norm_elementwise_affine, - norm_eps, - ada_norm_bias, - "rms_norm", - ) - else: - self.norm2 = nn.LayerNorm(dim, norm_eps, norm_elementwise_affine) - - self.attn2 = Attention( - query_dim=dim, - cross_attention_dim=cross_attention_dim if not double_self_attention else None, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - upcast_attention=upcast_attention, - out_bias=attention_out_bias, - ) # is self-attn if encoder_hidden_states is none - else: - if norm_type == "ada_norm_single": # For Latte - self.norm2 = nn.LayerNorm(dim, norm_eps, norm_elementwise_affine) - else: - self.norm2 = None - self.attn2 = None - - # 3. Feed-forward - if norm_type == "ada_norm_continuous": - self.norm3 = AdaLayerNormContinuous( - dim, - ada_norm_continous_conditioning_embedding_dim, - norm_elementwise_affine, - norm_eps, - ada_norm_bias, - "layer_norm", - ) - - elif norm_type in ["ada_norm_zero", "ada_norm", "layer_norm"]: - self.norm3 = nn.LayerNorm(dim, norm_eps, norm_elementwise_affine) - elif norm_type == "layer_norm_i2vgen": - self.norm3 = None - - self.ff = FeedForward( - dim, - dropout=dropout, - activation_fn=activation_fn, - final_dropout=final_dropout, - inner_dim=ff_inner_dim, - bias=ff_bias, - ) - - # 4. Fuser - if attention_type == "gated" or attention_type == "gated-text-image": - self.fuser = GatedSelfAttentionDense(dim, cross_attention_dim, num_attention_heads, attention_head_dim) - - # 5. Scale-shift for PixArt-Alpha. - if norm_type == "ada_norm_single": - self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5) - - # let chunk size default to None - self._chunk_size = None - self._chunk_dim = 0 - - def set_chunk_feed_forward(self, chunk_size: int | None, dim: int = 0): - # Sets chunk feed-forward - self._chunk_size = chunk_size - self._chunk_dim = dim - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - timestep: torch.LongTensor | None = None, - cross_attention_kwargs: dict[str, Any] = None, - class_labels: torch.LongTensor | None = None, - added_cond_kwargs: dict[str, torch.Tensor] | None = None, - ) -> torch.Tensor: - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - # Notice that normalization is always applied before the real computation in the following blocks. - # 0. Self-Attention - batch_size = hidden_states.shape[0] - - if self.norm_type == "ada_norm": - norm_hidden_states = self.norm1(hidden_states, timestep) - elif self.norm_type == "ada_norm_zero": - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1( - hidden_states, timestep, class_labels, hidden_dtype=hidden_states.dtype - ) - elif self.norm_type in ["layer_norm", "layer_norm_i2vgen"]: - norm_hidden_states = self.norm1(hidden_states) - elif self.norm_type == "ada_norm_continuous": - norm_hidden_states = self.norm1(hidden_states, added_cond_kwargs["pooled_text_emb"]) - elif self.norm_type == "ada_norm_single": - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( - self.scale_shift_table[None] + timestep.reshape(batch_size, 6, -1) - ).chunk(6, dim=1) - norm_hidden_states = self.norm1(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_msa) + shift_msa - else: - raise ValueError("Incorrect norm used") - - if self.pos_embed is not None: - norm_hidden_states = self.pos_embed(norm_hidden_states) - - # 1. Prepare GLIGEN inputs - cross_attention_kwargs = cross_attention_kwargs.copy() if cross_attention_kwargs is not None else {} - gligen_kwargs = cross_attention_kwargs.pop("gligen", None) - - attn_output = self.attn1( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states if self.only_cross_attention else None, - attention_mask=attention_mask, - **cross_attention_kwargs, - ) - - if self.norm_type == "ada_norm_zero": - attn_output = gate_msa.unsqueeze(1) * attn_output - elif self.norm_type == "ada_norm_single": - attn_output = gate_msa * attn_output - - hidden_states = attn_output + hidden_states - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - - # 1.2 GLIGEN Control - if gligen_kwargs is not None: - hidden_states = self.fuser(hidden_states, gligen_kwargs["objs"]) - - # 3. Cross-Attention - if self.attn2 is not None: - if self.norm_type == "ada_norm": - norm_hidden_states = self.norm2(hidden_states, timestep) - elif self.norm_type in ["ada_norm_zero", "layer_norm", "layer_norm_i2vgen"]: - norm_hidden_states = self.norm2(hidden_states) - elif self.norm_type == "ada_norm_single": - # For PixArt norm2 isn't applied here: - # https://github.com/PixArt-alpha/PixArt-alpha/blob/0f55e922376d8b797edd44d25d0e7464b260dcab/diffusion/model/nets/PixArtMS.py#L70C1-L76C103 - norm_hidden_states = hidden_states - elif self.norm_type == "ada_norm_continuous": - norm_hidden_states = self.norm2(hidden_states, added_cond_kwargs["pooled_text_emb"]) - else: - raise ValueError("Incorrect norm") - - if self.pos_embed is not None and self.norm_type != "ada_norm_single": - norm_hidden_states = self.pos_embed(norm_hidden_states) - - attn_output = self.attn2( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=encoder_attention_mask, - **cross_attention_kwargs, - ) - hidden_states = attn_output + hidden_states - - # 4. Feed-forward - # i2vgen doesn't have this norm 🤷‍♂️ - if self.norm_type == "ada_norm_continuous": - norm_hidden_states = self.norm3(hidden_states, added_cond_kwargs["pooled_text_emb"]) - elif not self.norm_type == "ada_norm_single": - norm_hidden_states = self.norm3(hidden_states) - - if self.norm_type == "ada_norm_zero": - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - - if self.norm_type == "ada_norm_single": - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp - - if self._chunk_size is not None: - # "feed_forward_chunk_size" can be used to save memory - ff_output = _chunked_feed_forward(self.ff, norm_hidden_states, self._chunk_dim, self._chunk_size) - else: - ff_output = self.ff(norm_hidden_states) - - if self.norm_type == "ada_norm_zero": - ff_output = gate_mlp.unsqueeze(1) * ff_output - elif self.norm_type == "ada_norm_single": - ff_output = gate_mlp * ff_output - - hidden_states = ff_output + hidden_states - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - - return hidden_states - - -class LuminaFeedForward(nn.Module): - r""" - A feed-forward layer. - - Parameters: - hidden_size (`int`): - The dimensionality of the hidden layers in the model. This parameter determines the width of the model's - hidden representations. - intermediate_size (`int`): The intermediate dimension of the feedforward layer. - multiple_of (`int`, *optional*): Value to ensure hidden dimension is a multiple - of this value. - ffn_dim_multiplier (float, *optional*): Custom multiplier for hidden - dimension. Defaults to None. - """ - - def __init__( - self, - dim: int, - inner_dim: int, - multiple_of: int | None = 256, - ffn_dim_multiplier: float | None = None, - ): - super().__init__() - # custom hidden_size factor multiplier - if ffn_dim_multiplier is not None: - inner_dim = int(ffn_dim_multiplier * inner_dim) - inner_dim = multiple_of * ((inner_dim + multiple_of - 1) // multiple_of) - - self.linear_1 = nn.Linear( - dim, - inner_dim, - bias=False, - ) - self.linear_2 = nn.Linear( - inner_dim, - dim, - bias=False, - ) - self.linear_3 = nn.Linear( - dim, - inner_dim, - bias=False, - ) - self.silu = FP32SiLU() - - def forward(self, x): - return self.linear_2(self.silu(self.linear_1(x)) * self.linear_3(x)) - - -@maybe_allow_in_graph -class TemporalBasicTransformerBlock(nn.Module): - r""" - A basic Transformer block for video like data. - - Parameters: - dim (`int`): The number of channels in the input and output. - time_mix_inner_dim (`int`): The number of channels for temporal attention. - num_attention_heads (`int`): The number of heads to use for multi-head attention. - attention_head_dim (`int`): The number of channels in each head. - cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention. - """ - - def __init__( - self, - dim: int, - time_mix_inner_dim: int, - num_attention_heads: int, - attention_head_dim: int, - cross_attention_dim: int | None = None, - ): - super().__init__() - self.is_res = dim == time_mix_inner_dim - - self.norm_in = nn.LayerNorm(dim) - - # Define 3 blocks. Each block has its own normalization layer. - # 1. Self-Attn - self.ff_in = FeedForward( - dim, - dim_out=time_mix_inner_dim, - activation_fn="geglu", - ) - - self.norm1 = nn.LayerNorm(time_mix_inner_dim) - self.attn1 = Attention( - query_dim=time_mix_inner_dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - cross_attention_dim=None, - ) - - # 2. Cross-Attn - if cross_attention_dim is not None: - # We currently only use AdaLayerNormZero for self attention where there will only be one attention block. - # I.e. the number of returned modulation chunks from AdaLayerZero would not make sense if returned during - # the second cross attention block. - self.norm2 = nn.LayerNorm(time_mix_inner_dim) - self.attn2 = Attention( - query_dim=time_mix_inner_dim, - cross_attention_dim=cross_attention_dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - ) # is self-attn if encoder_hidden_states is none - else: - self.norm2 = None - self.attn2 = None - - # 3. Feed-forward - self.norm3 = nn.LayerNorm(time_mix_inner_dim) - self.ff = FeedForward(time_mix_inner_dim, activation_fn="geglu") - - # let chunk size default to None - self._chunk_size = None - self._chunk_dim = None - - def set_chunk_feed_forward(self, chunk_size: int | None, **kwargs): - # Sets chunk feed-forward - self._chunk_size = chunk_size - # chunk dim should be hardcoded to 1 to have better speed vs. memory trade-off - self._chunk_dim = 1 - - def forward( - self, - hidden_states: torch.Tensor, - num_frames: int, - encoder_hidden_states: torch.Tensor | None = None, - ) -> torch.Tensor: - # Notice that normalization is always applied before the real computation in the following blocks. - # 0. Self-Attention - batch_size = hidden_states.shape[0] - - batch_frames, seq_length, channels = hidden_states.shape - batch_size = batch_frames // num_frames - - hidden_states = hidden_states[None, :].reshape(batch_size, num_frames, seq_length, channels) - hidden_states = hidden_states.permute(0, 2, 1, 3) - hidden_states = hidden_states.reshape(batch_size * seq_length, num_frames, channels) - - residual = hidden_states - hidden_states = self.norm_in(hidden_states) - - if self._chunk_size is not None: - hidden_states = _chunked_feed_forward(self.ff_in, hidden_states, self._chunk_dim, self._chunk_size) - else: - hidden_states = self.ff_in(hidden_states) - - if self.is_res: - hidden_states = hidden_states + residual - - norm_hidden_states = self.norm1(hidden_states) - attn_output = self.attn1(norm_hidden_states, encoder_hidden_states=None) - hidden_states = attn_output + hidden_states - - # 3. Cross-Attention - if self.attn2 is not None: - norm_hidden_states = self.norm2(hidden_states) - attn_output = self.attn2(norm_hidden_states, encoder_hidden_states=encoder_hidden_states) - hidden_states = attn_output + hidden_states - - # 4. Feed-forward - norm_hidden_states = self.norm3(hidden_states) - - if self._chunk_size is not None: - ff_output = _chunked_feed_forward(self.ff, norm_hidden_states, self._chunk_dim, self._chunk_size) - else: - ff_output = self.ff(norm_hidden_states) - - if self.is_res: - hidden_states = ff_output + hidden_states - else: - hidden_states = ff_output - - hidden_states = hidden_states[None, :].reshape(batch_size, seq_length, num_frames, channels) - hidden_states = hidden_states.permute(0, 2, 1, 3) - hidden_states = hidden_states.reshape(batch_size * num_frames, seq_length, channels) - - return hidden_states - - -class SkipFFTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - kv_input_dim: int, - kv_input_dim_proj_use_bias: bool, - dropout=0.0, - cross_attention_dim: int | None = None, - attention_bias: bool = False, - attention_out_bias: bool = True, - ): - super().__init__() - if kv_input_dim != dim: - self.kv_mapper = nn.Linear(kv_input_dim, dim, kv_input_dim_proj_use_bias) - else: - self.kv_mapper = None - - self.norm1 = RMSNorm(dim, 1e-06) - - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - cross_attention_dim=cross_attention_dim, - out_bias=attention_out_bias, - ) - - self.norm2 = RMSNorm(dim, 1e-06) - - self.attn2 = Attention( - query_dim=dim, - cross_attention_dim=cross_attention_dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - out_bias=attention_out_bias, - ) - - def forward(self, hidden_states, encoder_hidden_states, cross_attention_kwargs): - cross_attention_kwargs = cross_attention_kwargs.copy() if cross_attention_kwargs is not None else {} - - if self.kv_mapper is not None: - encoder_hidden_states = self.kv_mapper(F.silu(encoder_hidden_states)) - - norm_hidden_states = self.norm1(hidden_states) - - attn_output = self.attn1( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - **cross_attention_kwargs, - ) - - hidden_states = attn_output + hidden_states - - norm_hidden_states = self.norm2(hidden_states) - - attn_output = self.attn2( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - **cross_attention_kwargs, - ) - - hidden_states = attn_output + hidden_states - - return hidden_states - - -@maybe_allow_in_graph -class FreeNoiseTransformerBlock(nn.Module): - r""" - A FreeNoise Transformer block. - - Parameters: - dim (`int`): - The number of channels in the input and output. - num_attention_heads (`int`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`): - The number of channels in each head. - dropout (`float`, *optional*, defaults to 0.0): - The dropout probability to use. - cross_attention_dim (`int`, *optional*): - The size of the encoder_hidden_states vector for cross attention. - activation_fn (`str`, *optional*, defaults to `"geglu"`): - Activation function to be used in feed-forward. - num_embeds_ada_norm (`int`, *optional*): - The number of diffusion steps used during training. See `Transformer2DModel`. - attention_bias (`bool`, defaults to `False`): - Configure if the attentions should contain a bias parameter. - only_cross_attention (`bool`, defaults to `False`): - Whether to use only cross-attention layers. In this case two cross attention layers are used. - double_self_attention (`bool`, defaults to `False`): - Whether to use two self-attention layers. In this case no cross attention layers are used. - upcast_attention (`bool`, defaults to `False`): - Whether to upcast the attention computation to float32. This is useful for mixed precision training. - norm_elementwise_affine (`bool`, defaults to `True`): - Whether to use learnable elementwise affine parameters for normalization. - norm_type (`str`, defaults to `"layer_norm"`): - The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`. - final_dropout (`bool` defaults to `False`): - Whether to apply a final dropout after the last feed-forward layer. - attention_type (`str`, defaults to `"default"`): - The type of attention to use. Can be `"default"` or `"gated"` or `"gated-text-image"`. - positional_embeddings (`str`, *optional*): - The type of positional embeddings to apply to. - num_positional_embeddings (`int`, *optional*, defaults to `None`): - The maximum number of positional embeddings to apply. - ff_inner_dim (`int`, *optional*): - Hidden dimension of feed-forward MLP. - ff_bias (`bool`, defaults to `True`): - Whether or not to use bias in feed-forward MLP. - attention_out_bias (`bool`, defaults to `True`): - Whether or not to use bias in attention output project layer. - context_length (`int`, defaults to `16`): - The maximum number of frames that the FreeNoise block processes at once. - context_stride (`int`, defaults to `4`): - The number of frames to be skipped before starting to process a new batch of `context_length` frames. - weighting_scheme (`str`, defaults to `"pyramid"`): - The weighting scheme to use for weighting averaging of processed latent frames. As described in the - Equation 9. of the [FreeNoise](https://huggingface.co/papers/2310.15169) paper, "pyramid" is the default - setting used. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - dropout: float = 0.0, - cross_attention_dim: int | None = None, - activation_fn: str = "geglu", - num_embeds_ada_norm: int | None = None, - attention_bias: bool = False, - only_cross_attention: bool = False, - double_self_attention: bool = False, - upcast_attention: bool = False, - norm_elementwise_affine: bool = True, - norm_type: str = "layer_norm", - norm_eps: float = 1e-5, - final_dropout: bool = False, - positional_embeddings: str | None = None, - num_positional_embeddings: int | None = None, - ff_inner_dim: int | None = None, - ff_bias: bool = True, - attention_out_bias: bool = True, - context_length: int = 16, - context_stride: int = 4, - weighting_scheme: str = "pyramid", - ): - super().__init__() - self.dim = dim - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - self.dropout = dropout - self.cross_attention_dim = cross_attention_dim - self.activation_fn = activation_fn - self.attention_bias = attention_bias - self.double_self_attention = double_self_attention - self.norm_elementwise_affine = norm_elementwise_affine - self.positional_embeddings = positional_embeddings - self.num_positional_embeddings = num_positional_embeddings - self.only_cross_attention = only_cross_attention - - self.set_free_noise_properties(context_length, context_stride, weighting_scheme) - - # We keep these boolean flags for backward-compatibility. - self.use_ada_layer_norm_zero = (num_embeds_ada_norm is not None) and norm_type == "ada_norm_zero" - self.use_ada_layer_norm = (num_embeds_ada_norm is not None) and norm_type == "ada_norm" - self.use_ada_layer_norm_single = norm_type == "ada_norm_single" - self.use_layer_norm = norm_type == "layer_norm" - self.use_ada_layer_norm_continuous = norm_type == "ada_norm_continuous" - - if norm_type in ("ada_norm", "ada_norm_zero") and num_embeds_ada_norm is None: - raise ValueError( - f"`norm_type` is set to {norm_type}, but `num_embeds_ada_norm` is not defined. Please make sure to" - f" define `num_embeds_ada_norm` if setting `norm_type` to {norm_type}." - ) - - self.norm_type = norm_type - self.num_embeds_ada_norm = num_embeds_ada_norm - - if positional_embeddings and (num_positional_embeddings is None): - raise ValueError( - "If `positional_embedding` type is defined, `num_positition_embeddings` must also be defined." - ) - - if positional_embeddings == "sinusoidal": - self.pos_embed = SinusoidalPositionalEmbedding(dim, max_seq_length=num_positional_embeddings) - else: - self.pos_embed = None - - # Define 3 blocks. Each block has its own normalization layer. - # 1. Self-Attn - self.norm1 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps) - - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - cross_attention_dim=cross_attention_dim if only_cross_attention else None, - upcast_attention=upcast_attention, - out_bias=attention_out_bias, - ) - - # 2. Cross-Attn - if cross_attention_dim is not None or double_self_attention: - self.norm2 = nn.LayerNorm(dim, norm_eps, norm_elementwise_affine) - - self.attn2 = Attention( - query_dim=dim, - cross_attention_dim=cross_attention_dim if not double_self_attention else None, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - upcast_attention=upcast_attention, - out_bias=attention_out_bias, - ) # is self-attn if encoder_hidden_states is none - - # 3. Feed-forward - self.ff = FeedForward( - dim, - dropout=dropout, - activation_fn=activation_fn, - final_dropout=final_dropout, - inner_dim=ff_inner_dim, - bias=ff_bias, - ) - - self.norm3 = nn.LayerNorm(dim, norm_eps, norm_elementwise_affine) - - # let chunk size default to None - self._chunk_size = None - self._chunk_dim = 0 - - def _get_frame_indices(self, num_frames: int) -> list[tuple[int, int]]: - frame_indices = [] - for i in range(0, num_frames - self.context_length + 1, self.context_stride): - window_start = i - window_end = min(num_frames, i + self.context_length) - frame_indices.append((window_start, window_end)) - return frame_indices - - def _get_frame_weights(self, num_frames: int, weighting_scheme: str = "pyramid") -> list[float]: - if weighting_scheme == "flat": - weights = [1.0] * num_frames - - elif weighting_scheme == "pyramid": - if num_frames % 2 == 0: - # num_frames = 4 => [1, 2, 2, 1] - mid = num_frames // 2 - weights = list(range(1, mid + 1)) - weights = weights + weights[::-1] - else: - # num_frames = 5 => [1, 2, 3, 2, 1] - mid = (num_frames + 1) // 2 - weights = list(range(1, mid)) - weights = weights + [mid] + weights[::-1] - - elif weighting_scheme == "delayed_reverse_sawtooth": - if num_frames % 2 == 0: - # num_frames = 4 => [0.01, 2, 2, 1] - mid = num_frames // 2 - weights = [0.01] * (mid - 1) + [mid] - weights = weights + list(range(mid, 0, -1)) - else: - # num_frames = 5 => [0.01, 0.01, 3, 2, 1] - mid = (num_frames + 1) // 2 - weights = [0.01] * mid - weights = weights + list(range(mid, 0, -1)) - else: - raise ValueError(f"Unsupported value for weighting_scheme={weighting_scheme}") - - return weights - - def set_free_noise_properties( - self, context_length: int, context_stride: int, weighting_scheme: str = "pyramid" - ) -> None: - self.context_length = context_length - self.context_stride = context_stride - self.weighting_scheme = weighting_scheme - - def set_chunk_feed_forward(self, chunk_size: int | None, dim: int = 0) -> None: - # Sets chunk feed-forward - self._chunk_size = chunk_size - self._chunk_dim = dim - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] = None, - *args, - **kwargs, - ) -> torch.Tensor: - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - cross_attention_kwargs = cross_attention_kwargs.copy() if cross_attention_kwargs is not None else {} - - # hidden_states: [B x H x W, F, C] - device = hidden_states.device - dtype = hidden_states.dtype - - num_frames = hidden_states.size(1) - frame_indices = self._get_frame_indices(num_frames) - frame_weights = self._get_frame_weights(self.context_length, self.weighting_scheme) - frame_weights = torch.tensor(frame_weights, device=device, dtype=dtype).unsqueeze(0).unsqueeze(-1) - is_last_frame_batch_complete = frame_indices[-1][1] == num_frames - - # Handle out-of-bounds case if num_frames isn't perfectly divisible by context_length - # For example, num_frames=25, context_length=16, context_stride=4, then we expect the ranges: - # [(0, 16), (4, 20), (8, 24), (10, 26)] - if not is_last_frame_batch_complete: - if num_frames < self.context_length: - raise ValueError(f"Expected {num_frames=} to be greater or equal than {self.context_length=}") - last_frame_batch_length = num_frames - frame_indices[-1][1] - frame_indices.append((num_frames - self.context_length, num_frames)) - - num_times_accumulated = torch.zeros((1, num_frames, 1), device=device) - accumulated_values = torch.zeros_like(hidden_states) - - for i, (frame_start, frame_end) in enumerate(frame_indices): - # The reason for slicing here is to ensure that if (frame_end - frame_start) is to handle - # cases like frame_indices=[(0, 16), (16, 20)], if the user provided a video with 19 frames, or - # essentially a non-multiple of `context_length`. - weights = torch.ones_like(num_times_accumulated[:, frame_start:frame_end]) - weights *= frame_weights - - hidden_states_chunk = hidden_states[:, frame_start:frame_end] - - # Notice that normalization is always applied before the real computation in the following blocks. - # 1. Self-Attention - norm_hidden_states = self.norm1(hidden_states_chunk) - - if self.pos_embed is not None: - norm_hidden_states = self.pos_embed(norm_hidden_states) - - attn_output = self.attn1( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states if self.only_cross_attention else None, - attention_mask=attention_mask, - **cross_attention_kwargs, - ) - - hidden_states_chunk = attn_output + hidden_states_chunk - if hidden_states_chunk.ndim == 4: - hidden_states_chunk = hidden_states_chunk.squeeze(1) - - # 2. Cross-Attention - if self.attn2 is not None: - norm_hidden_states = self.norm2(hidden_states_chunk) - - if self.pos_embed is not None and self.norm_type != "ada_norm_single": - norm_hidden_states = self.pos_embed(norm_hidden_states) - - attn_output = self.attn2( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=encoder_attention_mask, - **cross_attention_kwargs, - ) - hidden_states_chunk = attn_output + hidden_states_chunk - - if i == len(frame_indices) - 1 and not is_last_frame_batch_complete: - accumulated_values[:, -last_frame_batch_length:] += ( - hidden_states_chunk[:, -last_frame_batch_length:] * weights[:, -last_frame_batch_length:] - ) - num_times_accumulated[:, -last_frame_batch_length:] += weights[:, -last_frame_batch_length] - else: - accumulated_values[:, frame_start:frame_end] += hidden_states_chunk * weights - num_times_accumulated[:, frame_start:frame_end] += weights - - # TODO(aryan): Maybe this could be done in a better way. - # - # Previously, this was: - # hidden_states = torch.where( - # num_times_accumulated > 0, accumulated_values / num_times_accumulated, accumulated_values - # ) - # - # The reasoning for the change here is `torch.where` became a bottleneck at some point when golfing memory - # spikes. It is particularly noticeable when the number of frames is high. My understanding is that this comes - # from tensors being copied - which is why we resort to spliting and concatenating here. I've not particularly - # looked into this deeply because other memory optimizations led to more pronounced reductions. - hidden_states = torch.cat( - [ - torch.where(num_times_split > 0, accumulated_split / num_times_split, accumulated_split) - for accumulated_split, num_times_split in zip( - accumulated_values.split(self.context_length, dim=1), - num_times_accumulated.split(self.context_length, dim=1), - ) - ], - dim=1, - ).to(dtype) - - # 3. Feed-forward - norm_hidden_states = self.norm3(hidden_states) - - if self._chunk_size is not None: - ff_output = _chunked_feed_forward(self.ff, norm_hidden_states, self._chunk_dim, self._chunk_size) - else: - ff_output = self.ff(norm_hidden_states) - - hidden_states = ff_output + hidden_states - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - - return hidden_states - - -class FeedForward(nn.Module): - r""" - A feed-forward layer. - - Parameters: - dim (`int`): The number of channels in the input. - dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`. - mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension. - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward. - final_dropout (`bool` *optional*, defaults to False): Apply a final dropout. - bias (`bool`, defaults to True): Whether to use a bias in the linear layer. - """ - - def __init__( - self, - dim: int, - dim_out: int | None = None, - mult: int = 4, - dropout: float = 0.0, - activation_fn: str = "geglu", - final_dropout: bool = False, - inner_dim=None, - bias: bool = True, - ): - super().__init__() - if inner_dim is None: - inner_dim = int(dim * mult) - dim_out = dim_out if dim_out is not None else dim - - if activation_fn == "gelu": - act_fn = GELU(dim, inner_dim, bias=bias) - if activation_fn == "gelu-approximate": - act_fn = GELU(dim, inner_dim, approximate="tanh", bias=bias) - elif activation_fn == "geglu": - act_fn = GEGLU(dim, inner_dim, bias=bias) - elif activation_fn == "geglu-approximate": - act_fn = ApproximateGELU(dim, inner_dim, bias=bias) - elif activation_fn == "swiglu": - act_fn = SwiGLU(dim, inner_dim, bias=bias) - elif activation_fn == "linear-silu": - act_fn = LinearActivation(dim, inner_dim, bias=bias, activation="silu") - - self.net = nn.ModuleList([]) - # project in - self.net.append(act_fn) - # project dropout - self.net.append(nn.Dropout(dropout)) - # project out - self.net.append(nn.Linear(inner_dim, dim_out, bias=bias)) - # FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout - if final_dropout: - self.net.append(nn.Dropout(dropout)) - - def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - for module in self.net: - hidden_states = module(hidden_states) - return hidden_states diff --git a/diffusers/models/attention_dispatch.py b/diffusers/models/attention_dispatch.py deleted file mode 100644 index 9414c151fd670c22fe1a144f2a6972ffea41970d..0000000000000000000000000000000000000000 --- a/diffusers/models/attention_dispatch.py +++ /dev/null @@ -1,4176 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import contextlib -import functools -import inspect -import math -from dataclasses import dataclass -from enum import Enum -from typing import TYPE_CHECKING, Any, Callable - -import torch -import torch.distributed as dist -import torch.nn.functional as F - - -if torch.distributed.is_available(): - import torch.distributed._functional_collectives as funcol - -from ..utils import ( - get_logger, - is_aiter_available, - is_aiter_version, - is_flash_attn_3_available, - is_flash_attn_available, - is_flash_attn_version, - is_kernels_available, - is_kernels_version, - is_sageattention_available, - is_sageattention_version, - is_torch_npu_available, - is_torch_version, - is_torch_xla_available, - is_torch_xla_version, - is_xformers_available, - is_xformers_version, -) -from ..utils.constants import DIFFUSERS_ATTN_BACKEND, DIFFUSERS_ATTN_CHECKS -from ..utils.torch_utils import lru_cache_unless_export, maybe_allow_in_graph -from ._modeling_parallel import gather_size_by_comm - - -if TYPE_CHECKING: - from ._modeling_parallel import ParallelConfig - -_REQUIRED_FLASH_VERSION = "2.6.3" -_REQUIRED_AITER_VERSION = "0.1.5" -_REQUIRED_SAGE_VERSION = "2.1.1" -_REQUIRED_FLEX_VERSION = "2.5.0" -_REQUIRED_XLA_VERSION = "2.2" -_REQUIRED_XFORMERS_VERSION = "0.0.29" - -logger = get_logger(__name__) # pylint: disable=invalid-name - -_CAN_USE_FLASH_ATTN = is_flash_attn_available() and is_flash_attn_version(">=", _REQUIRED_FLASH_VERSION) -_CAN_USE_FLASH_ATTN_3 = is_flash_attn_3_available() -_CAN_USE_AITER_ATTN = is_aiter_available() and is_aiter_version(">=", _REQUIRED_AITER_VERSION) -_CAN_USE_SAGE_ATTN = is_sageattention_available() and is_sageattention_version(">=", _REQUIRED_SAGE_VERSION) -_CAN_USE_FLEX_ATTN = is_torch_version(">=", _REQUIRED_FLEX_VERSION) -_CAN_USE_NPU_ATTN = is_torch_npu_available() -_CAN_USE_XLA_ATTN = is_torch_xla_available() and is_torch_xla_version(">=", _REQUIRED_XLA_VERSION) -_CAN_USE_XFORMERS_ATTN = is_xformers_available() and is_xformers_version(">=", _REQUIRED_XFORMERS_VERSION) - - -if _CAN_USE_FLASH_ATTN: - try: - from flash_attn import flash_attn_func, flash_attn_varlen_func - from flash_attn.flash_attn_interface import _wrapped_flash_attn_backward, _wrapped_flash_attn_forward - except (ImportError, OSError, RuntimeError) as e: - # Handle ABI mismatch or other import failures gracefully. - # This can happen when flash_attn was compiled against a different PyTorch version. - logger.warning(f"flash_attn is installed but failed to import: {e}. Falling back to native PyTorch attention.") - _CAN_USE_FLASH_ATTN = False - flash_attn_func = None - flash_attn_varlen_func = None - _wrapped_flash_attn_backward = None - _wrapped_flash_attn_forward = None -else: - flash_attn_func = None - flash_attn_varlen_func = None - _wrapped_flash_attn_backward = None - _wrapped_flash_attn_forward = None - - -if _CAN_USE_FLASH_ATTN_3: - try: - from flash_attn_interface import flash_attn_func as flash_attn_3_func - from flash_attn_interface import flash_attn_varlen_func as flash_attn_3_varlen_func - except (ImportError, OSError, RuntimeError) as e: - logger.warning(f"flash_attn_3 failed to import: {e}. Falling back to native attention.") - _CAN_USE_FLASH_ATTN_3 = False - flash_attn_3_func = None - flash_attn_3_varlen_func = None -else: - flash_attn_3_func = None - flash_attn_3_varlen_func = None - -if _CAN_USE_AITER_ATTN: - try: - from aiter import flash_attn_func as aiter_flash_attn_func - except (ImportError, OSError, RuntimeError) as e: - logger.warning(f"aiter failed to import: {e}. Falling back to native attention.") - _CAN_USE_AITER_ATTN = False - aiter_flash_attn_func = None -else: - aiter_flash_attn_func = None - -if _CAN_USE_SAGE_ATTN: - try: - from sageattention import ( - sageattn, - sageattn_qk_int8_pv_fp8_cuda, - sageattn_qk_int8_pv_fp8_cuda_sm90, - sageattn_qk_int8_pv_fp16_cuda, - sageattn_qk_int8_pv_fp16_triton, - sageattn_varlen, - ) - except (ImportError, OSError, RuntimeError) as e: - logger.warning(f"sageattention failed to import: {e}. Falling back to native attention.") - _CAN_USE_SAGE_ATTN = False - sageattn = None - sageattn_qk_int8_pv_fp8_cuda = None - sageattn_qk_int8_pv_fp8_cuda_sm90 = None - sageattn_qk_int8_pv_fp16_cuda = None - sageattn_qk_int8_pv_fp16_triton = None - sageattn_varlen = None -else: - sageattn = None - sageattn_qk_int8_pv_fp16_cuda = None - sageattn_qk_int8_pv_fp16_triton = None - sageattn_qk_int8_pv_fp8_cuda = None - sageattn_qk_int8_pv_fp8_cuda_sm90 = None - sageattn_varlen = None - - -if _CAN_USE_FLEX_ATTN: - try: - # We cannot import the flex_attention function from the package directly because it is expected (from the - # pytorch documentation) that the user may compile it. If we import directly, we will not have access to the - # compiled function. - import torch.nn.attention.flex_attention as flex_attention - except (ImportError, OSError, RuntimeError) as e: - logger.warning(f"flex_attention failed to import: {e}. Falling back to native attention.") - _CAN_USE_FLEX_ATTN = False - flex_attention = None -else: - flex_attention = None - - -if _CAN_USE_NPU_ATTN: - try: - from torch_npu import npu_fusion_attention - except (ImportError, OSError, RuntimeError) as e: - logger.warning(f"torch_npu failed to import: {e}. Falling back to native attention.") - _CAN_USE_NPU_ATTN = False - npu_fusion_attention = None -else: - npu_fusion_attention = None - - -if _CAN_USE_XLA_ATTN: - try: - from torch_xla.experimental.custom_kernel import flash_attention as xla_flash_attention - except (ImportError, OSError, RuntimeError) as e: - logger.warning(f"torch_xla failed to import: {e}. Falling back to native attention.") - _CAN_USE_XLA_ATTN = False - xla_flash_attention = None -else: - xla_flash_attention = None - - -if _CAN_USE_XFORMERS_ATTN: - try: - import xformers.ops as xops - except (ImportError, OSError, RuntimeError) as e: - logger.warning(f"xformers failed to import: {e}. Falling back to native attention.") - _CAN_USE_XFORMERS_ATTN = False - xops = None -else: - xops = None - -# Version guard for PyTorch compatibility - custom_op was added in PyTorch 2.4 -if torch.__version__ >= "2.4.0": - _custom_op = torch.library.custom_op - _register_fake = torch.library.register_fake -else: - - def custom_op_no_op(name, fn=None, /, *, mutates_args, device_types=None, schema=None): - def wrap(func): - return func - - return wrap if fn is None else fn - - def register_fake_no_op(op, fn=None, /, *, lib=None, _stacklevel=1): - def wrap(func): - return func - - return wrap if fn is None else fn - - _custom_op = custom_op_no_op - _register_fake = register_fake_no_op - - -# TODO(aryan): Add support for the following: -# - Sage Attention++ -# - block sparse, radial and other attention methods -# - CP with sage attention, flex, xformers, other missing backends -# - Add support for normal and CP training with backends that don't support it yet - - -class AttentionBackendName(str, Enum): - # EAGER = "eager" - - # `flash-attn` - FLASH = "flash" - FLASH_HUB = "flash_hub" - FLASH_VARLEN = "flash_varlen" - FLASH_VARLEN_HUB = "flash_varlen_hub" - FLASH_4_HUB = "flash_4_hub" - _FLASH_3 = "_flash_3" - _FLASH_VARLEN_3 = "_flash_varlen_3" - _FLASH_3_HUB = "_flash_3_hub" - _FLASH_3_VARLEN_HUB = "_flash_3_varlen_hub" - - # `aiter` - AITER = "aiter" - - # PyTorch native - FLEX = "flex" - NATIVE = "native" - _NATIVE_CUDNN = "_native_cudnn" - _NATIVE_EFFICIENT = "_native_efficient" - _NATIVE_FLASH = "_native_flash" - _NATIVE_MATH = "_native_math" - _NATIVE_NPU = "_native_npu" - _NATIVE_XLA = "_native_xla" - - # `sageattention` - SAGE = "sage" - SAGE_HUB = "sage_hub" - SAGE_VARLEN = "sage_varlen" - _SAGE_QK_INT8_PV_FP8_CUDA = "_sage_qk_int8_pv_fp8_cuda" - _SAGE_QK_INT8_PV_FP8_CUDA_SM90 = "_sage_qk_int8_pv_fp8_cuda_sm90" - _SAGE_QK_INT8_PV_FP16_CUDA = "_sage_qk_int8_pv_fp16_cuda" - _SAGE_QK_INT8_PV_FP16_TRITON = "_sage_qk_int8_pv_fp16_triton" - # TODO: let's not add support for Sparge Attention now because it requires tuning per model - # We can look into supporting something "autotune"-ing in the future - # SPARGE = "sparge" - - # `xformers` - XFORMERS = "xformers" - - -class _AttentionBackendRegistry: - _backends = {} - _constraints = {} - _supported_arg_names = {} - _supports_context_parallel = set() - _active_backend = AttentionBackendName(DIFFUSERS_ATTN_BACKEND) - _checks_enabled = DIFFUSERS_ATTN_CHECKS - - @classmethod - def register( - cls, - backend: AttentionBackendName, - constraints: list[Callable] | None = None, - supports_context_parallel: bool = False, - ): - logger.debug(f"Registering attention backend: {backend} with constraints: {constraints}") - - def decorator(func): - cls._backends[backend] = func - cls._constraints[backend] = constraints or [] - cls._supported_arg_names[backend] = set(inspect.signature(func).parameters.keys()) - if supports_context_parallel: - cls._supports_context_parallel.add(backend.value) - - return func - - return decorator - - @classmethod - def get_active_backend(cls): - return cls._active_backend, cls._backends[cls._active_backend] - - @classmethod - def set_active_backend(cls, backend: str): - cls._active_backend = backend - - @classmethod - def list_backends(cls): - return list(cls._backends.keys()) - - @classmethod - def _is_context_parallel_available( - cls, - backend: AttentionBackendName, - ) -> bool: - supports_context_parallel = backend.value in cls._supports_context_parallel - return supports_context_parallel - - -@dataclass -class _HubKernelConfig: - """Configuration for downloading and using a hub-based attention kernel.""" - - repo_id: str - function_attr: str - revision: str | None = None - version: int | None = None - kernel_fn: Callable | None = None - wrapped_forward_attr: str | None = None - wrapped_backward_attr: str | None = None - wrapped_forward_fn: Callable | None = None - wrapped_backward_fn: Callable | None = None - - -# Registry for hub-based attention kernels -_HUB_KERNELS_REGISTRY: dict["AttentionBackendName", _HubKernelConfig] = { - AttentionBackendName._FLASH_3_HUB: _HubKernelConfig( - repo_id="kernels-community/flash-attn3", - function_attr="flash_attn_func", - wrapped_forward_attr="flash_attn_interface._flash_attn_forward", - wrapped_backward_attr="flash_attn_interface._flash_attn_backward", - version=1, - ), - AttentionBackendName._FLASH_3_VARLEN_HUB: _HubKernelConfig( - repo_id="kernels-community/flash-attn3", - function_attr="flash_attn_varlen_func", - wrapped_forward_attr="flash_attn_interface._flash_attn_forward", - wrapped_backward_attr="flash_attn_interface._flash_attn_backward", - version=1, - ), - AttentionBackendName.FLASH_HUB: _HubKernelConfig( - repo_id="kernels-community/flash-attn2", - function_attr="flash_attn_func", - wrapped_forward_attr="flash_attn_interface._wrapped_flash_attn_forward", - wrapped_backward_attr="flash_attn_interface._wrapped_flash_attn_backward", - version=1, - ), - AttentionBackendName.FLASH_VARLEN_HUB: _HubKernelConfig( - repo_id="kernels-community/flash-attn2", - function_attr="flash_attn_varlen_func", - wrapped_forward_attr="flash_attn_interface._wrapped_flash_attn_varlen_forward", - wrapped_backward_attr="flash_attn_interface._wrapped_flash_attn_varlen_backward", - version=1, - ), - AttentionBackendName.SAGE_HUB: _HubKernelConfig( - repo_id="kernels-community/sage-attention", - function_attr="sageattn", - version=1, - ), - AttentionBackendName.FLASH_4_HUB: _HubKernelConfig( - repo_id="kernels-community/flash-attn4", - function_attr="flash_attn_func", - version=0, - ), -} - - -@contextlib.contextmanager -def attention_backend(backend: str | AttentionBackendName = AttentionBackendName.NATIVE): - """ - Context manager to set the active attention backend. - """ - if backend not in _AttentionBackendRegistry._backends: - raise ValueError(f"Backend {backend} is not registered.") - - backend = AttentionBackendName(backend) - _check_attention_backend_requirements(backend) - _maybe_download_kernel_for_backend(backend) - - old_backend = _AttentionBackendRegistry._active_backend - _AttentionBackendRegistry.set_active_backend(backend) - - try: - yield - finally: - _AttentionBackendRegistry.set_active_backend(old_backend) - - -def dispatch_attention_fn( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - attention_kwargs: dict[str, Any] | None = None, - *, - backend: AttentionBackendName | None = None, - parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - attention_kwargs = attention_kwargs or {} - - if backend is None: - # If no backend is specified, we either use the default backend (set via the DIFFUSERS_ATTN_BACKEND environment - # variable), or we use a custom backend based on whether user is using the `attention_backend` context manager - backend_name, backend_fn = _AttentionBackendRegistry.get_active_backend() - else: - backend_name = AttentionBackendName(backend) - backend_fn = _AttentionBackendRegistry._backends.get(backend_name) - - kwargs = { - "query": query, - "key": key, - "value": value, - "attn_mask": attn_mask, - "dropout_p": dropout_p, - "is_causal": is_causal, - "scale": scale, - **attention_kwargs, - "_parallel_config": parallel_config, - } - # Equivalent to `is_torch_version(">=", "2.5.0")` — use module-level constant to avoid - # Dynamo tracing into the lru_cache-wrapped `is_torch_version` during torch.compile. - if _CAN_USE_FLEX_ATTN: - kwargs["enable_gqa"] = enable_gqa - - if _AttentionBackendRegistry._checks_enabled: - removed_kwargs = set(kwargs) - set(_AttentionBackendRegistry._supported_arg_names[backend_name]) - if removed_kwargs: - logger.warning(f"Removing unsupported arguments for attention backend {backend_name}: {removed_kwargs}.") - for check in _AttentionBackendRegistry._constraints.get(backend_name): - check(**kwargs) - - kwargs = {k: v for k, v in kwargs.items() if k in _AttentionBackendRegistry._supported_arg_names[backend_name]} - - return backend_fn(**kwargs) - - -# ===== Checks ===== -# A list of very simple functions to catch common errors quickly when debugging. - - -def _check_attn_mask_or_causal(attn_mask: torch.Tensor | None, is_causal: bool, **kwargs) -> None: - if attn_mask is not None and is_causal: - raise ValueError("`is_causal` cannot be True when `attn_mask` is not None.") - - -def _check_device(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, **kwargs) -> None: - if query.device != key.device or query.device != value.device: - raise ValueError("Query, key, and value must be on the same device.") - if query.dtype != key.dtype or query.dtype != value.dtype: - raise ValueError("Query, key, and value must have the same dtype.") - - -def _check_device_cuda(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, **kwargs) -> None: - _check_device(query, key, value) - if query.device.type != "cuda": - raise ValueError("Query, key, and value must be on a CUDA device.") - - -def _check_device_cuda_atleast_smXY(major: int, minor: int) -> Callable: - def check_device_cuda(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, **kwargs) -> None: - _check_device_cuda(query, key, value) - if torch.cuda.get_device_capability(query.device) < (major, minor): - raise ValueError( - f"Query, key, and value must be on a CUDA device with compute capability >= {major}.{minor}." - ) - - return check_device_cuda - - -def _check_qkv_dtype_match(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, **kwargs) -> None: - if query.dtype != key.dtype: - raise ValueError("Query and key must have the same dtype.") - if query.dtype != value.dtype: - raise ValueError("Query and value must have the same dtype.") - - -def _check_qkv_dtype_bf16_or_fp16(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, **kwargs) -> None: - _check_qkv_dtype_match(query, key, value) - if query.dtype not in (torch.bfloat16, torch.float16): - raise ValueError("Query, key, and value must be either bfloat16 or float16.") - - -def _check_shape( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - **kwargs, -) -> None: - # Expected shapes: - # query: (batch_size, seq_len_q, num_heads, head_dim) - # key: (batch_size, seq_len_kv, num_heads, head_dim) - # value: (batch_size, seq_len_kv, num_heads, head_dim) - # attn_mask: (seq_len_q, seq_len_kv) or (batch_size, seq_len_q, seq_len_kv) - # or (batch_size, num_heads, seq_len_q, seq_len_kv) - if query.shape[-1] != key.shape[-1]: - raise ValueError("Query and key must have the same head dimension.") - if key.shape[-3] != value.shape[-3]: - raise ValueError("Key and value must have the same sequence length.") - if attn_mask is not None and attn_mask.shape[-1] != key.shape[-3]: - raise ValueError("Attention mask must match the key's sequence length.") - - -# ===== Helper functions ===== - - -def _check_attention_backend_requirements(backend: AttentionBackendName) -> None: - if backend in [AttentionBackendName.FLASH, AttentionBackendName.FLASH_VARLEN]: - if not _CAN_USE_FLASH_ATTN: - raise RuntimeError( - f"Flash Attention backend '{backend.value}' is not usable because of missing package or the version is too old. Please install `flash-attn>={_REQUIRED_FLASH_VERSION}`." - ) - - elif backend in [AttentionBackendName._FLASH_3, AttentionBackendName._FLASH_VARLEN_3]: - if not _CAN_USE_FLASH_ATTN_3: - raise RuntimeError( - f"Flash Attention 3 backend '{backend.value}' is not usable because of missing package or the version is too old. Please build FA3 beta release from source." - ) - - elif backend in [ - AttentionBackendName.FLASH_HUB, - AttentionBackendName.FLASH_VARLEN_HUB, - AttentionBackendName._FLASH_3_HUB, - AttentionBackendName._FLASH_3_VARLEN_HUB, - AttentionBackendName.SAGE_HUB, - AttentionBackendName.FLASH_4_HUB, - ]: - if not is_kernels_available(): - raise RuntimeError( - f"Backend '{backend.value}' is not usable because the `kernels` package isn't available. Please install it with `pip install kernels`." - ) - if not is_kernels_version(">=", "0.12"): - raise RuntimeError( - f"Backend '{backend.value}' needs to be used with a `kernels` version of at least 0.12. Please update with `pip install -U kernels`." - ) - - if backend == AttentionBackendName.FLASH_4_HUB and not is_kernels_version(">=", "0.12.3"): - raise RuntimeError( - f"Backend '{backend.value}' needs to be used with a `kernels` version of at least 0.12.3. Please update with `pip install -U kernels`." - ) - - elif backend == AttentionBackendName.AITER: - if not _CAN_USE_AITER_ATTN: - raise RuntimeError( - f"Aiter Attention backend '{backend.value}' is not usable because of missing package or the version is too old. Please install `aiter>={_REQUIRED_AITER_VERSION}`." - ) - - elif backend in [ - AttentionBackendName.SAGE, - AttentionBackendName.SAGE_VARLEN, - AttentionBackendName._SAGE_QK_INT8_PV_FP8_CUDA, - AttentionBackendName._SAGE_QK_INT8_PV_FP8_CUDA_SM90, - AttentionBackendName._SAGE_QK_INT8_PV_FP16_CUDA, - AttentionBackendName._SAGE_QK_INT8_PV_FP16_TRITON, - ]: - if not _CAN_USE_SAGE_ATTN: - raise RuntimeError( - f"Sage Attention backend '{backend.value}' is not usable because of missing package or the version is too old. Please install `sageattention>={_REQUIRED_SAGE_VERSION}`." - ) - - elif backend == AttentionBackendName.FLEX: - if not _CAN_USE_FLEX_ATTN: - raise RuntimeError( - f"Flex Attention backend '{backend.value}' is not usable because of missing package or the version is too old. Please install `torch>=2.5.0`." - ) - - elif backend == AttentionBackendName._NATIVE_NPU: - if not _CAN_USE_NPU_ATTN: - raise RuntimeError( - f"NPU Attention backend '{backend.value}' is not usable because of missing package or the version is too old. Please install `torch_npu`." - ) - - elif backend == AttentionBackendName._NATIVE_XLA: - if not _CAN_USE_XLA_ATTN: - raise RuntimeError( - f"XLA Attention backend '{backend.value}' is not usable because of missing package or the version is too old. Please install `torch_xla>={_REQUIRED_XLA_VERSION}`." - ) - - elif backend == AttentionBackendName.XFORMERS: - if not _CAN_USE_XFORMERS_ATTN: - raise RuntimeError( - f"Xformers Attention backend '{backend.value}' is not usable because of missing package or the version is too old. Please install `xformers>={_REQUIRED_XFORMERS_VERSION}`." - ) - - -@lru_cache_unless_export(maxsize=128) -def _prepare_for_flash_attn_or_sage_varlen_without_mask( - batch_size: int, - seq_len_q: int, - seq_len_kv: int, - device: torch.device | None = None, -): - seqlens_q = torch.full((batch_size,), seq_len_q, dtype=torch.int32, device=device) - seqlens_k = torch.full((batch_size,), seq_len_kv, dtype=torch.int32, device=device) - cu_seqlens_q = torch.zeros(batch_size + 1, dtype=torch.int32, device=device) - cu_seqlens_k = torch.zeros(batch_size + 1, dtype=torch.int32, device=device) - cu_seqlens_q[1:] = torch.cumsum(seqlens_q, dim=0) - cu_seqlens_k[1:] = torch.cumsum(seqlens_k, dim=0) - max_seqlen_q = seqlens_q.max().item() - max_seqlen_k = seqlens_k.max().item() - return (seqlens_q, seqlens_k), (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) - - -def _prepare_for_flash_attn_or_sage_varlen_with_mask( - batch_size: int, - seq_len_q: int, - attn_mask: torch.Tensor, - device: torch.device | None = None, -): - seqlens_q = torch.full((batch_size,), seq_len_q, dtype=torch.int32, device=device) - seqlens_k = attn_mask.sum(dim=1, dtype=torch.int32) - cu_seqlens_q = torch.zeros(batch_size + 1, dtype=torch.int32, device=device) - cu_seqlens_k = torch.zeros(batch_size + 1, dtype=torch.int32, device=device) - cu_seqlens_q[1:] = torch.cumsum(seqlens_q, dim=0) - cu_seqlens_k[1:] = torch.cumsum(seqlens_k, dim=0) - max_seqlen_q = seqlens_q.max().item() - max_seqlen_k = seqlens_k.max().item() - return (seqlens_q, seqlens_k), (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) - - -def _prepare_for_flash_attn_or_sage_varlen( - batch_size: int, - seq_len_q: int, - seq_len_kv: int, - attn_mask: torch.Tensor | None = None, - device: torch.device | None = None, -) -> None: - if attn_mask is None: - return _prepare_for_flash_attn_or_sage_varlen_without_mask(batch_size, seq_len_q, seq_len_kv, device) - return _prepare_for_flash_attn_or_sage_varlen_with_mask(batch_size, seq_len_q, attn_mask, device) - - -def _unpad_to_padded(packed: torch.Tensor, indices: torch.Tensor, batch_size: int, seq_len: int) -> torch.Tensor: - """scatter a packed `(nnz, ...)` tensor back to padded `(batch_size, seq_len, ...)`.""" - output = torch.zeros(batch_size * seq_len, *packed.shape[1:], dtype=packed.dtype, device=packed.device) - output[indices] = packed - return output.view(batch_size, seq_len, *packed.shape[1:]) - - -def _normalize_attn_mask(attn_mask: torch.Tensor, batch_size: int, seq_len_k: int) -> torch.Tensor: - """ - Normalize an attention mask to shape [batch_size, seq_len_k] (bool) suitable for inferring seqlens_[q|k] in - FlashAttention/Sage varlen. - - Supports 1D to 4D shapes and common broadcasting patterns. - """ - if attn_mask.dtype != torch.bool: - raise ValueError(f"Attention mask must be of type bool, got {attn_mask.dtype}.") - - if attn_mask.ndim == 1: - # [seq_len_k] -> broadcast across batch - attn_mask = attn_mask.unsqueeze(0).expand(batch_size, seq_len_k) - - elif attn_mask.ndim == 2: - # [batch_size, seq_len_k]. Maybe broadcast across batch - if attn_mask.size(0) not in [1, batch_size]: - raise ValueError( - f"attn_mask.shape[0] ({attn_mask.shape[0]}) must be 1 or {batch_size} for 2D attention mask." - ) - attn_mask = attn_mask.expand(batch_size, seq_len_k) - - elif attn_mask.ndim == 3: - # [batch_size, seq_len_q, seq_len_k] -> reduce over query dimension - # We do this reduction because we know that arbitrary QK masks is not supported in Flash/Sage varlen. - if attn_mask.size(0) not in [1, batch_size]: - raise ValueError( - f"attn_mask.shape[0] ({attn_mask.shape[0]}) must be 1 or {batch_size} for 3D attention mask." - ) - attn_mask = attn_mask.any(dim=1) - attn_mask = attn_mask.expand(batch_size, seq_len_k) - - elif attn_mask.ndim == 4: - # [batch_size, num_heads, seq_len_q, seq_len_k] or broadcastable versions - if attn_mask.size(0) not in [1, batch_size]: - raise ValueError( - f"attn_mask.shape[0] ({attn_mask.shape[0]}) must be 1 or {batch_size} for 4D attention mask." - ) - attn_mask = attn_mask.expand(batch_size, -1, -1, seq_len_k) # [B, H, Q, K] - attn_mask = attn_mask.any(dim=(1, 2)) # [B, K] - - else: - raise ValueError(f"Unsupported attention mask shape: {attn_mask.shape}") - - if attn_mask.shape != (batch_size, seq_len_k): - raise ValueError( - f"Normalized attention mask shape mismatch: got {attn_mask.shape}, expected ({batch_size}, {seq_len_k})" - ) - - return attn_mask - - -def _flex_attention_causal_mask_mod(batch_idx, head_idx, q_idx, kv_idx): - return q_idx >= kv_idx - - -# ===== Helpers for downloading kernels ===== -def _resolve_kernel_attr(module, attr_path: str): - target = module - for attr in attr_path.split("."): - if not hasattr(target, attr): - raise AttributeError(f"Kernel module '{module.__name__}' does not define attribute path '{attr_path}'.") - target = getattr(target, attr) - return target - - -def _maybe_download_kernel_for_backend(backend: AttentionBackendName) -> None: - if backend not in _HUB_KERNELS_REGISTRY: - return - config = _HUB_KERNELS_REGISTRY[backend] - - needs_kernel = config.kernel_fn is None - needs_wrapped_forward = config.wrapped_forward_attr is not None and config.wrapped_forward_fn is None - needs_wrapped_backward = config.wrapped_backward_attr is not None and config.wrapped_backward_fn is None - - if not (needs_kernel or needs_wrapped_forward or needs_wrapped_backward): - return - - try: - from kernels import get_kernel - - kernel_module = get_kernel(config.repo_id, revision=config.revision, version=config.version) - if needs_kernel: - config.kernel_fn = _resolve_kernel_attr(kernel_module, config.function_attr) - - if needs_wrapped_forward: - config.wrapped_forward_fn = _resolve_kernel_attr(kernel_module, config.wrapped_forward_attr) - - if needs_wrapped_backward: - config.wrapped_backward_fn = _resolve_kernel_attr(kernel_module, config.wrapped_backward_attr) - - except Exception as e: - logger.error(f"An error occurred while fetching kernel '{config.repo_id}' from the Hub: {e}") - raise - - -# ===== torch op registrations ===== -# Registrations are required for fullgraph tracing compatibility -# TODO: this is only required because the beta release FA3 does not have it. There is a PR adding -# this but it was never merged: https://github.com/Dao-AILab/flash-attention/pull/1590 -@_custom_op("_diffusers_flash_attn_3::_flash_attn_forward", mutates_args=(), device_types="cuda") -def _wrapped_flash_attn_3( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - softmax_scale: float | None = None, - causal: bool = False, - qv: torch.Tensor | None = None, - q_descale: torch.Tensor | None = None, - k_descale: torch.Tensor | None = None, - v_descale: torch.Tensor | None = None, - attention_chunk: int = 0, - softcap: float = 0.0, - num_splits: int = 1, - pack_gqa: bool | None = None, - deterministic: bool = False, - sm_margin: int = 0, -) -> tuple[torch.Tensor, torch.Tensor]: - # Hardcoded for now because pytorch does not support tuple/int type hints - window_size = (-1, -1) - result = flash_attn_3_func( - q=q, - k=k, - v=v, - softmax_scale=softmax_scale, - causal=causal, - qv=qv, - q_descale=q_descale, - k_descale=k_descale, - v_descale=v_descale, - window_size=window_size, - attention_chunk=attention_chunk, - softcap=softcap, - num_splits=num_splits, - pack_gqa=pack_gqa, - deterministic=deterministic, - sm_margin=sm_margin, - return_attn_probs=True, - ) - out, lse, *_ = result - lse = lse.permute(0, 2, 1) - return out, lse - - -@_register_fake("_diffusers_flash_attn_3::_flash_attn_forward") -def _( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - softmax_scale: float | None = None, - causal: bool = False, - qv: torch.Tensor | None = None, - q_descale: torch.Tensor | None = None, - k_descale: torch.Tensor | None = None, - v_descale: torch.Tensor | None = None, - attention_chunk: int = 0, - softcap: float = 0.0, - num_splits: int = 1, - pack_gqa: bool | None = None, - deterministic: bool = False, - sm_margin: int = 0, -) -> tuple[torch.Tensor, torch.Tensor]: - window_size = (-1, -1) # noqa: F841 - # A lot of the parameters here are not yet used in any way within diffusers. - # We can safely ignore for now and keep the fake op shape propagation simple. - batch_size, seq_len, num_heads, head_dim = q.shape - lse_shape = (batch_size, seq_len, num_heads) - return torch.empty_like(q), q.new_empty(lse_shape) - - -# ===== Helper functions to use attention backends with templated CP autograd functions ===== - - -def _native_attention_forward_op( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _save_ctx: bool = True, - _parallel_config: "ParallelConfig" | None = None, -): - # Native attention does not return_lse - if return_lse: - raise ValueError("Native attention does not support return_lse=True") - - # used for backward pass - if _save_ctx: - ctx.save_for_backward(query, key, value) - ctx.attn_mask = attn_mask - ctx.dropout_p = dropout_p - ctx.is_causal = is_causal - ctx.scale = scale - ctx.enable_gqa = enable_gqa - - query, key, value = (x.permute(0, 2, 1, 3) for x in (query, key, value)) - out = torch.nn.functional.scaled_dot_product_attention( - query=query, - key=key, - value=value, - attn_mask=attn_mask, - dropout_p=dropout_p, - is_causal=is_causal, - scale=scale, - enable_gqa=enable_gqa, - ) - out = out.permute(0, 2, 1, 3) - - return out - - -def _native_attention_backward_op( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - **kwargs, -): - query, key, value = ctx.saved_tensors - - query.requires_grad_(True) - key.requires_grad_(True) - value.requires_grad_(True) - - with torch.enable_grad(): - query_t, key_t, value_t = (x.permute(0, 2, 1, 3) for x in (query, key, value)) - out = torch.nn.functional.scaled_dot_product_attention( - query=query_t, - key=key_t, - value=value_t, - attn_mask=ctx.attn_mask, - dropout_p=ctx.dropout_p, - is_causal=ctx.is_causal, - scale=ctx.scale, - enable_gqa=ctx.enable_gqa, - ) - out = out.permute(0, 2, 1, 3) - - grad_query_t, grad_key_t, grad_value_t = torch.autograd.grad( - outputs=out, inputs=[query_t, key_t, value_t], grad_outputs=grad_out, retain_graph=False - ) - - grad_query = grad_query_t.permute(0, 2, 1, 3) - grad_key = grad_key_t.permute(0, 2, 1, 3) - grad_value = grad_value_t.permute(0, 2, 1, 3) - - return grad_query, grad_key, grad_value - - -# https://github.com/pytorch/pytorch/blob/8904ba638726f8c9a5aff5977c4aa76c9d2edfa6/aten/src/ATen/native/native_functions.yaml#L14958 -# forward declaration: -# aten::_scaled_dot_product_cudnn_attention(Tensor query, Tensor key, Tensor value, Tensor? attn_bias, bool compute_log_sumexp, float dropout_p=0., bool is_causal=False, bool return_debug_mask=False, *, float? scale=None) -> (Tensor output, Tensor logsumexp, Tensor cum_seq_q, Tensor cum_seq_k, SymInt max_q, SymInt max_k, Tensor philox_seed, Tensor philox_offset, Tensor debug_attn_mask) -def _cudnn_attention_forward_op( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _save_ctx: bool = True, - _parallel_config: "ParallelConfig" | None = None, -): - if enable_gqa: - raise ValueError("`enable_gqa` is not yet supported for cuDNN attention.") - - tensors_to_save = () - - # Contiguous is a must here! Calling cuDNN backend with aten ops produces incorrect results - # if the input tensors are not contiguous. - query = query.transpose(1, 2).contiguous() - key = key.transpose(1, 2).contiguous() - value = value.transpose(1, 2).contiguous() - tensors_to_save += (query, key, value) - - out, lse, cum_seq_q, cum_seq_k, max_q, max_k, philox_seed, philox_offset, debug_attn_mask = ( - torch.ops.aten._scaled_dot_product_cudnn_attention( - query=query, - key=key, - value=value, - attn_bias=attn_mask, - compute_log_sumexp=return_lse, - dropout_p=dropout_p, - is_causal=is_causal, - return_debug_mask=False, - scale=scale, - ) - ) - - tensors_to_save += (out, lse, cum_seq_q, cum_seq_k, philox_seed, philox_offset) - if _save_ctx: - ctx.save_for_backward(*tensors_to_save) - ctx.dropout_p = dropout_p - ctx.is_causal = is_causal - ctx.scale = scale - ctx.attn_mask = attn_mask - ctx.max_q = max_q - ctx.max_k = max_k - - out = out.transpose(1, 2).contiguous() - if lse is not None: - lse = lse.transpose(1, 2).contiguous() - return (out, lse) if return_lse else out - - -# backward declaration: -# aten::_scaled_dot_product_cudnn_attention_backward(Tensor grad_out, Tensor query, Tensor key, Tensor value, Tensor out, Tensor logsumexp, Tensor philox_seed, Tensor philox_offset, Tensor attn_bias, Tensor cum_seq_q, Tensor cum_seq_k, SymInt max_q, SymInt max_k, float dropout_p, bool is_causal, *, float? scale=None) -> (Tensor, Tensor, Tensor) -def _cudnn_attention_backward_op( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - **kwargs, -): - query, key, value, out, lse, cum_seq_q, cum_seq_k, philox_seed, philox_offset = ctx.saved_tensors - - grad_out = grad_out.transpose(1, 2).contiguous() - key = key.transpose(1, 2).contiguous() - value = value.transpose(1, 2).contiguous() - - # Cannot pass first 5 arguments as kwargs because: https://github.com/pytorch/pytorch/blob/d26ca5de058dbcf56ac52bb43e84dd98df2ace97/torch/_dynamo/variables/torch.py#L1341 - grad_query, grad_key, grad_value = torch.ops.aten._scaled_dot_product_cudnn_attention_backward( - grad_out, - query, - key, - value, - out, - logsumexp=lse, - philox_seed=philox_seed, - philox_offset=philox_offset, - attn_bias=ctx.attn_mask, - cum_seq_q=cum_seq_q, - cum_seq_k=cum_seq_k, - max_q=ctx.max_q, - max_k=ctx.max_k, - dropout_p=ctx.dropout_p, - is_causal=ctx.is_causal, - scale=ctx.scale, - ) - grad_query, grad_key, grad_value = (x.transpose(1, 2).contiguous() for x in (grad_query, grad_key, grad_value)) - - return grad_query, grad_key, grad_value - - -# https://github.com/pytorch/pytorch/blob/e33fa0ece36a93dbc8ff19b0251b8d99f8ae8668/aten/src/ATen/native/native_functions.yaml#L15135 -# forward declaration: -# aten::_scaled_dot_product_flash_attention(Tensor query, Tensor key, Tensor value, float dropout_p=0.0, bool is_causal=False, bool return_debug_mask=False, *, float? scale=None) -> (Tensor output, Tensor logsumexp, Tensor cum_seq_q, Tensor cum_seq_k, SymInt max_q, SymInt max_k, Tensor rng_state, Tensor unused, Tensor debug_attn_mask) -def _native_flash_attention_forward_op( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _save_ctx: bool = True, - _parallel_config: "ParallelConfig" | None = None, -): - if enable_gqa: - raise ValueError("`enable_gqa` is not yet supported for native flash attention.") - - tensors_to_save = () - - query = query.transpose(1, 2).contiguous() - key = key.transpose(1, 2).contiguous() - value = value.transpose(1, 2).contiguous() - tensors_to_save += (query, key, value) - - out, lse, cum_seq_q, cum_seq_k, max_q, max_k, philox_seed, philox_offset, debug_attn_mask = ( - torch.ops.aten._scaled_dot_product_flash_attention( - query=query, - key=key, - value=value, - dropout_p=dropout_p, - is_causal=is_causal, - return_debug_mask=False, - scale=scale, - ) - ) - - tensors_to_save += (out, lse, cum_seq_q, cum_seq_k, philox_seed, philox_offset) - if _save_ctx: - ctx.save_for_backward(*tensors_to_save) - ctx.dropout_p = dropout_p - ctx.is_causal = is_causal - ctx.scale = scale - ctx.max_q = max_q - ctx.max_k = max_k - - out = out.transpose(1, 2).contiguous() - if lse is not None: - lse = lse.transpose(1, 2).contiguous() - return (out, lse) if return_lse else out - - -# https://github.com/pytorch/pytorch/blob/e33fa0ece36a93dbc8ff19b0251b8d99f8ae8668/aten/src/ATen/native/native_functions.yaml#L15153 -# backward declaration: -# aten::_scaled_dot_product_flash_attention_backward(Tensor grad_out, Tensor query, Tensor key, Tensor value, Tensor out, Tensor logsumexp, Tensor cum_seq_q, Tensor cum_seq_k, SymInt max_q, SymInt max_k, float dropout_p, bool is_causal, Tensor philox_seed, Tensor philox_offset, *, float? scale=None) -> (Tensor grad_query, Tensor grad_key, Tensor grad_value) -def _native_flash_attention_backward_op( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - **kwargs, -): - query, key, value, out, lse, cum_seq_q, cum_seq_k, philox_seed, philox_offset = ctx.saved_tensors - - grad_out = grad_out.transpose(1, 2).contiguous() - key = key.transpose(1, 2).contiguous() - value = value.transpose(1, 2).contiguous() - - grad_query, grad_key, grad_value = torch.ops.aten._scaled_dot_product_flash_attention_backward( - grad_out, - query, - key, - value, - out, - logsumexp=lse, - philox_seed=philox_seed, - philox_offset=philox_offset, - cum_seq_q=cum_seq_q, - cum_seq_k=cum_seq_k, - max_q=ctx.max_q, - max_k=ctx.max_k, - dropout_p=ctx.dropout_p, - is_causal=ctx.is_causal, - scale=ctx.scale, - ) - grad_query, grad_key, grad_value = (x.transpose(1, 2).contiguous() for x in (grad_query, grad_key, grad_value)) - - return grad_query, grad_key, grad_value - - -# Adapted from: https://github.com/Dao-AILab/flash-attention/blob/fd2fc9d85c8e54e5c20436465bca709bc1a6c5a1/flash_attn/flash_attn_interface.py#L807 -def _flash_attention_forward_op( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _save_ctx: bool = True, - _parallel_config: "ParallelConfig" | None = None, - *, - window_size: tuple[int, int] = (-1, -1), -): - if attn_mask is not None: - raise ValueError("`attn_mask` is not yet supported for flash-attn 2.") - if enable_gqa: - raise ValueError("`enable_gqa` is not yet supported for flash-attn 2.") - - softcap = 0.0 - alibi_slopes = None - deterministic = False - grad_enabled = any(x.requires_grad for x in (query, key, value)) - - if scale is None: - scale = query.shape[-1] ** (-0.5) - - # flash-attn only returns LSE if dropout_p > 0. So, we need to workaround. - if grad_enabled or (_parallel_config is not None and _parallel_config.context_parallel_config._world_size > 1): - dropout_p = dropout_p if dropout_p > 0 else 1e-30 - - with torch.set_grad_enabled(grad_enabled): - out, lse, S_dmask, rng_state = _wrapped_flash_attn_forward( - query, - key, - value, - dropout_p, - scale, - is_causal, - window_size[0], - window_size[1], - softcap, - alibi_slopes, - return_lse, - ) - lse = lse.permute(0, 2, 1) - - if _save_ctx: - ctx.save_for_backward(query, key, value, out, lse, rng_state) - ctx.dropout_p = dropout_p - ctx.scale = scale - ctx.is_causal = is_causal - ctx.window_size = window_size - ctx.softcap = softcap - ctx.alibi_slopes = alibi_slopes - ctx.deterministic = deterministic - - return (out, lse) if return_lse else out - - -def _flash_attention_backward_op( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - **kwargs, -): - query, key, value, out, lse, rng_state = ctx.saved_tensors - grad_query, grad_key, grad_value = torch.empty_like(query), torch.empty_like(key), torch.empty_like(value) - - lse_d = _wrapped_flash_attn_backward( # noqa: F841 - grad_out, - query, - key, - value, - out, - lse, - grad_query, - grad_key, - grad_value, - ctx.dropout_p, - ctx.scale, - ctx.is_causal, - ctx.window_size[0], - ctx.window_size[1], - ctx.softcap, - ctx.alibi_slopes, - ctx.deterministic, - rng_state, - ) - - # Head dimension may have been padded - grad_query = grad_query[..., : grad_out.shape[-1]] - grad_key = grad_key[..., : grad_out.shape[-1]] - grad_value = grad_value[..., : grad_out.shape[-1]] - - return grad_query, grad_key, grad_value - - -def _flash_attention_hub_forward_op( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _save_ctx: bool = True, - _parallel_config: "ParallelConfig" | None = None, - *, - window_size: tuple[int, int] = (-1, -1), -): - if attn_mask is not None: - raise ValueError("`attn_mask` is not yet supported for flash-attn hub kernels.") - if enable_gqa: - raise ValueError("`enable_gqa` is not yet supported for flash-attn hub kernels.") - - config = _HUB_KERNELS_REGISTRY[AttentionBackendName.FLASH_HUB] - wrapped_forward_fn = config.wrapped_forward_fn - wrapped_backward_fn = config.wrapped_backward_fn - if wrapped_forward_fn is None or wrapped_backward_fn is None: - raise RuntimeError( - "Flash attention hub kernels must expose `_wrapped_flash_attn_forward` and `_wrapped_flash_attn_backward` " - "for context parallel execution." - ) - - if scale is None: - scale = query.shape[-1] ** (-0.5) - - softcap = 0.0 - alibi_slopes = None - deterministic = False - grad_enabled = any(x.requires_grad for x in (query, key, value)) - - if grad_enabled or (_parallel_config is not None and _parallel_config.context_parallel_config._world_size > 1): - dropout_p = dropout_p if dropout_p > 0 else 1e-30 - - with torch.set_grad_enabled(grad_enabled): - out, lse, S_dmask, rng_state = wrapped_forward_fn( - query, - key, - value, - dropout_p, - scale, - is_causal, - window_size[0], - window_size[1], - softcap, - alibi_slopes, - return_lse, - ) - lse = lse.permute(0, 2, 1).contiguous() - - if _save_ctx: - ctx.save_for_backward(query, key, value, out, lse, rng_state) - ctx.dropout_p = dropout_p - ctx.scale = scale - ctx.is_causal = is_causal - ctx.window_size = window_size - ctx.softcap = softcap - ctx.alibi_slopes = alibi_slopes - ctx.deterministic = deterministic - - return (out, lse) if return_lse else out - - -def _flash_attention_hub_backward_op( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - **kwargs, -): - config = _HUB_KERNELS_REGISTRY[AttentionBackendName.FLASH_HUB] - wrapped_backward_fn = config.wrapped_backward_fn - if wrapped_backward_fn is None: - raise RuntimeError( - "Flash attention hub kernels must expose `_wrapped_flash_attn_backward` for context parallel execution." - ) - - query, key, value, out, lse, rng_state = ctx.saved_tensors - grad_query, grad_key, grad_value = torch.empty_like(query), torch.empty_like(key), torch.empty_like(value) - - _ = wrapped_backward_fn( - grad_out, - query, - key, - value, - out, - lse, - grad_query, - grad_key, - grad_value, - ctx.dropout_p, - ctx.scale, - ctx.is_causal, - ctx.window_size[0], - ctx.window_size[1], - ctx.softcap, - ctx.alibi_slopes, - ctx.deterministic, - rng_state, - ) - - grad_query = grad_query[..., : grad_out.shape[-1]] - grad_key = grad_key[..., : grad_out.shape[-1]] - grad_value = grad_value[..., : grad_out.shape[-1]] - - return grad_query, grad_key, grad_value - - -def _flash_varlen_attention_hub_forward_op( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _save_ctx: bool = True, - _parallel_config: "ParallelConfig" | None = None, - *, - window_size: tuple[int, int] = (-1, -1), -): - if enable_gqa: - raise ValueError("`enable_gqa` is not yet supported for flash-attn varlen hub kernels.") - - config = _HUB_KERNELS_REGISTRY[AttentionBackendName.FLASH_VARLEN_HUB] - wrapped_forward_fn = config.wrapped_forward_fn - wrapped_backward_fn = config.wrapped_backward_fn - if wrapped_forward_fn is None or wrapped_backward_fn is None: - raise RuntimeError( - "Flash attention varlen hub kernels must expose `_wrapped_flash_attn_varlen_forward` and " - "`_wrapped_flash_attn_varlen_backward` for context parallel execution." - ) - - if scale is None: - scale = query.shape[-1] ** (-0.5) - - softcap = 0.0 - alibi_slopes = None - deterministic = False - grad_enabled = any(x.requires_grad for x in (query, key, value)) - - if grad_enabled or (_parallel_config is not None and _parallel_config.context_parallel_config._world_size > 1): - dropout_p = dropout_p if dropout_p > 0 else 1e-30 - - batch_size, seq_len_q, num_heads, _ = query.shape - _, seq_len_kv, _, _ = key.shape - - if attn_mask is not None: - attn_mask = _normalize_attn_mask(attn_mask, batch_size, seq_len_kv) - (_, seqlens_k), (cu_seqlens_q, cu_seqlens_k), (_, max_seqlen_k) = ( - _prepare_for_flash_attn_or_sage_varlen_with_mask(batch_size, seq_len_q, attn_mask, query.device) - ) - indices_k = attn_mask.flatten().nonzero(as_tuple=False).flatten() - query_packed = query.flatten(0, 1) - key_packed = key.reshape(-1, *key.shape[2:])[indices_k] - value_packed = value.reshape(-1, *value.shape[2:])[indices_k] - max_seqlen_q = seq_len_q - else: - (_, seqlens_k), (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = ( - _prepare_for_flash_attn_or_sage_varlen_without_mask(batch_size, seq_len_q, seq_len_kv, query.device) - ) - query_packed = query.flatten(0, 1) - key_packed = key.flatten(0, 1) - value_packed = value.flatten(0, 1) - seqlens_k = None - - with torch.set_grad_enabled(grad_enabled): - out_packed, lse, _, rng_state = wrapped_forward_fn( - query_packed, - key_packed, - value_packed, - cu_seqlens_q, - cu_seqlens_k, - max_seqlen_q, - max_seqlen_k, - dropout_p, - scale, - is_causal, - window_size[0], - window_size[1], - softcap, - alibi_slopes, - return_lse, - ) - - out = out_packed.view(batch_size, seq_len_q, *out_packed.shape[1:]) - - if _save_ctx: - ctx.save_for_backward( - query_packed, key_packed, value_packed, out_packed, lse, rng_state, cu_seqlens_q, cu_seqlens_k - ) - ctx.seqlens_k = seqlens_k # None if unmasked - ctx.indices_k = indices_k if attn_mask is not None else None - ctx.max_seqlen_q = max_seqlen_q - ctx.max_seqlen_k = max_seqlen_k - ctx.batch_size = batch_size - ctx.seq_len_q = seq_len_q - ctx.seq_len_kv = seq_len_kv - ctx.num_heads = num_heads - ctx.dropout_p = dropout_p - ctx.scale = scale - ctx.is_causal = is_causal - ctx.window_size = window_size - ctx.softcap = softcap - ctx.alibi_slopes = alibi_slopes - ctx.deterministic = deterministic - - # (num_heads, batch_size * seq_len_q) -> (batch_size, seq_len_q, num_heads) - lse_sp = lse.view(num_heads, batch_size, seq_len_q).permute(1, 2, 0).contiguous() - - return (out, lse_sp) if return_lse else out - - -def _flash_varlen_attention_hub_backward_op( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - **kwargs, -): - config = _HUB_KERNELS_REGISTRY[AttentionBackendName.FLASH_VARLEN_HUB] - wrapped_backward_fn = config.wrapped_backward_fn - if wrapped_backward_fn is None: - raise RuntimeError( - "Flash attention varlen hub kernels must expose `_wrapped_flash_attn_varlen_backward` " - "for context parallel execution." - ) - - query_packed, key_packed, value_packed, out_packed, lse, rng_state, cu_seqlens_q, cu_seqlens_k = ctx.saved_tensors - - grad_out_packed = grad_out.flatten(0, 1) - grad_query, grad_key, grad_value = ( - torch.empty_like(query_packed), - torch.empty_like(key_packed), - torch.empty_like(value_packed), - ) - - _ = wrapped_backward_fn( - grad_out_packed, - query_packed, - key_packed, - value_packed, - out_packed, - lse, - grad_query, - grad_key, - grad_value, - cu_seqlens_q, - cu_seqlens_k, - ctx.max_seqlen_q, - ctx.max_seqlen_k, - ctx.dropout_p, - ctx.scale, - ctx.is_causal, - ctx.window_size[0], - ctx.window_size[1], - ctx.softcap, - ctx.alibi_slopes, - ctx.deterministic, - rng_state, - ) - - grad_query = grad_query.view(ctx.batch_size, ctx.seq_len_q, *grad_query.shape[1:]) - - if ctx.seqlens_k is not None: - grad_key = _unpad_to_padded(grad_key, ctx.indices_k, ctx.batch_size, ctx.seq_len_kv) - grad_value = _unpad_to_padded(grad_value, ctx.indices_k, ctx.batch_size, ctx.seq_len_kv) - else: - grad_key = grad_key.view(ctx.batch_size, ctx.seq_len_kv, *grad_key.shape[1:]) - grad_value = grad_value.view(ctx.batch_size, ctx.seq_len_kv, *grad_value.shape[1:]) - - grad_query = grad_query[..., : grad_out.shape[-1]] - grad_key = grad_key[..., : grad_out.shape[-1]] - grad_value = grad_value[..., : grad_out.shape[-1]] - - return grad_query, grad_key, grad_value - - -def _flash_attention_3_hub_forward_op( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _save_ctx: bool = True, - _parallel_config: "ParallelConfig" | None = None, - *, - window_size: tuple[int, int] = (-1, -1), - softcap: float = 0.0, - num_splits: int = 1, - pack_gqa: bool | None = None, - deterministic: bool = False, - sm_margin: int = 0, -): - if attn_mask is not None: - raise ValueError("`attn_mask` is not yet supported for flash-attn 3 hub kernels.") - if dropout_p != 0.0: - raise ValueError("`dropout_p` is not yet supported for flash-attn 3 hub kernels.") - if enable_gqa: - raise ValueError("`enable_gqa` is not yet supported for flash-attn 3 hub kernels.") - - config = _HUB_KERNELS_REGISTRY[AttentionBackendName._FLASH_3_HUB] - wrapped_forward_fn = config.wrapped_forward_fn - if wrapped_forward_fn is None: - raise RuntimeError( - "Flash attention 3 hub kernels must expose `flash_attn_interface._flash_attn_forward` " - "for context parallel execution." - ) - - if scale is None: - scale = query.shape[-1] ** (-0.5) - - out, softmax_lse, *_ = wrapped_forward_fn( - query, - key, - value, - None, - None, # k_new, v_new - None, # qv - None, # out - None, - None, - None, # cu_seqlens_q/k/k_new - None, - None, # seqused_q/k - None, - None, # max_seqlen_q/k - None, - None, - None, # page_table, kv_batch_idx, leftpad_k - None, - None, - None, # rotary_cos/sin, seqlens_rotary - None, - None, - None, # q_descale, k_descale, v_descale - scale, - causal=is_causal, - window_size_left=window_size[0], - window_size_right=window_size[1], - attention_chunk=0, - softcap=softcap, - num_splits=num_splits, - pack_gqa=pack_gqa, - sm_margin=sm_margin, - ) - - lse = softmax_lse.permute(0, 2, 1).contiguous() if return_lse else None - - if _save_ctx: - ctx.save_for_backward(query, key, value, out, softmax_lse) - ctx.scale = scale - ctx.is_causal = is_causal - ctx.window_size = window_size - ctx.softcap = softcap - ctx.deterministic = deterministic - ctx.sm_margin = sm_margin - - return (out, lse) if return_lse else out - - -def _flash_attention_3_hub_backward_op( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - **kwargs, -): - config = _HUB_KERNELS_REGISTRY[AttentionBackendName._FLASH_3_HUB] - wrapped_backward_fn = config.wrapped_backward_fn - if wrapped_backward_fn is None: - raise RuntimeError( - "Flash attention 3 hub kernels must expose `flash_attn_interface._flash_attn_backward` " - "for context parallel execution." - ) - - query, key, value, out, softmax_lse = ctx.saved_tensors - grad_query = torch.empty_like(query) - grad_key = torch.empty_like(key) - grad_value = torch.empty_like(value) - - wrapped_backward_fn( - grad_out, - query, - key, - value, - out, - softmax_lse, - None, - None, # cu_seqlens_q, cu_seqlens_k - None, - None, # seqused_q, seqused_k - None, - None, # max_seqlen_q, max_seqlen_k - grad_query, - grad_key, - grad_value, - ctx.scale, - ctx.is_causal, - ctx.window_size[0], - ctx.window_size[1], - ctx.softcap, - ctx.deterministic, - ctx.sm_margin, - ) - - grad_query = grad_query[..., : grad_out.shape[-1]] - grad_key = grad_key[..., : grad_out.shape[-1]] - grad_value = grad_value[..., : grad_out.shape[-1]] - - return grad_query, grad_key, grad_value - - -def _flash_attention_3_varlen_hub_forward_op( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _save_ctx: bool = True, - _parallel_config: "ParallelConfig" | None = None, - *, - window_size: tuple[int, int] = (-1, -1), - softcap: float = 0.0, - num_splits: int = 1, - pack_gqa: bool | None = None, - deterministic: bool = False, - sm_margin: int = 0, -): - if dropout_p != 0.0: - raise ValueError("`dropout_p` is not yet supported for flash-attn 3 varlen hub kernels.") - if enable_gqa: - raise ValueError("`enable_gqa` is not yet supported for flash-attn 3 varlen hub kernels.") - - config = _HUB_KERNELS_REGISTRY[AttentionBackendName._FLASH_3_VARLEN_HUB] - wrapped_forward_fn = config.wrapped_forward_fn - wrapped_backward_fn = config.wrapped_backward_fn - if wrapped_forward_fn is None or wrapped_backward_fn is None: - raise RuntimeError( - "Flash attention 3 varlen hub kernels must expose `flash_attn_interface._flash_attn_forward` and " - "`flash_attn_interface._flash_attn_backward` for context parallel execution." - ) - - if scale is None: - scale = query.shape[-1] ** (-0.5) - - batch_size, seq_len_q, num_heads, _ = query.shape - _, seq_len_kv, _, _ = key.shape - - if attn_mask is not None: - attn_mask = _normalize_attn_mask(attn_mask, batch_size, seq_len_kv) - (_, seqlens_k), (cu_seqlens_q, cu_seqlens_k), (_, max_seqlen_k) = ( - _prepare_for_flash_attn_or_sage_varlen_with_mask(batch_size, seq_len_q, attn_mask, query.device) - ) - indices_k = attn_mask.flatten().nonzero(as_tuple=False).flatten() - query_packed = query.flatten(0, 1) - key_packed = key.reshape(-1, *key.shape[2:])[indices_k] - value_packed = value.reshape(-1, *value.shape[2:])[indices_k] - max_seqlen_q = seq_len_q - else: - (_, seqlens_k), (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = ( - _prepare_for_flash_attn_or_sage_varlen_without_mask(batch_size, seq_len_q, seq_len_kv, query.device) - ) - query_packed = query.flatten(0, 1) - key_packed = key.flatten(0, 1) - value_packed = value.flatten(0, 1) - seqlens_k = None - - out_packed, softmax_lse, *_ = wrapped_forward_fn( - query_packed, - key_packed, - value_packed, - None, # k_new - None, # v_new - None, # qv - None, # out_ - cu_seqlens_q, - cu_seqlens_k, - None, # cu_seqlens_k_new - None, # seqused_q - None, # seqused_k - max_seqlen_q, - max_seqlen_k, - None, # page_table - None, # kv_batch_idx - None, # leftpad_k - None, # rotary_cos - None, # rotary_sin - None, # seqlens_rotary - None, # q_descale - None, # k_descale - None, # v_descale - scale, - causal=is_causal, - window_size_left=window_size[0], - window_size_right=window_size[1], - attention_chunk=0, - softcap=softcap, - rotary_interleaved=True, - scheduler_metadata=None, - num_splits=num_splits, - pack_gqa=pack_gqa, - sm_margin=sm_margin, - ) - - out = out_packed.view(batch_size, seq_len_q, *out_packed.shape[1:]) - - if _save_ctx: - ctx.save_for_backward( - query_packed, key_packed, value_packed, out_packed, softmax_lse, cu_seqlens_q, cu_seqlens_k - ) - ctx.seqlens_k = seqlens_k # None if unmasked - ctx.indices_k = indices_k if attn_mask is not None else None - ctx.max_seqlen_q = max_seqlen_q - ctx.max_seqlen_k = max_seqlen_k - ctx.batch_size = batch_size - ctx.seq_len_q = seq_len_q - ctx.seq_len_kv = seq_len_kv - ctx.num_heads = num_heads - ctx.scale = scale - ctx.is_causal = is_causal - ctx.window_size = window_size - ctx.softcap = softcap - ctx.deterministic = deterministic - ctx.sm_margin = sm_margin - - # softmax_lse in varlen mode: (num_heads, total_q) -> (batch_size, seq_len_q, num_heads) - lse_sp = softmax_lse.view(num_heads, batch_size, seq_len_q).permute(1, 2, 0).contiguous() - - return (out, lse_sp) if return_lse else out - - -def _flash_attention_3_varlen_hub_backward_op( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - **kwargs, -): - config = _HUB_KERNELS_REGISTRY[AttentionBackendName._FLASH_3_VARLEN_HUB] - wrapped_backward_fn = config.wrapped_backward_fn - if wrapped_backward_fn is None: - raise RuntimeError( - "Flash attention 3 varlen hub kernels must expose `flash_attn_interface._flash_attn_backward` " - "for context parallel execution." - ) - - query_packed, key_packed, value_packed, out_packed, softmax_lse, cu_seqlens_q, cu_seqlens_k = ctx.saved_tensors - - grad_out_packed = grad_out.flatten(0, 1) - grad_query, grad_key, grad_value = ( - torch.empty_like(query_packed), - torch.empty_like(key_packed), - torch.empty_like(value_packed), - ) - - wrapped_backward_fn( - grad_out_packed, - query_packed, - key_packed, - value_packed, - out_packed, - softmax_lse, - cu_seqlens_q, - cu_seqlens_k, - None, - None, # seqused_q, seqused_k - ctx.max_seqlen_q, - ctx.max_seqlen_k, - grad_query, - grad_key, - grad_value, - ctx.scale, - ctx.is_causal, - ctx.window_size[0], - ctx.window_size[1], - ctx.softcap, - ctx.deterministic, - ctx.sm_margin, - ) - - grad_query = grad_query.view(ctx.batch_size, ctx.seq_len_q, *grad_query.shape[1:]) - - if ctx.seqlens_k is not None: - grad_key = _unpad_to_padded(grad_key, ctx.indices_k, ctx.batch_size, ctx.seq_len_kv) - grad_value = _unpad_to_padded(grad_value, ctx.indices_k, ctx.batch_size, ctx.seq_len_kv) - else: - grad_key = grad_key.view(ctx.batch_size, ctx.seq_len_kv, *grad_key.shape[1:]) - grad_value = grad_value.view(ctx.batch_size, ctx.seq_len_kv, *grad_value.shape[1:]) - - grad_query = grad_query[..., : grad_out.shape[-1]] - grad_key = grad_key[..., : grad_out.shape[-1]] - grad_value = grad_value[..., : grad_out.shape[-1]] - - return grad_query, grad_key, grad_value - - -def _sage_attention_forward_op( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _save_ctx: bool = True, - _parallel_config: "ParallelConfig" | None = None, -): - if attn_mask is not None: - raise ValueError("`attn_mask` is not yet supported for Sage attention.") - if dropout_p > 0.0: - raise ValueError("`dropout_p` is not yet supported for Sage attention.") - if enable_gqa: - raise ValueError("`enable_gqa` is not yet supported for Sage attention.") - - out = sageattn( - q=query, - k=key, - v=value, - tensor_layout="NHD", - is_causal=is_causal, - sm_scale=scale, - return_lse=return_lse, - ) - lse = None - if return_lse: - out, lse, *_ = out - lse = lse.permute(0, 2, 1) - - return (out, lse) if return_lse else out - - -def _sage_attention_hub_forward_op( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _save_ctx: bool = True, - _parallel_config: "ParallelConfig" | None = None, -): - if attn_mask is not None: - raise ValueError("`attn_mask` is not yet supported for Sage attention.") - if dropout_p > 0.0: - raise ValueError("`dropout_p` is not yet supported for Sage attention.") - if enable_gqa: - raise ValueError("`enable_gqa` is not yet supported for Sage attention.") - - func = _HUB_KERNELS_REGISTRY[AttentionBackendName.SAGE_HUB].kernel_fn - out = func( - q=query, - k=key, - v=value, - tensor_layout="NHD", - is_causal=is_causal, - sm_scale=scale, - return_lse=return_lse, - ) - - lse = None - if return_lse: - out, lse, *_ = out - lse = lse.permute(0, 2, 1).contiguous() - - return (out, lse) if return_lse else out - - -def _sage_attention_backward_op( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, -): - raise NotImplementedError("Backward pass is not implemented for Sage attention.") - - -def _maybe_modify_attn_mask_npu(query: torch.Tensor, key: torch.Tensor, attn_mask: torch.Tensor | None = None): - # Skip Attention Mask if all values are 1, `None` mask can speedup the computation - if attn_mask is not None and torch.all(attn_mask != 0): - attn_mask = None - - # Reshape Attention Mask: [batch_size, seq_len_k] or [batch_size, 1, 1, seq_len_k] -> [batch_size, 1, sqe_len_q, seq_len_k] - # https://www.hiascend.com/document/detail/zh/Pytorch/730/apiref/torchnpuCustomsapi/docs/context/torch_npu-npu_fusion_attention.md - if attn_mask is not None: - if attn_mask.ndim == 2 and attn_mask.shape[0] == query.shape[0] and attn_mask.shape[1] == key.shape[1]: - batch_size, seq_len_q, seq_len_kv = attn_mask.shape[0], query.shape[1], key.shape[1] - attn_mask = attn_mask.unsqueeze(1).expand(batch_size, seq_len_q, seq_len_kv).unsqueeze(1).contiguous() - elif attn_mask.ndim == 4 and attn_mask.shape[1:3] == (1, 1): - attn_mask = attn_mask.expand(-1, -1, query.shape[1], -1).contiguous() - - attn_mask = ~attn_mask.to(torch.bool) - - return attn_mask - - -def _npu_attention_forward_op( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _save_ctx: bool = True, - _parallel_config: "ParallelConfig" | None = None, -): - if return_lse: - raise ValueError("NPU attention backend does not support setting `return_lse=True`.") - - attn_mask = _maybe_modify_attn_mask_npu(query, key, attn_mask) - - out = npu_fusion_attention( - query, - key, - value, - query.size(2), # num_heads - atten_mask=attn_mask, - input_layout="BSND", - pse=None, - scale=1.0 / math.sqrt(query.shape[-1]) if scale is None else scale, - pre_tockens=65536, - next_tockens=65536, - keep_prob=1.0 - dropout_p, - sync=False, - inner_precise=0, - )[0] - - return out - - -# Not implemented yet. -def _npu_attention_backward_op( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - **kwargs, -): - raise NotImplementedError("Backward pass is not implemented for Npu Fusion Attention.") - - -# ===== Context parallel ===== - - -# Reference: -# - https://github.com/pytorch/pytorch/blob/f58a680d09e13658a52c6ba05c63c15759846bcc/torch/distributed/_functional_collectives.py#L827 -# - https://github.com/pytorch/pytorch/blob/f58a680d09e13658a52c6ba05c63c15759846bcc/torch/distributed/_functional_collectives.py#L246 -# For fullgraph=True tracing compatibility (since FakeTensor does not have a `wait` method): -def _wait_tensor(tensor): - if isinstance(tensor, funcol.AsyncCollectiveTensor): - tensor = tensor.wait() - return tensor - - -def _all_to_all_single(x: torch.Tensor, group) -> torch.Tensor: - shape = x.shape - # HACK: We need to flatten because despite making tensors contiguous, torch single-file-ization - # to benchmark triton codegen fails somewhere: - # buf25 = torch.ops._c10d_functional.all_to_all_single.default(buf24, [1, 1], [1, 1], '3') - # ValueError: Tensors must be contiguous - x = x.flatten() - x = funcol.all_to_all_single(x, None, None, group) - x = x.reshape(shape) - x = _wait_tensor(x) - return x - - -def _all_to_all_dim_exchange(x: torch.Tensor, scatter_idx: int = 2, gather_idx: int = 1, group=None) -> torch.Tensor: - """ - Perform dimension sharding / reassembly across processes using _all_to_all_single. - - This utility reshapes and redistributes tensor `x` across the given process group, across sequence dimension or - head dimension flexibly by accepting scatter_idx and gather_idx. - - Args: - x (torch.Tensor): - Input tensor. Expected shapes: - - When scatter_idx=2, gather_idx=1: (batch_size, seq_len_local, num_heads, head_dim) - - When scatter_idx=1, gather_idx=2: (batch_size, seq_len, num_heads_local, head_dim) - scatter_idx (int) : - Dimension along which the tensor is partitioned before all-to-all. - gather_idx (int): - Dimension along which the output is reassembled after all-to-all. - group : - Distributed process group for the Ulysses group. - - Returns: - torch.Tensor: Tensor with globally exchanged dimensions. - - For (scatter_idx=2 → gather_idx=1): (batch_size, seq_len, num_heads_local, head_dim) - - For (scatter_idx=1 → gather_idx=2): (batch_size, seq_len_local, num_heads, head_dim) - """ - group_world_size = torch.distributed.get_world_size(group) - - if scatter_idx == 2 and gather_idx == 1: - # Used before Ulysses sequence parallel (SP) attention. Scatters the gathers sequence - # dimension and scatters head dimension - batch_size, seq_len_local, num_heads, head_dim = x.shape - seq_len = seq_len_local * group_world_size - num_heads_local = num_heads // group_world_size - - # B, S_LOCAL, H, D -> group_world_size, S_LOCAL, B, H_LOCAL, D - x_temp = ( - x.reshape(batch_size, seq_len_local, group_world_size, num_heads_local, head_dim) - .transpose(0, 2) - .contiguous() - ) - - if group_world_size > 1: - out = _all_to_all_single(x_temp, group=group) - else: - out = x_temp - # group_world_size, S_LOCAL, B, H_LOCAL, D -> B, S, H_LOCAL, D - out = out.reshape(seq_len, batch_size, num_heads_local, head_dim).permute(1, 0, 2, 3).contiguous() - out = out.reshape(batch_size, seq_len, num_heads_local, head_dim) - return out - elif scatter_idx == 1 and gather_idx == 2: - # Used after ulysses sequence parallel in unified SP. gathers the head dimension - # scatters back the sequence dimension. - batch_size, seq_len, num_heads_local, head_dim = x.shape - num_heads = num_heads_local * group_world_size - seq_len_local = seq_len // group_world_size - - # B, S, H_LOCAL, D -> group_world_size, H_LOCAL, S_LOCAL, B, D - x_temp = ( - x.reshape(batch_size, group_world_size, seq_len_local, num_heads_local, head_dim) - .permute(1, 3, 2, 0, 4) - .reshape(group_world_size, num_heads_local, seq_len_local, batch_size, head_dim) - ) - - if group_world_size > 1: - output = _all_to_all_single(x_temp, group) - else: - output = x_temp - output = output.reshape(num_heads, seq_len_local, batch_size, head_dim).transpose(0, 2).contiguous() - output = output.reshape(batch_size, seq_len_local, num_heads, head_dim) - return output - else: - raise ValueError("Invalid scatter/gather indices for _all_to_all_dim_exchange.") - - -class SeqAllToAllDim(torch.autograd.Function): - """ - all_to_all operation for unified sequence parallelism. uses _all_to_all_dim_exchange, see _all_to_all_dim_exchange - for more info. - """ - - @staticmethod - def forward(ctx, group, input, scatter_id=2, gather_id=1): - ctx.group = group - ctx.scatter_id = scatter_id - ctx.gather_id = gather_id - return _all_to_all_dim_exchange(input, scatter_id, gather_id, group) - - @staticmethod - def backward(ctx, grad_outputs): - grad_input = SeqAllToAllDim.apply( - ctx.group, - grad_outputs, - ctx.gather_id, # reversed - ctx.scatter_id, # reversed - ) - return (None, grad_input, None, None) - - -# Below are helper functions to handle abritrary head num and abritrary sequence length for Ulysses Anything Attention. -def _maybe_pad_qkv_head(x: torch.Tensor, H: int, group: dist.ProcessGroup) -> tuple[torch.Tensor, int]: - r"""Maybe pad the head dimension to be divisible by world_size. - x: torch.Tensor, shape (B, S_LOCAL, H, D) H: int, original global head num return: tuple[torch.Tensor, int], padded - tensor (B, S_LOCAL, H + H_PAD, D) and H_PAD - """ - world_size = dist.get_world_size(group=group) - H_PAD = 0 - if H % world_size != 0: - H_PAD = world_size - (H % world_size) - NEW_H_LOCAL = (H + H_PAD) // world_size - # e.g., Allow: H=30, world_size=8 -> NEW_H_LOCAL=4, H_PAD=2. - # NOT ALLOW: H=30, world_size=16 -> NEW_H_LOCAL=2, H_PAD=14. - assert H_PAD < NEW_H_LOCAL, f"Padding head num {H_PAD} should be less than new local head num {NEW_H_LOCAL}" - x = F.pad(x, (0, 0, 0, H_PAD)).contiguous() - return x, H_PAD - - -def _maybe_unpad_qkv_head(x: torch.Tensor, H_PAD: int, group: dist.ProcessGroup) -> torch.Tensor: - r"""Maybe unpad the head dimension. - x: torch.Tensor, shape (B, S_GLOBAL, H_LOCAL + H_PAD, D) H_PAD: int, head padding num return: torch.Tensor, - unpadded tensor (B, S_GLOBAL, H_LOCAL, D) - """ - rank = dist.get_rank(group=group) - world_size = dist.get_world_size(group=group) - # Only the last rank may have padding - if H_PAD > 0 and rank == world_size - 1: - x = x[:, :, :-H_PAD, :] - return x.contiguous() - - -def _maybe_pad_o_head(x: torch.Tensor, H: int, group: dist.ProcessGroup) -> tuple[torch.Tensor, int]: - r"""Maybe pad the head dimension to be divisible by world_size. - x: torch.Tensor, shape (B, S_GLOBAL, H_LOCAL, D) H: int, original global head num return: tuple[torch.Tensor, int], - padded tensor (B, S_GLOBAL, H_LOCAL + H_PAD, D) and H_PAD - """ - if H is None: - return x, 0 - - rank = dist.get_rank(group=group) - world_size = dist.get_world_size(group=group) - H_PAD = 0 - # Only the last rank may need padding - if H % world_size != 0: - # We need to broadcast H_PAD to all ranks to keep consistency - # in unpadding step later for all ranks. - H_PAD = world_size - (H % world_size) - NEW_H_LOCAL = (H + H_PAD) // world_size - assert H_PAD < NEW_H_LOCAL, f"Padding head num {H_PAD} should be less than new local head num {NEW_H_LOCAL}" - if rank == world_size - 1: - x = F.pad(x, (0, 0, 0, H_PAD)).contiguous() - return x, H_PAD - - -def _maybe_unpad_o_head(x: torch.Tensor, H_PAD: int, group: dist.ProcessGroup) -> torch.Tensor: - r"""Maybe unpad the head dimension. - x: torch.Tensor, shape (B, S_LOCAL, H_GLOBAL + H_PAD, D) H_PAD: int, head padding num return: torch.Tensor, - unpadded tensor (B, S_LOCAL, H_GLOBAL, D) - """ - if H_PAD > 0: - x = x[:, :, :-H_PAD, :] - return x.contiguous() - - -def ulysses_anything_metadata(query: torch.Tensor, **kwargs) -> dict: - # query: (B, S_LOCAL, H_GLOBAL, D) - assert len(query.shape) == 4, "Query tensor must be 4-dimensional of shape (B, S_LOCAL, H_GLOBAL, D)" - extra_kwargs = {} - extra_kwargs["NUM_QO_HEAD"] = query.shape[2] - extra_kwargs["Q_S_LOCAL"] = query.shape[1] - # Add other kwargs if needed in future - return extra_kwargs - - -@maybe_allow_in_graph -def all_to_all_single_any_qkv_async( - x: torch.Tensor, group: dist.ProcessGroup, **kwargs -) -> Callable[..., torch.Tensor]: - r""" - x: torch.Tensor, shape (B, S_LOCAL, H, D) return: Callable that returns (B, S_GLOBAL, H_LOCAL, D) - """ - world_size = dist.get_world_size(group=group) - B, S_LOCAL, H, D = x.shape - x, H_PAD = _maybe_pad_qkv_head(x, H, group) - H_LOCAL = (H + H_PAD) // world_size - # (world_size, S_LOCAL, B, H_LOCAL, D) - x = x.reshape(B, S_LOCAL, world_size, H_LOCAL, D).permute(2, 1, 0, 3, 4).contiguous() - - input_split_sizes = [S_LOCAL] * world_size - # S_LOCAL maybe not equal for all ranks in dynamic shape case, - # since we don't know the actual shape before this timing, thus, - # we have to use all gather to collect the S_LOCAL first. - output_split_sizes = gather_size_by_comm(S_LOCAL, group) - x = x.flatten(0, 1) # (world_size * S_LOCAL, B, H_LOCAL, D) - x = funcol.all_to_all_single(x, output_split_sizes, input_split_sizes, group) - - def wait() -> torch.Tensor: - nonlocal x, H_PAD - x = _wait_tensor(x) # (S_GLOBAL, B, H_LOCAL, D) - # (S_GLOBAL, B, H_LOCAL, D) - # -> (B, S_GLOBAL, H_LOCAL, D) - x = x.permute(1, 0, 2, 3).contiguous() - x = _maybe_unpad_qkv_head(x, H_PAD, group) - return x - - return wait - - -@maybe_allow_in_graph -def all_to_all_single_any_o_async(x: torch.Tensor, group: dist.ProcessGroup, **kwargs) -> Callable[..., torch.Tensor]: - r""" - x: torch.Tensor, shape (B, S_GLOBAL, H_LOCAL, D) return: Callable that returns (B, S_LOCAL, H_GLOBAL, D) - """ - # Assume H is provided in kwargs, since we can't infer H from x's shape. - # The padding logic needs H to determine if padding is necessary. - H = kwargs.get("NUM_QO_HEAD", None) - world_size = dist.get_world_size(group=group) - - x, H_PAD = _maybe_pad_o_head(x, H, group) - shape = x.shape # (B, S_GLOBAL, H_LOCAL, D) - (B, S_GLOBAL, H_LOCAL, D) = shape - - # input_split: e.g, S_GLOBAL=9 input splits across ranks [[5,4], [5,4],..] - # output_split: e.g, S_GLOBAL=9 output splits across ranks [[5,5], [4,4],..] - - # WARN: In some cases, e.g, joint attn in Qwen-Image, the S_LOCAL can not infer - # from tensor split due to: if c = torch.cat((a, b)), world_size=4, then, - # c.tensor_split(4)[0].shape[1] may != to (a.tensor_split(4)[0].shape[1] + - # b.tensor_split(4)[0].shape[1]) - - S_LOCAL = kwargs.get("Q_S_LOCAL") - input_split_sizes = gather_size_by_comm(S_LOCAL, group) - x = x.permute(1, 0, 2, 3).contiguous() # (S_GLOBAL, B, H_LOCAL, D) - output_split_sizes = [S_LOCAL] * world_size - x = funcol.all_to_all_single(x, output_split_sizes, input_split_sizes, group) - - def wait() -> torch.Tensor: - nonlocal x, H_PAD - x = _wait_tensor(x) # (S_GLOBAL, B, H_LOCAL, D) - x = x.reshape(world_size, S_LOCAL, B, H_LOCAL, D) - x = x.permute(2, 1, 0, 3, 4).contiguous() - x = x.reshape(B, S_LOCAL, world_size * H_LOCAL, D) - x = _maybe_unpad_o_head(x, H_PAD, group) - return x - - return wait - - -class TemplatedRingAttention(torch.autograd.Function): - @staticmethod - def forward( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None, - dropout_p: float, - is_causal: bool, - scale: float | None, - enable_gqa: bool, - return_lse: bool, - forward_op, - backward_op, - _parallel_config: "ParallelConfig" | None = None, - ): - ring_mesh = _parallel_config.context_parallel_config._ring_mesh - rank = _parallel_config.context_parallel_config._ring_local_rank - world_size = _parallel_config.context_parallel_config.ring_degree - next_rank = (rank + 1) % world_size - prev_out = prev_lse = None - - ctx.forward_op = forward_op - ctx.backward_op = backward_op - ctx.q_shape = query.shape - ctx.kv_shape = key.shape - ctx._parallel_config = _parallel_config - - kv_buffer = torch.cat([key.flatten(), value.flatten()]).contiguous() - kv_buffer = funcol.all_gather_tensor(kv_buffer, gather_dim=0, group=ring_mesh.get_group()) - kv_buffer = kv_buffer.chunk(world_size) - - for i in range(world_size): - if i > 0: - kv = kv_buffer[next_rank] - key_numel = key.numel() - key = kv[:key_numel].reshape_as(key) - value = kv[key_numel:].reshape_as(value) - next_rank = (next_rank + 1) % world_size - - out, lse = forward_op( - ctx, - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - True, - _save_ctx=i == 0, - _parallel_config=_parallel_config, - ) - - if _parallel_config.context_parallel_config.convert_to_fp32: - out = out.to(torch.float32) - lse = lse.to(torch.float32) - - # lse must be 4-D to broadcast with out (B, S, H, D). - # Some backends (e.g. cuDNN on torch>=2.9) already return a - # trailing-1 dim; others (e.g. flash-hub / native-flash) always - # return 3-D lse, so we add the dim here when needed. - # See: https://github.com/huggingface/diffusers/pull/12693#issuecomment-3627519544 - if lse.ndim == 3: - lse = lse.unsqueeze(-1) - if prev_out is not None: - out = prev_out - torch.nn.functional.sigmoid(lse - prev_lse) * (prev_out - out) - lse = prev_lse - torch.nn.functional.logsigmoid(prev_lse - lse) - prev_out = out - prev_lse = lse - - out = out.to(query.dtype) - lse = lse.squeeze(-1) - - return (out, lse) if return_lse else out - - @staticmethod - def backward( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - ): - ring_mesh = ctx._parallel_config.context_parallel_config._ring_mesh - rank = ctx._parallel_config.context_parallel_config._ring_local_rank - world_size = ctx._parallel_config.context_parallel_config.ring_degree - next_rank = (rank + 1) % world_size - next_ranks = list(range(1, world_size)) + [0] - - accum_dtype = torch.float32 if ctx._parallel_config.context_parallel_config.convert_to_fp32 else grad_out.dtype - grad_query = torch.zeros(ctx.q_shape, dtype=accum_dtype, device=grad_out.device) - grad_key = torch.zeros(ctx.kv_shape, dtype=accum_dtype, device=grad_out.device) - grad_value = torch.zeros(ctx.kv_shape, dtype=accum_dtype, device=grad_out.device) - next_grad_kv = None - - query, key, value, *_ = ctx.saved_tensors - kv_buffer = torch.cat([key.flatten(), value.flatten()]).contiguous() - kv_buffer = funcol.all_gather_tensor(kv_buffer, gather_dim=0, group=ring_mesh.get_group()) - kv_buffer = kv_buffer.chunk(world_size) - - for i in range(world_size): - if i > 0: - kv = kv_buffer[next_rank] - key_numel = key.numel() - key = kv[:key_numel].reshape_as(key) - value = kv[key_numel:].reshape_as(value) - next_rank = (next_rank + 1) % world_size - - grad_query_op, grad_key_op, grad_value_op, *_ = ctx.backward_op(ctx, grad_out) - - if i > 0: - grad_kv_buffer = _wait_tensor(next_grad_kv) - grad_key_numel = grad_key.numel() - grad_key = grad_kv_buffer[:grad_key_numel].reshape_as(grad_key) - grad_value = grad_kv_buffer[grad_key_numel:].reshape_as(grad_value) - - grad_query += grad_query_op - grad_key += grad_key_op - grad_value += grad_value_op - - if i < world_size - 1: - grad_kv_buffer = torch.cat([grad_key.flatten(), grad_value.flatten()]).contiguous() - next_grad_kv = funcol.permute_tensor(grad_kv_buffer, next_ranks, group=ring_mesh.get_group()) - - grad_query, grad_key, grad_value = (x.to(grad_out.dtype) for x in (grad_query, grad_key, grad_value)) - - return grad_query, grad_key, grad_value, None, None, None, None, None, None, None, None, None - - -class TemplatedUlyssesAttention(torch.autograd.Function): - @staticmethod - def forward( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None, - dropout_p: float, - is_causal: bool, - scale: float | None, - enable_gqa: bool, - return_lse: bool, - forward_op, - backward_op, - _parallel_config: "ParallelConfig" | None = None, - ): - ulysses_mesh = _parallel_config.context_parallel_config._ulysses_mesh - world_size = _parallel_config.context_parallel_config.ulysses_degree - group = ulysses_mesh.get_group() - - ctx.forward_op = forward_op - ctx.backward_op = backward_op - ctx._parallel_config = _parallel_config - - B, S_Q_LOCAL, H, D = query.shape - _, S_KV_LOCAL, _, _ = key.shape - H_LOCAL = H // world_size - query = query.reshape(B, S_Q_LOCAL, world_size, H_LOCAL, D).permute(2, 1, 0, 3, 4).contiguous() - key = key.reshape(B, S_KV_LOCAL, world_size, H_LOCAL, D).permute(2, 1, 0, 3, 4).contiguous() - value = value.reshape(B, S_KV_LOCAL, world_size, H_LOCAL, D).permute(2, 1, 0, 3, 4).contiguous() - query, key, value = (_all_to_all_single(x, group) for x in (query, key, value)) - query, key, value = (x.flatten(0, 1).permute(1, 0, 2, 3).contiguous() for x in (query, key, value)) - - if attn_mask is not None and attn_mask.shape[-1] == S_KV_LOCAL: - # All-gather a local mask so its layout matches the QKV layout after all-to-all. - mask_list = [torch.empty_like(attn_mask) for _ in range(world_size)] - dist.all_gather(mask_list, attn_mask, group=group) - attn_mask = torch.cat(mask_list, dim=-1) - - out = forward_op( - ctx, - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - return_lse, - _save_ctx=True, - _parallel_config=_parallel_config, - ) - if return_lse: - out, lse, *_ = out - - out = out.reshape(B, world_size, S_Q_LOCAL, H_LOCAL, D).permute(1, 3, 0, 2, 4).contiguous() - out = _all_to_all_single(out, group) - out = out.flatten(0, 1).permute(1, 2, 0, 3).contiguous() - - if return_lse: - lse = lse.reshape(B, world_size, S_Q_LOCAL, H_LOCAL).permute(1, 3, 0, 2).contiguous() - lse = _all_to_all_single(lse, group) - lse = lse.flatten(0, 1).permute(1, 2, 0).contiguous() - else: - lse = None - - return (out, lse) if return_lse else out - - @staticmethod - def backward( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - ): - ulysses_mesh = ctx._parallel_config.context_parallel_config._ulysses_mesh - world_size = ctx._parallel_config.context_parallel_config.ulysses_degree - group = ulysses_mesh.get_group() - - B, S_LOCAL, H, D = grad_out.shape - H_LOCAL = H // world_size - - grad_out = grad_out.reshape(B, S_LOCAL, world_size, H_LOCAL, D).permute(2, 1, 0, 3, 4).contiguous() - grad_out = _all_to_all_single(grad_out, group) - grad_out = grad_out.flatten(0, 1).permute(1, 0, 2, 3).contiguous() - - grad_query_op, grad_key_op, grad_value_op, *_ = ctx.backward_op(ctx, grad_out) - - grad_query, grad_key, grad_value = ( - x.reshape(B, world_size, S_LOCAL, H_LOCAL, D).permute(1, 3, 0, 2, 4).contiguous() - for x in (grad_query_op, grad_key_op, grad_value_op) - ) - grad_query, grad_key, grad_value = (_all_to_all_single(x, group) for x in (grad_query, grad_key, grad_value)) - grad_query, grad_key, grad_value = ( - x.flatten(0, 1).permute(1, 2, 0, 3).contiguous() for x in (grad_query, grad_key, grad_value) - ) - - return grad_query, grad_key, grad_value, None, None, None, None, None, None, None, None, None - - -class TemplatedRingAnythingAttention(torch.autograd.Function): - @staticmethod - def forward( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None, - dropout_p: float, - is_causal: bool, - scale: float | None, - enable_gqa: bool, - return_lse: bool, - forward_op, - backward_op, - _parallel_config: "ParallelConfig" | None = None, - ): - # Ring attention for arbitrary sequence lengths. - if attn_mask is not None: - raise ValueError( - "TemplatedRingAnythingAttention does not support non-None attn_mask: " - "non-uniform sequence lengths across ranks make cross-rank mask slicing ambiguous." - ) - ring_mesh = _parallel_config.context_parallel_config._ring_mesh - group = ring_mesh.get_group() - rank = _parallel_config.context_parallel_config._ring_local_rank - world_size = _parallel_config.context_parallel_config.ring_degree - next_rank = (rank + 1) % world_size - prev_out = prev_lse = None - - ctx.forward_op = forward_op - ctx.backward_op = backward_op - ctx.q_shape = query.shape - ctx.kv_shape = key.shape - ctx._parallel_config = _parallel_config - - kv_seq_len = key.shape[1] # local S_KV (may differ across ranks) - all_kv_seq_lens = gather_size_by_comm(kv_seq_len, group) - s_max = max(all_kv_seq_lens) - - # Padding is applied on the sequence dimension (dim=1) at the end. - def pad_to_s_max(t: torch.Tensor) -> torch.Tensor: - pad_len = s_max - t.shape[1] - if pad_len == 0: - return t - pad_shape = (t.shape[0], pad_len, *t.shape[2:]) - return torch.cat([t, t.new_zeros(pad_shape)], dim=1) - - # Pad each local KV to the maximum local sequence length so all ranks can all-gather same-sized buffers. - key_padded = pad_to_s_max(key) - value_padded = pad_to_s_max(value) - - kv_buffer = torch.cat([key_padded.flatten(), value_padded.flatten()]).contiguous() - kv_buffer = funcol.all_gather_tensor(kv_buffer, gather_dim=0, group=group) - kv_buffer = kv_buffer.chunk(world_size) - - # numel per-rank in the padded layout - kv_padded_numel = key_padded.numel() - - for i in range(world_size): - if i > 0: - true_seq_len = all_kv_seq_lens[next_rank] - kv = kv_buffer[next_rank] - # Reshape to padded shape, then slice to true sequence length - key = kv[:kv_padded_numel].reshape_as(key_padded)[:, :true_seq_len] - value = kv[kv_padded_numel:].reshape_as(value_padded)[:, :true_seq_len] - next_rank = (next_rank + 1) % world_size - else: - # i == 0: use local (unpadded) key/value - key = key_padded[:, :kv_seq_len] - value = value_padded[:, :kv_seq_len] - - out, lse = forward_op( - ctx, - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - True, - _save_ctx=i == 0, - _parallel_config=_parallel_config, - ) - - if _parallel_config.context_parallel_config.convert_to_fp32: - out = out.to(torch.float32) - lse = lse.to(torch.float32) - - if is_torch_version("<", "2.9.0"): - lse = lse.unsqueeze(-1) - if prev_out is not None: - out = prev_out - torch.nn.functional.sigmoid(lse - prev_lse) * (prev_out - out) - lse = prev_lse - torch.nn.functional.logsigmoid(prev_lse - lse) - prev_out = out - prev_lse = lse - - out = out.to(query.dtype) - lse = lse.squeeze(-1) - - return (out, lse) if return_lse else out - - @staticmethod - def backward( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - ): - raise NotImplementedError("Backward pass for Ring Anything Attention in diffusers is not implemented yet.") - - -class TemplatedUlyssesAnythingAttention(torch.autograd.Function): - @staticmethod - def forward( - ctx: torch.autograd.function.FunctionCtx, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor, - dropout_p: float, - is_causal: bool, - scale: float, - enable_gqa: bool, - return_lse: bool, - forward_op, - backward_op, - _parallel_config: "ParallelConfig" | None = None, - **kwargs, - ): - ulysses_mesh = _parallel_config.context_parallel_config._ulysses_mesh - group = ulysses_mesh.get_group() - - ctx.forward_op = forward_op - ctx.backward_op = backward_op - ctx._parallel_config = _parallel_config - - _, S_KV_LOCAL, _, _ = key.shape - - metadata = ulysses_anything_metadata(query) - query_wait = all_to_all_single_any_qkv_async(query, group, **metadata) - key_wait = all_to_all_single_any_qkv_async(key, group, **metadata) - value_wait = all_to_all_single_any_qkv_async(value, group, **metadata) - - query = query_wait() # type: torch.Tensor - key = key_wait() # type: torch.Tensor - value = value_wait() # type: torch.Tensor - - if attn_mask is not None and attn_mask.shape[-1] == S_KV_LOCAL: - # All-gather a local mask to match the post-all-to-all global sequence. - # The "anything" path allows unequal local sizes, so we pad to the - # maximum across ranks before all-gathering, then trim back. - mask_local_sizes = gather_size_by_comm(attn_mask.shape[-1], group) - max_local = max(mask_local_sizes) - if attn_mask.shape[-1] < max_local: - attn_mask = F.pad(attn_mask, (0, max_local - attn_mask.shape[-1])) - mask_list = [torch.empty_like(attn_mask) for _ in range(dist.get_world_size(group=group))] - dist.all_gather(mask_list, attn_mask, group=group) - attn_mask = torch.cat(mask_list, dim=-1) - attn_mask = attn_mask[..., : sum(mask_local_sizes)] - - out = forward_op( - ctx, - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - return_lse, - _save_ctx=False, # ulysses anything only support forward pass now. - _parallel_config=_parallel_config, - ) - if return_lse: - out, lse, *_ = out - - # out: (B, S_Q_GLOBAL, H_LOCAL, D) -> (B, S_Q_LOCAL, H_GLOBAL, D) - out_wait = all_to_all_single_any_o_async(out, group, **metadata) - - if return_lse: - # lse: (B, S_Q_GLOBAL, H_LOCAL) - lse = lse.unsqueeze(-1) # (B, S_Q_GLOBAL, H_LOCAL, D=1) - lse_wait = all_to_all_single_any_o_async(lse, group, **metadata) - out = out_wait() # type: torch.Tensor - lse = lse_wait() # type: torch.Tensor - lse = lse.squeeze(-1).contiguous() # (B, S_Q_LOCAL, H_GLOBAL) - else: - out = out_wait() # type: torch.Tensor - lse = None - - return (out, lse) if return_lse else out - - @staticmethod - def backward( - ctx: torch.autograd.function.FunctionCtx, - grad_out: torch.Tensor, - *args, - ): - raise NotImplementedError("Backward pass for Ulysses Anything Attention in diffusers is not implemented yet.") - - -def _templated_unified_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor, - dropout_p: float, - is_causal: bool, - scale: float, - enable_gqa: bool, - return_lse: bool, - forward_op, - backward_op, - _parallel_config: "ParallelConfig" | None = None, - scatter_idx: int = 2, - gather_idx: int = 1, -): - """ - Unified Sequence Parallelism attention combining Ulysses and ring attention. See: https://arxiv.org/abs/2405.07719 - """ - ulysses_mesh = _parallel_config.context_parallel_config._ulysses_mesh - ulysses_group = ulysses_mesh.get_group() - - query = SeqAllToAllDim.apply(ulysses_group, query, scatter_idx, gather_idx) - key = SeqAllToAllDim.apply(ulysses_group, key, scatter_idx, gather_idx) - value = SeqAllToAllDim.apply(ulysses_group, value, scatter_idx, gather_idx) - out = TemplatedRingAttention.apply( - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - return_lse, - forward_op, - backward_op, - _parallel_config, - ) - if return_lse: - context_layer, lse, *_ = out - else: - context_layer = out - # context_layer is of shape (B, S, H_LOCAL, D) - output = SeqAllToAllDim.apply( - ulysses_group, - context_layer, - gather_idx, - scatter_idx, - ) - if return_lse: - # lse from TemplatedRingAttention is 3-D (B, S, H_LOCAL) after its - # final squeeze(-1). SeqAllToAllDim requires a 4-D input, so we add - # the trailing dim here and remove it after the collective. - # See: https://github.com/huggingface/diffusers/pull/12693#issuecomment-3627519544 - if lse.ndim == 3: - lse = lse.unsqueeze(-1) # (B, S, H_LOCAL, 1) - lse = SeqAllToAllDim.apply(ulysses_group, lse, gather_idx, scatter_idx) - lse = lse.squeeze(-1) - return (output, lse) - return output - - -def _templated_context_parallel_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - *, - forward_op, - backward_op, - _parallel_config: "ParallelConfig" | None = None, -): - if is_causal: - raise ValueError("Causal attention is not yet supported for templated attention.") - if enable_gqa: - raise ValueError("GQA is not yet supported for templated attention.") - - # TODO: add support for unified attention with ring/ulysses degree both being > 1 - if ( - _parallel_config.context_parallel_config.ring_degree > 1 - and _parallel_config.context_parallel_config.ulysses_degree > 1 - ): - return _templated_unified_attention( - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - return_lse, - forward_op, - backward_op, - _parallel_config, - ) - elif _parallel_config.context_parallel_config.ring_degree > 1: - if _parallel_config.context_parallel_config.ring_anything: - return TemplatedRingAnythingAttention.apply( - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - return_lse, - forward_op, - backward_op, - _parallel_config, - ) - else: - return TemplatedRingAttention.apply( - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - return_lse, - forward_op, - backward_op, - _parallel_config, - ) - elif _parallel_config.context_parallel_config.ulysses_degree > 1: - if _parallel_config.context_parallel_config.ulysses_anything: - # For Any sequence lengths and Any head num support - return TemplatedUlyssesAnythingAttention.apply( - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - return_lse, - forward_op, - backward_op, - _parallel_config, - ) - else: - return TemplatedUlyssesAttention.apply( - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - return_lse, - forward_op, - backward_op, - _parallel_config, - ) - else: - raise ValueError("Reaching this branch of code is unexpected. Please report a bug.") - - -# ===== Attention backends ===== - - -@_AttentionBackendRegistry.register( - AttentionBackendName.FLASH, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], - supports_context_parallel=True, -) -def _flash_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - window_size: tuple[int, int] = (-1, -1), - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - lse = None - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for flash-attn 2.") - - if _parallel_config is None: - out = flash_attn_func( - q=query, - k=key, - v=value, - dropout_p=dropout_p, - softmax_scale=scale, - causal=is_causal, - window_size=window_size, - return_attn_probs=return_lse, - ) - if return_lse: - out, lse, *_ = out - else: - forward_op = functools.partial(_flash_attention_forward_op, window_size=window_size) - out = _templated_context_parallel_attention( - query, - key, - value, - None, - dropout_p, - is_causal, - scale, - False, - return_lse, - forward_op=forward_op, - backward_op=_flash_attention_backward_op, - _parallel_config=_parallel_config, - ) - if return_lse: - out, lse = out - - return (out, lse) if return_lse else out - - -@_AttentionBackendRegistry.register( - AttentionBackendName.FLASH_HUB, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], - supports_context_parallel=True, -) -def _flash_attention_hub( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - window_size: tuple[int, int] = (-1, -1), - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - lse = None - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for flash-attn 2.") - - func = _HUB_KERNELS_REGISTRY[AttentionBackendName.FLASH_HUB].kernel_fn - if _parallel_config is None: - out = func( - q=query, - k=key, - v=value, - dropout_p=dropout_p, - softmax_scale=scale, - causal=is_causal, - window_size=window_size, - return_attn_probs=return_lse, - ) - if return_lse: - out, lse, *_ = out - else: - forward_op = functools.partial(_flash_attention_hub_forward_op, window_size=window_size) - out = _templated_context_parallel_attention( - query, - key, - value, - None, - dropout_p, - is_causal, - scale, - False, - return_lse, - forward_op=forward_op, - backward_op=_flash_attention_hub_backward_op, - _parallel_config=_parallel_config, - ) - if return_lse: - out, lse = out - - return (out, lse) if return_lse else out - - -@_AttentionBackendRegistry.register( - AttentionBackendName.FLASH_VARLEN_HUB, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], - supports_context_parallel=True, -) -def _flash_varlen_attention_hub( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - scale: float | None = None, - is_causal: bool = False, - window_size: tuple[int, int] = (-1, -1), - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if _parallel_config is not None and _parallel_config.context_parallel_config.ring_degree > 1: - raise NotImplementedError("`ring_degree > 1` is not yet supported for the FLASH_VARLEN_HUB backend.") - - lse = None - batch_size, seq_len_q, _, _ = query.shape - _, seq_len_kv, _, _ = key.shape - - if _parallel_config is None: - if attn_mask is not None: - attn_mask = _normalize_attn_mask(attn_mask, batch_size, seq_len_kv) - (_, _), (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = ( - _prepare_for_flash_attn_or_sage_varlen_with_mask(batch_size, seq_len_q, attn_mask, query.device) - ) - indices_k = attn_mask.flatten().nonzero(as_tuple=False).flatten() - key_packed = key.reshape(-1, *key.shape[2:])[indices_k] - value_packed = value.reshape(-1, *value.shape[2:])[indices_k] - else: - (_, _), (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = ( - _prepare_for_flash_attn_or_sage_varlen_without_mask(batch_size, seq_len_q, seq_len_kv, query.device) - ) - key_packed = key.flatten(0, 1) - value_packed = value.flatten(0, 1) - - query_packed = query.flatten(0, 1) - - func = _HUB_KERNELS_REGISTRY[AttentionBackendName.FLASH_VARLEN_HUB].kernel_fn - out = func( - q=query_packed, - k=key_packed, - v=value_packed, - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_k=cu_seqlens_k, - max_seqlen_q=max_seqlen_q, - max_seqlen_k=max_seqlen_k, - dropout_p=dropout_p, - softmax_scale=scale, - causal=is_causal, - window_size=window_size, - return_attn_probs=return_lse, - ) - if return_lse: - out, lse, *_ = out - out = out.unflatten(0, (batch_size, -1)) - else: - forward_op = functools.partial(_flash_varlen_attention_hub_forward_op, window_size=window_size) - out = _templated_context_parallel_attention( - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - False, - return_lse, - forward_op=forward_op, - backward_op=_flash_varlen_attention_hub_backward_op, - _parallel_config=_parallel_config, - ) - if return_lse: - out, lse = out - - return (out, lse) if return_lse else out - - -@_AttentionBackendRegistry.register( - AttentionBackendName.FLASH_VARLEN, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], -) -def _flash_varlen_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - scale: float | None = None, - is_causal: bool = False, - window_size: tuple[int, int] = (-1, -1), - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - batch_size, seq_len_q, _, _ = query.shape - _, seq_len_kv, _, _ = key.shape - - if attn_mask is not None: - attn_mask = _normalize_attn_mask(attn_mask, batch_size, seq_len_kv) - - (_, seqlens_k), (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = ( - _prepare_for_flash_attn_or_sage_varlen( - batch_size, seq_len_q, seq_len_kv, attn_mask=attn_mask, device=query.device - ) - ) - - key_valid, value_valid = [], [] - for b in range(batch_size): - valid_len = seqlens_k[b] - key_valid.append(key[b, :valid_len]) - value_valid.append(value[b, :valid_len]) - - query_packed = query.flatten(0, 1) - key_packed = torch.cat(key_valid, dim=0) - value_packed = torch.cat(value_valid, dim=0) - - out = flash_attn_varlen_func( - q=query_packed, - k=key_packed, - v=value_packed, - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_k=cu_seqlens_k, - max_seqlen_q=max_seqlen_q, - max_seqlen_k=max_seqlen_k, - dropout_p=dropout_p, - softmax_scale=scale, - causal=is_causal, - window_size=window_size, - return_attn_probs=return_lse, - ) - out = out.unflatten(0, (batch_size, -1)) - - return out - - -@_AttentionBackendRegistry.register( - AttentionBackendName._FLASH_3, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], -) -def _flash_attention_3( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - scale: float | None = None, - is_causal: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for flash-attn 3.") - - out, lse = _wrapped_flash_attn_3( - q=query, - k=key, - v=value, - softmax_scale=scale, - causal=is_causal, - ) - return (out, lse) if return_lse else out - - -@_AttentionBackendRegistry.register( - AttentionBackendName._FLASH_3_HUB, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], - supports_context_parallel=True, -) -def _flash_attention_3_hub( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - scale: float | None = None, - is_causal: bool = False, - window_size: tuple[int, int] = (-1, -1), - softcap: float = 0.0, - deterministic: bool = False, - return_attn_probs: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for flash-attn 3.") - - func = _HUB_KERNELS_REGISTRY[AttentionBackendName._FLASH_3_HUB].kernel_fn - if _parallel_config is None: - out = func( - q=query, - k=key, - v=value, - softmax_scale=scale, - causal=is_causal, - qv=None, - q_descale=None, - k_descale=None, - v_descale=None, - window_size=window_size, - softcap=softcap, - num_splits=1, - pack_gqa=None, - deterministic=deterministic, - sm_margin=0, - return_attn_probs=return_attn_probs, - ) - return (out[0], out[1]) if return_attn_probs else out - - forward_op = functools.partial( - _flash_attention_3_hub_forward_op, - window_size=window_size, - softcap=softcap, - num_splits=1, - pack_gqa=None, - deterministic=deterministic, - sm_margin=0, - ) - backward_op = functools.partial( - _flash_attention_3_hub_backward_op, - window_size=window_size, - softcap=softcap, - num_splits=1, - pack_gqa=None, - deterministic=deterministic, - sm_margin=0, - ) - out = _templated_context_parallel_attention( - query, - key, - value, - None, - 0.0, - is_causal, - scale, - False, - return_attn_probs, - forward_op=forward_op, - backward_op=backward_op, - _parallel_config=_parallel_config, - ) - if return_attn_probs: - out, lse = out - return out, lse - - return out - - -@_AttentionBackendRegistry.register( - AttentionBackendName._FLASH_3_VARLEN_HUB, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], - supports_context_parallel=True, -) -def _flash_attention_3_varlen_hub( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - scale: float | None = None, - is_causal: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if _parallel_config is not None and _parallel_config.context_parallel_config.ring_degree > 1: - raise NotImplementedError("`ring_degree > 1` is not yet supported for the _FLASH_3_VARLEN_HUB backend.") - - batch_size, seq_len_q, _, _ = query.shape - _, seq_len_kv, _, _ = key.shape - - if _parallel_config is None: - if attn_mask is not None: - attn_mask = _normalize_attn_mask(attn_mask, batch_size, seq_len_kv) - (_, _), (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = ( - _prepare_for_flash_attn_or_sage_varlen_with_mask(batch_size, seq_len_q, attn_mask, query.device) - ) - indices_k = attn_mask.flatten().nonzero(as_tuple=False).flatten() - key_packed = key.reshape(-1, *key.shape[2:])[indices_k] - value_packed = value.reshape(-1, *value.shape[2:])[indices_k] - else: - (_, _), (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = ( - _prepare_for_flash_attn_or_sage_varlen_without_mask(batch_size, seq_len_q, seq_len_kv, query.device) - ) - key_packed = key.flatten(0, 1) - value_packed = value.flatten(0, 1) - - query_packed = query.flatten(0, 1) - - func = _HUB_KERNELS_REGISTRY[AttentionBackendName._FLASH_3_VARLEN_HUB].kernel_fn - result = func( - q=query_packed, - k=key_packed, - v=value_packed, - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_k=cu_seqlens_k, - max_seqlen_q=max_seqlen_q, - max_seqlen_k=max_seqlen_k, - softmax_scale=scale, - causal=is_causal, - ) - if isinstance(result, tuple): - out, lse, *_ = result - else: - out = result - lse = None - out = out.unflatten(0, (batch_size, -1)) - else: - forward_op = functools.partial( - _flash_attention_3_varlen_hub_forward_op, - window_size=(-1, -1), - softcap=0.0, - num_splits=1, - pack_gqa=None, - deterministic=False, - sm_margin=0, - ) - out = _templated_context_parallel_attention( - query, - key, - value, - attn_mask, - 0.0, - is_causal, - scale, - False, - return_lse, - forward_op=forward_op, - backward_op=_flash_attention_3_varlen_hub_backward_op, - _parallel_config=_parallel_config, - ) - if return_lse: - out, lse = out - - return (out, lse) if return_lse else out - - -@_AttentionBackendRegistry.register( - AttentionBackendName.FLASH_4_HUB, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], - supports_context_parallel=False, -) -def _flash_attention_4_hub( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - scale: float | None = None, - is_causal: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for flash-attn 4.") - - func = _HUB_KERNELS_REGISTRY[AttentionBackendName.FLASH_4_HUB].kernel_fn - out = func( - q=query, - k=key, - v=value, - softmax_scale=scale, - causal=is_causal, - ) - if isinstance(out, tuple): - return (out[0], out[1]) if return_lse else out[0] - return out - - -@_AttentionBackendRegistry.register( - AttentionBackendName._FLASH_VARLEN_3, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], -) -def _flash_varlen_attention_3( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - scale: float | None = None, - is_causal: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - batch_size, seq_len_q, _, _ = query.shape - _, seq_len_kv, _, _ = key.shape - - if attn_mask is not None: - attn_mask = _normalize_attn_mask(attn_mask, batch_size, seq_len_kv) - - (_, seqlens_k), (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = ( - _prepare_for_flash_attn_or_sage_varlen( - batch_size, seq_len_q, seq_len_kv, attn_mask=attn_mask, device=query.device - ) - ) - - key_valid, value_valid = [], [] - for b in range(batch_size): - valid_len = seqlens_k[b] - key_valid.append(key[b, :valid_len]) - value_valid.append(value[b, :valid_len]) - - query_packed = query.flatten(0, 1) - key_packed = torch.cat(key_valid, dim=0) - value_packed = torch.cat(value_valid, dim=0) - - result = flash_attn_3_varlen_func( - q=query_packed, - k=key_packed, - v=value_packed, - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_k=cu_seqlens_k, - max_seqlen_q=max_seqlen_q, - max_seqlen_k=max_seqlen_k, - softmax_scale=scale, - causal=is_causal, - return_attn_probs=return_lse, - ) - if isinstance(result, tuple): - out, lse, *_ = result - else: - out = result - lse = None - out = out.unflatten(0, (batch_size, -1)) - - return (out, lse) if return_lse else out - - -@_AttentionBackendRegistry.register( - AttentionBackendName.AITER, - constraints=[_check_device_cuda, _check_qkv_dtype_bf16_or_fp16, _check_shape], -) -def _aiter_flash_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for aiter attention") - - if not return_lse and torch.is_grad_enabled(): - # aiter requires return_lse=True by assertion when gradients are enabled. - out, lse, *_ = aiter_flash_attn_func( - q=query, - k=key, - v=value, - dropout_p=dropout_p, - softmax_scale=scale, - causal=is_causal, - return_lse=True, - ) - else: - out = aiter_flash_attn_func( - q=query, - k=key, - v=value, - dropout_p=dropout_p, - softmax_scale=scale, - causal=is_causal, - return_lse=return_lse, - ) - if return_lse: - out, lse, *_ = out - - return (out, lse) if return_lse else out - - -@_AttentionBackendRegistry.register( - AttentionBackendName.FLEX, - constraints=[_check_attn_mask_or_causal, _check_device, _check_shape], -) -def _native_flex_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | "flex_attention.BlockMask" | None = None, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - # TODO: should we LRU cache the block mask creation? - score_mod = None - block_mask = None - batch_size, seq_len_q, num_heads, _ = query.shape - _, seq_len_kv, _, _ = key.shape - - if attn_mask is None or isinstance(attn_mask, flex_attention.BlockMask): - block_mask = attn_mask - elif is_causal: - block_mask = flex_attention.create_block_mask( - _flex_attention_causal_mask_mod, batch_size, num_heads, seq_len_q, seq_len_kv, query.device - ) - elif torch.is_tensor(attn_mask): - if attn_mask.ndim == 2: - attn_mask = attn_mask.view(attn_mask.size(0), 1, attn_mask.size(1), 1) - - attn_mask = attn_mask.expand(batch_size, num_heads, seq_len_q, seq_len_kv) - - if attn_mask.dtype == torch.bool: - # TODO: this probably does not work but verify! - def mask_mod(batch_idx, head_idx, q_idx, kv_idx): - return attn_mask[batch_idx, head_idx, q_idx, kv_idx] - - block_mask = flex_attention.create_block_mask( - mask_mod, batch_size, None, seq_len_q, seq_len_kv, query.device - ) - else: - - def score_mod(score, batch_idx, head_idx, q_idx, kv_idx): - return score + attn_mask[batch_idx, head_idx, q_idx, kv_idx] - else: - raise ValueError("Attention mask must be either None, a BlockMask, or a 2D/4D tensor.") - - query, key, value = (x.permute(0, 2, 1, 3) for x in (query, key, value)) - out = flex_attention.flex_attention( - query=query, - key=key, - value=value, - score_mod=score_mod, - block_mask=block_mask, - scale=scale, - enable_gqa=enable_gqa, - return_lse=return_lse, - ) - out = out.permute(0, 2, 1, 3) - return out - - -def _prepare_additive_attn_mask( - attn_mask: torch.Tensor, target_dtype: torch.dtype, reshape_4d: bool = True -) -> torch.Tensor: - """ - Convert a 2D attention mask to an additive mask, optionally reshaping to 4D for SDPA. - - This helper is used by both native SDPA and xformers backends to handle both boolean and additive masks. - - Args: - attn_mask: 2D tensor [batch_size, seq_len_k] - - Boolean: True means attend, False means mask out - - Additive: 0.0 means attend, -inf means mask out - target_dtype: The dtype to convert the mask to (usually query.dtype) - reshape_4d: If True, reshape from [batch_size, seq_len_k] to [batch_size, 1, 1, seq_len_k] for broadcasting - - Returns: - Additive mask tensor where 0.0 means attend and -inf means mask out. Shape is [batch_size, seq_len_k] if - reshape_4d=False, or [batch_size, 1, 1, seq_len_k] if reshape_4d=True. - """ - # Check if the mask is boolean or already additive - if attn_mask.dtype == torch.bool: - # Convert boolean to additive: True -> 0.0, False -> -inf - attn_mask = torch.where(attn_mask, 0.0, float("-inf")) - # Convert to target dtype - attn_mask = attn_mask.to(dtype=target_dtype) - else: - # Already additive mask - just ensure correct dtype - attn_mask = attn_mask.to(dtype=target_dtype) - - # Optionally reshape to 4D for broadcasting in attention mechanisms - if reshape_4d: - batch_size, seq_len_k = attn_mask.shape - attn_mask = attn_mask.view(batch_size, 1, 1, seq_len_k) - - return attn_mask - - -@_AttentionBackendRegistry.register( - AttentionBackendName.NATIVE, - constraints=[_check_device, _check_shape], - supports_context_parallel=True, -) -def _native_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if return_lse: - raise ValueError("Native attention backend does not support setting `return_lse=True`.") - - # Reshape 2D mask to 4D for SDPA - # SDPA accepts both boolean masks (torch.bool) and additive masks (float) - if ( - attn_mask is not None - and attn_mask.ndim == 2 - and attn_mask.shape[0] == query.shape[0] - and attn_mask.shape[1] == key.shape[1] - ): - # Just reshape [batch_size, seq_len_k] -> [batch_size, 1, 1, seq_len_k] - # SDPA handles both boolean and additive masks correctly - attn_mask = attn_mask.unsqueeze(1).unsqueeze(1) - - if _parallel_config is None: - query, key, value = (x.permute(0, 2, 1, 3) for x in (query, key, value)) - out = torch.nn.functional.scaled_dot_product_attention( - query=query, - key=key, - value=value, - attn_mask=attn_mask, - dropout_p=dropout_p, - is_causal=is_causal, - scale=scale, - enable_gqa=enable_gqa, - ) - out = out.permute(0, 2, 1, 3) - else: - out = _templated_context_parallel_attention( - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - return_lse, - forward_op=_native_attention_forward_op, - backward_op=_native_attention_backward_op, - _parallel_config=_parallel_config, - ) - - return out - - -@_AttentionBackendRegistry.register( - AttentionBackendName._NATIVE_CUDNN, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], - supports_context_parallel=True, -) -def _native_cudnn_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - lse = None - if _parallel_config is None and not return_lse: - query, key, value = (x.permute(0, 2, 1, 3).contiguous() for x in (query, key, value)) - with torch.nn.attention.sdpa_kernel(torch.nn.attention.SDPBackend.CUDNN_ATTENTION): - out = torch.nn.functional.scaled_dot_product_attention( - query=query, - key=key, - value=value, - attn_mask=attn_mask, - dropout_p=dropout_p, - is_causal=is_causal, - scale=scale, - enable_gqa=enable_gqa, - ) - out = out.permute(0, 2, 1, 3) - else: - out = _templated_context_parallel_attention( - query, - key, - value, - attn_mask, - dropout_p, - is_causal, - scale, - enable_gqa, - return_lse, - forward_op=_cudnn_attention_forward_op, - backward_op=_cudnn_attention_backward_op, - _parallel_config=_parallel_config, - ) - if return_lse: - out, lse = out - - return (out, lse) if return_lse else out - - -@_AttentionBackendRegistry.register( - AttentionBackendName._NATIVE_EFFICIENT, - constraints=[_check_device, _check_shape], -) -def _native_efficient_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if return_lse: - raise ValueError("Native efficient attention backend does not support setting `return_lse=True`.") - query, key, value = (x.permute(0, 2, 1, 3) for x in (query, key, value)) - with torch.nn.attention.sdpa_kernel(torch.nn.attention.SDPBackend.EFFICIENT_ATTENTION): - out = torch.nn.functional.scaled_dot_product_attention( - query=query, - key=key, - value=value, - attn_mask=attn_mask, - dropout_p=dropout_p, - is_causal=is_causal, - scale=scale, - enable_gqa=enable_gqa, - ) - out = out.permute(0, 2, 1, 3) - return out - - -@_AttentionBackendRegistry.register( - AttentionBackendName._NATIVE_FLASH, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], - supports_context_parallel=True, -) -def _native_flash_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for aiter attention") - - lse = None - if _parallel_config is None and not return_lse: - query, key, value = (x.permute(0, 2, 1, 3) for x in (query, key, value)) - with torch.nn.attention.sdpa_kernel(torch.nn.attention.SDPBackend.FLASH_ATTENTION): - out = torch.nn.functional.scaled_dot_product_attention( - query=query, - key=key, - value=value, - attn_mask=None, # not supported - dropout_p=dropout_p, - is_causal=is_causal, - scale=scale, - enable_gqa=enable_gqa, - ) - out = out.permute(0, 2, 1, 3) - else: - out = _templated_context_parallel_attention( - query, - key, - value, - None, - dropout_p, - is_causal, - scale, - enable_gqa, - return_lse, - forward_op=_native_flash_attention_forward_op, - backward_op=_native_flash_attention_backward_op, - _parallel_config=_parallel_config, - ) - if return_lse: - out, lse = out - - return (out, lse) if return_lse else out - - -@_AttentionBackendRegistry.register( - AttentionBackendName._NATIVE_MATH, - constraints=[_check_device, _check_shape], -) -def _native_math_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if return_lse: - raise ValueError("Native math attention backend does not support setting `return_lse=True`.") - query, key, value = (x.permute(0, 2, 1, 3) for x in (query, key, value)) - with torch.nn.attention.sdpa_kernel(torch.nn.attention.SDPBackend.MATH): - out = torch.nn.functional.scaled_dot_product_attention( - query=query, - key=key, - value=value, - attn_mask=attn_mask, - dropout_p=dropout_p, - is_causal=is_causal, - scale=scale, - enable_gqa=enable_gqa, - ) - out = out.permute(0, 2, 1, 3) - return out - - -@_AttentionBackendRegistry.register( - AttentionBackendName._NATIVE_NPU, - constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], - supports_context_parallel=True, -) -def _native_npu_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - scale: float | None = None, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if return_lse: - raise ValueError("NPU attention backend does not support setting `return_lse=True`.") - if _parallel_config is None: - attn_mask = _maybe_modify_attn_mask_npu(query, key, attn_mask) - - out = npu_fusion_attention( - query, - key, - value, - query.size(2), # num_heads - atten_mask=attn_mask, - input_layout="BSND", - pse=None, - scale=1.0 / math.sqrt(query.shape[-1]) if scale is None else scale, - pre_tockens=65536, - next_tockens=65536, - keep_prob=1.0 - dropout_p, - sync=False, - inner_precise=0, - )[0] - else: - out = _templated_context_parallel_attention( - query, - key, - value, - attn_mask, - dropout_p, - None, - scale, - None, - return_lse, - forward_op=_npu_attention_forward_op, - backward_op=_npu_attention_backward_op, - _parallel_config=_parallel_config, - ) - return out - - -# Reference: https://github.com/pytorch/xla/blob/06c5533de6588f6b90aa1655d9850bcf733b90b4/torch_xla/experimental/custom_kernel.py#L853 -@_AttentionBackendRegistry.register( - AttentionBackendName._NATIVE_XLA, - constraints=[_check_device, _check_shape], -) -def _native_xla_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - is_causal: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for XLA attention") - if return_lse: - raise ValueError("XLA attention backend does not support setting `return_lse=True`.") - query, key, value = (x.permute(0, 2, 1, 3) for x in (query, key, value)) - query = query / math.sqrt(query.shape[-1]) - out = xla_flash_attention( - q=query, - k=key, - v=value, - causal=is_causal, - ) - out = out.permute(0, 2, 1, 3) - return out - - -@_AttentionBackendRegistry.register( - AttentionBackendName.SAGE, - constraints=[_check_device_cuda, _check_qkv_dtype_bf16_or_fp16, _check_shape], - supports_context_parallel=True, -) -def _sage_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - is_causal: bool = False, - scale: float | None = None, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for sage attention") - lse = None - if _parallel_config is None: - out = sageattn( - q=query, - k=key, - v=value, - tensor_layout="NHD", - is_causal=is_causal, - sm_scale=scale, - return_lse=return_lse, - ) - if return_lse: - out, lse, *_ = out - else: - out = _templated_context_parallel_attention( - query, - key, - value, - None, - 0.0, - is_causal, - scale, - False, - return_lse, - forward_op=_sage_attention_forward_op, - backward_op=_sage_attention_backward_op, - _parallel_config=_parallel_config, - ) - if return_lse: - out, lse = out - - return (out, lse) if return_lse else out - - -@_AttentionBackendRegistry.register( - AttentionBackendName.SAGE_HUB, - constraints=[_check_device_cuda, _check_qkv_dtype_bf16_or_fp16, _check_shape], - supports_context_parallel=True, -) -def _sage_attention_hub( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - is_causal: bool = False, - scale: float | None = None, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for sage attention") - lse = None - func = _HUB_KERNELS_REGISTRY[AttentionBackendName.SAGE_HUB].kernel_fn - if _parallel_config is None: - out = func( - q=query, - k=key, - v=value, - tensor_layout="NHD", - is_causal=is_causal, - sm_scale=scale, - return_lse=return_lse, - ) - if return_lse: - out, lse, *_ = out - else: - out = _templated_context_parallel_attention( - query, - key, - value, - None, - 0.0, - is_causal, - scale, - False, - return_lse, - forward_op=_sage_attention_hub_forward_op, - backward_op=_sage_attention_backward_op, - _parallel_config=_parallel_config, - ) - if return_lse: - out, lse = out - - return (out, lse) if return_lse else out - - -@_AttentionBackendRegistry.register( - AttentionBackendName.SAGE_VARLEN, - constraints=[_check_device_cuda, _check_qkv_dtype_bf16_or_fp16, _check_shape], -) -def _sage_varlen_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - is_causal: bool = False, - scale: float | None = None, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if return_lse: - raise ValueError("Sage varlen backend does not support setting `return_lse=True`.") - - batch_size, seq_len_q, _, _ = query.shape - _, seq_len_kv, _, _ = key.shape - - if attn_mask is not None: - attn_mask = _normalize_attn_mask(attn_mask, batch_size, seq_len_kv) - - (_, seqlens_k), (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = ( - _prepare_for_flash_attn_or_sage_varlen( - batch_size, seq_len_q, seq_len_kv, attn_mask=attn_mask, device=query.device - ) - ) - - key_valid, value_valid = [], [] - for b in range(batch_size): - valid_len = seqlens_k[b] - key_valid.append(key[b, :valid_len]) - value_valid.append(value[b, :valid_len]) - - query_packed = query.flatten(0, 1) - key_packed = torch.cat(key_valid, dim=0) - value_packed = torch.cat(value_valid, dim=0) - - out = sageattn_varlen( - q=query_packed, - k=key_packed, - v=value_packed, - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_k=cu_seqlens_k, - max_seqlen_q=max_seqlen_q, - max_seqlen_k=max_seqlen_k, - is_causal=is_causal, - sm_scale=scale, - ) - out = out.unflatten(0, (batch_size, -1)) - - return out - - -@_AttentionBackendRegistry.register( - AttentionBackendName._SAGE_QK_INT8_PV_FP8_CUDA, - constraints=[_check_device_cuda_atleast_smXY(9, 0), _check_shape], -) -def _sage_qk_int8_pv_fp8_cuda_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - is_causal: bool = False, - scale: float | None = None, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for sage attention") - return sageattn_qk_int8_pv_fp8_cuda( - q=query, - k=key, - v=value, - tensor_layout="NHD", - is_causal=is_causal, - sm_scale=scale, - return_lse=return_lse, - ) - - -@_AttentionBackendRegistry.register( - AttentionBackendName._SAGE_QK_INT8_PV_FP8_CUDA_SM90, - constraints=[_check_device_cuda_atleast_smXY(9, 0), _check_shape], -) -def _sage_qk_int8_pv_fp8_cuda_sm90_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - is_causal: bool = False, - scale: float | None = None, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for sage attention") - return sageattn_qk_int8_pv_fp8_cuda_sm90( - q=query, - k=key, - v=value, - tensor_layout="NHD", - is_causal=is_causal, - sm_scale=scale, - return_lse=return_lse, - ) - - -@_AttentionBackendRegistry.register( - AttentionBackendName._SAGE_QK_INT8_PV_FP16_CUDA, - constraints=[_check_device_cuda_atleast_smXY(8, 0), _check_shape], -) -def _sage_qk_int8_pv_fp16_cuda_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - is_causal: bool = False, - scale: float | None = None, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for sage attention") - return sageattn_qk_int8_pv_fp16_cuda( - q=query, - k=key, - v=value, - tensor_layout="NHD", - is_causal=is_causal, - sm_scale=scale, - return_lse=return_lse, - ) - - -@_AttentionBackendRegistry.register( - AttentionBackendName._SAGE_QK_INT8_PV_FP16_TRITON, - constraints=[_check_device_cuda_atleast_smXY(8, 0), _check_shape], -) -def _sage_qk_int8_pv_fp16_triton_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - is_causal: bool = False, - scale: float | None = None, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if attn_mask is not None: - raise ValueError("`attn_mask` is not supported for sage attention") - return sageattn_qk_int8_pv_fp16_triton( - q=query, - k=key, - v=value, - tensor_layout="NHD", - is_causal=is_causal, - sm_scale=scale, - return_lse=return_lse, - ) - - -@_AttentionBackendRegistry.register( - AttentionBackendName.XFORMERS, - constraints=[_check_attn_mask_or_causal, _check_device, _check_shape], -) -def _xformers_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attn_mask: torch.Tensor | None = None, - dropout_p: float = 0.0, - is_causal: bool = False, - scale: float | None = None, - enable_gqa: bool = False, - return_lse: bool = False, - _parallel_config: "ParallelConfig" | None = None, -) -> torch.Tensor: - if return_lse: - raise ValueError("xformers attention backend does not support setting `return_lse=True`.") - - batch_size, seq_len_q, num_heads_q, _ = query.shape - _, seq_len_kv, num_heads_kv, _ = key.shape - - if is_causal: - attn_mask = xops.LowerTriangularMask() - elif attn_mask is not None: - if attn_mask.ndim == 2: - # Convert 2D mask to 4D for xformers - # Mask can be boolean (True=attend, False=mask) or additive (0.0=attend, -inf=mask) - # xformers requires 4D additive masks [batch, heads, seq_q, seq_k] - # Need memory alignment - create larger tensor and slice for alignment - original_seq_len = attn_mask.size(1) - aligned_seq_len = ((original_seq_len + 7) // 8) * 8 # Round up to multiple of 8 - - # Create aligned 4D tensor and slice to ensure proper memory layout - aligned_mask = torch.zeros( - (batch_size, num_heads_q, seq_len_q, aligned_seq_len), - dtype=query.dtype, - device=query.device, - ) - # Convert to 4D additive mask (handles both boolean and additive inputs) - mask_additive = _prepare_additive_attn_mask( - attn_mask, target_dtype=query.dtype - ) # [batch, 1, 1, seq_len_k] - # Broadcast to [batch, heads, seq_q, seq_len_k] - aligned_mask[:, :, :, :original_seq_len] = mask_additive - # Mask out the padding (already -inf from zeros -> where with default) - aligned_mask[:, :, :, original_seq_len:] = float("-inf") - - # Slice to actual size with proper alignment - attn_mask = aligned_mask[:, :, :, :seq_len_kv] - elif attn_mask.ndim != 4: - raise ValueError("Only 2D and 4D attention masks are supported for xformers attention.") - elif attn_mask.ndim == 4: - attn_mask = attn_mask.expand(batch_size, num_heads_q, seq_len_q, seq_len_kv).type_as(query) - - if enable_gqa: - if num_heads_q % num_heads_kv != 0: - raise ValueError("Number of heads in query must be divisible by number of heads in key/value.") - num_heads_per_group = num_heads_q // num_heads_kv - query = query.unflatten(2, (num_heads_kv, -1)) - key = key.unflatten(2, (num_heads_kv, -1)).expand(-1, -1, -1, num_heads_per_group, -1) - value = value.unflatten(2, (num_heads_kv, -1)).expand(-1, -1, -1, num_heads_per_group, -1) - - out = xops.memory_efficient_attention(query, key, value, attn_mask, dropout_p, scale) - - if enable_gqa: - out = out.flatten(2, 3) - - return out diff --git a/diffusers/models/attention_processor.py b/diffusers/models/attention_processor.py deleted file mode 100644 index 1b923e7496639eaf94dba248967dd6500878b6b4..0000000000000000000000000000000000000000 --- a/diffusers/models/attention_processor.py +++ /dev/null @@ -1,5679 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from __future__ import annotations - -import inspect -import math -from typing import Callable - -import torch -import torch.nn.functional as F -from torch import nn - -from ..image_processor import IPAdapterMaskProcessor -from ..utils import deprecate, is_torch_xla_available, logging -from ..utils.import_utils import is_torch_npu_available, is_torch_xla_version, is_xformers_available -from ..utils.torch_utils import is_torch_version, maybe_allow_in_graph - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - -if is_torch_npu_available(): - import torch_npu - -if is_xformers_available(): - import xformers - import xformers.ops -else: - xformers = None - -if is_torch_xla_available(): - # flash attention pallas kernel is introduced in the torch_xla 2.3 release. - if is_torch_xla_version(">", "2.2"): - from torch_xla.experimental.custom_kernel import flash_attention - from torch_xla.runtime import is_spmd - XLA_AVAILABLE = True -else: - XLA_AVAILABLE = False - - -@maybe_allow_in_graph -class Attention(nn.Module): - r""" - A cross attention layer. - - Parameters: - query_dim (`int`): - The number of channels in the query. - cross_attention_dim (`int`, *optional*): - The number of channels in the encoder_hidden_states. If not given, defaults to `query_dim`. - heads (`int`, *optional*, defaults to 8): - The number of heads to use for multi-head attention. - kv_heads (`int`, *optional*, defaults to `None`): - The number of key and value heads to use for multi-head attention. Defaults to `heads`. If - `kv_heads=heads`, the model will use Multi Head Attention (MHA), if `kv_heads=1` the model will use Multi - Query Attention (MQA) otherwise GQA is used. - dim_head (`int`, *optional*, defaults to 64): - The number of channels in each head. - dropout (`float`, *optional*, defaults to 0.0): - The dropout probability to use. - bias (`bool`, *optional*, defaults to False): - Set to `True` for the query, key, and value linear layers to contain a bias parameter. - upcast_attention (`bool`, *optional*, defaults to False): - Set to `True` to upcast the attention computation to `float32`. - upcast_softmax (`bool`, *optional*, defaults to False): - Set to `True` to upcast the softmax computation to `float32`. - cross_attention_norm (`str`, *optional*, defaults to `None`): - The type of normalization to use for the cross attention. Can be `None`, `layer_norm`, or `group_norm`. - cross_attention_norm_num_groups (`int`, *optional*, defaults to 32): - The number of groups to use for the group norm in the cross attention. - added_kv_proj_dim (`int`, *optional*, defaults to `None`): - The number of channels to use for the added key and value projections. If `None`, no projection is used. - norm_num_groups (`int`, *optional*, defaults to `None`): - The number of groups to use for the group norm in the attention. - spatial_norm_dim (`int`, *optional*, defaults to `None`): - The number of channels to use for the spatial normalization. - out_bias (`bool`, *optional*, defaults to `True`): - Set to `True` to use a bias in the output linear layer. - scale_qk (`bool`, *optional*, defaults to `True`): - Set to `True` to scale the query and key by `1 / sqrt(dim_head)`. - only_cross_attention (`bool`, *optional*, defaults to `False`): - Set to `True` to only use cross attention and not added_kv_proj_dim. Can only be set to `True` if - `added_kv_proj_dim` is not `None`. - eps (`float`, *optional*, defaults to 1e-5): - An additional value added to the denominator in group normalization that is used for numerical stability. - rescale_output_factor (`float`, *optional*, defaults to 1.0): - A factor to rescale the output by dividing it with this value. - residual_connection (`bool`, *optional*, defaults to `False`): - Set to `True` to add the residual connection to the output. - _from_deprecated_attn_block (`bool`, *optional*, defaults to `False`): - Set to `True` if the attention block is loaded from a deprecated state dict. - processor (`AttnProcessor`, *optional*, defaults to `None`): - The attention processor to use. If `None`, defaults to `AttnProcessor2_0` if `torch 2.x` is used and - `AttnProcessor` otherwise. - """ - - def __init__( - self, - query_dim: int, - cross_attention_dim: int | None = None, - heads: int = 8, - kv_heads: int | None = None, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - upcast_attention: bool = False, - upcast_softmax: bool = False, - cross_attention_norm: str | None = None, - cross_attention_norm_num_groups: int = 32, - qk_norm: str | None = None, - added_kv_proj_dim: int | None = None, - added_proj_bias: bool | None = True, - norm_num_groups: int | None = None, - spatial_norm_dim: int | None = None, - out_bias: bool = True, - scale_qk: bool = True, - only_cross_attention: bool = False, - eps: float = 1e-5, - rescale_output_factor: float = 1.0, - residual_connection: bool = False, - _from_deprecated_attn_block: bool = False, - processor: "AttnProcessor" | None = None, - out_dim: int = None, - out_context_dim: int = None, - context_pre_only=None, - pre_only=False, - elementwise_affine: bool = True, - is_causal: bool = False, - ): - super().__init__() - - # To prevent circular import. - from .normalization import FP32LayerNorm, LpNorm, RMSNorm - - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.inner_kv_dim = self.inner_dim if kv_heads is None else dim_head * kv_heads - self.query_dim = query_dim - self.use_bias = bias - self.is_cross_attention = cross_attention_dim is not None - self.cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim - self.upcast_attention = upcast_attention - self.upcast_softmax = upcast_softmax - self.rescale_output_factor = rescale_output_factor - self.residual_connection = residual_connection - self.dropout = dropout - self.fused_projections = False - self.out_dim = out_dim if out_dim is not None else query_dim - self.out_context_dim = out_context_dim if out_context_dim is not None else query_dim - self.context_pre_only = context_pre_only - self.pre_only = pre_only - self.is_causal = is_causal - - # we make use of this private variable to know whether this class is loaded - # with an deprecated state dict so that we can convert it on the fly - self._from_deprecated_attn_block = _from_deprecated_attn_block - - self.scale_qk = scale_qk - self.scale = dim_head**-0.5 if self.scale_qk else 1.0 - - self.heads = out_dim // dim_head if out_dim is not None else heads - # for slice_size > 0 the attention score computation - # is split across the batch axis to save memory - # You can set slice_size with `set_attention_slice` - self.sliceable_head_dim = heads - - self.added_kv_proj_dim = added_kv_proj_dim - self.only_cross_attention = only_cross_attention - - if self.added_kv_proj_dim is None and self.only_cross_attention: - raise ValueError( - "`only_cross_attention` can only be set to True if `added_kv_proj_dim` is not None. Make sure to set either `only_cross_attention=False` or define `added_kv_proj_dim`." - ) - - if norm_num_groups is not None: - self.group_norm = nn.GroupNorm(num_channels=query_dim, num_groups=norm_num_groups, eps=eps, affine=True) - else: - self.group_norm = None - - if spatial_norm_dim is not None: - self.spatial_norm = SpatialNorm(f_channels=query_dim, zq_channels=spatial_norm_dim) - else: - self.spatial_norm = None - - if qk_norm is None: - self.norm_q = None - self.norm_k = None - elif qk_norm == "layer_norm": - self.norm_q = nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - elif qk_norm == "fp32_layer_norm": - self.norm_q = FP32LayerNorm(dim_head, elementwise_affine=False, bias=False, eps=eps) - self.norm_k = FP32LayerNorm(dim_head, elementwise_affine=False, bias=False, eps=eps) - elif qk_norm == "layer_norm_across_heads": - # Lumina applies qk norm across all heads - self.norm_q = nn.LayerNorm(dim_head * heads, eps=eps) - self.norm_k = nn.LayerNorm(dim_head * kv_heads, eps=eps) - elif qk_norm == "rms_norm": - self.norm_q = RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - elif qk_norm == "rms_norm_across_heads": - # LTX applies qk norm across all heads - self.norm_q = RMSNorm(dim_head * heads, eps=eps) - self.norm_k = RMSNorm(dim_head * kv_heads, eps=eps) - elif qk_norm == "l2": - self.norm_q = LpNorm(p=2, dim=-1, eps=eps) - self.norm_k = LpNorm(p=2, dim=-1, eps=eps) - else: - raise ValueError( - f"unknown qk_norm: {qk_norm}. Should be one of None, 'layer_norm', 'fp32_layer_norm', 'layer_norm_across_heads', 'rms_norm', 'rms_norm_across_heads', 'l2'." - ) - - if cross_attention_norm is None: - self.norm_cross = None - elif cross_attention_norm == "layer_norm": - self.norm_cross = nn.LayerNorm(self.cross_attention_dim) - elif cross_attention_norm == "group_norm": - if self.added_kv_proj_dim is not None: - # The given `encoder_hidden_states` are initially of shape - # (batch_size, seq_len, added_kv_proj_dim) before being projected - # to (batch_size, seq_len, cross_attention_dim). The norm is applied - # before the projection, so we need to use `added_kv_proj_dim` as - # the number of channels for the group norm. - norm_cross_num_channels = added_kv_proj_dim - else: - norm_cross_num_channels = self.cross_attention_dim - - self.norm_cross = nn.GroupNorm( - num_channels=norm_cross_num_channels, num_groups=cross_attention_norm_num_groups, eps=1e-5, affine=True - ) - else: - raise ValueError( - f"unknown cross_attention_norm: {cross_attention_norm}. Should be None, 'layer_norm' or 'group_norm'" - ) - - self.to_q = nn.Linear(query_dim, self.inner_dim, bias=bias) - - if not self.only_cross_attention: - # only relevant for the `AddedKVProcessor` classes - self.to_k = nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias) - self.to_v = nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias) - else: - self.to_k = None - self.to_v = None - - self.added_proj_bias = added_proj_bias - if self.added_kv_proj_dim is not None: - self.add_k_proj = nn.Linear(added_kv_proj_dim, self.inner_kv_dim, bias=added_proj_bias) - self.add_v_proj = nn.Linear(added_kv_proj_dim, self.inner_kv_dim, bias=added_proj_bias) - if self.context_pre_only is not None: - self.add_q_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - else: - self.add_q_proj = None - self.add_k_proj = None - self.add_v_proj = None - - if not self.pre_only: - self.to_out = nn.ModuleList([]) - self.to_out.append(nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(nn.Dropout(dropout)) - else: - self.to_out = None - - if self.context_pre_only is not None and not self.context_pre_only: - self.to_add_out = nn.Linear(self.inner_dim, self.out_context_dim, bias=out_bias) - else: - self.to_add_out = None - - if qk_norm is not None and added_kv_proj_dim is not None: - if qk_norm == "layer_norm": - self.norm_added_q = nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_added_k = nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - elif qk_norm == "fp32_layer_norm": - self.norm_added_q = FP32LayerNorm(dim_head, elementwise_affine=False, bias=False, eps=eps) - self.norm_added_k = FP32LayerNorm(dim_head, elementwise_affine=False, bias=False, eps=eps) - elif qk_norm == "rms_norm": - self.norm_added_q = RMSNorm(dim_head, eps=eps) - self.norm_added_k = RMSNorm(dim_head, eps=eps) - elif qk_norm == "rms_norm_across_heads": - # Wan applies qk norm across all heads - # Wan also doesn't apply a q norm - self.norm_added_q = None - self.norm_added_k = RMSNorm(dim_head * kv_heads, eps=eps) - else: - raise ValueError( - f"unknown qk_norm: {qk_norm}. Should be one of `None,'layer_norm','fp32_layer_norm','rms_norm'`" - ) - else: - self.norm_added_q = None - self.norm_added_k = None - - # set attention processor - # We use the AttnProcessor2_0 by default when torch 2.x is used which uses - # torch.nn.functional.scaled_dot_product_attention for native Flash/memory_efficient_attention - # but only if it has the default `scale` argument. TODO remove scale_qk check when we move to torch 2.1 - if processor is None: - processor = ( - AttnProcessor2_0() if hasattr(F, "scaled_dot_product_attention") and self.scale_qk else AttnProcessor() - ) - self.set_processor(processor) - - def set_use_xla_flash_attention( - self, - use_xla_flash_attention: bool, - partition_spec: tuple[str | None, ...] | None = None, - is_flux=False, - ) -> None: - r""" - Set whether to use xla flash attention from `torch_xla` or not. - - Args: - use_xla_flash_attention (`bool`): - Whether to use pallas flash attention kernel from `torch_xla` or not. - partition_spec (`tuple[]`, *optional*): - Specify the partition specification if using SPMD. Otherwise None. - """ - if use_xla_flash_attention: - if not is_torch_xla_available: - raise "torch_xla is not available" - elif is_torch_xla_version("<", "2.3"): - raise "flash attention pallas kernel is supported from torch_xla version 2.3" - elif is_spmd() and is_torch_xla_version("<", "2.4"): - raise "flash attention pallas kernel using SPMD is supported from torch_xla version 2.4" - else: - if is_flux: - processor = XLAFluxFlashAttnProcessor2_0(partition_spec) - else: - processor = XLAFlashAttnProcessor2_0(partition_spec) - else: - processor = ( - AttnProcessor2_0() if hasattr(F, "scaled_dot_product_attention") and self.scale_qk else AttnProcessor() - ) - self.set_processor(processor) - - def set_use_npu_flash_attention(self, use_npu_flash_attention: bool) -> None: - r""" - Set whether to use npu flash attention from `torch_npu` or not. - - """ - if use_npu_flash_attention: - processor = AttnProcessorNPU() - else: - # set attention processor - # We use the AttnProcessor2_0 by default when torch 2.x is used which uses - # torch.nn.functional.scaled_dot_product_attention for native Flash/memory_efficient_attention - # but only if it has the default `scale` argument. TODO remove scale_qk check when we move to torch 2.1 - processor = ( - AttnProcessor2_0() if hasattr(F, "scaled_dot_product_attention") and self.scale_qk else AttnProcessor() - ) - self.set_processor(processor) - - def set_use_memory_efficient_attention_xformers( - self, use_memory_efficient_attention_xformers: bool, attention_op: Callable | None = None - ) -> None: - r""" - Set whether to use memory efficient attention from `xformers` or not. - - Args: - use_memory_efficient_attention_xformers (`bool`): - Whether to use memory efficient attention from `xformers` or not. - attention_op (`Callable`, *optional*): - The attention operation to use. Defaults to `None` which uses the default attention operation from - `xformers`. - """ - is_custom_diffusion = hasattr(self, "processor") and isinstance( - self.processor, - (CustomDiffusionAttnProcessor, CustomDiffusionXFormersAttnProcessor, CustomDiffusionAttnProcessor2_0), - ) - is_added_kv_processor = hasattr(self, "processor") and isinstance( - self.processor, - ( - AttnAddedKVProcessor, - AttnAddedKVProcessor2_0, - SlicedAttnAddedKVProcessor, - XFormersAttnAddedKVProcessor, - ), - ) - is_ip_adapter = hasattr(self, "processor") and isinstance( - self.processor, - (IPAdapterAttnProcessor, IPAdapterAttnProcessor2_0, IPAdapterXFormersAttnProcessor), - ) - is_joint_processor = hasattr(self, "processor") and isinstance( - self.processor, - ( - JointAttnProcessor2_0, - XFormersJointAttnProcessor, - ), - ) - - if use_memory_efficient_attention_xformers: - if is_added_kv_processor and is_custom_diffusion: - raise NotImplementedError( - f"Memory efficient attention is currently not supported for custom diffusion for attention processor type {self.processor}" - ) - if not is_xformers_available(): - raise ModuleNotFoundError( - ( - "Refer to https://github.com/facebookresearch/xformers for more information on how to install" - " xformers" - ), - name="xformers", - ) - elif not torch.cuda.is_available(): - raise ValueError( - "torch.cuda.is_available() should be True but is False. xformers' memory efficient attention is" - " only available for GPU " - ) - else: - try: - # Make sure we can run the memory efficient attention - dtype = None - if attention_op is not None: - op_fw, op_bw = attention_op - dtype, *_ = op_fw.SUPPORTED_DTYPES - q = torch.randn((1, 2, 40), device="cuda", dtype=dtype) - _ = xformers.ops.memory_efficient_attention(q, q, q) - except Exception as e: - raise e - - if is_custom_diffusion: - processor = CustomDiffusionXFormersAttnProcessor( - train_kv=self.processor.train_kv, - train_q_out=self.processor.train_q_out, - hidden_size=self.processor.hidden_size, - cross_attention_dim=self.processor.cross_attention_dim, - attention_op=attention_op, - ) - processor.load_state_dict(self.processor.state_dict()) - if hasattr(self.processor, "to_k_custom_diffusion"): - processor.to(self.processor.to_k_custom_diffusion.weight.device) - elif is_added_kv_processor: - # TODO(Patrick, Suraj, William) - currently xformers doesn't work for UnCLIP - # which uses this type of cross attention ONLY because the attention mask of format - # [0, ..., -10.000, ..., 0, ...,] is not supported - # throw warning - logger.info( - "Memory efficient attention with `xformers` might currently not work correctly if an attention mask is required for the attention operation." - ) - processor = XFormersAttnAddedKVProcessor(attention_op=attention_op) - elif is_ip_adapter: - processor = IPAdapterXFormersAttnProcessor( - hidden_size=self.processor.hidden_size, - cross_attention_dim=self.processor.cross_attention_dim, - num_tokens=self.processor.num_tokens, - scale=self.processor.scale, - attention_op=attention_op, - ) - processor.load_state_dict(self.processor.state_dict()) - if hasattr(self.processor, "to_k_ip"): - processor.to( - device=self.processor.to_k_ip[0].weight.device, dtype=self.processor.to_k_ip[0].weight.dtype - ) - elif is_joint_processor: - processor = XFormersJointAttnProcessor(attention_op=attention_op) - else: - processor = XFormersAttnProcessor(attention_op=attention_op) - else: - if is_custom_diffusion: - attn_processor_class = ( - CustomDiffusionAttnProcessor2_0 - if hasattr(F, "scaled_dot_product_attention") - else CustomDiffusionAttnProcessor - ) - processor = attn_processor_class( - train_kv=self.processor.train_kv, - train_q_out=self.processor.train_q_out, - hidden_size=self.processor.hidden_size, - cross_attention_dim=self.processor.cross_attention_dim, - ) - processor.load_state_dict(self.processor.state_dict()) - if hasattr(self.processor, "to_k_custom_diffusion"): - processor.to(self.processor.to_k_custom_diffusion.weight.device) - elif is_ip_adapter: - processor = IPAdapterAttnProcessor2_0( - hidden_size=self.processor.hidden_size, - cross_attention_dim=self.processor.cross_attention_dim, - num_tokens=self.processor.num_tokens, - scale=self.processor.scale, - ) - processor.load_state_dict(self.processor.state_dict()) - if hasattr(self.processor, "to_k_ip"): - processor.to( - device=self.processor.to_k_ip[0].weight.device, dtype=self.processor.to_k_ip[0].weight.dtype - ) - else: - # set attention processor - # We use the AttnProcessor2_0 by default when torch 2.x is used which uses - # torch.nn.functional.scaled_dot_product_attention for native Flash/memory_efficient_attention - # but only if it has the default `scale` argument. TODO remove scale_qk check when we move to torch 2.1 - processor = ( - AttnProcessor2_0() - if hasattr(F, "scaled_dot_product_attention") and self.scale_qk - else AttnProcessor() - ) - - self.set_processor(processor) - - def set_attention_slice(self, slice_size: int) -> None: - r""" - Set the slice size for attention computation. - - Args: - slice_size (`int`): - The slice size for attention computation. - """ - if slice_size is not None and slice_size > self.sliceable_head_dim: - raise ValueError(f"slice_size {slice_size} has to be smaller or equal to {self.sliceable_head_dim}.") - - if slice_size is not None and self.added_kv_proj_dim is not None: - processor = SlicedAttnAddedKVProcessor(slice_size) - elif slice_size is not None: - processor = SlicedAttnProcessor(slice_size) - elif self.added_kv_proj_dim is not None: - processor = AttnAddedKVProcessor() - else: - # set attention processor - # We use the AttnProcessor2_0 by default when torch 2.x is used which uses - # torch.nn.functional.scaled_dot_product_attention for native Flash/memory_efficient_attention - # but only if it has the default `scale` argument. TODO remove scale_qk check when we move to torch 2.1 - processor = ( - AttnProcessor2_0() if hasattr(F, "scaled_dot_product_attention") and self.scale_qk else AttnProcessor() - ) - - self.set_processor(processor) - - def set_processor(self, processor: "AttnProcessor") -> None: - r""" - Set the attention processor to use. - - Args: - processor (`AttnProcessor`): - The attention processor to use. - """ - # if current processor is in `self._modules` and if passed `processor` is not, we need to - # pop `processor` from `self._modules` - if ( - hasattr(self, "processor") - and isinstance(self.processor, torch.nn.Module) - and not isinstance(processor, torch.nn.Module) - ): - logger.info(f"You are removing possibly trained weights of {self.processor} with {processor}") - self._modules.pop("processor") - - self.processor = processor - - def get_processor(self, return_deprecated_lora: bool = False) -> "AttentionProcessor": - r""" - Get the attention processor in use. - - Args: - return_deprecated_lora (`bool`, *optional*, defaults to `False`): - Set to `True` to return the deprecated LoRA attention processor. - - Returns: - "AttentionProcessor": The attention processor in use. - """ - if not return_deprecated_lora: - return self.processor - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - **cross_attention_kwargs, - ) -> torch.Tensor: - r""" - The forward method of the `Attention` class. - - Args: - hidden_states (`torch.Tensor`): - The hidden states of the query. - encoder_hidden_states (`torch.Tensor`, *optional*): - The hidden states of the encoder. - attention_mask (`torch.Tensor`, *optional*): - The attention mask to use. If `None`, no mask is applied. - **cross_attention_kwargs: - Additional keyword arguments to pass along to the cross attention. - - Returns: - `torch.Tensor`: The output of the attention layer. - """ - # The `Attention` class can call different attention processors / attention functions - # here we simply pass along all tensors to the selected processor class - # For standard processors that are defined here, `**cross_attention_kwargs` is empty - - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - quiet_attn_parameters = {"ip_adapter_masks", "ip_hidden_states"} - unused_kwargs = [ - k for k, _ in cross_attention_kwargs.items() if k not in attn_parameters and k not in quiet_attn_parameters - ] - if len(unused_kwargs) > 0: - logger.warning( - f"cross_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - cross_attention_kwargs = {k: w for k, w in cross_attention_kwargs.items() if k in attn_parameters} - - return self.processor( - self, - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - **cross_attention_kwargs, - ) - - def batch_to_head_dim(self, tensor: torch.Tensor) -> torch.Tensor: - r""" - Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size // heads, seq_len, dim * heads]`. `heads` - is the number of heads initialized while constructing the `Attention` class. - - Args: - tensor (`torch.Tensor`): The tensor to reshape. - - Returns: - `torch.Tensor`: The reshaped tensor. - """ - head_size = self.heads - batch_size, seq_len, dim = tensor.shape - tensor = tensor.reshape(batch_size // head_size, head_size, seq_len, dim) - tensor = tensor.permute(0, 2, 1, 3).reshape(batch_size // head_size, seq_len, dim * head_size) - return tensor - - def head_to_batch_dim(self, tensor: torch.Tensor, out_dim: int = 3) -> torch.Tensor: - r""" - Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size, seq_len, heads, dim // heads]` `heads` is - the number of heads initialized while constructing the `Attention` class. - - Args: - tensor (`torch.Tensor`): The tensor to reshape. - out_dim (`int`, *optional*, defaults to `3`): The output dimension of the tensor. If `3`, the tensor is - reshaped to `[batch_size * heads, seq_len, dim // heads]`. - - Returns: - `torch.Tensor`: The reshaped tensor. - """ - head_size = self.heads - if tensor.ndim == 3: - batch_size, seq_len, dim = tensor.shape - extra_dim = 1 - else: - batch_size, extra_dim, seq_len, dim = tensor.shape - tensor = tensor.reshape(batch_size, seq_len * extra_dim, head_size, dim // head_size) - tensor = tensor.permute(0, 2, 1, 3) - - if out_dim == 3: - tensor = tensor.reshape(batch_size * head_size, seq_len * extra_dim, dim // head_size) - - return tensor - - def get_attention_scores( - self, query: torch.Tensor, key: torch.Tensor, attention_mask: torch.Tensor | None = None - ) -> torch.Tensor: - r""" - Compute the attention scores. - - Args: - query (`torch.Tensor`): The query tensor. - key (`torch.Tensor`): The key tensor. - attention_mask (`torch.Tensor`, *optional*): The attention mask to use. If `None`, no mask is applied. - - Returns: - `torch.Tensor`: The attention probabilities/scores. - """ - dtype = query.dtype - if self.upcast_attention: - query = query.float() - key = key.float() - - if attention_mask is None: - baddbmm_input = torch.empty( - query.shape[0], query.shape[1], key.shape[1], dtype=query.dtype, device=query.device - ) - beta = 0 - else: - baddbmm_input = attention_mask - beta = 1 - - attention_scores = torch.baddbmm( - baddbmm_input, - query, - key.transpose(-1, -2), - beta=beta, - alpha=self.scale, - ) - del baddbmm_input - - if self.upcast_softmax: - attention_scores = attention_scores.float() - - attention_probs = attention_scores.softmax(dim=-1) - del attention_scores - - attention_probs = attention_probs.to(dtype) - - return attention_probs - - def prepare_attention_mask( - self, attention_mask: torch.Tensor, target_length: int, batch_size: int, out_dim: int = 3 - ) -> torch.Tensor: - r""" - Prepare the attention mask for the attention computation. - - Args: - attention_mask (`torch.Tensor`): - The attention mask to prepare. - target_length (`int`): - The target length of the attention mask. This is the length of the attention mask after padding. - batch_size (`int`): - The batch size, which is used to repeat the attention mask. - out_dim (`int`, *optional*, defaults to `3`): - The output dimension of the attention mask. Can be either `3` or `4`. - - Returns: - `torch.Tensor`: The prepared attention mask. - """ - head_size = self.heads - if attention_mask is None: - return attention_mask - - current_length: int = attention_mask.shape[-1] - if current_length != target_length: - if attention_mask.device.type == "mps": - # HACK: MPS: Does not support padding by greater than dimension of input tensor. - # Instead, we can manually construct the padding tensor. - padding_shape = (attention_mask.shape[0], attention_mask.shape[1], target_length) - padding = torch.zeros(padding_shape, dtype=attention_mask.dtype, device=attention_mask.device) - attention_mask = torch.cat([attention_mask, padding], dim=2) - else: - # TODO: for pipelines such as stable-diffusion, padding cross-attn mask: - # we want to instead pad by (0, remaining_length), where remaining_length is: - # remaining_length: int = target_length - current_length - # TODO: re-enable tests/models/test_models_unet_2d_condition.py#test_model_xattn_padding - attention_mask = F.pad(attention_mask, (0, target_length), value=0.0) - - if out_dim == 3: - if attention_mask.shape[0] < batch_size * head_size: - attention_mask = attention_mask.repeat_interleave( - head_size, dim=0, output_size=attention_mask.shape[0] * head_size - ) - elif out_dim == 4: - attention_mask = attention_mask.unsqueeze(1) - attention_mask = attention_mask.repeat_interleave( - head_size, dim=1, output_size=attention_mask.shape[1] * head_size - ) - - return attention_mask - - def norm_encoder_hidden_states(self, encoder_hidden_states: torch.Tensor) -> torch.Tensor: - r""" - Normalize the encoder hidden states. Requires `self.norm_cross` to be specified when constructing the - `Attention` class. - - Args: - encoder_hidden_states (`torch.Tensor`): Hidden states of the encoder. - - Returns: - `torch.Tensor`: The normalized encoder hidden states. - """ - assert self.norm_cross is not None, "self.norm_cross must be defined to call self.norm_encoder_hidden_states" - - if isinstance(self.norm_cross, nn.LayerNorm): - encoder_hidden_states = self.norm_cross(encoder_hidden_states) - elif isinstance(self.norm_cross, nn.GroupNorm): - # Group norm norms along the channels dimension and expects - # input to be in the shape of (N, C, *). In this case, we want - # to norm along the hidden dimension, so we need to move - # (batch_size, sequence_length, hidden_size) -> - # (batch_size, hidden_size, sequence_length) - encoder_hidden_states = encoder_hidden_states.transpose(1, 2) - encoder_hidden_states = self.norm_cross(encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states.transpose(1, 2) - else: - assert False - - return encoder_hidden_states - - @torch.no_grad() - def fuse_projections(self, fuse=True): - device = self.to_q.weight.data.device - dtype = self.to_q.weight.data.dtype - - if not self.is_cross_attention: - # fetch weight matrices. - concatenated_weights = torch.cat([self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data]) - in_features = concatenated_weights.shape[1] - out_features = concatenated_weights.shape[0] - - # create a new single projection layer and copy over the weights. - self.to_qkv = nn.Linear(in_features, out_features, bias=self.use_bias, device=device, dtype=dtype) - self.to_qkv.weight.copy_(concatenated_weights) - if self.use_bias: - concatenated_bias = torch.cat([self.to_q.bias.data, self.to_k.bias.data, self.to_v.bias.data]) - self.to_qkv.bias.copy_(concatenated_bias) - - else: - concatenated_weights = torch.cat([self.to_k.weight.data, self.to_v.weight.data]) - in_features = concatenated_weights.shape[1] - out_features = concatenated_weights.shape[0] - - self.to_kv = nn.Linear(in_features, out_features, bias=self.use_bias, device=device, dtype=dtype) - self.to_kv.weight.copy_(concatenated_weights) - if self.use_bias: - concatenated_bias = torch.cat([self.to_k.bias.data, self.to_v.bias.data]) - self.to_kv.bias.copy_(concatenated_bias) - - # handle added projections for SD3 and others. - if ( - getattr(self, "add_q_proj", None) is not None - and getattr(self, "add_k_proj", None) is not None - and getattr(self, "add_v_proj", None) is not None - ): - concatenated_weights = torch.cat( - [self.add_q_proj.weight.data, self.add_k_proj.weight.data, self.add_v_proj.weight.data] - ) - in_features = concatenated_weights.shape[1] - out_features = concatenated_weights.shape[0] - - self.to_added_qkv = nn.Linear( - in_features, out_features, bias=self.added_proj_bias, device=device, dtype=dtype - ) - self.to_added_qkv.weight.copy_(concatenated_weights) - if self.added_proj_bias: - concatenated_bias = torch.cat( - [self.add_q_proj.bias.data, self.add_k_proj.bias.data, self.add_v_proj.bias.data] - ) - self.to_added_qkv.bias.copy_(concatenated_bias) - - self.fused_projections = fuse - - -class SanaMultiscaleAttentionProjection(nn.Module): - def __init__( - self, - in_channels: int, - num_attention_heads: int, - kernel_size: int, - ) -> None: - super().__init__() - - channels = 3 * in_channels - self.proj_in = nn.Conv2d( - channels, - channels, - kernel_size, - padding=kernel_size // 2, - groups=channels, - bias=False, - ) - self.proj_out = nn.Conv2d(channels, channels, 1, 1, 0, groups=3 * num_attention_heads, bias=False) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.proj_in(hidden_states) - hidden_states = self.proj_out(hidden_states) - return hidden_states - - -class SanaMultiscaleLinearAttention(nn.Module): - r"""Lightweight multi-scale linear attention""" - - def __init__( - self, - in_channels: int, - out_channels: int, - num_attention_heads: int | None = None, - attention_head_dim: int = 8, - mult: float = 1.0, - norm_type: str = "batch_norm", - kernel_sizes: tuple[int, ...] = (5,), - eps: float = 1e-15, - residual_connection: bool = False, - ): - super().__init__() - - # To prevent circular import - from .normalization import get_normalization - - self.eps = eps - self.attention_head_dim = attention_head_dim - self.norm_type = norm_type - self.residual_connection = residual_connection - - num_attention_heads = ( - int(in_channels // attention_head_dim * mult) if num_attention_heads is None else num_attention_heads - ) - inner_dim = num_attention_heads * attention_head_dim - - self.to_q = nn.Linear(in_channels, inner_dim, bias=False) - self.to_k = nn.Linear(in_channels, inner_dim, bias=False) - self.to_v = nn.Linear(in_channels, inner_dim, bias=False) - - self.to_qkv_multiscale = nn.ModuleList() - for kernel_size in kernel_sizes: - self.to_qkv_multiscale.append( - SanaMultiscaleAttentionProjection(inner_dim, num_attention_heads, kernel_size) - ) - - self.nonlinearity = nn.ReLU() - self.to_out = nn.Linear(inner_dim * (1 + len(kernel_sizes)), out_channels, bias=False) - self.norm_out = get_normalization(norm_type, num_features=out_channels) - - self.processor = SanaMultiscaleAttnProcessor2_0() - - def apply_linear_attention(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor) -> torch.Tensor: - value = F.pad(value, (0, 0, 0, 1), mode="constant", value=1) # Adds padding - scores = torch.matmul(value, key.transpose(-1, -2)) - hidden_states = torch.matmul(scores, query) - - hidden_states = hidden_states.to(dtype=torch.float32) - hidden_states = hidden_states[:, :, :-1] / (hidden_states[:, :, -1:] + self.eps) - return hidden_states - - def apply_quadratic_attention(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor) -> torch.Tensor: - scores = torch.matmul(key.transpose(-1, -2), query) - scores = scores.to(dtype=torch.float32) - scores = scores / (torch.sum(scores, dim=2, keepdim=True) + self.eps) - hidden_states = torch.matmul(value, scores.to(value.dtype)) - return hidden_states - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return self.processor(self, hidden_states) - - -class MochiAttention(nn.Module): - def __init__( - self, - query_dim: int, - added_kv_proj_dim: int, - processor: "MochiAttnProcessor2_0", - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - added_proj_bias: bool = True, - out_dim: int | None = None, - out_context_dim: int | None = None, - out_bias: bool = True, - context_pre_only: bool = False, - eps: float = 1e-5, - ): - super().__init__() - from .normalization import MochiRMSNorm - - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.out_dim = out_dim if out_dim is not None else query_dim - self.out_context_dim = out_context_dim if out_context_dim else query_dim - self.context_pre_only = context_pre_only - - self.heads = out_dim // dim_head if out_dim is not None else heads - - self.norm_q = MochiRMSNorm(dim_head, eps, True) - self.norm_k = MochiRMSNorm(dim_head, eps, True) - self.norm_added_q = MochiRMSNorm(dim_head, eps, True) - self.norm_added_k = MochiRMSNorm(dim_head, eps, True) - - self.to_q = nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_v = nn.Linear(query_dim, self.inner_dim, bias=bias) - - self.add_k_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_v_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - if self.context_pre_only is not None: - self.add_q_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - - self.to_out = nn.ModuleList([]) - self.to_out.append(nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(nn.Dropout(dropout)) - - if not self.context_pre_only: - self.to_add_out = nn.Linear(self.inner_dim, self.out_context_dim, bias=out_bias) - - self.processor = processor - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - **kwargs, - ): - return self.processor( - self, - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - **kwargs, - ) - - -class MochiAttnProcessor2_0: - """Attention processor used in Mochi.""" - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("MochiAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: "MochiAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: torch.Tensor, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - encoder_query = encoder_query.unflatten(2, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(2, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(2, (attn.heads, -1)) - - if attn.norm_added_q is not None: - encoder_query = attn.norm_added_q(encoder_query) - if attn.norm_added_k is not None: - encoder_key = attn.norm_added_k(encoder_key) - - if image_rotary_emb is not None: - - def apply_rotary_emb(x, freqs_cos, freqs_sin): - x_even = x[..., 0::2].float() - x_odd = x[..., 1::2].float() - - cos = (x_even * freqs_cos - x_odd * freqs_sin).to(x.dtype) - sin = (x_even * freqs_sin + x_odd * freqs_cos).to(x.dtype) - - return torch.stack([cos, sin], dim=-1).flatten(-2) - - query = apply_rotary_emb(query, *image_rotary_emb) - key = apply_rotary_emb(key, *image_rotary_emb) - - query, key, value = query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2) - encoder_query, encoder_key, encoder_value = ( - encoder_query.transpose(1, 2), - encoder_key.transpose(1, 2), - encoder_value.transpose(1, 2), - ) - - sequence_length = query.size(2) - encoder_sequence_length = encoder_query.size(2) - total_length = sequence_length + encoder_sequence_length - - batch_size, heads, _, dim = query.shape - attn_outputs = [] - for idx in range(batch_size): - mask = attention_mask[idx][None, :] - valid_prompt_token_indices = torch.nonzero(mask.flatten(), as_tuple=False).flatten() - - valid_encoder_query = encoder_query[idx : idx + 1, :, valid_prompt_token_indices, :] - valid_encoder_key = encoder_key[idx : idx + 1, :, valid_prompt_token_indices, :] - valid_encoder_value = encoder_value[idx : idx + 1, :, valid_prompt_token_indices, :] - - valid_query = torch.cat([query[idx : idx + 1], valid_encoder_query], dim=2) - valid_key = torch.cat([key[idx : idx + 1], valid_encoder_key], dim=2) - valid_value = torch.cat([value[idx : idx + 1], valid_encoder_value], dim=2) - - attn_output = F.scaled_dot_product_attention( - valid_query, valid_key, valid_value, dropout_p=0.0, is_causal=False - ) - valid_sequence_length = attn_output.size(2) - attn_output = F.pad(attn_output, (0, 0, 0, total_length - valid_sequence_length)) - attn_outputs.append(attn_output) - - hidden_states = torch.cat(attn_outputs, dim=0) - hidden_states = hidden_states.transpose(1, 2).flatten(2, 3) - - hidden_states, encoder_hidden_states = hidden_states.split_with_sizes( - (sequence_length, encoder_sequence_length), dim=1 - ) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if hasattr(attn, "to_add_out"): - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - return hidden_states, encoder_hidden_states - - -class AttnProcessor: - r""" - Default processor for performing attention-related computations. - """ - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - residual = hidden_states - - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - query = attn.head_to_batch_dim(query) - key = attn.head_to_batch_dim(key) - value = attn.head_to_batch_dim(value) - - attention_probs = attn.get_attention_scores(query, key, attention_mask) - hidden_states = torch.bmm(attention_probs, value) - hidden_states = attn.batch_to_head_dim(hidden_states) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class CustomDiffusionAttnProcessor(nn.Module): - r""" - Processor for implementing attention for the Custom Diffusion method. - - Args: - train_kv (`bool`, defaults to `True`): - Whether to newly train the key and value matrices corresponding to the text features. - train_q_out (`bool`, defaults to `True`): - Whether to newly train query matrices corresponding to the latent image features. - hidden_size (`int`, *optional*, defaults to `None`): - The hidden size of the attention layer. - cross_attention_dim (`int`, *optional*, defaults to `None`): - The number of channels in the `encoder_hidden_states`. - out_bias (`bool`, defaults to `True`): - Whether to include the bias parameter in `train_q_out`. - dropout (`float`, *optional*, defaults to 0.0): - The dropout probability to use. - """ - - def __init__( - self, - train_kv: bool = True, - train_q_out: bool = True, - hidden_size: int | None = None, - cross_attention_dim: int | None = None, - out_bias: bool = True, - dropout: float = 0.0, - ): - super().__init__() - self.train_kv = train_kv - self.train_q_out = train_q_out - - self.hidden_size = hidden_size - self.cross_attention_dim = cross_attention_dim - - # `_custom_diffusion` id for easy serialization and loading. - if self.train_kv: - self.to_k_custom_diffusion = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) - self.to_v_custom_diffusion = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) - if self.train_q_out: - self.to_q_custom_diffusion = nn.Linear(hidden_size, hidden_size, bias=False) - self.to_out_custom_diffusion = nn.ModuleList([]) - self.to_out_custom_diffusion.append(nn.Linear(hidden_size, hidden_size, bias=out_bias)) - self.to_out_custom_diffusion.append(nn.Dropout(dropout)) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - batch_size, sequence_length, _ = hidden_states.shape - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - if self.train_q_out: - query = self.to_q_custom_diffusion(hidden_states).to(attn.to_q.weight.dtype) - else: - query = attn.to_q(hidden_states.to(attn.to_q.weight.dtype)) - - if encoder_hidden_states is None: - crossattn = False - encoder_hidden_states = hidden_states - else: - crossattn = True - if attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - if self.train_kv: - key = self.to_k_custom_diffusion(encoder_hidden_states.to(self.to_k_custom_diffusion.weight.dtype)) - value = self.to_v_custom_diffusion(encoder_hidden_states.to(self.to_v_custom_diffusion.weight.dtype)) - key = key.to(attn.to_q.weight.dtype) - value = value.to(attn.to_q.weight.dtype) - else: - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if crossattn: - detach = torch.ones_like(key) - detach[:, :1, :] = detach[:, :1, :] * 0.0 - key = detach * key + (1 - detach) * key.detach() - value = detach * value + (1 - detach) * value.detach() - - query = attn.head_to_batch_dim(query) - key = attn.head_to_batch_dim(key) - value = attn.head_to_batch_dim(value) - - attention_probs = attn.get_attention_scores(query, key, attention_mask) - hidden_states = torch.bmm(attention_probs, value) - hidden_states = attn.batch_to_head_dim(hidden_states) - - if self.train_q_out: - # linear proj - hidden_states = self.to_out_custom_diffusion[0](hidden_states) - # dropout - hidden_states = self.to_out_custom_diffusion[1](hidden_states) - else: - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - return hidden_states - - -class AttnAddedKVProcessor: - r""" - Processor for performing attention-related computations with extra learnable key and value matrices for the text - encoder. - """ - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - residual = hidden_states - - hidden_states = hidden_states.view(hidden_states.shape[0], hidden_states.shape[1], -1).transpose(1, 2) - batch_size, sequence_length, _ = hidden_states.shape - - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - query = attn.head_to_batch_dim(query) - - encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) - encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) - encoder_hidden_states_key_proj = attn.head_to_batch_dim(encoder_hidden_states_key_proj) - encoder_hidden_states_value_proj = attn.head_to_batch_dim(encoder_hidden_states_value_proj) - - if not attn.only_cross_attention: - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - key = attn.head_to_batch_dim(key) - value = attn.head_to_batch_dim(value) - key = torch.cat([encoder_hidden_states_key_proj, key], dim=1) - value = torch.cat([encoder_hidden_states_value_proj, value], dim=1) - else: - key = encoder_hidden_states_key_proj - value = encoder_hidden_states_value_proj - - attention_probs = attn.get_attention_scores(query, key, attention_mask) - hidden_states = torch.bmm(attention_probs, value) - hidden_states = attn.batch_to_head_dim(hidden_states) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - hidden_states = hidden_states.transpose(-1, -2).reshape(residual.shape) - hidden_states = hidden_states + residual - - return hidden_states - - -class AttnAddedKVProcessor2_0: - r""" - Processor for performing scaled dot-product attention (enabled by default if you're using PyTorch 2.0), with extra - learnable key and value matrices for the text encoder. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "AttnAddedKVProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - residual = hidden_states - - hidden_states = hidden_states.view(hidden_states.shape[0], hidden_states.shape[1], -1).transpose(1, 2) - batch_size, sequence_length, _ = hidden_states.shape - - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size, out_dim=4) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - query = attn.head_to_batch_dim(query, out_dim=4) - - encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) - encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) - encoder_hidden_states_key_proj = attn.head_to_batch_dim(encoder_hidden_states_key_proj, out_dim=4) - encoder_hidden_states_value_proj = attn.head_to_batch_dim(encoder_hidden_states_value_proj, out_dim=4) - - if not attn.only_cross_attention: - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - key = attn.head_to_batch_dim(key, out_dim=4) - value = attn.head_to_batch_dim(value, out_dim=4) - key = torch.cat([encoder_hidden_states_key_proj, key], dim=2) - value = torch.cat([encoder_hidden_states_value_proj, value], dim=2) - else: - key = encoder_hidden_states_key_proj - value = encoder_hidden_states_value_proj - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, residual.shape[1]) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - hidden_states = hidden_states.transpose(-1, -2).reshape(residual.shape) - hidden_states = hidden_states + residual - - return hidden_states - - -class JointAttnProcessor2_0: - """Attention processor used typically in processing the SD3-like self-attention projections.""" - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("JointAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor = None, - attention_mask: torch.FloatTensor | None = None, - *args, - **kwargs, - ) -> torch.FloatTensor: - residual = hidden_states - - batch_size = hidden_states.shape[0] - - # `sample` projections. - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # `context` projections. - if encoder_hidden_states is not None: - encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states) - encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) - encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) - - encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view( - batch_size, -1, attn.heads, head_dim - ).transpose(1, 2) - encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view( - batch_size, -1, attn.heads, head_dim - ).transpose(1, 2) - encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view( - batch_size, -1, attn.heads, head_dim - ).transpose(1, 2) - - if attn.norm_added_q is not None: - encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj) - if attn.norm_added_k is not None: - encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj) - - query = torch.cat([query, encoder_hidden_states_query_proj], dim=2) - key = torch.cat([key, encoder_hidden_states_key_proj], dim=2) - value = torch.cat([value, encoder_hidden_states_value_proj], dim=2) - - hidden_states = F.scaled_dot_product_attention(query, key, value, dropout_p=0.0, is_causal=False) - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - if encoder_hidden_states is not None: - # Split the attention outputs. - hidden_states, encoder_hidden_states = ( - hidden_states[:, : residual.shape[1]], - hidden_states[:, residual.shape[1] :], - ) - if not attn.context_pre_only: - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if encoder_hidden_states is not None: - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -class PAGJointAttnProcessor2_0: - """Attention processor used typically in processing the SD3-like self-attention projections.""" - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "PAGJointAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor = None, - attention_mask: torch.FloatTensor | None = None, - ) -> torch.FloatTensor: - residual = hidden_states - - input_ndim = hidden_states.ndim - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - context_input_ndim = encoder_hidden_states.ndim - if context_input_ndim == 4: - batch_size, channel, height, width = encoder_hidden_states.shape - encoder_hidden_states = encoder_hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - # store the length of image patch sequences to create a mask that prevents interaction between patches - # similar to making the self-attention map an identity matrix - identity_block_size = hidden_states.shape[1] - - # chunk - hidden_states_org, hidden_states_ptb = hidden_states.chunk(2) - encoder_hidden_states_org, encoder_hidden_states_ptb = encoder_hidden_states.chunk(2) - - ################## original path ################## - batch_size = encoder_hidden_states_org.shape[0] - - # `sample` projections. - query_org = attn.to_q(hidden_states_org) - key_org = attn.to_k(hidden_states_org) - value_org = attn.to_v(hidden_states_org) - - # `context` projections. - encoder_hidden_states_org_query_proj = attn.add_q_proj(encoder_hidden_states_org) - encoder_hidden_states_org_key_proj = attn.add_k_proj(encoder_hidden_states_org) - encoder_hidden_states_org_value_proj = attn.add_v_proj(encoder_hidden_states_org) - - # attention - query_org = torch.cat([query_org, encoder_hidden_states_org_query_proj], dim=1) - key_org = torch.cat([key_org, encoder_hidden_states_org_key_proj], dim=1) - value_org = torch.cat([value_org, encoder_hidden_states_org_value_proj], dim=1) - - inner_dim = key_org.shape[-1] - head_dim = inner_dim // attn.heads - query_org = query_org.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key_org = key_org.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value_org = value_org.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - hidden_states_org = F.scaled_dot_product_attention( - query_org, key_org, value_org, dropout_p=0.0, is_causal=False - ) - hidden_states_org = hidden_states_org.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states_org = hidden_states_org.to(query_org.dtype) - - # Split the attention outputs. - hidden_states_org, encoder_hidden_states_org = ( - hidden_states_org[:, : residual.shape[1]], - hidden_states_org[:, residual.shape[1] :], - ) - - # linear proj - hidden_states_org = attn.to_out[0](hidden_states_org) - # dropout - hidden_states_org = attn.to_out[1](hidden_states_org) - if not attn.context_pre_only: - encoder_hidden_states_org = attn.to_add_out(encoder_hidden_states_org) - - if input_ndim == 4: - hidden_states_org = hidden_states_org.transpose(-1, -2).reshape(batch_size, channel, height, width) - if context_input_ndim == 4: - encoder_hidden_states_org = encoder_hidden_states_org.transpose(-1, -2).reshape( - batch_size, channel, height, width - ) - - ################## perturbed path ################## - - batch_size = encoder_hidden_states_ptb.shape[0] - - # `sample` projections. - query_ptb = attn.to_q(hidden_states_ptb) - key_ptb = attn.to_k(hidden_states_ptb) - value_ptb = attn.to_v(hidden_states_ptb) - - # `context` projections. - encoder_hidden_states_ptb_query_proj = attn.add_q_proj(encoder_hidden_states_ptb) - encoder_hidden_states_ptb_key_proj = attn.add_k_proj(encoder_hidden_states_ptb) - encoder_hidden_states_ptb_value_proj = attn.add_v_proj(encoder_hidden_states_ptb) - - # attention - query_ptb = torch.cat([query_ptb, encoder_hidden_states_ptb_query_proj], dim=1) - key_ptb = torch.cat([key_ptb, encoder_hidden_states_ptb_key_proj], dim=1) - value_ptb = torch.cat([value_ptb, encoder_hidden_states_ptb_value_proj], dim=1) - - inner_dim = key_ptb.shape[-1] - head_dim = inner_dim // attn.heads - query_ptb = query_ptb.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key_ptb = key_ptb.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value_ptb = value_ptb.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # create a full mask with all entries set to 0 - seq_len = query_ptb.size(2) - full_mask = torch.zeros((seq_len, seq_len), device=query_ptb.device, dtype=query_ptb.dtype) - - # set the attention value between image patches to -inf - full_mask[:identity_block_size, :identity_block_size] = float("-inf") - - # set the diagonal of the attention value between image patches to 0 - full_mask[:identity_block_size, :identity_block_size].fill_diagonal_(0) - - # expand the mask to match the attention weights shape - full_mask = full_mask.unsqueeze(0).unsqueeze(0) # Add batch and num_heads dimensions - - hidden_states_ptb = F.scaled_dot_product_attention( - query_ptb, key_ptb, value_ptb, attn_mask=full_mask, dropout_p=0.0, is_causal=False - ) - hidden_states_ptb = hidden_states_ptb.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states_ptb = hidden_states_ptb.to(query_ptb.dtype) - - # split the attention outputs. - hidden_states_ptb, encoder_hidden_states_ptb = ( - hidden_states_ptb[:, : residual.shape[1]], - hidden_states_ptb[:, residual.shape[1] :], - ) - - # linear proj - hidden_states_ptb = attn.to_out[0](hidden_states_ptb) - # dropout - hidden_states_ptb = attn.to_out[1](hidden_states_ptb) - if not attn.context_pre_only: - encoder_hidden_states_ptb = attn.to_add_out(encoder_hidden_states_ptb) - - if input_ndim == 4: - hidden_states_ptb = hidden_states_ptb.transpose(-1, -2).reshape(batch_size, channel, height, width) - if context_input_ndim == 4: - encoder_hidden_states_ptb = encoder_hidden_states_ptb.transpose(-1, -2).reshape( - batch_size, channel, height, width - ) - - ################ concat ############### - hidden_states = torch.cat([hidden_states_org, hidden_states_ptb]) - encoder_hidden_states = torch.cat([encoder_hidden_states_org, encoder_hidden_states_ptb]) - - return hidden_states, encoder_hidden_states - - -class PAGCFGJointAttnProcessor2_0: - """Attention processor used typically in processing the SD3-like self-attention projections.""" - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "PAGCFGJointAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor = None, - attention_mask: torch.FloatTensor | None = None, - *args, - **kwargs, - ) -> torch.FloatTensor: - residual = hidden_states - - input_ndim = hidden_states.ndim - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - context_input_ndim = encoder_hidden_states.ndim - if context_input_ndim == 4: - batch_size, channel, height, width = encoder_hidden_states.shape - encoder_hidden_states = encoder_hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - identity_block_size = hidden_states.shape[ - 1 - ] # patch embeddings width * height (correspond to self-attention map width or height) - - # chunk - hidden_states_uncond, hidden_states_org, hidden_states_ptb = hidden_states.chunk(3) - hidden_states_org = torch.cat([hidden_states_uncond, hidden_states_org]) - - ( - encoder_hidden_states_uncond, - encoder_hidden_states_org, - encoder_hidden_states_ptb, - ) = encoder_hidden_states.chunk(3) - encoder_hidden_states_org = torch.cat([encoder_hidden_states_uncond, encoder_hidden_states_org]) - - ################## original path ################## - batch_size = encoder_hidden_states_org.shape[0] - - # `sample` projections. - query_org = attn.to_q(hidden_states_org) - key_org = attn.to_k(hidden_states_org) - value_org = attn.to_v(hidden_states_org) - - # `context` projections. - encoder_hidden_states_org_query_proj = attn.add_q_proj(encoder_hidden_states_org) - encoder_hidden_states_org_key_proj = attn.add_k_proj(encoder_hidden_states_org) - encoder_hidden_states_org_value_proj = attn.add_v_proj(encoder_hidden_states_org) - - # attention - query_org = torch.cat([query_org, encoder_hidden_states_org_query_proj], dim=1) - key_org = torch.cat([key_org, encoder_hidden_states_org_key_proj], dim=1) - value_org = torch.cat([value_org, encoder_hidden_states_org_value_proj], dim=1) - - inner_dim = key_org.shape[-1] - head_dim = inner_dim // attn.heads - query_org = query_org.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key_org = key_org.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value_org = value_org.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - hidden_states_org = F.scaled_dot_product_attention( - query_org, key_org, value_org, dropout_p=0.0, is_causal=False - ) - hidden_states_org = hidden_states_org.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states_org = hidden_states_org.to(query_org.dtype) - - # Split the attention outputs. - hidden_states_org, encoder_hidden_states_org = ( - hidden_states_org[:, : residual.shape[1]], - hidden_states_org[:, residual.shape[1] :], - ) - - # linear proj - hidden_states_org = attn.to_out[0](hidden_states_org) - # dropout - hidden_states_org = attn.to_out[1](hidden_states_org) - if not attn.context_pre_only: - encoder_hidden_states_org = attn.to_add_out(encoder_hidden_states_org) - - if input_ndim == 4: - hidden_states_org = hidden_states_org.transpose(-1, -2).reshape(batch_size, channel, height, width) - if context_input_ndim == 4: - encoder_hidden_states_org = encoder_hidden_states_org.transpose(-1, -2).reshape( - batch_size, channel, height, width - ) - - ################## perturbed path ################## - - batch_size = encoder_hidden_states_ptb.shape[0] - - # `sample` projections. - query_ptb = attn.to_q(hidden_states_ptb) - key_ptb = attn.to_k(hidden_states_ptb) - value_ptb = attn.to_v(hidden_states_ptb) - - # `context` projections. - encoder_hidden_states_ptb_query_proj = attn.add_q_proj(encoder_hidden_states_ptb) - encoder_hidden_states_ptb_key_proj = attn.add_k_proj(encoder_hidden_states_ptb) - encoder_hidden_states_ptb_value_proj = attn.add_v_proj(encoder_hidden_states_ptb) - - # attention - query_ptb = torch.cat([query_ptb, encoder_hidden_states_ptb_query_proj], dim=1) - key_ptb = torch.cat([key_ptb, encoder_hidden_states_ptb_key_proj], dim=1) - value_ptb = torch.cat([value_ptb, encoder_hidden_states_ptb_value_proj], dim=1) - - inner_dim = key_ptb.shape[-1] - head_dim = inner_dim // attn.heads - query_ptb = query_ptb.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key_ptb = key_ptb.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value_ptb = value_ptb.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # create a full mask with all entries set to 0 - seq_len = query_ptb.size(2) - full_mask = torch.zeros((seq_len, seq_len), device=query_ptb.device, dtype=query_ptb.dtype) - - # set the attention value between image patches to -inf - full_mask[:identity_block_size, :identity_block_size] = float("-inf") - - # set the diagonal of the attention value between image patches to 0 - full_mask[:identity_block_size, :identity_block_size].fill_diagonal_(0) - - # expand the mask to match the attention weights shape - full_mask = full_mask.unsqueeze(0).unsqueeze(0) # Add batch and num_heads dimensions - - hidden_states_ptb = F.scaled_dot_product_attention( - query_ptb, key_ptb, value_ptb, attn_mask=full_mask, dropout_p=0.0, is_causal=False - ) - hidden_states_ptb = hidden_states_ptb.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states_ptb = hidden_states_ptb.to(query_ptb.dtype) - - # split the attention outputs. - hidden_states_ptb, encoder_hidden_states_ptb = ( - hidden_states_ptb[:, : residual.shape[1]], - hidden_states_ptb[:, residual.shape[1] :], - ) - - # linear proj - hidden_states_ptb = attn.to_out[0](hidden_states_ptb) - # dropout - hidden_states_ptb = attn.to_out[1](hidden_states_ptb) - if not attn.context_pre_only: - encoder_hidden_states_ptb = attn.to_add_out(encoder_hidden_states_ptb) - - if input_ndim == 4: - hidden_states_ptb = hidden_states_ptb.transpose(-1, -2).reshape(batch_size, channel, height, width) - if context_input_ndim == 4: - encoder_hidden_states_ptb = encoder_hidden_states_ptb.transpose(-1, -2).reshape( - batch_size, channel, height, width - ) - - ################ concat ############### - hidden_states = torch.cat([hidden_states_org, hidden_states_ptb]) - encoder_hidden_states = torch.cat([encoder_hidden_states_org, encoder_hidden_states_ptb]) - - return hidden_states, encoder_hidden_states - - -class FusedJointAttnProcessor2_0: - """Attention processor used typically in processing the SD3-like self-attention projections.""" - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor = None, - attention_mask: torch.FloatTensor | None = None, - *args, - **kwargs, - ) -> torch.FloatTensor: - residual = hidden_states - - input_ndim = hidden_states.ndim - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - context_input_ndim = encoder_hidden_states.ndim - if context_input_ndim == 4: - batch_size, channel, height, width = encoder_hidden_states.shape - encoder_hidden_states = encoder_hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size = encoder_hidden_states.shape[0] - - # `sample` projections. - qkv = attn.to_qkv(hidden_states) - split_size = qkv.shape[-1] // 3 - query, key, value = torch.split(qkv, split_size, dim=-1) - - # `context` projections. - encoder_qkv = attn.to_added_qkv(encoder_hidden_states) - split_size = encoder_qkv.shape[-1] // 3 - ( - encoder_hidden_states_query_proj, - encoder_hidden_states_key_proj, - encoder_hidden_states_value_proj, - ) = torch.split(encoder_qkv, split_size, dim=-1) - - # attention - query = torch.cat([query, encoder_hidden_states_query_proj], dim=1) - key = torch.cat([key, encoder_hidden_states_key_proj], dim=1) - value = torch.cat([value, encoder_hidden_states_value_proj], dim=1) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - hidden_states = F.scaled_dot_product_attention(query, key, value, dropout_p=0.0, is_causal=False) - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # Split the attention outputs. - hidden_states, encoder_hidden_states = ( - hidden_states[:, : residual.shape[1]], - hidden_states[:, residual.shape[1] :], - ) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - if not attn.context_pre_only: - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - if context_input_ndim == 4: - encoder_hidden_states = encoder_hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - return hidden_states, encoder_hidden_states - - -class XFormersJointAttnProcessor: - r""" - Processor for implementing memory efficient attention using xFormers. - - Args: - attention_op (`Callable`, *optional*, defaults to `None`): - The base - [operator](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.AttentionOpBase) to - use as the attention operator. It is recommended to set to `None`, and allow xFormers to choose the best - operator. - """ - - def __init__(self, attention_op: Callable | None = None): - self.attention_op = attention_op - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor = None, - attention_mask: torch.FloatTensor | None = None, - *args, - **kwargs, - ) -> torch.FloatTensor: - residual = hidden_states - - # `sample` projections. - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = attn.head_to_batch_dim(query).contiguous() - key = attn.head_to_batch_dim(key).contiguous() - value = attn.head_to_batch_dim(value).contiguous() - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # `context` projections. - if encoder_hidden_states is not None: - encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states) - encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) - encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) - - encoder_hidden_states_query_proj = attn.head_to_batch_dim(encoder_hidden_states_query_proj).contiguous() - encoder_hidden_states_key_proj = attn.head_to_batch_dim(encoder_hidden_states_key_proj).contiguous() - encoder_hidden_states_value_proj = attn.head_to_batch_dim(encoder_hidden_states_value_proj).contiguous() - - if attn.norm_added_q is not None: - encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj) - if attn.norm_added_k is not None: - encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj) - - query = torch.cat([query, encoder_hidden_states_query_proj], dim=1) - key = torch.cat([key, encoder_hidden_states_key_proj], dim=1) - value = torch.cat([value, encoder_hidden_states_value_proj], dim=1) - - hidden_states = xformers.ops.memory_efficient_attention( - query, key, value, attn_bias=attention_mask, op=self.attention_op, scale=attn.scale - ) - hidden_states = hidden_states.to(query.dtype) - hidden_states = attn.batch_to_head_dim(hidden_states) - - if encoder_hidden_states is not None: - # Split the attention outputs. - hidden_states, encoder_hidden_states = ( - hidden_states[:, : residual.shape[1]], - hidden_states[:, residual.shape[1] :], - ) - if not attn.context_pre_only: - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if encoder_hidden_states is not None: - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -class AllegroAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is - used in the Allegro model. It applies a normalization layer and rotary embedding on the query and key vector. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "AllegroAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - residual = hidden_states - - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # Apply RoPE if needed - if image_rotary_emb is not None and not attn.is_cross_attention: - from .embeddings import apply_rotary_emb_allegro - - query = apply_rotary_emb_allegro(query, image_rotary_emb[0], image_rotary_emb[1]) - key = apply_rotary_emb_allegro(key, image_rotary_emb[0], image_rotary_emb[1]) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class AuraFlowAttnProcessor2_0: - """Attention processor used typically in processing Aura Flow.""" - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention") and is_torch_version("<", "2.1"): - raise ImportError( - "AuraFlowAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to at least 2.1 or above as we use `scale` in `F.scaled_dot_product_attention()`. " - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor = None, - *args, - **kwargs, - ) -> torch.FloatTensor: - batch_size = hidden_states.shape[0] - - # `sample` projections. - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - # `context` projections. - if encoder_hidden_states is not None: - encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states) - encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) - encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) - - # Reshape. - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - query = query.view(batch_size, -1, attn.heads, head_dim) - key = key.view(batch_size, -1, attn.heads, head_dim) - value = value.view(batch_size, -1, attn.heads, head_dim) - - # Apply QK norm. - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Concatenate the projections. - if encoder_hidden_states is not None: - encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view( - batch_size, -1, attn.heads, head_dim - ) - encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view(batch_size, -1, attn.heads, head_dim) - encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view( - batch_size, -1, attn.heads, head_dim - ) - - if attn.norm_added_q is not None: - encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj) - if attn.norm_added_k is not None: - encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj) - - query = torch.cat([encoder_hidden_states_query_proj, query], dim=1) - key = torch.cat([encoder_hidden_states_key_proj, key], dim=1) - value = torch.cat([encoder_hidden_states_value_proj, value], dim=1) - - query = query.transpose(1, 2) - key = key.transpose(1, 2) - value = value.transpose(1, 2) - - # Attention. - hidden_states = F.scaled_dot_product_attention( - query, key, value, dropout_p=0.0, scale=attn.scale, is_causal=False - ) - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # Split the attention outputs. - if encoder_hidden_states is not None: - hidden_states, encoder_hidden_states = ( - hidden_states[:, encoder_hidden_states.shape[1] :], - hidden_states[:, : encoder_hidden_states.shape[1]], - ) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - if encoder_hidden_states is not None: - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - if encoder_hidden_states is not None: - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -class FusedAuraFlowAttnProcessor2_0: - """Attention processor used typically in processing Aura Flow with fused projections.""" - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention") and is_torch_version("<", "2.1"): - raise ImportError( - "FusedAuraFlowAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to at least 2.1 or above as we use `scale` in `F.scaled_dot_product_attention()`. " - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor = None, - *args, - **kwargs, - ) -> torch.FloatTensor: - batch_size = hidden_states.shape[0] - - # `sample` projections. - qkv = attn.to_qkv(hidden_states) - split_size = qkv.shape[-1] // 3 - query, key, value = torch.split(qkv, split_size, dim=-1) - - # `context` projections. - if encoder_hidden_states is not None: - encoder_qkv = attn.to_added_qkv(encoder_hidden_states) - split_size = encoder_qkv.shape[-1] // 3 - ( - encoder_hidden_states_query_proj, - encoder_hidden_states_key_proj, - encoder_hidden_states_value_proj, - ) = torch.split(encoder_qkv, split_size, dim=-1) - - # Reshape. - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - query = query.view(batch_size, -1, attn.heads, head_dim) - key = key.view(batch_size, -1, attn.heads, head_dim) - value = value.view(batch_size, -1, attn.heads, head_dim) - - # Apply QK norm. - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Concatenate the projections. - if encoder_hidden_states is not None: - encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view( - batch_size, -1, attn.heads, head_dim - ) - encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view(batch_size, -1, attn.heads, head_dim) - encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view( - batch_size, -1, attn.heads, head_dim - ) - - if attn.norm_added_q is not None: - encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj) - if attn.norm_added_k is not None: - encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj) - - query = torch.cat([encoder_hidden_states_query_proj, query], dim=1) - key = torch.cat([encoder_hidden_states_key_proj, key], dim=1) - value = torch.cat([encoder_hidden_states_value_proj, value], dim=1) - - query = query.transpose(1, 2) - key = key.transpose(1, 2) - value = value.transpose(1, 2) - - # Attention. - hidden_states = F.scaled_dot_product_attention( - query, key, value, dropout_p=0.0, scale=attn.scale, is_causal=False - ) - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # Split the attention outputs. - if encoder_hidden_states is not None: - hidden_states, encoder_hidden_states = ( - hidden_states[:, encoder_hidden_states.shape[1] :], - hidden_states[:, : encoder_hidden_states.shape[1]], - ) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - if encoder_hidden_states is not None: - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - if encoder_hidden_states is not None: - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -class CogVideoXAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention for the CogVideoX model. It applies a rotary embedding on - query and key vectors, but does not include spatial normalization. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - text_seq_length = encoder_hidden_states.size(1) - - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - batch_size, sequence_length, _ = hidden_states.shape - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE if needed - if image_rotary_emb is not None: - from .embeddings import apply_rotary_emb - - query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb) - if not attn.is_cross_attention: - key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb) - - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - encoder_hidden_states, hidden_states = hidden_states.split( - [text_seq_length, hidden_states.size(1) - text_seq_length], dim=1 - ) - return hidden_states, encoder_hidden_states - - -class FusedCogVideoXAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention for the CogVideoX model. It applies a rotary embedding on - query and key vectors, but does not include spatial normalization. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - text_seq_length = encoder_hidden_states.size(1) - - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - qkv = attn.to_qkv(hidden_states) - split_size = qkv.shape[-1] // 3 - query, key, value = torch.split(qkv, split_size, dim=-1) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE if needed - if image_rotary_emb is not None: - from .embeddings import apply_rotary_emb - - query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb) - if not attn.is_cross_attention: - key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb) - - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - encoder_hidden_states, hidden_states = hidden_states.split( - [text_seq_length, hidden_states.size(1) - text_seq_length], dim=1 - ) - return hidden_states, encoder_hidden_states - - -class XFormersAttnAddedKVProcessor: - r""" - Processor for implementing memory efficient attention using xFormers. - - Args: - attention_op (`Callable`, *optional*, defaults to `None`): - The base - [operator](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.AttentionOpBase) to - use as the attention operator. It is recommended to set to `None`, and allow xFormers to choose the best - operator. - """ - - def __init__(self, attention_op: Callable | None = None): - self.attention_op = attention_op - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - residual = hidden_states - hidden_states = hidden_states.view(hidden_states.shape[0], hidden_states.shape[1], -1).transpose(1, 2) - batch_size, sequence_length, _ = hidden_states.shape - - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - query = attn.head_to_batch_dim(query) - - encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) - encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) - encoder_hidden_states_key_proj = attn.head_to_batch_dim(encoder_hidden_states_key_proj) - encoder_hidden_states_value_proj = attn.head_to_batch_dim(encoder_hidden_states_value_proj) - - if not attn.only_cross_attention: - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - key = attn.head_to_batch_dim(key) - value = attn.head_to_batch_dim(value) - key = torch.cat([encoder_hidden_states_key_proj, key], dim=1) - value = torch.cat([encoder_hidden_states_value_proj, value], dim=1) - else: - key = encoder_hidden_states_key_proj - value = encoder_hidden_states_value_proj - - hidden_states = xformers.ops.memory_efficient_attention( - query, key, value, attn_bias=attention_mask, op=self.attention_op, scale=attn.scale - ) - hidden_states = hidden_states.to(query.dtype) - hidden_states = attn.batch_to_head_dim(hidden_states) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - hidden_states = hidden_states.transpose(-1, -2).reshape(residual.shape) - hidden_states = hidden_states + residual - - return hidden_states - - -class XFormersAttnProcessor: - r""" - Processor for implementing memory efficient attention using xFormers. - - Args: - attention_op (`Callable`, *optional*, defaults to `None`): - The base - [operator](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.AttentionOpBase) to - use as the attention operator. It is recommended to set to `None`, and allow xFormers to choose the best - operator. - """ - - def __init__(self, attention_op: Callable | None = None): - self.attention_op = attention_op - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - residual = hidden_states - - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, key_tokens, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - attention_mask = attn.prepare_attention_mask(attention_mask, key_tokens, batch_size) - if attention_mask is not None: - # expand our mask's singleton query_tokens dimension: - # [batch*heads, 1, key_tokens] -> - # [batch*heads, query_tokens, key_tokens] - # so that it can be added as a bias onto the attention scores that xformers computes: - # [batch*heads, query_tokens, key_tokens] - # we do this explicitly because xformers doesn't broadcast the singleton dimension for us. - _, query_tokens, _ = hidden_states.shape - attention_mask = attention_mask.expand(-1, query_tokens, -1) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - query = attn.head_to_batch_dim(query).contiguous() - key = attn.head_to_batch_dim(key).contiguous() - value = attn.head_to_batch_dim(value).contiguous() - - hidden_states = xformers.ops.memory_efficient_attention( - query, key, value, attn_bias=attention_mask, op=self.attention_op, scale=attn.scale - ) - hidden_states = hidden_states.to(query.dtype) - hidden_states = attn.batch_to_head_dim(hidden_states) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class AttnProcessorNPU: - r""" - Processor for implementing flash attention using torch_npu. Torch_npu supports only fp16 and bf16 data types. If - fp32 is used, F.scaled_dot_product_attention will be used for computation, but the acceleration effect on NPU is - not significant. - - """ - - def __init__(self): - if not is_torch_npu_available(): - raise ImportError("AttnProcessorNPU requires torch_npu extensions and is supported only on npu devices.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - attention_mask = attention_mask.repeat(1, 1, hidden_states.shape[1], 1) - if attention_mask.dtype == torch.bool: - attention_mask = torch.logical_not(attention_mask.bool()) - else: - attention_mask = attention_mask.bool() - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - if query.dtype in (torch.float16, torch.bfloat16): - hidden_states = torch_npu.npu_fusion_attention( - query, - key, - value, - attn.heads, - input_layout="BNSD", - pse=None, - atten_mask=attention_mask, - scale=1.0 / math.sqrt(query.shape[-1]), - pre_tockens=65536, - next_tockens=65536, - keep_prob=1.0, - sync=False, - inner_precise=0, - )[0] - else: - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class AttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class XLAFlashAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention with pallas flash attention kernel if using `torch_xla`. - """ - - def __init__(self, partition_spec: tuple[str | None, ...] | None = None): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "XLAFlashAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - if is_torch_xla_version("<", "2.3"): - raise ImportError("XLA flash attention requires torch_xla version >= 2.3.") - if is_spmd() and is_torch_xla_version("<", "2.4"): - raise ImportError("SPMD support for XLA flash attention needs torch_xla version >= 2.4.") - self.partition_spec = partition_spec - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - if all(tensor.shape[2] >= 4096 for tensor in [query, key, value]): - if attention_mask is not None: - attention_mask = attention_mask.view(batch_size, 1, 1, attention_mask.shape[-1]) - # Convert mask to float and replace 0s with -inf and 1s with 0 - attention_mask = ( - attention_mask.float() - .masked_fill(attention_mask == 0, float("-inf")) - .masked_fill(attention_mask == 1, float(0.0)) - ) - - # Apply attention mask to key - key = key + attention_mask - query /= math.sqrt(query.shape[3]) - partition_spec = self.partition_spec if is_spmd() else None - hidden_states = flash_attention(query, key, value, causal=False, partition_spec=partition_spec) - else: - logger.warning( - "Unable to use the flash attention pallas kernel API call due to QKV sequence length < 4096." - ) - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class MochiVaeAttnProcessor2_0: - r""" - Attention processor used in Mochi VAE. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - residual = hidden_states - is_single_frame = hidden_states.shape[1] == 1 - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if is_single_frame: - hidden_states = attn.to_v(hidden_states) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - return hidden_states - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=attn.is_causal - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class StableAudioAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is - used in the Stable Audio model. It applies rotary embedding on query and key vector, and allows MHA, GQA or MQA. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "StableAudioAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def apply_partial_rotary_emb( - self, - x: torch.Tensor, - freqs_cis: tuple[torch.Tensor], - ) -> torch.Tensor: - from .embeddings import apply_rotary_emb - - rot_dim = freqs_cis[0].shape[-1] - x_to_rotate, x_unrotated = x[..., :rot_dim], x[..., rot_dim:] - - x_rotated = apply_rotary_emb(x_to_rotate, freqs_cis, use_real=True, use_real_unbind_dim=-2) - - out = torch.cat((x_rotated, x_unrotated), dim=-1) - return out - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - from .embeddings import apply_rotary_emb - - residual = hidden_states - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - head_dim = query.shape[-1] // attn.heads - kv_heads = key.shape[-1] // head_dim - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, kv_heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, kv_heads, head_dim).transpose(1, 2) - - if kv_heads != attn.heads: - # if GQA or MQA, repeat the key/value heads to reach the number of query heads. - heads_per_kv_head = attn.heads // kv_heads - key = torch.repeat_interleave(key, heads_per_kv_head, dim=1, output_size=key.shape[1] * heads_per_kv_head) - value = torch.repeat_interleave( - value, heads_per_kv_head, dim=1, output_size=value.shape[1] * heads_per_kv_head - ) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE if needed - if rotary_emb is not None: - query_dtype = query.dtype - key_dtype = key.dtype - query = query.to(torch.float32) - key = key.to(torch.float32) - - rot_dim = rotary_emb[0].shape[-1] - query_to_rotate, query_unrotated = query[..., :rot_dim], query[..., rot_dim:] - query_rotated = apply_rotary_emb(query_to_rotate, rotary_emb, use_real=True, use_real_unbind_dim=-2) - - query = torch.cat((query_rotated, query_unrotated), dim=-1) - - if not attn.is_cross_attention: - key_to_rotate, key_unrotated = key[..., :rot_dim], key[..., rot_dim:] - key_rotated = apply_rotary_emb(key_to_rotate, rotary_emb, use_real=True, use_real_unbind_dim=-2) - - key = torch.cat((key_rotated, key_unrotated), dim=-1) - - query = query.to(query_dtype) - key = key.to(key_dtype) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class HunyuanAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is - used in the HunyuanDiT model. It applies a s normalization layer and rotary embedding on query and key vector. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - from .embeddings import apply_rotary_emb - - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE if needed - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb) - if not attn.is_cross_attention: - key = apply_rotary_emb(key, image_rotary_emb) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class FusedHunyuanAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0) with fused - projection layers. This is used in the HunyuanDiT model. It applies a s normalization layer and rotary embedding on - query and key vector. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "FusedHunyuanAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - from .embeddings import apply_rotary_emb - - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - if encoder_hidden_states is None: - qkv = attn.to_qkv(hidden_states) - split_size = qkv.shape[-1] // 3 - query, key, value = torch.split(qkv, split_size, dim=-1) - else: - if attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - query = attn.to_q(hidden_states) - - kv = attn.to_kv(encoder_hidden_states) - split_size = kv.shape[-1] // 2 - key, value = torch.split(kv, split_size, dim=-1) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE if needed - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb) - if not attn.is_cross_attention: - key = apply_rotary_emb(key, image_rotary_emb) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class PAGHunyuanAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is - used in the HunyuanDiT model. It applies a normalization layer and rotary embedding on query and key vector. This - variant of the processor employs [Pertubed Attention Guidance](https://huggingface.co/papers/2403.17377). - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "PAGHunyuanAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - from .embeddings import apply_rotary_emb - - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - # chunk - hidden_states_org, hidden_states_ptb = hidden_states.chunk(2) - - # 1. Original Path - batch_size, sequence_length, _ = ( - hidden_states_org.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states_org = attn.group_norm(hidden_states_org.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states_org) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states_org - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE if needed - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb) - if not attn.is_cross_attention: - key = apply_rotary_emb(key, image_rotary_emb) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states_org = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states_org = hidden_states_org.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states_org = hidden_states_org.to(query.dtype) - - # linear proj - hidden_states_org = attn.to_out[0](hidden_states_org) - # dropout - hidden_states_org = attn.to_out[1](hidden_states_org) - - if input_ndim == 4: - hidden_states_org = hidden_states_org.transpose(-1, -2).reshape(batch_size, channel, height, width) - - # 2. Perturbed Path - if attn.group_norm is not None: - hidden_states_ptb = attn.group_norm(hidden_states_ptb.transpose(1, 2)).transpose(1, 2) - - hidden_states_ptb = attn.to_v(hidden_states_ptb) - hidden_states_ptb = hidden_states_ptb.to(query.dtype) - - # linear proj - hidden_states_ptb = attn.to_out[0](hidden_states_ptb) - # dropout - hidden_states_ptb = attn.to_out[1](hidden_states_ptb) - - if input_ndim == 4: - hidden_states_ptb = hidden_states_ptb.transpose(-1, -2).reshape(batch_size, channel, height, width) - - # cat - hidden_states = torch.cat([hidden_states_org, hidden_states_ptb]) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class PAGCFGHunyuanAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is - used in the HunyuanDiT model. It applies a normalization layer and rotary embedding on query and key vector. This - variant of the processor employs [Pertubed Attention Guidance](https://huggingface.co/papers/2403.17377). - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "PAGCFGHunyuanAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - from .embeddings import apply_rotary_emb - - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - # chunk - hidden_states_uncond, hidden_states_org, hidden_states_ptb = hidden_states.chunk(3) - hidden_states_org = torch.cat([hidden_states_uncond, hidden_states_org]) - - # 1. Original Path - batch_size, sequence_length, _ = ( - hidden_states_org.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states_org = attn.group_norm(hidden_states_org.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states_org) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states_org - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE if needed - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb) - if not attn.is_cross_attention: - key = apply_rotary_emb(key, image_rotary_emb) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states_org = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states_org = hidden_states_org.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states_org = hidden_states_org.to(query.dtype) - - # linear proj - hidden_states_org = attn.to_out[0](hidden_states_org) - # dropout - hidden_states_org = attn.to_out[1](hidden_states_org) - - if input_ndim == 4: - hidden_states_org = hidden_states_org.transpose(-1, -2).reshape(batch_size, channel, height, width) - - # 2. Perturbed Path - if attn.group_norm is not None: - hidden_states_ptb = attn.group_norm(hidden_states_ptb.transpose(1, 2)).transpose(1, 2) - - hidden_states_ptb = attn.to_v(hidden_states_ptb) - hidden_states_ptb = hidden_states_ptb.to(query.dtype) - - # linear proj - hidden_states_ptb = attn.to_out[0](hidden_states_ptb) - # dropout - hidden_states_ptb = attn.to_out[1](hidden_states_ptb) - - if input_ndim == 4: - hidden_states_ptb = hidden_states_ptb.transpose(-1, -2).reshape(batch_size, channel, height, width) - - # cat - hidden_states = torch.cat([hidden_states_org, hidden_states_ptb]) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class LuminaAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is - used in the LuminaNextDiT model. It applies a s normalization layer and rotary embedding on query and key vector. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - query_rotary_emb: torch.Tensor | None = None, - key_rotary_emb: torch.Tensor | None = None, - base_sequence_length: int | None = None, - ) -> torch.Tensor: - from .embeddings import apply_rotary_emb - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = hidden_states.shape - - # Get Query-Key-Value Pair - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - query_dim = query.shape[-1] - inner_dim = key.shape[-1] - head_dim = query_dim // attn.heads - dtype = query.dtype - - # Get key-value heads - kv_heads = inner_dim // head_dim - - # Apply Query-Key Norm if needed - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - query = query.view(batch_size, -1, attn.heads, head_dim) - - key = key.view(batch_size, -1, kv_heads, head_dim) - value = value.view(batch_size, -1, kv_heads, head_dim) - - # Apply RoPE if needed - if query_rotary_emb is not None: - query = apply_rotary_emb(query, query_rotary_emb, use_real=False) - if key_rotary_emb is not None: - key = apply_rotary_emb(key, key_rotary_emb, use_real=False) - - query, key = query.to(dtype), key.to(dtype) - - # Apply proportional attention if true - if key_rotary_emb is None: - softmax_scale = None - else: - if base_sequence_length is not None: - softmax_scale = math.sqrt(math.log(sequence_length, base_sequence_length)) * attn.scale - else: - softmax_scale = attn.scale - - # perform Grouped-qurey Attention (GQA) - n_rep = attn.heads // kv_heads - if n_rep >= 1: - key = key.unsqueeze(3).repeat(1, 1, 1, n_rep, 1).flatten(2, 3) - value = value.unsqueeze(3).repeat(1, 1, 1, n_rep, 1).flatten(2, 3) - - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.bool().view(batch_size, 1, 1, -1) - attention_mask = attention_mask.expand(-1, attn.heads, sequence_length, -1) - - query = query.transpose(1, 2) - key = key.transpose(1, 2) - value = value.transpose(1, 2) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, scale=softmax_scale - ) - hidden_states = hidden_states.transpose(1, 2).to(dtype) - - return hidden_states - - -class FusedAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). It uses - fused projection layers. For self-attention modules, all projection matrices (i.e., query, key, value) are fused. - For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is currently 🧪 experimental in nature and can change in future. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "FusedAttnProcessor2_0 requires at least PyTorch 2.0, to use it. Please upgrade PyTorch to > 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - if encoder_hidden_states is None: - qkv = attn.to_qkv(hidden_states) - split_size = qkv.shape[-1] // 3 - query, key, value = torch.split(qkv, split_size, dim=-1) - else: - if attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - query = attn.to_q(hidden_states) - - kv = attn.to_kv(encoder_hidden_states) - split_size = kv.shape[-1] // 2 - key, value = torch.split(kv, split_size, dim=-1) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class CustomDiffusionXFormersAttnProcessor(nn.Module): - r""" - Processor for implementing memory efficient attention using xFormers for the Custom Diffusion method. - - Args: - train_kv (`bool`, defaults to `True`): - Whether to newly train the key and value matrices corresponding to the text features. - train_q_out (`bool`, defaults to `True`): - Whether to newly train query matrices corresponding to the latent image features. - hidden_size (`int`, *optional*, defaults to `None`): - The hidden size of the attention layer. - cross_attention_dim (`int`, *optional*, defaults to `None`): - The number of channels in the `encoder_hidden_states`. - out_bias (`bool`, defaults to `True`): - Whether to include the bias parameter in `train_q_out`. - dropout (`float`, *optional*, defaults to 0.0): - The dropout probability to use. - attention_op (`Callable`, *optional*, defaults to `None`): - The base - [operator](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.AttentionOpBase) to use - as the attention operator. It is recommended to set to `None`, and allow xFormers to choose the best operator. - """ - - def __init__( - self, - train_kv: bool = True, - train_q_out: bool = False, - hidden_size: int | None = None, - cross_attention_dim: int | None = None, - out_bias: bool = True, - dropout: float = 0.0, - attention_op: Callable | None = None, - ): - super().__init__() - self.train_kv = train_kv - self.train_q_out = train_q_out - - self.hidden_size = hidden_size - self.cross_attention_dim = cross_attention_dim - self.attention_op = attention_op - - # `_custom_diffusion` id for easy serialization and loading. - if self.train_kv: - self.to_k_custom_diffusion = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) - self.to_v_custom_diffusion = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) - if self.train_q_out: - self.to_q_custom_diffusion = nn.Linear(hidden_size, hidden_size, bias=False) - self.to_out_custom_diffusion = nn.ModuleList([]) - self.to_out_custom_diffusion.append(nn.Linear(hidden_size, hidden_size, bias=out_bias)) - self.to_out_custom_diffusion.append(nn.Dropout(dropout)) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - - if self.train_q_out: - query = self.to_q_custom_diffusion(hidden_states).to(attn.to_q.weight.dtype) - else: - query = attn.to_q(hidden_states.to(attn.to_q.weight.dtype)) - - if encoder_hidden_states is None: - crossattn = False - encoder_hidden_states = hidden_states - else: - crossattn = True - if attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - if self.train_kv: - key = self.to_k_custom_diffusion(encoder_hidden_states.to(self.to_k_custom_diffusion.weight.dtype)) - value = self.to_v_custom_diffusion(encoder_hidden_states.to(self.to_v_custom_diffusion.weight.dtype)) - key = key.to(attn.to_q.weight.dtype) - value = value.to(attn.to_q.weight.dtype) - else: - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if crossattn: - detach = torch.ones_like(key) - detach[:, :1, :] = detach[:, :1, :] * 0.0 - key = detach * key + (1 - detach) * key.detach() - value = detach * value + (1 - detach) * value.detach() - - query = attn.head_to_batch_dim(query).contiguous() - key = attn.head_to_batch_dim(key).contiguous() - value = attn.head_to_batch_dim(value).contiguous() - - hidden_states = xformers.ops.memory_efficient_attention( - query, key, value, attn_bias=attention_mask, op=self.attention_op, scale=attn.scale - ) - hidden_states = hidden_states.to(query.dtype) - hidden_states = attn.batch_to_head_dim(hidden_states) - - if self.train_q_out: - # linear proj - hidden_states = self.to_out_custom_diffusion[0](hidden_states) - # dropout - hidden_states = self.to_out_custom_diffusion[1](hidden_states) - else: - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - return hidden_states - - -class CustomDiffusionAttnProcessor2_0(nn.Module): - r""" - Processor for implementing attention for the Custom Diffusion method using PyTorch 2.0’s memory-efficient scaled - dot-product attention. - - Args: - train_kv (`bool`, defaults to `True`): - Whether to newly train the key and value matrices corresponding to the text features. - train_q_out (`bool`, defaults to `True`): - Whether to newly train query matrices corresponding to the latent image features. - hidden_size (`int`, *optional*, defaults to `None`): - The hidden size of the attention layer. - cross_attention_dim (`int`, *optional*, defaults to `None`): - The number of channels in the `encoder_hidden_states`. - out_bias (`bool`, defaults to `True`): - Whether to include the bias parameter in `train_q_out`. - dropout (`float`, *optional*, defaults to 0.0): - The dropout probability to use. - """ - - def __init__( - self, - train_kv: bool = True, - train_q_out: bool = True, - hidden_size: int | None = None, - cross_attention_dim: int | None = None, - out_bias: bool = True, - dropout: float = 0.0, - ): - super().__init__() - self.train_kv = train_kv - self.train_q_out = train_q_out - - self.hidden_size = hidden_size - self.cross_attention_dim = cross_attention_dim - - # `_custom_diffusion` id for easy serialization and loading. - if self.train_kv: - self.to_k_custom_diffusion = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) - self.to_v_custom_diffusion = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) - if self.train_q_out: - self.to_q_custom_diffusion = nn.Linear(hidden_size, hidden_size, bias=False) - self.to_out_custom_diffusion = nn.ModuleList([]) - self.to_out_custom_diffusion.append(nn.Linear(hidden_size, hidden_size, bias=out_bias)) - self.to_out_custom_diffusion.append(nn.Dropout(dropout)) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - batch_size, sequence_length, _ = hidden_states.shape - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - if self.train_q_out: - query = self.to_q_custom_diffusion(hidden_states) - else: - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - crossattn = False - encoder_hidden_states = hidden_states - else: - crossattn = True - if attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - if self.train_kv: - key = self.to_k_custom_diffusion(encoder_hidden_states.to(self.to_k_custom_diffusion.weight.dtype)) - value = self.to_v_custom_diffusion(encoder_hidden_states.to(self.to_v_custom_diffusion.weight.dtype)) - key = key.to(attn.to_q.weight.dtype) - value = value.to(attn.to_q.weight.dtype) - - else: - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if crossattn: - detach = torch.ones_like(key) - detach[:, :1, :] = detach[:, :1, :] * 0.0 - key = detach * key + (1 - detach) * key.detach() - value = detach * value + (1 - detach) * value.detach() - - inner_dim = hidden_states.shape[-1] - - head_dim = inner_dim // attn.heads - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - if self.train_q_out: - # linear proj - hidden_states = self.to_out_custom_diffusion[0](hidden_states) - # dropout - hidden_states = self.to_out_custom_diffusion[1](hidden_states) - else: - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - return hidden_states - - -class SlicedAttnProcessor: - r""" - Processor for implementing sliced attention. - - Args: - slice_size (`int`, *optional*): - The number of steps to compute attention. Uses as many slices as `attention_head_dim // slice_size`, and - `attention_head_dim` must be a multiple of the `slice_size`. - """ - - def __init__(self, slice_size: int): - self.slice_size = slice_size - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - residual = hidden_states - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - dim = query.shape[-1] - query = attn.head_to_batch_dim(query) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - key = attn.head_to_batch_dim(key) - value = attn.head_to_batch_dim(value) - - batch_size_attention, query_tokens, _ = query.shape - hidden_states = torch.zeros( - (batch_size_attention, query_tokens, dim // attn.heads), device=query.device, dtype=query.dtype - ) - - for i in range((batch_size_attention - 1) // self.slice_size + 1): - start_idx = i * self.slice_size - end_idx = (i + 1) * self.slice_size - - query_slice = query[start_idx:end_idx] - key_slice = key[start_idx:end_idx] - attn_mask_slice = attention_mask[start_idx:end_idx] if attention_mask is not None else None - - attn_slice = attn.get_attention_scores(query_slice, key_slice, attn_mask_slice) - - attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx]) - - hidden_states[start_idx:end_idx] = attn_slice - - hidden_states = attn.batch_to_head_dim(hidden_states) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class SlicedAttnAddedKVProcessor: - r""" - Processor for implementing sliced attention with extra learnable key and value matrices for the text encoder. - - Args: - slice_size (`int`, *optional*): - The number of steps to compute attention. Uses as many slices as `attention_head_dim // slice_size`, and - `attention_head_dim` must be a multiple of the `slice_size`. - """ - - def __init__(self, slice_size): - self.slice_size = slice_size - - def __call__( - self, - attn: "Attention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - ) -> torch.Tensor: - residual = hidden_states - - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - hidden_states = hidden_states.view(hidden_states.shape[0], hidden_states.shape[1], -1).transpose(1, 2) - - batch_size, sequence_length, _ = hidden_states.shape - - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - dim = query.shape[-1] - query = attn.head_to_batch_dim(query) - - encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) - encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) - - encoder_hidden_states_key_proj = attn.head_to_batch_dim(encoder_hidden_states_key_proj) - encoder_hidden_states_value_proj = attn.head_to_batch_dim(encoder_hidden_states_value_proj) - - if not attn.only_cross_attention: - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - key = attn.head_to_batch_dim(key) - value = attn.head_to_batch_dim(value) - key = torch.cat([encoder_hidden_states_key_proj, key], dim=1) - value = torch.cat([encoder_hidden_states_value_proj, value], dim=1) - else: - key = encoder_hidden_states_key_proj - value = encoder_hidden_states_value_proj - - batch_size_attention, query_tokens, _ = query.shape - hidden_states = torch.zeros( - (batch_size_attention, query_tokens, dim // attn.heads), device=query.device, dtype=query.dtype - ) - - for i in range((batch_size_attention - 1) // self.slice_size + 1): - start_idx = i * self.slice_size - end_idx = (i + 1) * self.slice_size - - query_slice = query[start_idx:end_idx] - key_slice = key[start_idx:end_idx] - attn_mask_slice = attention_mask[start_idx:end_idx] if attention_mask is not None else None - - attn_slice = attn.get_attention_scores(query_slice, key_slice, attn_mask_slice) - - attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx]) - - hidden_states[start_idx:end_idx] = attn_slice - - hidden_states = attn.batch_to_head_dim(hidden_states) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - hidden_states = hidden_states.transpose(-1, -2).reshape(residual.shape) - hidden_states = hidden_states + residual - - return hidden_states - - -class SpatialNorm(nn.Module): - """ - Spatially conditioned normalization as defined in https://huggingface.co/papers/2209.09002. - - Args: - f_channels (`int`): - The number of channels for input to group normalization layer, and output of the spatial norm layer. - zq_channels (`int`): - The number of channels for the quantized vector as described in the paper. - """ - - def __init__( - self, - f_channels: int, - zq_channels: int, - ): - super().__init__() - self.norm_layer = nn.GroupNorm(num_channels=f_channels, num_groups=32, eps=1e-6, affine=True) - self.conv_y = nn.Conv2d(zq_channels, f_channels, kernel_size=1, stride=1, padding=0) - self.conv_b = nn.Conv2d(zq_channels, f_channels, kernel_size=1, stride=1, padding=0) - - def forward(self, f: torch.Tensor, zq: torch.Tensor) -> torch.Tensor: - f_size = f.shape[-2:] - zq = F.interpolate(zq, size=f_size, mode="nearest") - norm_f = self.norm_layer(f) - new_f = norm_f * self.conv_y(zq) + self.conv_b(zq) - return new_f - - -class IPAdapterAttnProcessor(nn.Module): - r""" - Attention processor for Multiple IP-Adapters. - - Args: - hidden_size (`int`): - The hidden size of the attention layer. - cross_attention_dim (`int`): - The number of channels in the `encoder_hidden_states`. - num_tokens (`int`, `tuple[int]` or `list[int]`, defaults to `(4,)`): - The context length of the image features. - scale (`float` or list[`float`], defaults to 1.0): - the weight scale of image prompt. - """ - - def __init__(self, hidden_size, cross_attention_dim=None, num_tokens=(4,), scale=1.0): - super().__init__() - - self.hidden_size = hidden_size - self.cross_attention_dim = cross_attention_dim - - if not isinstance(num_tokens, (tuple, list)): - num_tokens = [num_tokens] - self.num_tokens = num_tokens - - if not isinstance(scale, list): - scale = [scale] * len(num_tokens) - if len(scale) != len(num_tokens): - raise ValueError("`scale` should be a list of integers with the same length as `num_tokens`.") - self.scale = scale - - self.to_k_ip = nn.ModuleList( - [nn.Linear(cross_attention_dim, hidden_size, bias=False) for _ in range(len(num_tokens))] - ) - self.to_v_ip = nn.ModuleList( - [nn.Linear(cross_attention_dim, hidden_size, bias=False) for _ in range(len(num_tokens))] - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - scale: float = 1.0, - ip_adapter_masks: torch.Tensor | None = None, - ): - residual = hidden_states - - # separate ip_hidden_states from encoder_hidden_states - if encoder_hidden_states is not None: - if isinstance(encoder_hidden_states, tuple): - encoder_hidden_states, ip_hidden_states = encoder_hidden_states - else: - deprecation_message = ( - "You have passed a tensor as `encoder_hidden_states`. This is deprecated and will be removed in a future release." - " Please make sure to update your script to pass `encoder_hidden_states` as a tuple to suppress this warning." - ) - deprecate("encoder_hidden_states not a tuple", "1.0.0", deprecation_message, standard_warn=False) - end_pos = encoder_hidden_states.shape[1] - self.num_tokens[0] - encoder_hidden_states, ip_hidden_states = ( - encoder_hidden_states[:, :end_pos, :], - [encoder_hidden_states[:, end_pos:, :]], - ) - - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - query = attn.head_to_batch_dim(query) - key = attn.head_to_batch_dim(key) - value = attn.head_to_batch_dim(value) - - attention_probs = attn.get_attention_scores(query, key, attention_mask) - hidden_states = torch.bmm(attention_probs, value) - hidden_states = attn.batch_to_head_dim(hidden_states) - - if ip_adapter_masks is not None: - if not isinstance(ip_adapter_masks, list): - # for backward compatibility, we accept `ip_adapter_mask` as a tensor of shape [num_ip_adapter, 1, height, width] - ip_adapter_masks = list(ip_adapter_masks.unsqueeze(1)) - if not (len(ip_adapter_masks) == len(self.scale) == len(ip_hidden_states)): - raise ValueError( - f"Length of ip_adapter_masks array ({len(ip_adapter_masks)}) must match " - f"length of self.scale array ({len(self.scale)}) and number of ip_hidden_states " - f"({len(ip_hidden_states)})" - ) - else: - for index, (mask, scale, ip_state) in enumerate(zip(ip_adapter_masks, self.scale, ip_hidden_states)): - if mask is None: - continue - if not isinstance(mask, torch.Tensor) or mask.ndim != 4: - raise ValueError( - "Each element of the ip_adapter_masks array should be a tensor with shape " - "[1, num_images_for_ip_adapter, height, width]." - " Please use `IPAdapterMaskProcessor` to preprocess your mask" - ) - if mask.shape[1] != ip_state.shape[1]: - raise ValueError( - f"Number of masks ({mask.shape[1]}) does not match " - f"number of ip images ({ip_state.shape[1]}) at index {index}" - ) - if isinstance(scale, list) and not len(scale) == mask.shape[1]: - raise ValueError( - f"Number of masks ({mask.shape[1]}) does not match " - f"number of scales ({len(scale)}) at index {index}" - ) - else: - ip_adapter_masks = [None] * len(self.scale) - - # for ip-adapter - for current_ip_hidden_states, scale, to_k_ip, to_v_ip, mask in zip( - ip_hidden_states, self.scale, self.to_k_ip, self.to_v_ip, ip_adapter_masks - ): - skip = False - if isinstance(scale, list): - if all(s == 0 for s in scale): - skip = True - elif scale == 0: - skip = True - if not skip: - if mask is not None: - if not isinstance(scale, list): - scale = [scale] * mask.shape[1] - - current_num_images = mask.shape[1] - for i in range(current_num_images): - ip_key = to_k_ip(current_ip_hidden_states[:, i, :, :]) - ip_value = to_v_ip(current_ip_hidden_states[:, i, :, :]) - - ip_key = attn.head_to_batch_dim(ip_key) - ip_value = attn.head_to_batch_dim(ip_value) - - ip_attention_probs = attn.get_attention_scores(query, ip_key, None) - _current_ip_hidden_states = torch.bmm(ip_attention_probs, ip_value) - _current_ip_hidden_states = attn.batch_to_head_dim(_current_ip_hidden_states) - - mask_downsample = IPAdapterMaskProcessor.downsample( - mask[:, i, :, :], - batch_size, - _current_ip_hidden_states.shape[1], - _current_ip_hidden_states.shape[2], - ) - - mask_downsample = mask_downsample.to(dtype=query.dtype, device=query.device) - - hidden_states = hidden_states + scale[i] * (_current_ip_hidden_states * mask_downsample) - else: - ip_key = to_k_ip(current_ip_hidden_states) - ip_value = to_v_ip(current_ip_hidden_states) - - ip_key = attn.head_to_batch_dim(ip_key) - ip_value = attn.head_to_batch_dim(ip_value) - - ip_attention_probs = attn.get_attention_scores(query, ip_key, None) - current_ip_hidden_states = torch.bmm(ip_attention_probs, ip_value) - current_ip_hidden_states = attn.batch_to_head_dim(current_ip_hidden_states) - - hidden_states = hidden_states + scale * current_ip_hidden_states - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class IPAdapterAttnProcessor2_0(torch.nn.Module): - r""" - Attention processor for IP-Adapter for PyTorch 2.0. - - Args: - hidden_size (`int`): - The hidden size of the attention layer. - cross_attention_dim (`int`): - The number of channels in the `encoder_hidden_states`. - num_tokens (`int`, `tuple[int]` or `list[int]`, defaults to `(4,)`): - The context length of the image features. - scale (`float` or `list[float]`, defaults to 1.0): - the weight scale of image prompt. - """ - - def __init__(self, hidden_size, cross_attention_dim=None, num_tokens=(4,), scale=1.0): - super().__init__() - - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - f"{self.__class__.__name__} requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - self.hidden_size = hidden_size - self.cross_attention_dim = cross_attention_dim - - if not isinstance(num_tokens, (tuple, list)): - num_tokens = [num_tokens] - self.num_tokens = num_tokens - - if not isinstance(scale, list): - scale = [scale] * len(num_tokens) - if len(scale) != len(num_tokens): - raise ValueError("`scale` should be a list of integers with the same length as `num_tokens`.") - self.scale = scale - - self.to_k_ip = nn.ModuleList( - [nn.Linear(cross_attention_dim, hidden_size, bias=False) for _ in range(len(num_tokens))] - ) - self.to_v_ip = nn.ModuleList( - [nn.Linear(cross_attention_dim, hidden_size, bias=False) for _ in range(len(num_tokens))] - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - scale: float = 1.0, - ip_adapter_masks: torch.Tensor | None = None, - ): - residual = hidden_states - - # separate ip_hidden_states from encoder_hidden_states - if encoder_hidden_states is not None: - if isinstance(encoder_hidden_states, tuple): - encoder_hidden_states, ip_hidden_states = encoder_hidden_states - else: - deprecation_message = ( - "You have passed a tensor as `encoder_hidden_states`. This is deprecated and will be removed in a future release." - " Please make sure to update your script to pass `encoder_hidden_states` as a tuple to suppress this warning." - ) - deprecate("encoder_hidden_states not a tuple", "1.0.0", deprecation_message, standard_warn=False) - end_pos = encoder_hidden_states.shape[1] - self.num_tokens[0] - encoder_hidden_states, ip_hidden_states = ( - encoder_hidden_states[:, :end_pos, :], - [encoder_hidden_states[:, end_pos:, :]], - ) - - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - if ip_adapter_masks is not None: - if not isinstance(ip_adapter_masks, list): - # for backward compatibility, we accept `ip_adapter_mask` as a tensor of shape [num_ip_adapter, 1, height, width] - ip_adapter_masks = list(ip_adapter_masks.unsqueeze(1)) - if not (len(ip_adapter_masks) == len(self.scale) == len(ip_hidden_states)): - raise ValueError( - f"Length of ip_adapter_masks array ({len(ip_adapter_masks)}) must match " - f"length of self.scale array ({len(self.scale)}) and number of ip_hidden_states " - f"({len(ip_hidden_states)})" - ) - else: - for index, (mask, scale, ip_state) in enumerate(zip(ip_adapter_masks, self.scale, ip_hidden_states)): - if mask is None: - continue - if not isinstance(mask, torch.Tensor) or mask.ndim != 4: - raise ValueError( - "Each element of the ip_adapter_masks array should be a tensor with shape " - "[1, num_images_for_ip_adapter, height, width]." - " Please use `IPAdapterMaskProcessor` to preprocess your mask" - ) - if mask.shape[1] != ip_state.shape[1]: - raise ValueError( - f"Number of masks ({mask.shape[1]}) does not match " - f"number of ip images ({ip_state.shape[1]}) at index {index}" - ) - if isinstance(scale, list) and not len(scale) == mask.shape[1]: - raise ValueError( - f"Number of masks ({mask.shape[1]}) does not match " - f"number of scales ({len(scale)}) at index {index}" - ) - else: - ip_adapter_masks = [None] * len(self.scale) - - # for ip-adapter - for current_ip_hidden_states, scale, to_k_ip, to_v_ip, mask in zip( - ip_hidden_states, self.scale, self.to_k_ip, self.to_v_ip, ip_adapter_masks - ): - skip = False - if isinstance(scale, list): - if all(s == 0 for s in scale): - skip = True - elif scale == 0: - skip = True - if not skip: - if mask is not None: - if not isinstance(scale, list): - scale = [scale] * mask.shape[1] - - current_num_images = mask.shape[1] - for i in range(current_num_images): - ip_key = to_k_ip(current_ip_hidden_states[:, i, :, :]) - ip_value = to_v_ip(current_ip_hidden_states[:, i, :, :]) - - ip_key = ip_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - ip_value = ip_value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - _current_ip_hidden_states = F.scaled_dot_product_attention( - query, ip_key, ip_value, attn_mask=None, dropout_p=0.0, is_causal=False - ) - - _current_ip_hidden_states = _current_ip_hidden_states.transpose(1, 2).reshape( - batch_size, -1, attn.heads * head_dim - ) - _current_ip_hidden_states = _current_ip_hidden_states.to(query.dtype) - - mask_downsample = IPAdapterMaskProcessor.downsample( - mask[:, i, :, :], - batch_size, - _current_ip_hidden_states.shape[1], - _current_ip_hidden_states.shape[2], - ) - - mask_downsample = mask_downsample.to(dtype=query.dtype, device=query.device) - hidden_states = hidden_states + scale[i] * (_current_ip_hidden_states * mask_downsample) - else: - ip_key = to_k_ip(current_ip_hidden_states) - ip_value = to_v_ip(current_ip_hidden_states) - - ip_key = ip_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - ip_value = ip_value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - current_ip_hidden_states = F.scaled_dot_product_attention( - query, ip_key, ip_value, attn_mask=None, dropout_p=0.0, is_causal=False - ) - - current_ip_hidden_states = current_ip_hidden_states.transpose(1, 2).reshape( - batch_size, -1, attn.heads * head_dim - ) - current_ip_hidden_states = current_ip_hidden_states.to(query.dtype) - - hidden_states = hidden_states + scale * current_ip_hidden_states - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class IPAdapterXFormersAttnProcessor(torch.nn.Module): - r""" - Attention processor for IP-Adapter using xFormers. - - Args: - hidden_size (`int`): - The hidden size of the attention layer. - cross_attention_dim (`int`): - The number of channels in the `encoder_hidden_states`. - num_tokens (`int`, `tuple[int]` or `list[int]`, defaults to `(4,)`): - The context length of the image features. - scale (`float` or `list[float]`, defaults to 1.0): - the weight scale of image prompt. - attention_op (`Callable`, *optional*, defaults to `None`): - The base - [operator](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.AttentionOpBase) to - use as the attention operator. It is recommended to set to `None`, and allow xFormers to choose the best - operator. - """ - - def __init__( - self, - hidden_size, - cross_attention_dim=None, - num_tokens=(4,), - scale=1.0, - attention_op: Callable | None = None, - ): - super().__init__() - - self.hidden_size = hidden_size - self.cross_attention_dim = cross_attention_dim - self.attention_op = attention_op - - if not isinstance(num_tokens, (tuple, list)): - num_tokens = [num_tokens] - self.num_tokens = num_tokens - - if not isinstance(scale, list): - scale = [scale] * len(num_tokens) - if len(scale) != len(num_tokens): - raise ValueError("`scale` should be a list of integers with the same length as `num_tokens`.") - self.scale = scale - - self.to_k_ip = nn.ModuleList( - [nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) for _ in range(len(num_tokens))] - ) - self.to_v_ip = nn.ModuleList( - [nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) for _ in range(len(num_tokens))] - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor | None = None, - attention_mask: torch.FloatTensor | None = None, - temb: torch.FloatTensor | None = None, - scale: float = 1.0, - ip_adapter_masks: torch.FloatTensor | None = None, - ): - residual = hidden_states - - # separate ip_hidden_states from encoder_hidden_states - if encoder_hidden_states is not None: - if isinstance(encoder_hidden_states, tuple): - encoder_hidden_states, ip_hidden_states = encoder_hidden_states - else: - deprecation_message = ( - "You have passed a tensor as `encoder_hidden_states`. This is deprecated and will be removed in a future release." - " Please make sure to update your script to pass `encoder_hidden_states` as a tuple to suppress this warning." - ) - deprecate("encoder_hidden_states not a tuple", "1.0.0", deprecation_message, standard_warn=False) - end_pos = encoder_hidden_states.shape[1] - self.num_tokens[0] - encoder_hidden_states, ip_hidden_states = ( - encoder_hidden_states[:, :end_pos, :], - [encoder_hidden_states[:, end_pos:, :]], - ) - - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # expand our mask's singleton query_tokens dimension: - # [batch*heads, 1, key_tokens] -> - # [batch*heads, query_tokens, key_tokens] - # so that it can be added as a bias onto the attention scores that xformers computes: - # [batch*heads, query_tokens, key_tokens] - # we do this explicitly because xformers doesn't broadcast the singleton dimension for us. - _, query_tokens, _ = hidden_states.shape - attention_mask = attention_mask.expand(-1, query_tokens, -1) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - query = attn.head_to_batch_dim(query).contiguous() - key = attn.head_to_batch_dim(key).contiguous() - value = attn.head_to_batch_dim(value).contiguous() - - hidden_states = xformers.ops.memory_efficient_attention( - query, key, value, attn_bias=attention_mask, op=self.attention_op - ) - hidden_states = hidden_states.to(query.dtype) - hidden_states = attn.batch_to_head_dim(hidden_states) - - if ip_hidden_states: - if ip_adapter_masks is not None: - if not isinstance(ip_adapter_masks, list): - # for backward compatibility, we accept `ip_adapter_mask` as a tensor of shape [num_ip_adapter, 1, height, width] - ip_adapter_masks = list(ip_adapter_masks.unsqueeze(1)) - if not (len(ip_adapter_masks) == len(self.scale) == len(ip_hidden_states)): - raise ValueError( - f"Length of ip_adapter_masks array ({len(ip_adapter_masks)}) must match " - f"length of self.scale array ({len(self.scale)}) and number of ip_hidden_states " - f"({len(ip_hidden_states)})" - ) - else: - for index, (mask, scale, ip_state) in enumerate( - zip(ip_adapter_masks, self.scale, ip_hidden_states) - ): - if mask is None: - continue - if not isinstance(mask, torch.Tensor) or mask.ndim != 4: - raise ValueError( - "Each element of the ip_adapter_masks array should be a tensor with shape " - "[1, num_images_for_ip_adapter, height, width]." - " Please use `IPAdapterMaskProcessor` to preprocess your mask" - ) - if mask.shape[1] != ip_state.shape[1]: - raise ValueError( - f"Number of masks ({mask.shape[1]}) does not match " - f"number of ip images ({ip_state.shape[1]}) at index {index}" - ) - if isinstance(scale, list) and not len(scale) == mask.shape[1]: - raise ValueError( - f"Number of masks ({mask.shape[1]}) does not match " - f"number of scales ({len(scale)}) at index {index}" - ) - else: - ip_adapter_masks = [None] * len(self.scale) - - # for ip-adapter - for current_ip_hidden_states, scale, to_k_ip, to_v_ip, mask in zip( - ip_hidden_states, self.scale, self.to_k_ip, self.to_v_ip, ip_adapter_masks - ): - skip = False - if isinstance(scale, list): - if all(s == 0 for s in scale): - skip = True - elif scale == 0: - skip = True - if not skip: - if mask is not None: - mask = mask.to(torch.float16) - if not isinstance(scale, list): - scale = [scale] * mask.shape[1] - - current_num_images = mask.shape[1] - for i in range(current_num_images): - ip_key = to_k_ip(current_ip_hidden_states[:, i, :, :]) - ip_value = to_v_ip(current_ip_hidden_states[:, i, :, :]) - - ip_key = attn.head_to_batch_dim(ip_key).contiguous() - ip_value = attn.head_to_batch_dim(ip_value).contiguous() - - _current_ip_hidden_states = xformers.ops.memory_efficient_attention( - query, ip_key, ip_value, op=self.attention_op - ) - _current_ip_hidden_states = _current_ip_hidden_states.to(query.dtype) - _current_ip_hidden_states = attn.batch_to_head_dim(_current_ip_hidden_states) - - mask_downsample = IPAdapterMaskProcessor.downsample( - mask[:, i, :, :], - batch_size, - _current_ip_hidden_states.shape[1], - _current_ip_hidden_states.shape[2], - ) - - mask_downsample = mask_downsample.to(dtype=query.dtype, device=query.device) - hidden_states = hidden_states + scale[i] * (_current_ip_hidden_states * mask_downsample) - else: - ip_key = to_k_ip(current_ip_hidden_states) - ip_value = to_v_ip(current_ip_hidden_states) - - ip_key = attn.head_to_batch_dim(ip_key).contiguous() - ip_value = attn.head_to_batch_dim(ip_value).contiguous() - - current_ip_hidden_states = xformers.ops.memory_efficient_attention( - query, ip_key, ip_value, op=self.attention_op - ) - current_ip_hidden_states = current_ip_hidden_states.to(query.dtype) - current_ip_hidden_states = attn.batch_to_head_dim(current_ip_hidden_states) - - hidden_states = hidden_states + scale * current_ip_hidden_states - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class SD3IPAdapterJointAttnProcessor2_0(torch.nn.Module): - """ - Attention processor for IP-Adapter used typically in processing the SD3-like self-attention projections, with - additional image-based information and timestep embeddings. - - Args: - hidden_size (`int`): - The number of hidden channels. - ip_hidden_states_dim (`int`): - The image feature dimension. - head_dim (`int`): - The number of head channels. - timesteps_emb_dim (`int`, defaults to 1280): - The number of input channels for timestep embedding. - scale (`float`, defaults to 0.5): - IP-Adapter scale. - """ - - def __init__( - self, - hidden_size: int, - ip_hidden_states_dim: int, - head_dim: int, - timesteps_emb_dim: int = 1280, - scale: float = 0.5, - ): - super().__init__() - - # To prevent circular import - from .normalization import AdaLayerNorm, RMSNorm - - self.norm_ip = AdaLayerNorm(timesteps_emb_dim, output_dim=ip_hidden_states_dim * 2, norm_eps=1e-6, chunk_dim=1) - self.to_k_ip = nn.Linear(ip_hidden_states_dim, hidden_size, bias=False) - self.to_v_ip = nn.Linear(ip_hidden_states_dim, hidden_size, bias=False) - self.norm_q = RMSNorm(head_dim, 1e-6) - self.norm_k = RMSNorm(head_dim, 1e-6) - self.norm_ip_k = RMSNorm(head_dim, 1e-6) - self.scale = scale - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor = None, - attention_mask: torch.FloatTensor | None = None, - ip_hidden_states: torch.FloatTensor = None, - temb: torch.FloatTensor = None, - ) -> torch.FloatTensor: - """ - Perform the attention computation, integrating image features (if provided) and timestep embeddings. - - If `ip_hidden_states` is `None`, this is equivalent to using JointAttnProcessor2_0. - - Args: - attn (`Attention`): - Attention instance. - hidden_states (`torch.FloatTensor`): - Input `hidden_states`. - encoder_hidden_states (`torch.FloatTensor`, *optional*): - The encoder hidden states. - attention_mask (`torch.FloatTensor`, *optional*): - Attention mask. - ip_hidden_states (`torch.FloatTensor`, *optional*): - Image embeddings. - temb (`torch.FloatTensor`, *optional*): - Timestep embeddings. - - Returns: - `torch.FloatTensor`: Output hidden states. - """ - residual = hidden_states - - batch_size = hidden_states.shape[0] - - # `sample` projections. - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - img_query = query - img_key = key - img_value = value - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # `context` projections. - if encoder_hidden_states is not None: - encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states) - encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) - encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) - - encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view( - batch_size, -1, attn.heads, head_dim - ).transpose(1, 2) - encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view( - batch_size, -1, attn.heads, head_dim - ).transpose(1, 2) - encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view( - batch_size, -1, attn.heads, head_dim - ).transpose(1, 2) - - if attn.norm_added_q is not None: - encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj) - if attn.norm_added_k is not None: - encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj) - - query = torch.cat([query, encoder_hidden_states_query_proj], dim=2) - key = torch.cat([key, encoder_hidden_states_key_proj], dim=2) - value = torch.cat([value, encoder_hidden_states_value_proj], dim=2) - - hidden_states = F.scaled_dot_product_attention(query, key, value, dropout_p=0.0, is_causal=False) - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - if encoder_hidden_states is not None: - # Split the attention outputs. - hidden_states, encoder_hidden_states = ( - hidden_states[:, : residual.shape[1]], - hidden_states[:, residual.shape[1] :], - ) - if not attn.context_pre_only: - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - # IP Adapter - if self.scale != 0 and ip_hidden_states is not None: - # Norm image features - norm_ip_hidden_states = self.norm_ip(ip_hidden_states, temb=temb) - - # To k and v - ip_key = self.to_k_ip(norm_ip_hidden_states) - ip_value = self.to_v_ip(norm_ip_hidden_states) - - # Reshape - ip_key = ip_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - ip_value = ip_value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # Norm - query = self.norm_q(img_query) - img_key = self.norm_k(img_key) - ip_key = self.norm_ip_k(ip_key) - - # cat img - key = torch.cat([img_key, ip_key], dim=2) - value = torch.cat([img_value, ip_value], dim=2) - - ip_hidden_states = F.scaled_dot_product_attention(query, key, value, dropout_p=0.0, is_causal=False) - ip_hidden_states = ip_hidden_states.transpose(1, 2).view(batch_size, -1, attn.heads * head_dim) - ip_hidden_states = ip_hidden_states.to(query.dtype) - - hidden_states = hidden_states + ip_hidden_states * self.scale - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if encoder_hidden_states is not None: - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -class PAGIdentitySelfAttnProcessor2_0: - r""" - Processor for implementing PAG using scaled dot-product attention (enabled by default if you're using PyTorch 2.0). - PAG reference: https://huggingface.co/papers/2403.17377 - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "PAGIdentitySelfAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor | None = None, - attention_mask: torch.FloatTensor | None = None, - temb: torch.FloatTensor | None = None, - ) -> torch.Tensor: - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - # chunk - hidden_states_org, hidden_states_ptb = hidden_states.chunk(2) - - # original path - batch_size, sequence_length, _ = hidden_states_org.shape - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states_org = attn.group_norm(hidden_states_org.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states_org) - key = attn.to_k(hidden_states_org) - value = attn.to_v(hidden_states_org) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states_org = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - hidden_states_org = hidden_states_org.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states_org = hidden_states_org.to(query.dtype) - - # linear proj - hidden_states_org = attn.to_out[0](hidden_states_org) - # dropout - hidden_states_org = attn.to_out[1](hidden_states_org) - - if input_ndim == 4: - hidden_states_org = hidden_states_org.transpose(-1, -2).reshape(batch_size, channel, height, width) - - # perturbed path (identity attention) - batch_size, sequence_length, _ = hidden_states_ptb.shape - - if attn.group_norm is not None: - hidden_states_ptb = attn.group_norm(hidden_states_ptb.transpose(1, 2)).transpose(1, 2) - - hidden_states_ptb = attn.to_v(hidden_states_ptb) - hidden_states_ptb = hidden_states_ptb.to(query.dtype) - - # linear proj - hidden_states_ptb = attn.to_out[0](hidden_states_ptb) - # dropout - hidden_states_ptb = attn.to_out[1](hidden_states_ptb) - - if input_ndim == 4: - hidden_states_ptb = hidden_states_ptb.transpose(-1, -2).reshape(batch_size, channel, height, width) - - # cat - hidden_states = torch.cat([hidden_states_org, hidden_states_ptb]) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class PAGCFGIdentitySelfAttnProcessor2_0: - r""" - Processor for implementing PAG using scaled dot-product attention (enabled by default if you're using PyTorch 2.0). - PAG reference: https://huggingface.co/papers/2403.17377 - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "PAGCFGIdentitySelfAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor | None = None, - attention_mask: torch.FloatTensor | None = None, - temb: torch.FloatTensor | None = None, - ) -> torch.Tensor: - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - # chunk - hidden_states_uncond, hidden_states_org, hidden_states_ptb = hidden_states.chunk(3) - hidden_states_org = torch.cat([hidden_states_uncond, hidden_states_org]) - - # original path - batch_size, sequence_length, _ = hidden_states_org.shape - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states_org = attn.group_norm(hidden_states_org.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states_org) - key = attn.to_k(hidden_states_org) - value = attn.to_v(hidden_states_org) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states_org = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states_org = hidden_states_org.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states_org = hidden_states_org.to(query.dtype) - - # linear proj - hidden_states_org = attn.to_out[0](hidden_states_org) - # dropout - hidden_states_org = attn.to_out[1](hidden_states_org) - - if input_ndim == 4: - hidden_states_org = hidden_states_org.transpose(-1, -2).reshape(batch_size, channel, height, width) - - # perturbed path (identity attention) - batch_size, sequence_length, _ = hidden_states_ptb.shape - - if attn.group_norm is not None: - hidden_states_ptb = attn.group_norm(hidden_states_ptb.transpose(1, 2)).transpose(1, 2) - - value = attn.to_v(hidden_states_ptb) - hidden_states_ptb = value - hidden_states_ptb = hidden_states_ptb.to(query.dtype) - - # linear proj - hidden_states_ptb = attn.to_out[0](hidden_states_ptb) - # dropout - hidden_states_ptb = attn.to_out[1](hidden_states_ptb) - - if input_ndim == 4: - hidden_states_ptb = hidden_states_ptb.transpose(-1, -2).reshape(batch_size, channel, height, width) - - # cat - hidden_states = torch.cat([hidden_states_org, hidden_states_ptb]) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class SanaMultiscaleAttnProcessor2_0: - r""" - Processor for implementing multiscale quadratic attention. - """ - - def __call__(self, attn: SanaMultiscaleLinearAttention, hidden_states: torch.Tensor) -> torch.Tensor: - height, width = hidden_states.shape[-2:] - if height * width > attn.attention_head_dim: - use_linear_attention = True - else: - use_linear_attention = False - - residual = hidden_states - - batch_size, _, height, width = list(hidden_states.size()) - original_dtype = hidden_states.dtype - - hidden_states = hidden_states.movedim(1, -1) - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - hidden_states = torch.cat([query, key, value], dim=3) - hidden_states = hidden_states.movedim(-1, 1) - - multi_scale_qkv = [hidden_states] - for block in attn.to_qkv_multiscale: - multi_scale_qkv.append(block(hidden_states)) - - hidden_states = torch.cat(multi_scale_qkv, dim=1) - - if use_linear_attention: - # for linear attention upcast hidden_states to float32 - hidden_states = hidden_states.to(dtype=torch.float32) - - hidden_states = hidden_states.reshape(batch_size, -1, 3 * attn.attention_head_dim, height * width) - - query, key, value = hidden_states.chunk(3, dim=2) - query = attn.nonlinearity(query) - key = attn.nonlinearity(key) - - if use_linear_attention: - hidden_states = attn.apply_linear_attention(query, key, value) - hidden_states = hidden_states.to(dtype=original_dtype) - else: - hidden_states = attn.apply_quadratic_attention(query, key, value) - - hidden_states = torch.reshape(hidden_states, (batch_size, -1, height, width)) - hidden_states = attn.to_out(hidden_states.movedim(1, -1)).movedim(-1, 1) - - if attn.norm_type == "rms_norm": - hidden_states = attn.norm_out(hidden_states.movedim(1, -1)).movedim(-1, 1) - else: - hidden_states = attn.norm_out(hidden_states) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - return hidden_states - - -class LoRAAttnProcessor: - r""" - Processor for implementing attention with LoRA. - """ - - def __init__(self): - pass - - -class LoRAAttnProcessor2_0: - r""" - Processor for implementing attention with LoRA (enabled by default if you're using PyTorch 2.0). - """ - - def __init__(self): - pass - - -class LoRAXFormersAttnProcessor: - r""" - Processor for implementing attention with LoRA using xFormers. - """ - - def __init__(self): - pass - - -class LoRAAttnAddedKVProcessor: - r""" - Processor for implementing attention with LoRA with extra learnable key and value matrices for the text encoder. - """ - - def __init__(self): - pass - - -class SanaLinearAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product linear attention. - """ - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - original_dtype = hidden_states.dtype - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - query = query.transpose(1, 2).unflatten(1, (attn.heads, -1)) - key = key.transpose(1, 2).unflatten(1, (attn.heads, -1)).transpose(2, 3) - value = value.transpose(1, 2).unflatten(1, (attn.heads, -1)) - - query = F.relu(query) - key = F.relu(key) - - query, key, value = query.float(), key.float(), value.float() - - value = F.pad(value, (0, 0, 0, 1), mode="constant", value=1.0) - scores = torch.matmul(value, key) - hidden_states = torch.matmul(scores, query) - - hidden_states = hidden_states[:, :, :-1] / (hidden_states[:, :, -1:] + 1e-15) - hidden_states = hidden_states.flatten(1, 2).transpose(1, 2) - hidden_states = hidden_states.to(original_dtype) - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - if original_dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - return hidden_states - - -class PAGCFGSanaLinearAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product linear attention. - """ - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - original_dtype = hidden_states.dtype - - hidden_states_uncond, hidden_states_org, hidden_states_ptb = hidden_states.chunk(3) - hidden_states_org = torch.cat([hidden_states_uncond, hidden_states_org]) - - query = attn.to_q(hidden_states_org) - key = attn.to_k(hidden_states_org) - value = attn.to_v(hidden_states_org) - - query = query.transpose(1, 2).unflatten(1, (attn.heads, -1)) - key = key.transpose(1, 2).unflatten(1, (attn.heads, -1)).transpose(2, 3) - value = value.transpose(1, 2).unflatten(1, (attn.heads, -1)) - - query = F.relu(query) - key = F.relu(key) - - query, key, value = query.float(), key.float(), value.float() - - value = F.pad(value, (0, 0, 0, 1), mode="constant", value=1.0) - scores = torch.matmul(value, key) - hidden_states_org = torch.matmul(scores, query) - - hidden_states_org = hidden_states_org[:, :, :-1] / (hidden_states_org[:, :, -1:] + 1e-15) - hidden_states_org = hidden_states_org.flatten(1, 2).transpose(1, 2) - hidden_states_org = hidden_states_org.to(original_dtype) - - hidden_states_org = attn.to_out[0](hidden_states_org) - hidden_states_org = attn.to_out[1](hidden_states_org) - - # perturbed path (identity attention) - hidden_states_ptb = attn.to_v(hidden_states_ptb).to(original_dtype) - - hidden_states_ptb = attn.to_out[0](hidden_states_ptb) - hidden_states_ptb = attn.to_out[1](hidden_states_ptb) - - hidden_states = torch.cat([hidden_states_org, hidden_states_ptb]) - - if original_dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - return hidden_states - - -class PAGIdentitySanaLinearAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product linear attention. - """ - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - original_dtype = hidden_states.dtype - - hidden_states_org, hidden_states_ptb = hidden_states.chunk(2) - - query = attn.to_q(hidden_states_org) - key = attn.to_k(hidden_states_org) - value = attn.to_v(hidden_states_org) - - query = query.transpose(1, 2).unflatten(1, (attn.heads, -1)) - key = key.transpose(1, 2).unflatten(1, (attn.heads, -1)).transpose(2, 3) - value = value.transpose(1, 2).unflatten(1, (attn.heads, -1)) - - query = F.relu(query) - key = F.relu(key) - - query, key, value = query.float(), key.float(), value.float() - - value = F.pad(value, (0, 0, 0, 1), mode="constant", value=1.0) - scores = torch.matmul(value, key) - hidden_states_org = torch.matmul(scores, query) - - if hidden_states_org.dtype in [torch.float16, torch.bfloat16]: - hidden_states_org = hidden_states_org.float() - - hidden_states_org = hidden_states_org[:, :, :-1] / (hidden_states_org[:, :, -1:] + 1e-15) - hidden_states_org = hidden_states_org.flatten(1, 2).transpose(1, 2) - hidden_states_org = hidden_states_org.to(original_dtype) - - hidden_states_org = attn.to_out[0](hidden_states_org) - hidden_states_org = attn.to_out[1](hidden_states_org) - - # perturbed path (identity attention) - hidden_states_ptb = attn.to_v(hidden_states_ptb).to(original_dtype) - - hidden_states_ptb = attn.to_out[0](hidden_states_ptb) - hidden_states_ptb = attn.to_out[1](hidden_states_ptb) - - hidden_states = torch.cat([hidden_states_org, hidden_states_ptb]) - - if original_dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - return hidden_states - - -class FluxAttnProcessor2_0: - def __new__(cls, *args, **kwargs): - deprecation_message = "`FluxAttnProcessor2_0` is deprecated and this will be removed in a future version. Please use `FluxAttnProcessor`" - deprecate("FluxAttnProcessor2_0", "1.0.0", deprecation_message) - - from .transformers.transformer_flux import FluxAttnProcessor - - return FluxAttnProcessor(*args, **kwargs) - - -class FluxSingleAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). - """ - - def __new__(cls, *args, **kwargs): - deprecation_message = "`FluxSingleAttnProcessor` is deprecated and will be removed in a future version. Please use `FluxAttnProcessorSDPA` instead." - deprecate("FluxSingleAttnProcessor2_0", "1.0.0", deprecation_message) - - from .transformers.transformer_flux import FluxAttnProcessor - - return FluxAttnProcessor(*args, **kwargs) - - -class FusedFluxAttnProcessor2_0: - def __new__(cls, *args, **kwargs): - deprecation_message = "`FusedFluxAttnProcessor2_0` is deprecated and this will be removed in a future version. Please use `FluxAttnProcessor`" - deprecate("FusedFluxAttnProcessor2_0", "1.0.0", deprecation_message) - - from .transformers.transformer_flux import FluxAttnProcessor - - return FluxAttnProcessor(*args, **kwargs) - - -class FluxIPAdapterJointAttnProcessor2_0: - def __new__(cls, *args, **kwargs): - deprecation_message = "`FluxIPAdapterJointAttnProcessor2_0` is deprecated and this will be removed in a future version. Please use `FluxIPAdapterAttnProcessor`" - deprecate("FluxIPAdapterJointAttnProcessor2_0", "1.0.0", deprecation_message) - - from .transformers.transformer_flux import FluxIPAdapterAttnProcessor - - return FluxIPAdapterAttnProcessor(*args, **kwargs) - - -class FluxAttnProcessor2_0_NPU: - def __new__(cls, *args, **kwargs): - deprecation_message = ( - "FluxAttnProcessor2_0_NPU is deprecated and will be removed in a future version. An " - "alternative solution to use NPU Flash Attention will be provided in the future." - ) - deprecate("FluxAttnProcessor2_0_NPU", "1.0.0", deprecation_message, standard_warn=False) - - from .transformers.transformer_flux import FluxAttnProcessor - - processor = FluxAttnProcessor() - processor._attention_backend = "_native_npu" - return processor - - -class FusedFluxAttnProcessor2_0_NPU: - def __new__(self): - deprecation_message = ( - "FusedFluxAttnProcessor2_0_NPU is deprecated and will be removed in a future version. An " - "alternative solution to use NPU Flash Attention will be provided in the future." - ) - deprecate("FusedFluxAttnProcessor2_0_NPU", "1.0.0", deprecation_message, standard_warn=False) - - from .transformers.transformer_flux import FluxAttnProcessor - - processor = FluxAttnProcessor() - processor._attention_backend = "_fused_npu" - return processor - - -class XLAFluxFlashAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention with pallas flash attention kernel if using `torch_xla`. - """ - - def __new__(cls, *args, **kwargs): - deprecation_message = ( - "XLAFluxFlashAttnProcessor2_0 is deprecated and will be removed in diffusers 1.0.0. An " - "alternative solution to using XLA Flash Attention will be provided in the future." - ) - deprecate("XLAFluxFlashAttnProcessor2_0", "1.0.0", deprecation_message, standard_warn=False) - - if is_torch_xla_version("<", "2.3"): - raise ImportError("XLA flash attention requires torch_xla version >= 2.3.") - if is_spmd() and is_torch_xla_version("<", "2.4"): - raise ImportError("SPMD support for XLA flash attention needs torch_xla version >= 2.4.") - - from .transformers.transformer_flux import FluxAttnProcessor - - if len(args) > 0 or kwargs.get("partition_spec", None) is not None: - deprecation_message = ( - "partition_spec was not used in the processor implementation when it was added. Passing it " - "is a no-op and support for it will be removed." - ) - deprecate("partition_spec", "1.0.0", deprecation_message) - - processor = FluxAttnProcessor(*args, **kwargs) - processor._attention_backend = "_native_xla" - return processor - - -ADDED_KV_ATTENTION_PROCESSORS = ( - AttnAddedKVProcessor, - SlicedAttnAddedKVProcessor, - AttnAddedKVProcessor2_0, - XFormersAttnAddedKVProcessor, -) - -CROSS_ATTENTION_PROCESSORS = ( - AttnProcessor, - AttnProcessor2_0, - XFormersAttnProcessor, - SlicedAttnProcessor, - IPAdapterAttnProcessor, - IPAdapterAttnProcessor2_0, - FluxIPAdapterJointAttnProcessor2_0, -) - -AttentionProcessor = ( - AttnProcessor - | CustomDiffusionAttnProcessor - | AttnAddedKVProcessor - | AttnAddedKVProcessor2_0 - | JointAttnProcessor2_0 - | PAGJointAttnProcessor2_0 - | PAGCFGJointAttnProcessor2_0 - | FusedJointAttnProcessor2_0 - | AllegroAttnProcessor2_0 - | AuraFlowAttnProcessor2_0 - | FusedAuraFlowAttnProcessor2_0 - | FluxAttnProcessor2_0 - | FluxAttnProcessor2_0_NPU - | FusedFluxAttnProcessor2_0 - | FusedFluxAttnProcessor2_0_NPU - | CogVideoXAttnProcessor2_0 - | FusedCogVideoXAttnProcessor2_0 - | XFormersAttnAddedKVProcessor - | XFormersAttnProcessor - | XLAFlashAttnProcessor2_0 - | AttnProcessorNPU - | AttnProcessor2_0 - | MochiVaeAttnProcessor2_0 - | MochiAttnProcessor2_0 - | StableAudioAttnProcessor2_0 - | HunyuanAttnProcessor2_0 - | FusedHunyuanAttnProcessor2_0 - | PAGHunyuanAttnProcessor2_0 - | PAGCFGHunyuanAttnProcessor2_0 - | LuminaAttnProcessor2_0 - | FusedAttnProcessor2_0 - | CustomDiffusionXFormersAttnProcessor - | CustomDiffusionAttnProcessor2_0 - | SlicedAttnProcessor - | SlicedAttnAddedKVProcessor - | SanaLinearAttnProcessor2_0 - | PAGCFGSanaLinearAttnProcessor2_0 - | PAGIdentitySanaLinearAttnProcessor2_0 - | SanaMultiscaleLinearAttention - | SanaMultiscaleAttnProcessor2_0 - | SanaMultiscaleAttentionProjection - | IPAdapterAttnProcessor - | IPAdapterAttnProcessor2_0 - | IPAdapterXFormersAttnProcessor - | SD3IPAdapterJointAttnProcessor2_0 - | PAGIdentitySelfAttnProcessor2_0 - | PAGCFGIdentitySelfAttnProcessor2_0 - | LoRAAttnProcessor - | LoRAAttnProcessor2_0 - | LoRAXFormersAttnProcessor - | LoRAAttnAddedKVProcessor -) diff --git a/diffusers/models/auto_model.py b/diffusers/models/auto_model.py deleted file mode 100644 index 336650ef5703fc76eaa863aa60201b33eb28a7e7..0000000000000000000000000000000000000000 --- a/diffusers/models/auto_model.py +++ /dev/null @@ -1,345 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import os - -from huggingface_hub.utils import validate_hf_hub_args - -from ..configuration_utils import ConfigMixin -from ..utils import DIFFUSERS_LOAD_ID_FIELDS, logging -from ..utils.dynamic_modules_utils import get_class_from_dynamic_module, resolve_trust_remote_code - - -logger = logging.get_logger(__name__) - - -class AutoModel(ConfigMixin): - config_name = "config.json" - - def __init__(self, *args, **kwargs): - raise EnvironmentError( - f"{self.__class__.__name__} is designed to be instantiated " - f"using the `{self.__class__.__name__}.from_pretrained(pretrained_model_name_or_path)`, " - f"`{self.__class__.__name__}.from_config(config)`, or " - f"`{self.__class__.__name__}.from_pipe(pipeline)` methods." - ) - - @classmethod - def from_config(cls, pretrained_model_name_or_path_or_dict: str | os.PathLike | dict | None = None, **kwargs): - r""" - Instantiate a model from a config dictionary or a pretrained model configuration file with random weights (no - pretrained weights are loaded). - - Parameters: - pretrained_model_name_or_path_or_dict (`str`, `os.PathLike`, or `dict`): - Can be either: - - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model - configuration hosted on the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing a model configuration - file. - - A config dictionary. - - cache_dir (`Union[str, os.PathLike]`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model configuration, overriding the cached version if - it exists. - proxies (`Dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint. - local_files_only(`bool`, *optional*, defaults to `False`): - Whether to only load local model configuration files or not. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. - trust_remote_code (`bool`, *optional*, defaults to `False`): - Whether to trust remote code. - subfolder (`str`, *optional*, defaults to `""`): - The subfolder location of a model file within a larger model repository on the Hub or locally. - - Returns: - A model object instantiated from the config with random weights. - - Example: - - ```py - from diffusers import AutoModel - - model = AutoModel.from_config("stable-diffusion-v1-5/stable-diffusion-v1-5", subfolder="unet") - ``` - """ - subfolder = kwargs.pop("subfolder", None) - trust_remote_code = kwargs.pop("trust_remote_code", False) - - hub_kwargs_names = [ - "cache_dir", - "force_download", - "local_files_only", - "proxies", - "revision", - "token", - ] - hub_kwargs = {name: kwargs.pop(name, None) for name in hub_kwargs_names} - - if pretrained_model_name_or_path_or_dict is None: - raise ValueError( - "Please provide a `pretrained_model_name_or_path_or_dict` as the first positional argument." - ) - - if isinstance(pretrained_model_name_or_path_or_dict, (str, os.PathLike)): - pretrained_model_name_or_path = pretrained_model_name_or_path_or_dict - config = cls.load_config(pretrained_model_name_or_path, subfolder=subfolder, **hub_kwargs) - else: - config = pretrained_model_name_or_path_or_dict - pretrained_model_name_or_path = config.get("_name_or_path", None) - - has_remote_code = "auto_map" in config and cls.__name__ in config["auto_map"] - trust_remote_code = resolve_trust_remote_code( - trust_remote_code, pretrained_model_name_or_path, has_remote_code - ) - - if has_remote_code and trust_remote_code: - class_ref = config["auto_map"][cls.__name__] - module_file, class_name = class_ref.split(".") - module_file = module_file + ".py" - model_cls = get_class_from_dynamic_module( - pretrained_model_name_or_path, - subfolder=subfolder, - module_file=module_file, - class_name=class_name, - trust_remote_code=trust_remote_code, - **hub_kwargs, - ) - else: - if "_class_name" in config: - class_name = config["_class_name"] - library = "diffusers" - elif "model_type" in config: - class_name = "AutoModel" - library = "transformers" - else: - raise ValueError( - f"Couldn't find a model class associated with the config: {config}. Make sure the config " - "contains a `_class_name` or `model_type` key." - ) - - from ..pipelines.pipeline_loading_utils import ALL_IMPORTABLE_CLASSES, get_class_obj_and_candidates - - model_cls, _ = get_class_obj_and_candidates( - library_name=library, - class_name=class_name, - importable_classes=ALL_IMPORTABLE_CLASSES, - pipelines=None, - is_pipeline_module=False, - trust_remote_code=trust_remote_code, - ) - - if model_cls is None: - raise ValueError(f"AutoModel can't find a model linked to {class_name}.") - - return model_cls.from_config(config, **kwargs) - - @classmethod - @validate_hf_hub_args - def from_pretrained(cls, pretrained_model_or_path: str | os.PathLike | None = None, **kwargs): - r""" - Instantiate a pretrained PyTorch model from a pretrained model configuration. - - The model is set in evaluation mode - `model.eval()` - by default, and dropout modules are deactivated. To - train the model, set it back in training mode with `model.train()`. - - Parameters: - pretrained_model_name_or_path (`str` or `os.PathLike`, *optional*): - Can be either: - - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved - with [`~ModelMixin.save_pretrained`]. - - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - dtype (`torch.dtype`, *optional*): - Override the default `torch.dtype` and load the model with another dtype. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - output_loading_info (`bool`, *optional*, defaults to `False`): - Whether or not to also return a dictionary containing missing keys, unexpected keys and error messages. - local_files_only(`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - subfolder (`str`, *optional*, defaults to `""`): - The subfolder location of a model file within a larger model repository on the Hub or locally. - mirror (`str`, *optional*): - Mirror source to resolve accessibility issues if you're downloading a model in China. We do not - guarantee the timeliness or safety of the source, and you should refer to the mirror site for more - information. - device_map (`str` or `dict[str, int | str | torch.device]`, *optional*): - A map that specifies where each submodule should go. It doesn't need to be defined for each - parameter/buffer name; once a given module name is inside, every submodule of it will be sent to the - same device. Defaults to `None`, meaning that the model will be loaded on CPU. - - Set `device_map="auto"` to have 🤗 Accelerate automatically compute the most optimized `device_map`. For - more information about each option see [designing a device - map](https://hf.co/docs/accelerate/main/en/usage_guides/big_modeling#designing-a-device-map). - max_memory (`Dict`, *optional*): - A dictionary device identifier for the maximum memory. Will default to the maximum memory available for - each GPU and the available CPU RAM if unset. - offload_folder (`str` or `os.PathLike`, *optional*): - The path to offload weights if `device_map` contains the value `"disk"`. - offload_state_dict (`bool`, *optional*): - If `True`, temporarily offloads the CPU state dict to the hard drive to avoid running out of CPU RAM if - the weight of the CPU state dict + the biggest shard of the checkpoint does not fit. Defaults to `True` - when there is some disk offload. - low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 else `False`): - Speed up model loading only loading the pretrained weights and not initializing the weights. This also - tries to not use more than 1x model size in CPU memory (including peak memory) while loading the model. - Only supported for PyTorch >= 1.9.0. If you are using an older version of PyTorch, setting this - argument to `True` will raise an error. - variant (`str`, *optional*): - Load weights from a specified `variant` filename such as `"fp16"` or `"ema"`. - use_safetensors (`bool`, *optional*, defaults to `None`): - If set to `None`, the `safetensors` weights are downloaded if they're available **and** if the - `safetensors` library is installed. If set to `True`, the model is forcibly loaded from `safetensors` - weights. If set to `False`, `safetensors` weights are not loaded. - disable_mmap ('bool', *optional*, defaults to 'False'): - Whether to disable mmap when loading a Safetensors model. This option can perform better when the model - is on a network mount or hard drive, which may not handle the seeky-ness of mmap very well. - trust_remote_cocde (`bool`, *optional*, defaults to `False`): - Whether to trust remote code - - > [!TIP] > To use private or [gated models](https://huggingface.co/docs/hub/models-gated#gated-models), log-in - with `hf > auth login`. You can also activate the special > - ["offline-mode"](https://huggingface.co/diffusers/installation.html#offline-mode) to use this method in a > - firewalled environment. - - Example: - - ```py - from diffusers import AutoModel - - unet = AutoModel.from_pretrained("stable-diffusion-v1-5/stable-diffusion-v1-5", subfolder="unet") - ``` - - If you get the error message below, you need to finetune the weights for your downstream task: - - ```bash - Some weights of UNet2DConditionModel were not initialized from the model checkpoint at stable-diffusion-v1-5/stable-diffusion-v1-5 and are newly initialized because the shapes did not match: - - conv_in.weight: found shape torch.Size([320, 4, 3, 3]) in the checkpoint and torch.Size([320, 9, 3, 3]) in the model instantiated - You should probably TRAIN this model on a down-stream task to be able to use it for predictions and inference. - ``` - """ - subfolder = kwargs.pop("subfolder", None) - trust_remote_code = kwargs.pop("trust_remote_code", False) - - hub_kwargs_names = [ - "cache_dir", - "force_download", - "local_files_only", - "proxies", - "revision", - "token", - ] - hub_kwargs = {name: kwargs.pop(name, None) for name in hub_kwargs_names} - - # load_config_kwargs uses the same hub kwargs minus subfolder and resume_download - load_config_kwargs = {k: v for k, v in hub_kwargs.items() if k not in ["subfolder"]} - - library = None - orig_class_name = None - - # Always attempt to fetch model_index.json first - try: - cls.config_name = "model_index.json" - config = cls.load_config(pretrained_model_or_path, **load_config_kwargs) - - if subfolder is not None and subfolder in config: - library, orig_class_name = config[subfolder] - load_config_kwargs.update({"subfolder": subfolder}) - - except EnvironmentError as e: - logger.debug(e) - - # Unable to load from model_index.json so fallback to loading from config - if library is None and orig_class_name is None: - cls.config_name = "config.json" - config = cls.load_config(pretrained_model_or_path, subfolder=subfolder, **load_config_kwargs) - - if "_class_name" in config: - # If we find a class name in the config, we can try to load the model as a diffusers model - orig_class_name = config["_class_name"] - library = "diffusers" - load_config_kwargs.update({"subfolder": subfolder}) - elif "model_type" in config: - orig_class_name = "AutoModel" - library = "transformers" - load_config_kwargs.update({"subfolder": "" if subfolder is None else subfolder}) - else: - raise ValueError(f"Couldn't find model associated with the config file at {pretrained_model_or_path}.") - - has_remote_code = "auto_map" in config and cls.__name__ in config["auto_map"] - trust_remote_code = resolve_trust_remote_code(trust_remote_code, pretrained_model_or_path, has_remote_code) - if not has_remote_code and trust_remote_code: - raise ValueError( - "Selected model repository does not appear to have any custom code or does not have a valid `config.json` file." - ) - - if has_remote_code and trust_remote_code: - class_ref = config["auto_map"][cls.__name__] - module_file, class_name = class_ref.split(".") - module_file = module_file + ".py" - model_cls = get_class_from_dynamic_module( - pretrained_model_or_path, - subfolder=subfolder, - module_file=module_file, - class_name=class_name, - trust_remote_code=trust_remote_code, - **hub_kwargs, - ) - else: - from ..pipelines.pipeline_loading_utils import ALL_IMPORTABLE_CLASSES, get_class_obj_and_candidates - - model_cls, _ = get_class_obj_and_candidates( - library_name=library, - class_name=orig_class_name, - importable_classes=ALL_IMPORTABLE_CLASSES, - pipelines=None, - is_pipeline_module=False, - ) - - if model_cls is None: - raise ValueError(f"AutoModel can't find a model linked to {orig_class_name}.") - - kwargs = {**load_config_kwargs, **kwargs} - model = model_cls.from_pretrained(pretrained_model_or_path, **kwargs) - - load_id_kwargs = {"pretrained_model_name_or_path": pretrained_model_or_path, **kwargs} - parts = [load_id_kwargs.get(field, "null") for field in DIFFUSERS_LOAD_ID_FIELDS] - load_id = "|".join("null" if p is None else p for p in parts) - model._diffusers_load_id = load_id - - return model diff --git a/diffusers/models/autoencoders/__init__.py b/diffusers/models/autoencoders/__init__.py deleted file mode 100644 index dc481370204b6676c9963d99af4366c6e4b1562b..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/__init__.py +++ /dev/null @@ -1,31 +0,0 @@ -from .autoencoder_asym_kl import AsymmetricAutoencoderKL -from .autoencoder_cosmos3_audio import Cosmos3AVAEAudioTokenizer -from .autoencoder_dc import AutoencoderDC -from .autoencoder_kl import AutoencoderKL -from .autoencoder_kl_allegro import AutoencoderKLAllegro -from .autoencoder_kl_cogvideox import AutoencoderKLCogVideoX -from .autoencoder_kl_cosmos import AutoencoderKLCosmos -from .autoencoder_kl_flux2 import AutoencoderKLFlux2 -from .autoencoder_kl_hunyuan_video import AutoencoderKLHunyuanVideo -from .autoencoder_kl_hunyuanimage import AutoencoderKLHunyuanImage -from .autoencoder_kl_hunyuanimage_refiner import AutoencoderKLHunyuanImageRefiner -from .autoencoder_kl_hunyuanvideo15 import AutoencoderKLHunyuanVideo15 -from .autoencoder_kl_kvae import AutoencoderKLKVAE -from .autoencoder_kl_kvae_video import AutoencoderKLKVAEVideo -from .autoencoder_kl_ltx import AutoencoderKLLTXVideo -from .autoencoder_kl_ltx2 import AutoencoderKLLTX2Video -from .autoencoder_kl_ltx2_audio import AutoencoderKLLTX2Audio -from .autoencoder_kl_magvit import AutoencoderKLMagvit -from .autoencoder_kl_minimax_h3 import AutoencoderKLMiniMaxH3 -from .autoencoder_kl_minimax_h3_audio import AutoencoderKLMiniMaxH3Audio -from .autoencoder_kl_mochi import AutoencoderKLMochi -from .autoencoder_kl_qwenimage import AutoencoderKLQwenImage -from .autoencoder_kl_temporal_decoder import AutoencoderKLTemporalDecoder -from .autoencoder_kl_wan import AutoencoderKLWan -from .autoencoder_longcat_audio_dit import LongCatAudioDiTVae -from .autoencoder_oobleck import AutoencoderOobleck -from .autoencoder_rae import AutoencoderRAE -from .autoencoder_tiny import AutoencoderTiny -from .autoencoder_vidtok import AutoencoderVidTok -from .consistency_decoder_vae import ConsistencyDecoderVAE -from .vq_model import VQModel diff --git a/diffusers/models/autoencoders/autoencoder_asym_kl.py b/diffusers/models/autoencoders/autoencoder_asym_kl.py deleted file mode 100644 index bf13a4b3b134b58929a6e75dc5fd496ae86b9345..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_asym_kl.py +++ /dev/null @@ -1,188 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils.accelerate_utils import apply_forward_hook -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution, Encoder, MaskConditionDecoder - - -class AsymmetricAutoencoderKL(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - Designing a Better Asymmetric VQGAN for StableDiffusion https://huggingface.co/papers/2306.04632 . A VAE model with - KL loss for encoding images into latents and decoding latent representations into images. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - in_channels (int, *optional*, defaults to 3): Number of channels in the input image. - out_channels (int, *optional*, defaults to 3): Number of channels in the output. - down_block_types (`tuple[str]`, *optional*, defaults to `("DownEncoderBlock2D",)`): - tuple of downsample block types. - down_block_out_channels (`tuple[int]`, *optional*, defaults to `(64,)`): - tuple of down block output channels. - layers_per_down_block (`int`, *optional*, defaults to `1`): - Number layers for down block. - up_block_types (`tuple[str]`, *optional*, defaults to `("UpDecoderBlock2D",)`): - tuple of upsample block types. - up_block_out_channels (`tuple[int]`, *optional*, defaults to `(64,)`): - tuple of up block output channels. - layers_per_up_block (`int`, *optional*, defaults to `1`): - Number layers for up block. - act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use. - latent_channels (`int`, *optional*, defaults to 4): Number of channels in the latent space. - sample_size (`int`, *optional*, defaults to `32`): Sample input size. - norm_num_groups (`int`, *optional*, defaults to `32`): - Number of groups to use for the first normalization layer in ResNet blocks. - scaling_factor (`float`, *optional*, defaults to 0.18215): - The component-wise standard deviation of the trained latent space computed using the first batch of the - training set. This is used to scale the latent space to have unit variance when training the diffusion - model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the - diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 - / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image - Synthesis with Latent Diffusion Models](https://huggingface.co/papers/2112.10752) paper. - """ - - _skip_layerwise_casting_patterns = ["decoder"] - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - down_block_types: tuple[str, ...] = ("DownEncoderBlock2D",), - down_block_out_channels: tuple[int, ...] = (64,), - layers_per_down_block: int = 1, - up_block_types: tuple[str, ...] = ("UpDecoderBlock2D",), - up_block_out_channels: tuple[int, ...] = (64,), - layers_per_up_block: int = 1, - act_fn: str = "silu", - latent_channels: int = 4, - norm_num_groups: int = 32, - sample_size: int = 32, - scaling_factor: float = 0.18215, - ) -> None: - super().__init__() - - # pass init params to Encoder - self.encoder = Encoder( - in_channels=in_channels, - out_channels=latent_channels, - down_block_types=down_block_types, - block_out_channels=down_block_out_channels, - layers_per_block=layers_per_down_block, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - double_z=True, - ) - - # pass init params to Decoder - self.decoder = MaskConditionDecoder( - in_channels=latent_channels, - out_channels=out_channels, - up_block_types=up_block_types, - block_out_channels=up_block_out_channels, - layers_per_block=layers_per_up_block, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - ) - - self.quant_conv = nn.Conv2d(2 * latent_channels, 2 * latent_channels, 1) - self.post_quant_conv = nn.Conv2d(latent_channels, latent_channels, 1) - - self.register_to_config(block_out_channels=up_block_out_channels) - self.register_to_config(force_upcast=False) - - @apply_forward_hook - def encode(self, x: torch.Tensor, return_dict: bool = True) -> AutoencoderKLOutput | tuple[torch.Tensor]: - h = self.encoder(x) - moments = self.quant_conv(h) - posterior = DiagonalGaussianDistribution(moments) - - if not return_dict: - return (posterior,) - - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode( - self, - z: torch.Tensor, - image: torch.Tensor | None = None, - mask: torch.Tensor | None = None, - return_dict: bool = True, - ) -> DecoderOutput | tuple[torch.Tensor]: - z = self.post_quant_conv(z) - dec = self.decoder(z, image, mask) - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode( - self, - z: torch.Tensor, - generator: torch.Generator | None = None, - image: torch.Tensor | None = None, - mask: torch.Tensor | None = None, - return_dict: bool = True, - ) -> DecoderOutput | tuple[torch.Tensor]: - decoded = self._decode(z, image, mask).sample - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def forward( - self, - sample: torch.Tensor, - mask: torch.Tensor | None = None, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | tuple[torch.Tensor]: - r""" - Args: - sample (`torch.Tensor`): Input sample. - mask (`torch.Tensor`, *optional*, defaults to `None`): Optional inpainting mask. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`DecoderOutput`] is returned, otherwise a plain `tuple` is returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z, generator, sample, mask).sample - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) diff --git a/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py b/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py deleted file mode 100644 index e5549a47e9f151250d673c1f3cd4678ce1e31c58..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py +++ /dev/null @@ -1,657 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Cosmos3 AVAE Audio Tokenizer. - -The decoder reuses the Oobleck architecture (Snake1d activations + weight-norm convs + residual units), inlined here -instead of imported so the audio module is self-contained. The encoder is the Cosmos3 SpecConvNeXt audio encoder used -by AVAE checkpoints; it is intentionally separate from Oobleck's waveform encoder because the tensor layouts and -bottleneck semantics are different. -""" - -import math -from collections import OrderedDict -from dataclasses import dataclass - -import torch -import torch.nn as nn -import torch.nn.functional as F -from torch.nn.utils import weight_norm - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput -from ...utils.accelerate_utils import apply_forward_hook -from ..modeling_utils import ModelMixin, get_parameter_dtype -from ..normalization import FP32LayerNorm -from .autoencoder_oobleck import OobleckDiagonalGaussianDistribution - - -# Copied from diffusers.models.autoencoders.autoencoder_oobleck.Snake1d -class Snake1d(nn.Module): - """ - A 1-dimensional Snake activation function module. - """ - - def __init__(self, hidden_dim, logscale=True): - super().__init__() - self.alpha = nn.Parameter(torch.zeros(1, hidden_dim, 1)) - self.beta = nn.Parameter(torch.zeros(1, hidden_dim, 1)) - - self.alpha.requires_grad = True - self.beta.requires_grad = True - self.logscale = logscale - - def forward(self, hidden_states): - shape = hidden_states.shape - - alpha = self.alpha if not self.logscale else torch.exp(self.alpha) - beta = self.beta if not self.logscale else torch.exp(self.beta) - - hidden_states = hidden_states.reshape(shape[0], shape[1], -1) - hidden_states = hidden_states + (beta + 1e-9).reciprocal() * torch.sin(alpha * hidden_states).pow(2) - hidden_states = hidden_states.reshape(shape) - return hidden_states - - -class Cosmos3AudioConvNeXtBlock(nn.Module): - """1D ConvNeXt block used by the Cosmos3 SpecConvNeXt encoder.""" - - def __init__( - self, - hidden_dim: int, - intermediate_dim: int, - identity_init: bool = False, - use_snake: bool = True, - causal: bool = False, - ): - super().__init__() - self.causal = causal - - if causal: - self.dwconv = nn.Sequential( - nn.ConstantPad1d((6, 0), 0), - nn.Conv1d(hidden_dim, hidden_dim, kernel_size=7, groups=hidden_dim), - ) - else: - self.dwconv = nn.Sequential( - nn.ConstantPad1d((3, 3), 0), - nn.Conv1d(hidden_dim, hidden_dim, kernel_size=7, groups=hidden_dim), - ) - - self.norm = FP32LayerNorm(hidden_dim, eps=1e-5, bias=False) - self.pwconv1 = nn.Conv1d(hidden_dim, intermediate_dim, kernel_size=1) - self.act = Snake1d(intermediate_dim) if use_snake else nn.GELU() - self.pwconv2 = nn.Conv1d(intermediate_dim, hidden_dim, kernel_size=1) - if identity_init: - nn.init.zeros_(self.pwconv2.weight) - if self.pwconv2.bias is not None: - nn.init.zeros_(self.pwconv2.bias) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - residual = hidden_states - hidden_states = self.dwconv(hidden_states) - hidden_states = self.norm(hidden_states.permute(0, 2, 1)).permute(0, 2, 1) - hidden_states = self.pwconv1(hidden_states) - hidden_states = self.act(hidden_states) - hidden_states = self.pwconv2(hidden_states) - return residual + hidden_states - - -class Cosmos3AudioSpectrogramConvNeXtEncoder(nn.Module): - """Cosmos3 waveform-to-latent encoder using STFT features and ConvNeXt blocks.""" - - def __init__( - self, - input_channels: int, - stereo: bool, - channels: int, - latent_dim: int, - channel_multiples: tuple[int, ...], - strides: tuple[int, ...], - num_blocks: int, - n_fft: int, - hop_length: int, - identity_init: bool, - use_snake: bool, - causal: bool, - padding_mode: str, - ): - super().__init__() - - if causal: - raise NotImplementedError("Cosmos3 AVAE causal audio encoder is not supported yet.") - if len(channel_multiples) != len(strides): - raise ValueError( - "`enc_c_mults` and `enc_strides` must have the same length, got " - f"{len(channel_multiples)} and {len(strides)}." - ) - - self.input_channels = input_channels * (2 if stereo else 1) - self.channels = channels - self.latent_dim = latent_dim - self.channel_multiples = tuple(channel_multiples) - self.strides = tuple(strides) - self.num_blocks = num_blocks - self.n_fft = n_fft - self.hop_length = hop_length - self.causal = causal - - layers: list[nn.Module] = [ - weight_norm( - nn.Conv1d( - (n_fft + 2) * self.input_channels, - self.channel_multiples[0] * channels, - kernel_size=1, - bias=False, - ) - ) - ] - - for index, stride in enumerate(self.strides): - input_dim = self.channel_multiples[index] * channels - output_dim = ( - self.channel_multiples[index + 1] * channels - if index < len(self.channel_multiples) - 1 - else self.channel_multiples[-1] * channels - ) - - for _ in range(num_blocks): - layers.append( - Cosmos3AudioConvNeXtBlock( - hidden_dim=input_dim, - intermediate_dim=input_dim * 4, - identity_init=identity_init, - use_snake=use_snake, - causal=causal, - ) - ) - - layers.append( - weight_norm( - nn.Conv1d( - input_dim, - output_dim, - kernel_size=2 * stride, - stride=stride, - padding=math.ceil(stride / 2), - padding_mode=padding_mode, - ) - ) - ) - - layers.append( - weight_norm(nn.Conv1d(self.channel_multiples[-1] * channels, latent_dim, kernel_size=1, bias=False)) - ) - self.layers = nn.Sequential(*layers) - - def _spectrogram(self, waveform: torch.Tensor) -> torch.Tensor: - pad_left = (self.n_fft - self.hop_length) // 2 - pad_right = (self.n_fft - self.hop_length) - pad_left - waveform = F.pad(waveform, (pad_left, pad_right)).float() - window = torch.hann_window(self.n_fft, device=waveform.device, dtype=waveform.dtype) - return torch.stft( - waveform, - n_fft=self.n_fft, - hop_length=self.hop_length, - win_length=self.n_fft, - window=window, - center=False, - normalized=False, - onesided=True, - return_complex=True, - ) - - def forward(self, audio: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_samples = audio.shape - if num_channels != self.input_channels: - raise ValueError( - f"Cosmos3 AVAE encoder expected {self.input_channels} audio channels, got {num_channels}." - ) - - if num_channels > 1: - audio = audio.reshape(batch_size * num_channels, 1, num_samples) - - spectrogram = self._spectrogram(audio.squeeze(1)) - real, imaginary = torch.view_as_real(spectrogram).chunk(2, dim=-1) - spectrogram = torch.cat([real, imaginary], dim=1).squeeze(-1) - - spectrogram = spectrogram.to(audio.dtype) - if num_channels > 1: - spectrogram = spectrogram.reshape(batch_size, num_channels * spectrogram.shape[1], spectrogram.shape[2]) - - hidden_states = self.layers(spectrogram) - return hidden_states.transpose(1, 2) - - -# Copied from diffusers.models.autoencoders.autoencoder_oobleck.OobleckResidualUnit with Oobleck->Cosmos3Audio -class Cosmos3AudioResidualUnit(nn.Module): - """ - A residual unit composed of Snake1d and weight-normalized Conv1d layers with dilations. - """ - - def __init__(self, dimension: int = 16, dilation: int = 1): - super().__init__() - pad = ((7 - 1) * dilation) // 2 - - self.snake1 = Snake1d(dimension) - self.conv1 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=7, dilation=dilation, padding=pad)) - self.snake2 = Snake1d(dimension) - self.conv2 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=1)) - - def forward(self, hidden_state): - """ - Forward pass through the residual unit. - - Args: - hidden_state (`torch.Tensor` of shape `(batch_size, channels, time_steps)`): - Input tensor . - - Returns: - output_tensor (`torch.Tensor` of shape `(batch_size, channels, time_steps)`) - Input tensor after passing through the residual unit. - """ - output_tensor = hidden_state - output_tensor = self.conv1(self.snake1(output_tensor)) - output_tensor = self.conv2(self.snake2(output_tensor)) - - padding = (hidden_state.shape[-1] - output_tensor.shape[-1]) // 2 - if padding > 0: - hidden_state = hidden_state[..., padding:-padding] - output_tensor = hidden_state + output_tensor - return output_tensor - - -""" -Copied from diffusers.models.autoencoders.autoencoder_oobleck.OobleckDecoderBlock with Oobleck->Cosmos3Audio with -output_padding enabled. -""" - - -class Cosmos3AudioDecoderBlock(nn.Module): - """Decoder block used in Cosmos3Audio decoder.""" - - def __init__(self, input_dim, output_dim, stride: int = 1, output_padding: int = 0): - super().__init__() - - self.snake1 = Snake1d(input_dim) - self.conv_t1 = weight_norm( - nn.ConvTranspose1d( - input_dim, - output_dim, - kernel_size=2 * stride, - stride=stride, - padding=math.ceil(stride / 2), - output_padding=output_padding, - ) - ) - self.res_unit1 = Cosmos3AudioResidualUnit(output_dim, dilation=1) - self.res_unit2 = Cosmos3AudioResidualUnit(output_dim, dilation=3) - self.res_unit3 = Cosmos3AudioResidualUnit(output_dim, dilation=9) - - def forward(self, hidden_state): - hidden_state = self.snake1(hidden_state) - hidden_state = self.conv_t1(hidden_state) - hidden_state = self.res_unit1(hidden_state) - hidden_state = self.res_unit2(hidden_state) - hidden_state = self.res_unit3(hidden_state) - - return hidden_state - - -""" -Copied from diffusers.models.autoencoders.autoencoder_oobleck.OobleckDecoder with Oobleck->Cosmos3Audio and one change -of adding "output_padding=stride % 2," -""" - - -class Cosmos3AudioDecoder(nn.Module): - """Cosmos3Audio Decoder""" - - def __init__(self, channels, input_channels, audio_channels, upsampling_ratios, channel_multiples): - super().__init__() - - strides = upsampling_ratios - channel_multiples = [1] + channel_multiples - - # Add first conv layer - self.conv1 = weight_norm(nn.Conv1d(input_channels, channels * channel_multiples[-1], kernel_size=7, padding=3)) - - # Add upsampling + MRF blocks - block = [] - for stride_index, stride in enumerate(strides): - block += [ - Cosmos3AudioDecoderBlock( - input_dim=channels * channel_multiples[len(strides) - stride_index], - output_dim=channels * channel_multiples[len(strides) - stride_index - 1], - stride=stride, - output_padding=stride % 2, - ) - ] - - self.block = nn.ModuleList(block) - output_dim = channels - self.snake1 = Snake1d(output_dim) - self.conv2 = weight_norm(nn.Conv1d(channels, audio_channels, kernel_size=7, padding=3, bias=False)) - - def forward(self, hidden_state): - hidden_state = self.conv1(hidden_state) - - for layer in self.block: - hidden_state = layer(hidden_state) - - hidden_state = self.snake1(hidden_state) - hidden_state = self.conv2(hidden_state) - - return hidden_state - - -@dataclass -class Cosmos3AudioEncoderOutput(BaseOutput): - """Output of `Cosmos3AVAEAudioTokenizer.encode`.""" - - latent_dist: OobleckDiagonalGaussianDistribution - - -@dataclass -class Cosmos3AudioDecoderOutput(BaseOutput): - """Output of `Cosmos3AVAEAudioTokenizer.forward`.""" - - sample: torch.Tensor - - -class Cosmos3AVAEAudioTokenizer(ModelMixin, ConfigMixin): - """Audio tokenizer for Cosmos3 sound generation. - - Wraps the Cosmos3 AVAE SpecConvNeXt encoder and Oobleck-style decoder used by the Cosmos3 omni model. The decoder - API stays tensor-returning because ``Cosmos3OmniPipeline`` calls it directly when ``enable_sound=True``. - - Only the shipped AVAE configuration (``model_type="autoencoder_v2"``, waveform input, ``spec_convnext`` encoder, - ``vae`` bottleneck, ``oobleck`` decoder, log-scale SnakeBeta, no latent normalization) is supported; any other - value raises ``NotImplementedError``. - - Parameters: - model_type (`str`, defaults to `"autoencoder_v2"`): AVAE model variant; only `"autoencoder_v2"` is supported. - sampling_rate (`int`, defaults to `48000`): Audio sample rate in Hz. - vocoder_input_dim (`int`, defaults to `64`): Latent channel count fed into the decoder - (``== transformer sound_dim``). - dec_dim (`int`, defaults to `320`): Base decoder channel count. - dec_c_mults (`tuple[int, ...]`, defaults to `(1, 2, 4, 8, 16)`): Decoder channel multipliers. - dec_strides (`tuple[int, ...]`, defaults to `(2, 4, 5, 6, 8)`): Decoder upsampling strides. - dec_out_channels (`int`, defaults to `2`): Output audio channels (2 = stereo). - stereo (`bool`, defaults to `True`): - Whether the audio is stereo; doubles the encoder's effective channel count. - use_wav_as_input (`bool`, defaults to `True`): Whether the encoder consumes raw waveforms; only `True` is - supported. - normalize_volume (`bool`, defaults to `True`): Whether `encode` peak-normalizes the waveform before encoding. - hop_size (`int`, *optional*): Waveform→latent temporal compression factor used for `encode` padding. Defaults - to `prod(dec_strides)` when `None`. - input_channels (`int`, defaults to `1`): Per-channel encoder input count before the `stereo` doubling. - enc_type (`str`, defaults to `"spec_convnext"`): Encoder type; only `"spec_convnext"` is supported. - enc_dim (`int`, defaults to `192`): Base encoder channel count. - enc_intermediate_dim (`int`, defaults to `768`): Unused; kept for config fidelity (ConvNeXt blocks use - ``input_dim * 4``). - enc_num_layers (`int`, defaults to `12`): - Unused; kept for config fidelity (depth derives from `enc_num_blocks`). - enc_num_blocks (`int`, defaults to `2`): ConvNeXt blocks per encoder downsampling stage. - enc_n_fft (`int`, defaults to `64`): STFT FFT size for the encoder spectrogram front-end. - enc_hop_length (`int`, defaults to `16`): STFT hop length for the encoder spectrogram front-end. - enc_latent_dim (`int`, defaults to `128`): - Encoder output channels; split into mean/scale by the VAE bottleneck (so ``enc_latent_dim == 2 * - vocoder_input_dim``). - enc_c_mults (`tuple[int, ...]`, defaults to `(1, 2, 4)`): Encoder channel multipliers per stage. - enc_strides (`tuple[int, ...]`, defaults to `(4, 5, 6)`): Encoder downsampling strides per stage. - enc_identity_init (`bool`, defaults to `False`): Whether to zero-init the ConvNeXt residual 1x1 convs. - enc_use_snake (`bool`, defaults to `True`): Whether ConvNeXt blocks use SnakeBeta (else GELU). - dec_type (`str`, defaults to `"oobleck"`): Decoder type; only `"oobleck"` is supported. - dec_use_snake (`bool`, defaults to `True`): Whether the decoder uses SnakeBeta; only `True` is supported. - dec_final_tanh (`bool`, defaults to `False`): Vestigial decoder tanh flag; only `False` is supported. - dec_anti_aliasing (`bool`, defaults to `False`): Decoder anti-aliasing flag; only `False` is supported. - dec_use_nearest_upsample (`bool`, defaults to `False`): Decoder upsample mode flag; only `False` is supported. - dec_use_tanh_at_final (`bool`, defaults to `False`): Decoder final-tanh flag; only `False` is supported. - bottleneck_type (`str`, defaults to `"vae"`): Bottleneck type; only `"vae"` is supported. - bottleneck (`dict`, *optional*): Bottleneck config; if given, its `"type"` must be `"vae"`. - activation (`str`, defaults to `"snakebeta"`): Activation family; only `"snakebeta"` is supported. - snake_logscale (`bool`, defaults to `True`): Whether SnakeBeta parameters are log-scaled; only `True` is - supported. - anti_aliasing (`bool`, defaults to `False`): Global anti-aliasing flag; only `False` is supported. - use_cuda_kernel (`bool`, defaults to `False`): Whether to use fused CUDA kernels; only `False` is supported. - causal (`bool`, defaults to `False`): - Whether convolutions are causal; only `False` is supported by the encoder. - padding_mode (`str`, defaults to `"zeros"`): Convolution padding mode. - latent_mean (`float` or `list[float]`, *optional*): Latent normalization mean; latent normalization is not - implemented, so a non-`None` value raises ``NotImplementedError``. - latent_std (`float` or `list[float]`, *optional*): Latent normalization std; latent normalization is not - implemented, so a non-`None` value raises ``NotImplementedError``. - encoder_enabled (`bool`, defaults to `True`): Whether to instantiate the encoder. Set to `False` (or - auto-disabled on load) for decoder-only checkpoints, which cannot `encode`. - """ - - _supports_gradient_checkpointing = False - _supports_group_offloading = False - - @register_to_config - def __init__( - self, - model_type: str = "autoencoder_v2", - sampling_rate: int = 48000, - vocoder_input_dim: int = 64, - dec_dim: int = 320, - dec_c_mults: tuple = (1, 2, 4, 8, 16), - dec_strides: tuple = (2, 4, 5, 6, 8), - dec_out_channels: int = 2, - stereo: bool = True, - use_wav_as_input: bool = True, - normalize_volume: bool = True, - hop_size: int | None = None, - input_channels: int = 1, - enc_type: str = "spec_convnext", - enc_dim: int = 192, - enc_intermediate_dim: int = 768, - enc_num_layers: int = 12, - enc_num_blocks: int = 2, - enc_n_fft: int = 64, - enc_hop_length: int = 16, - enc_latent_dim: int = 128, - enc_c_mults: tuple = (1, 2, 4), - enc_strides: tuple = (4, 5, 6), - enc_identity_init: bool = False, - enc_use_snake: bool = True, - dec_type: str = "oobleck", - dec_use_snake: bool = True, - dec_final_tanh: bool = False, - dec_anti_aliasing: bool = False, - dec_use_nearest_upsample: bool = False, - dec_use_tanh_at_final: bool = False, - bottleneck_type: str = "vae", - bottleneck: dict | None = None, - activation: str = "snakebeta", - snake_logscale: bool = True, - anti_aliasing: bool = False, - use_cuda_kernel: bool = False, - causal: bool = False, - padding_mode: str = "zeros", - latent_mean: float | list[float] | None = None, - latent_std: float | list[float] | None = None, - encoder_enabled: bool = True, - ): - super().__init__() - - if model_type != "autoencoder_v2": - raise NotImplementedError(f"Cosmos3 AVAE model type {model_type!r} is not supported.") - if not use_wav_as_input: - raise NotImplementedError("Cosmos3 AVAE tokenizer only supports waveform input.") - if enc_type != "spec_convnext": - raise NotImplementedError(f"Cosmos3 AVAE encoder type {enc_type!r} is not supported.") - if bottleneck is not None and bottleneck.get("type", bottleneck_type) != "vae": - raise NotImplementedError("Cosmos3 AVAE tokenizer only supports the VAE bottleneck.") - if bottleneck_type != "vae": - raise NotImplementedError("Cosmos3 AVAE tokenizer only supports the VAE bottleneck.") - if dec_type != "oobleck": - raise NotImplementedError(f"Cosmos3 AVAE decoder type {dec_type!r} is not supported.") - if ( - not dec_use_snake - or dec_final_tanh - or dec_anti_aliasing - or dec_use_nearest_upsample - or dec_use_tanh_at_final - ): - raise NotImplementedError("Cosmos3 AVAE decoder only supports the shipped Oobleck decoder configuration.") - if activation != "snakebeta" or not snake_logscale or anti_aliasing or use_cuda_kernel: - raise NotImplementedError("Cosmos3 AVAE tokenizer only supports the shipped SnakeBeta configuration.") - if latent_mean is not None or latent_std is not None: - raise NotImplementedError( - "Cosmos3 AVAE tokenizer does not apply latent normalization; `latent_mean`/`latent_std` must be None." - ) - - self.encoder = None - self._encoder_available = False - if encoder_enabled: - self.encoder = Cosmos3AudioSpectrogramConvNeXtEncoder( - input_channels=input_channels, - stereo=stereo, - channels=enc_dim, - latent_dim=enc_latent_dim, - channel_multiples=tuple(enc_c_mults), - strides=tuple(enc_strides), - num_blocks=enc_num_blocks, - n_fft=enc_n_fft, - hop_length=enc_hop_length, - identity_init=enc_identity_init, - use_snake=enc_use_snake, - causal=causal, - padding_mode=padding_mode, - ) - self._encoder_available = True - - self.decoder = Cosmos3AudioDecoder( - channels=dec_dim, - input_channels=vocoder_input_dim, - audio_channels=dec_out_channels, - upsampling_ratios=list(reversed(dec_strides)), - channel_multiples=list(dec_c_mults), - ) - - self._hop_size: int = int(hop_size) if hop_size is not None else math.prod(dec_strides) - - def _disable_encoder(self): - self.encoder = None - self._encoder_available = False - self.register_to_config(encoder_enabled=False) - - def _fix_state_dict_keys_on_load(self, state_dict: OrderedDict) -> None: - super()._fix_state_dict_keys_on_load(state_dict) - if self.encoder is not None and not any(key.startswith("encoder.") for key in state_dict): - self._disable_encoder() - - def _encode(self, sample: torch.Tensor) -> torch.Tensor: - return self.encoder(sample).transpose(1, 2) - - @apply_forward_hook - def encode( - self, - sample: torch.Tensor, - return_dict: bool = True, - force_pad: bool = False, - ) -> Cosmos3AudioEncoderOutput | tuple[OobleckDiagonalGaussianDistribution]: - """Encode a waveform into a VAE latent distribution. - - Args: - sample: Audio waveform tensor with shape ``[B, C, T]``. - return_dict: Whether to return a ``Cosmos3AudioEncoderOutput``. - force_pad: Whether to right-pad to ``hop_size`` even when the model is in training mode. - """ - if sample.ndim != 3: - raise ValueError(f"`sample` must have shape [B, C, T], got {tuple(sample.shape)}.") - - if self.encoder is None or not self._encoder_available: - raise ValueError( - "This Cosmos3 AVAE sound tokenizer was loaded from decoder-only weights and cannot encode audio. " - "Re-convert the AVAE checkpoint with encoder weights to use `encode()`." - ) - - hidden_states = sample - if self.config.normalize_volume: - hidden_states = hidden_states / (hidden_states.abs().max() + 1e-5) * 0.95 - - if force_pad or not self.training: - sample_length = hidden_states.shape[-1] - padding = (self._hop_size - (sample_length % self._hop_size)) % self._hop_size - if padding > 0: - hidden_states = F.pad(hidden_states, (0, padding), mode="constant", value=0) - - encoder_dtype = get_parameter_dtype(self.encoder) - moments = self._encode(hidden_states.to(dtype=encoder_dtype)) - posterior = OobleckDiagonalGaussianDistribution(moments) - - if not return_dict: - return (posterior,) - - return Cosmos3AudioEncoderOutput(latent_dist=posterior) - - @apply_forward_hook - def decode(self, latents: torch.Tensor) -> torch.Tensor: - """Decode sound latents into an audio waveform. - - Args: - latents: ``[B, C, T]`` or ``[C, T]`` tensor of diffusion-model latents. - - Returns: - Waveform tensor ``[B, audio_channels, N]`` or ``[audio_channels, N]``. - """ - squeeze = latents.ndim == 2 - if squeeze: - latents = latents.unsqueeze(0) - audio = self.decoder(latents).clamp(-1.0, 1.0) - return audio.squeeze(0) if squeeze else audio - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - force_pad: bool = False, - ) -> Cosmos3AudioDecoderOutput | tuple[torch.Tensor]: - r""" - Encode then decode a waveform. `sample_posterior=False` (default) decodes the distribution mode (mean), whereas - the upstream Cosmos3 AVAE always samples; pass `sample_posterior=True` for reference-equivalent behavior. - - Args: - sample (`torch.Tensor`): - Input waveform sample with shape `(batch_size, audio_channels, num_samples)`. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior instead of decoding the distribution mode. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`Cosmos3AudioDecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - force_pad (`bool`, *optional*, defaults to `False`): - Whether to right-pad the waveform to `hop_size` before encoding even when the model is in training - mode. - - Returns: - [`Cosmos3AudioDecoderOutput`] or `tuple`: - If `return_dict` is True, a [`Cosmos3AudioDecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - posterior = self.encode(sample, force_pad=force_pad).latent_dist - latents = posterior.sample(generator=generator) if sample_posterior else posterior.mode() - decoded = self.decode(latents) - - if not return_dict: - return (decoded,) - - return Cosmos3AudioDecoderOutput(sample=decoded) diff --git a/diffusers/models/autoencoders/autoencoder_dc.py b/diffusers/models/autoencoders/autoencoder_dc.py deleted file mode 100644 index 859a4a6850b28b3e5b1b9098836c388d2e401847..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_dc.py +++ /dev/null @@ -1,724 +0,0 @@ -# Copyright 2025 MIT, Tsinghua University, NVIDIA CORPORATION and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..attention_processor import SanaMultiscaleLinearAttention -from ..modeling_utils import ModelMixin -from ..normalization import RMSNorm, get_normalization -from ..transformers.sana_transformer import GLUMBConv -from .vae import AutoencoderMixin, DecoderOutput, EncoderOutput - - -class ResBlock(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - norm_type: str = "batch_norm", - act_fn: str = "relu6", - ) -> None: - super().__init__() - - self.norm_type = norm_type - - self.nonlinearity = get_activation(act_fn) if act_fn is not None else nn.Identity() - self.conv1 = nn.Conv2d(in_channels, in_channels, 3, 1, 1) - self.conv2 = nn.Conv2d(in_channels, out_channels, 3, 1, 1, bias=False) - self.norm = get_normalization(norm_type, out_channels) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - residual = hidden_states - hidden_states = self.conv1(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv2(hidden_states) - - if self.norm_type == "rms_norm": - # move channel to the last dimension so we apply RMSnorm across channel dimension - hidden_states = self.norm(hidden_states.movedim(1, -1)).movedim(-1, 1) - else: - hidden_states = self.norm(hidden_states) - - return hidden_states + residual - - -class EfficientViTBlock(nn.Module): - def __init__( - self, - in_channels: int, - mult: float = 1.0, - attention_head_dim: int = 32, - qkv_multiscales: tuple[int, ...] = (5,), - norm_type: str = "batch_norm", - ) -> None: - super().__init__() - - self.attn = SanaMultiscaleLinearAttention( - in_channels=in_channels, - out_channels=in_channels, - mult=mult, - attention_head_dim=attention_head_dim, - norm_type=norm_type, - kernel_sizes=qkv_multiscales, - residual_connection=True, - ) - - self.conv_out = GLUMBConv( - in_channels=in_channels, - out_channels=in_channels, - norm_type="rms_norm", - ) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = self.attn(x) - x = self.conv_out(x) - return x - - -def get_block( - block_type: str, - in_channels: int, - out_channels: int, - attention_head_dim: int, - norm_type: str, - act_fn: str, - qkv_multiscales: tuple[int, ...] = (), -): - if block_type == "ResBlock": - block = ResBlock(in_channels, out_channels, norm_type, act_fn) - - elif block_type == "EfficientViTBlock": - block = EfficientViTBlock( - in_channels, attention_head_dim=attention_head_dim, norm_type=norm_type, qkv_multiscales=qkv_multiscales - ) - - else: - raise ValueError(f"Block with {block_type=} is not supported.") - - return block - - -class DCDownBlock2d(nn.Module): - def __init__(self, in_channels: int, out_channels: int, downsample: bool = False, shortcut: bool = True) -> None: - super().__init__() - - self.downsample = downsample - self.factor = 2 - self.stride = 1 if downsample else 2 - self.group_size = in_channels * self.factor**2 // out_channels - self.shortcut = shortcut - - out_ratio = self.factor**2 - if downsample: - assert out_channels % out_ratio == 0 - out_channels = out_channels // out_ratio - - self.conv = nn.Conv2d( - in_channels, - out_channels, - kernel_size=3, - stride=self.stride, - padding=1, - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - x = self.conv(hidden_states) - if self.downsample: - x = F.pixel_unshuffle(x, self.factor) - - if self.shortcut: - y = F.pixel_unshuffle(hidden_states, self.factor) - y = y.unflatten(1, (-1, self.group_size)) - y = y.mean(dim=2) - hidden_states = x + y - else: - hidden_states = x - - return hidden_states - - -class DCUpBlock2d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - interpolate: bool = False, - shortcut: bool = True, - interpolation_mode: str = "nearest", - ) -> None: - super().__init__() - - self.interpolate = interpolate - self.interpolation_mode = interpolation_mode - self.shortcut = shortcut - self.factor = 2 - self.repeats = out_channels * self.factor**2 // in_channels - - out_ratio = self.factor**2 - - if not interpolate: - out_channels = out_channels * out_ratio - - self.conv = nn.Conv2d(in_channels, out_channels, 3, 1, 1) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.interpolate: - x = F.interpolate(hidden_states, scale_factor=self.factor, mode=self.interpolation_mode) - x = self.conv(x) - else: - x = self.conv(hidden_states) - x = F.pixel_shuffle(x, self.factor) - - if self.shortcut: - y = hidden_states.repeat_interleave(self.repeats, dim=1, output_size=hidden_states.shape[1] * self.repeats) - y = F.pixel_shuffle(y, self.factor) - hidden_states = x + y - else: - hidden_states = x - - return hidden_states - - -class Encoder(nn.Module): - def __init__( - self, - in_channels: int, - latent_channels: int, - attention_head_dim: int = 32, - block_type: str | tuple[str] = "ResBlock", - block_out_channels: tuple[int, ...] = (128, 256, 512, 512, 1024, 1024), - layers_per_block: tuple[int, ...] = (2, 2, 2, 2, 2, 2), - qkv_multiscales: tuple[tuple[int, ...], ...] = ((), (), (), (5,), (5,), (5,)), - downsample_block_type: str = "pixel_unshuffle", - out_shortcut: bool = True, - ): - super().__init__() - - num_blocks = len(block_out_channels) - - if isinstance(block_type, str): - block_type = (block_type,) * num_blocks - - if layers_per_block[0] > 0: - self.conv_in = nn.Conv2d( - in_channels, - block_out_channels[0] if layers_per_block[0] > 0 else block_out_channels[1], - kernel_size=3, - stride=1, - padding=1, - ) - else: - self.conv_in = DCDownBlock2d( - in_channels=in_channels, - out_channels=block_out_channels[0] if layers_per_block[0] > 0 else block_out_channels[1], - downsample=downsample_block_type == "pixel_unshuffle", - shortcut=False, - ) - - down_blocks = [] - for i, (out_channel, num_layers) in enumerate(zip(block_out_channels, layers_per_block)): - down_block_list = [] - - for _ in range(num_layers): - block = get_block( - block_type[i], - out_channel, - out_channel, - attention_head_dim=attention_head_dim, - norm_type="rms_norm", - act_fn="silu", - qkv_multiscales=qkv_multiscales[i], - ) - down_block_list.append(block) - - if i < num_blocks - 1 and num_layers > 0: - downsample_block = DCDownBlock2d( - in_channels=out_channel, - out_channels=block_out_channels[i + 1], - downsample=downsample_block_type == "pixel_unshuffle", - shortcut=True, - ) - down_block_list.append(downsample_block) - - down_blocks.append(nn.Sequential(*down_block_list)) - - self.down_blocks = nn.ModuleList(down_blocks) - - self.conv_out = nn.Conv2d(block_out_channels[-1], latent_channels, 3, 1, 1) - - self.out_shortcut = out_shortcut - if out_shortcut: - self.out_shortcut_average_group_size = block_out_channels[-1] // latent_channels - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv_in(hidden_states) - for down_block in self.down_blocks: - hidden_states = down_block(hidden_states) - - if self.out_shortcut: - x = hidden_states.unflatten(1, (-1, self.out_shortcut_average_group_size)) - x = x.mean(dim=2) - hidden_states = self.conv_out(hidden_states) + x - else: - hidden_states = self.conv_out(hidden_states) - - return hidden_states - - -class Decoder(nn.Module): - def __init__( - self, - in_channels: int, - latent_channels: int, - attention_head_dim: int = 32, - block_type: str | tuple[str] = "ResBlock", - block_out_channels: tuple[int, ...] = (128, 256, 512, 512, 1024, 1024), - layers_per_block: tuple[int, ...] = (2, 2, 2, 2, 2, 2), - qkv_multiscales: tuple[tuple[int, ...], ...] = ((), (), (), (5,), (5,), (5,)), - norm_type: str | tuple[str] = "rms_norm", - act_fn: str | tuple[str] = "silu", - upsample_block_type: str = "pixel_shuffle", - in_shortcut: bool = True, - conv_act_fn: str = "relu", - ): - super().__init__() - - num_blocks = len(block_out_channels) - - if isinstance(block_type, str): - block_type = (block_type,) * num_blocks - if isinstance(norm_type, str): - norm_type = (norm_type,) * num_blocks - if isinstance(act_fn, str): - act_fn = (act_fn,) * num_blocks - - self.conv_in = nn.Conv2d(latent_channels, block_out_channels[-1], 3, 1, 1) - - self.in_shortcut = in_shortcut - if in_shortcut: - self.in_shortcut_repeats = block_out_channels[-1] // latent_channels - - up_blocks = [] - for i, (out_channel, num_layers) in reversed(list(enumerate(zip(block_out_channels, layers_per_block)))): - up_block_list = [] - - if i < num_blocks - 1 and num_layers > 0: - upsample_block = DCUpBlock2d( - block_out_channels[i + 1], - out_channel, - interpolate=upsample_block_type == "interpolate", - shortcut=True, - ) - up_block_list.append(upsample_block) - - for _ in range(num_layers): - block = get_block( - block_type[i], - out_channel, - out_channel, - attention_head_dim=attention_head_dim, - norm_type=norm_type[i], - act_fn=act_fn[i], - qkv_multiscales=qkv_multiscales[i], - ) - up_block_list.append(block) - - up_blocks.insert(0, nn.Sequential(*up_block_list)) - - self.up_blocks = nn.ModuleList(up_blocks) - - channels = block_out_channels[0] if layers_per_block[0] > 0 else block_out_channels[1] - - self.norm_out = RMSNorm(channels, 1e-5, elementwise_affine=True, bias=True) - self.conv_act = get_activation(conv_act_fn) - self.conv_out = None - - if layers_per_block[0] > 0: - self.conv_out = nn.Conv2d(channels, in_channels, 3, 1, 1) - else: - self.conv_out = DCUpBlock2d( - channels, in_channels, interpolate=upsample_block_type == "interpolate", shortcut=False - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.in_shortcut: - x = hidden_states.repeat_interleave( - self.in_shortcut_repeats, dim=1, output_size=hidden_states.shape[1] * self.in_shortcut_repeats - ) - hidden_states = self.conv_in(hidden_states) + x - else: - hidden_states = self.conv_in(hidden_states) - - for up_block in reversed(self.up_blocks): - hidden_states = up_block(hidden_states) - - hidden_states = self.norm_out(hidden_states.movedim(1, -1)).movedim(-1, 1) - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states) - return hidden_states - - -class AutoencoderDC(ModelMixin, AutoencoderMixin, ConfigMixin, FromOriginalModelMixin): - r""" - An Autoencoder model introduced in [DCAE](https://huggingface.co/papers/2410.10733) and used in - [SANA](https://huggingface.co/papers/2410.10629). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Args: - in_channels (`int`, defaults to `3`): - The number of input channels in samples. - latent_channels (`int`, defaults to `32`): - The number of channels in the latent space representation. - encoder_block_types (`str | tuple[str]`, defaults to `"ResBlock"`): - The type(s) of block to use in the encoder. - decoder_block_types (`str | tuple[str]`, defaults to `"ResBlock"`): - The type(s) of block to use in the decoder. - encoder_block_out_channels (`tuple[int, ...]`, defaults to `(128, 256, 512, 512, 1024, 1024)`): - The number of output channels for each block in the encoder. - decoder_block_out_channels (`tuple[int, ...]`, defaults to `(128, 256, 512, 512, 1024, 1024)`): - The number of output channels for each block in the decoder. - encoder_layers_per_block (`tuple[int]`, defaults to `(2, 2, 2, 3, 3, 3)`): - The number of layers per block in the encoder. - decoder_layers_per_block (`tuple[int]`, defaults to `(3, 3, 3, 3, 3, 3)`): - The number of layers per block in the decoder. - encoder_qkv_multiscales (`tuple[tuple[int, ...], ...]`, defaults to `((), (), (), (5,), (5,), (5,))`): - Multi-scale configurations for the encoder's QKV (query-key-value) transformations. - decoder_qkv_multiscales (`tuple[tuple[int, ...], ...]`, defaults to `((), (), (), (5,), (5,), (5,))`): - Multi-scale configurations for the decoder's QKV (query-key-value) transformations. - upsample_block_type (`str`, defaults to `"pixel_shuffle"`): - The type of block to use for upsampling in the decoder. - downsample_block_type (`str`, defaults to `"pixel_unshuffle"`): - The type of block to use for downsampling in the encoder. - decoder_norm_types (`str | tuple[str]`, defaults to `"rms_norm"`): - The normalization type(s) to use in the decoder. - decoder_act_fns (`str | tuple[str]`, defaults to `"silu"`): - The activation function(s) to use in the decoder. - encoder_out_shortcut (`bool`, defaults to `True`): - Whether to use shortcut at the end of the encoder. - decoder_in_shortcut (`bool`, defaults to `True`): - Whether to use shortcut at the beginning of the decoder. - decoder_conv_act_fn (`str`, defaults to `"relu"`): - The activation function to use at the end of the decoder. - scaling_factor (`float`, defaults to `1.0`): - The multiplicative inverse of the root mean square of the latent features. This is used to scale the latent - space to have unit variance when training the diffusion model. The latents are scaled with the formula `z = - z * scaling_factor` before being passed to the diffusion model. When decoding, the latents are scaled back - to the original scale with the formula: `z = 1 / scaling_factor * z`. - """ - - _supports_gradient_checkpointing = False - - @register_to_config - def __init__( - self, - in_channels: int = 3, - latent_channels: int = 32, - attention_head_dim: int = 32, - encoder_block_types: str | tuple[str] = "ResBlock", - decoder_block_types: str | tuple[str] = "ResBlock", - encoder_block_out_channels: tuple[int, ...] = (128, 256, 512, 512, 1024, 1024), - decoder_block_out_channels: tuple[int, ...] = (128, 256, 512, 512, 1024, 1024), - encoder_layers_per_block: tuple[int, ...] = (2, 2, 2, 3, 3, 3), - decoder_layers_per_block: tuple[int, ...] = (3, 3, 3, 3, 3, 3), - encoder_qkv_multiscales: tuple[tuple[int, ...], ...] = ((), (), (), (5,), (5,), (5,)), - decoder_qkv_multiscales: tuple[tuple[int, ...], ...] = ((), (), (), (5,), (5,), (5,)), - upsample_block_type: str = "pixel_shuffle", - downsample_block_type: str = "pixel_unshuffle", - decoder_norm_types: str | tuple[str] = "rms_norm", - decoder_act_fns: str | tuple[str] = "silu", - encoder_out_shortcut: bool = True, - decoder_in_shortcut: bool = True, - decoder_conv_act_fn: str = "relu", - scaling_factor: float = 1.0, - ) -> None: - super().__init__() - - self.encoder = Encoder( - in_channels=in_channels, - latent_channels=latent_channels, - attention_head_dim=attention_head_dim, - block_type=encoder_block_types, - block_out_channels=encoder_block_out_channels, - layers_per_block=encoder_layers_per_block, - qkv_multiscales=encoder_qkv_multiscales, - downsample_block_type=downsample_block_type, - out_shortcut=encoder_out_shortcut, - ) - self.decoder = Decoder( - in_channels=in_channels, - latent_channels=latent_channels, - attention_head_dim=attention_head_dim, - block_type=decoder_block_types, - block_out_channels=decoder_block_out_channels, - layers_per_block=decoder_layers_per_block, - qkv_multiscales=decoder_qkv_multiscales, - norm_type=decoder_norm_types, - act_fn=decoder_act_fns, - upsample_block_type=upsample_block_type, - in_shortcut=decoder_in_shortcut, - conv_act_fn=decoder_conv_act_fn, - ) - - self.spatial_compression_ratio = 2 ** (len(encoder_block_out_channels) - 1) - self.temporal_compression_ratio = 1 - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent - # frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the - # intermediate tiles together, the memory requirement can be lowered. - self.use_tiling = False - - # The minimal tile height and width for spatial tiling to be used - self.tile_sample_min_height = 512 - self.tile_sample_min_width = 512 - - # The minimal distance between two spatial tiles - self.tile_sample_stride_height = 448 - self.tile_sample_stride_width = 448 - - self.tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - self.tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_sample_stride_height: float | None = None, - tile_sample_stride_width: float | None = None, - ) -> None: - r""" - Enable tiled AE decoding. When this option is enabled, the AE will split the input tensor into tiles to compute - decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_sample_stride_height (`int`, *optional*): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. - tile_sample_stride_width (`int`, *optional*): - The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling - artifacts produced across the width dimension. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height - self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width - self.tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - self.tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, height, width = x.shape - - if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): - return self.tiled_encode(x, return_dict=False)[0] - - encoded = self.encoder(x) - - return encoded - - @apply_forward_hook - def encode(self, x: torch.Tensor, return_dict: bool = True) -> EncoderOutput | tuple[torch.Tensor]: - r""" - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, defaults to `True`): - Whether to return a [`~models.vae.EncoderOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.vae.EncoderOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - encoded = torch.cat(encoded_slices) - else: - encoded = self._encode(x) - - if not return_dict: - return (encoded,) - return EncoderOutput(latent=encoded) - - def _decode(self, z: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, height, width = z.shape - - if self.use_tiling and (width > self.tile_latent_min_width or height > self.tile_latent_min_height): - return self.tiled_decode(z, return_dict=False)[0] - - decoded = self.decoder(z) - - return decoded - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | tuple[torch.Tensor]: - r""" - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.size(0) > 1: - decoded_slices = [self._decode(z_slice) for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z) - - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[2], b.shape[2], blend_extent) - for y in range(blend_extent): - b[:, :, y, :] = a[:, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, y, :] * (y / blend_extent) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[3], b.shape[3], blend_extent) - for x in range(blend_extent): - b[:, :, :, x] = a[:, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, x] * (x / blend_extent) - return b - - def tiled_encode(self, x: torch.Tensor, return_dict: bool = True) -> torch.Tensor: - batch_size, num_channels, height, width = x.shape - latent_height = height // self.spatial_compression_ratio - latent_width = width // self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - blend_height = tile_latent_min_height - tile_latent_stride_height - blend_width = tile_latent_min_width - tile_latent_stride_width - - # Split x into overlapping tiles and encode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, x.shape[2], self.tile_sample_stride_height): - row = [] - for j in range(0, x.shape[3], self.tile_sample_stride_width): - tile = x[:, :, i : i + self.tile_sample_min_height, j : j + self.tile_sample_min_width] - if ( - tile.shape[2] % self.spatial_compression_ratio != 0 - or tile.shape[3] % self.spatial_compression_ratio != 0 - ): - pad_h = (self.spatial_compression_ratio - tile.shape[2]) % self.spatial_compression_ratio - pad_w = (self.spatial_compression_ratio - tile.shape[3]) % self.spatial_compression_ratio - tile = F.pad(tile, (0, pad_w, 0, pad_h)) - tile = self.encoder(tile) - row.append(tile) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :tile_latent_stride_height, :tile_latent_stride_width]) - result_rows.append(torch.cat(result_row, dim=3)) - - encoded = torch.cat(result_rows, dim=2)[:, :, :latent_height, :latent_width] - - if not return_dict: - return (encoded,) - return EncoderOutput(latent=encoded) - - def tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - batch_size, num_channels, height, width = z.shape - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = self.tile_sample_min_height - self.tile_sample_stride_height - blend_width = self.tile_sample_min_width - self.tile_sample_stride_width - - # Split z into overlapping tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, tile_latent_stride_height): - row = [] - for j in range(0, width, tile_latent_stride_width): - tile = z[:, :, i : i + tile_latent_min_height, j : j + tile_latent_min_width] - decoded = self.decoder(tile) - row.append(decoded) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, : self.tile_sample_stride_height, : self.tile_sample_stride_width]) - result_rows.append(torch.cat(result_row, dim=3)) - - decoded = torch.cat(result_rows, dim=2) - - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def forward(self, sample: torch.Tensor, return_dict: bool = True) -> torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - encoded = self.encode(sample, return_dict=False)[0] - decoded = self.decode(encoded, return_dict=False)[0] - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) diff --git a/diffusers/models/autoencoders/autoencoder_kl.py b/diffusers/models/autoencoders/autoencoder_kl.py deleted file mode 100644 index 4434b735d949a05d36a6a0116e1c60138ff1ec6e..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl.py +++ /dev/null @@ -1,479 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...loaders.single_file_model import FromOriginalModelMixin -from ...utils import deprecate -from ...utils.accelerate_utils import apply_forward_hook -from ..attention import AttentionMixin -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - Attention, - AttnAddedKVProcessor, - AttnProcessor, - FusedAttnProcessor2_0, -) -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, Decoder, DecoderOutput, DiagonalGaussianDistribution, Encoder - - -class AutoencoderKL( - ModelMixin, AttentionMixin, AutoencoderMixin, ConfigMixin, FromOriginalModelMixin, PeftAdapterMixin -): - r""" - A VAE model with KL loss for encoding images into latents and decoding latent representations into images. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - in_channels (int, *optional*, defaults to 3): Number of channels in the input image. - out_channels (int, *optional*, defaults to 3): Number of channels in the output. - down_block_types (`tuple[str]`, *optional*, defaults to `("DownEncoderBlock2D",)`): - tuple of downsample block types. - up_block_types (`tuple[str]`, *optional*, defaults to `("UpDecoderBlock2D",)`): - tuple of upsample block types. - block_out_channels (`tuple[int]`, *optional*, defaults to `(64,)`): - tuple of block output channels. - act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use. - latent_channels (`int`, *optional*, defaults to 4): Number of channels in the latent space. - sample_size (`int`, *optional*, defaults to `32`): Sample input size. - scaling_factor (`float`, *optional*, defaults to 0.18215): - The component-wise standard deviation of the trained latent space computed using the first batch of the - training set. This is used to scale the latent space to have unit variance when training the diffusion - model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the - diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 - / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image - Synthesis with Latent Diffusion Models](https://huggingface.co/papers/2112.10752) paper. - force_upcast (`bool`, *optional*, default to `True`): - If enabled it will force the VAE to run in float32 for high image resolution pipelines, such as SD-XL. VAE - can be fine-tuned / trained to a lower range without losing too much precision in which case `force_upcast` - can be set to `False` - see: https://huggingface.co/madebyollin/sdxl-vae-fp16-fix - mid_block_add_attention (`bool`, *optional*, default to `True`): - If enabled, the mid_block of the Encoder and Decoder will have attention blocks. If set to false, the - mid_block will only have resnet blocks - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["BasicTransformerBlock", "ResnetBlock2D"] - _group_offload_block_modules = ["quant_conv", "post_quant_conv", "encoder", "decoder"] - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - down_block_types: tuple[str] = ("DownEncoderBlock2D",), - up_block_types: tuple[str] = ("UpDecoderBlock2D",), - block_out_channels: tuple[int] = (64,), - layers_per_block: int = 1, - act_fn: str = "silu", - latent_channels: int = 4, - norm_num_groups: int = 32, - sample_size: int = 32, - scaling_factor: float = 0.18215, - shift_factor: float | None = None, - latents_mean: tuple[float] | None = None, - latents_std: tuple[float] | None = None, - force_upcast: bool = True, - use_quant_conv: bool = True, - use_post_quant_conv: bool = True, - mid_block_add_attention: bool = True, - ): - super().__init__() - - # pass init params to Encoder - self.encoder = Encoder( - in_channels=in_channels, - out_channels=latent_channels, - down_block_types=down_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - double_z=True, - mid_block_add_attention=mid_block_add_attention, - ) - - # pass init params to Decoder - self.decoder = Decoder( - in_channels=latent_channels, - out_channels=out_channels, - up_block_types=up_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - norm_num_groups=norm_num_groups, - act_fn=act_fn, - mid_block_add_attention=mid_block_add_attention, - ) - - self.quant_conv = nn.Conv2d(2 * latent_channels, 2 * latent_channels, 1) if use_quant_conv else None - self.post_quant_conv = nn.Conv2d(latent_channels, latent_channels, 1) if use_post_quant_conv else None - - self.use_slicing = False - self.use_tiling = False - - # only relevant if vae tiling is enabled - self.tile_sample_min_size = self.config.sample_size - sample_size = ( - self.config.sample_size[0] - if isinstance(self.config.sample_size, (list, tuple)) - else self.config.sample_size - ) - self.tile_latent_min_size = int(sample_size / (2 ** (len(self.config.block_out_channels) - 1))) - self.tile_overlap_factor = 0.25 - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, height, width = x.shape - - if self.use_tiling and (width > self.tile_sample_min_size or height > self.tile_sample_min_size): - return self._tiled_encode(x) - - enc = self.encoder(x) - if self.quant_conv is not None: - enc = self.quant_conv(enc) - - return enc - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded images. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - if self.use_tiling and (z.shape[-1] > self.tile_latent_min_size or z.shape[-2] > self.tile_latent_min_size): - return self.tiled_decode(z, return_dict=return_dict) - - if self.post_quant_conv is not None: - z = self.post_quant_conv(z) - - dec = self.decoder(z) - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode( - self, z: torch.FloatTensor, return_dict: bool = True, generator=None - ) -> DecoderOutput | torch.FloatTensor: - """ - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[2], b.shape[2], blend_extent) - for y in range(blend_extent): - b[:, :, y, :] = a[:, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, y, :] * (y / blend_extent) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[3], b.shape[3], blend_extent) - for x in range(blend_extent): - b[:, :, :, x] = a[:, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, x] * (x / blend_extent) - return b - - def _tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - r"""Encode a batch of images using a tiled encoder. - - When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several - steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is - different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the - tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the - output, but they should be much less noticeable. - - Args: - x (`torch.Tensor`): Input batch of images. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - - overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor)) - blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor) - row_limit = self.tile_latent_min_size - blend_extent - - # Split the image into 512x512 tiles and encode them separately. - rows = [] - for i in range(0, x.shape[2], overlap_size): - row = [] - for j in range(0, x.shape[3], overlap_size): - tile = x[:, :, i : i + self.tile_sample_min_size, j : j + self.tile_sample_min_size] - tile = self.encoder(tile) - if self.config.use_quant_conv: - tile = self.quant_conv(tile) - row.append(tile) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=3)) - - enc = torch.cat(result_rows, dim=2) - return enc - - def tiled_encode(self, x: torch.Tensor, return_dict: bool = True) -> AutoencoderKLOutput: - r"""Encode a batch of images using a tiled encoder. - - When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several - steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is - different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the - tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the - output, but they should be much less noticeable. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - [`~models.autoencoder_kl.AutoencoderKLOutput`] or `tuple`: - If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain - `tuple` is returned. - """ - deprecation_message = ( - "The tiled_encode implementation supporting the `return_dict` parameter is deprecated. In the future, the " - "implementation of this method will be replaced with that of `_tiled_encode` and you will no longer be able " - "to pass `return_dict`. You will also have to create a `DiagonalGaussianDistribution()` from the returned value." - ) - deprecate("tiled_encode", "1.0.0", deprecation_message, standard_warn=False) - - overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor)) - blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor) - row_limit = self.tile_latent_min_size - blend_extent - - # Split the image into 512x512 tiles and encode them separately. - rows = [] - for i in range(0, x.shape[2], overlap_size): - row = [] - for j in range(0, x.shape[3], overlap_size): - tile = x[:, :, i : i + self.tile_sample_min_size, j : j + self.tile_sample_min_size] - tile = self.encoder(tile) - if self.config.use_quant_conv: - tile = self.quant_conv(tile) - row.append(tile) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=3)) - - moments = torch.cat(result_rows, dim=2) - posterior = DiagonalGaussianDistribution(moments) - - if not return_dict: - return (posterior,) - - return AutoencoderKLOutput(latent_dist=posterior) - - def tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor)) - blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor) - row_limit = self.tile_sample_min_size - blend_extent - - # Split z into overlapping 64x64 tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, z.shape[2], overlap_size): - row = [] - for j in range(0, z.shape[3], overlap_size): - tile = z[:, :, i : i + self.tile_latent_min_size, j : j + self.tile_latent_min_size] - if self.config.use_post_quant_conv: - tile = self.post_quant_conv(tile) - decoded = self.decoder(tile) - row.append(decoded) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=3)) - - dec = torch.cat(result_rows, dim=2) - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z).sample - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) diff --git a/diffusers/models/autoencoders/autoencoder_kl_allegro.py b/diffusers/models/autoencoders/autoencoder_kl_allegro.py deleted file mode 100644 index 5983c08a6f8660db13d49f07190b3f800868caf7..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_allegro.py +++ /dev/null @@ -1,1107 +0,0 @@ -# Copyright 2025 The RhymesAI and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils.accelerate_utils import apply_forward_hook -from ..attention_processor import Attention, SpatialNorm -from ..autoencoders.vae import DecoderOutput, DiagonalGaussianDistribution -from ..downsampling import Downsample2D -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from ..resnet import ResnetBlock2D -from ..upsampling import Upsample2D -from .vae import AutoencoderMixin - - -class AllegroTemporalConvLayer(nn.Module): - r""" - Temporal convolutional layer that can be used for video (sequence of images) input. Code adapted from: - https://github.com/modelscope/modelscope/blob/1509fdb973e5871f37148a4b5e5964cafd43e64d/modelscope/models/multi_modal/video_synthesis/unet_sd.py#L1016 - """ - - def __init__( - self, - in_dim: int, - out_dim: int | None = None, - dropout: float = 0.0, - norm_num_groups: int = 32, - up_sample: bool = False, - down_sample: bool = False, - stride: int = 1, - ) -> None: - super().__init__() - - out_dim = out_dim or in_dim - pad_h = pad_w = int((stride - 1) * 0.5) - pad_t = 0 - - self.down_sample = down_sample - self.up_sample = up_sample - - if down_sample: - self.conv1 = nn.Sequential( - nn.GroupNorm(norm_num_groups, in_dim), - nn.SiLU(), - nn.Conv3d(in_dim, out_dim, (2, stride, stride), stride=(2, 1, 1), padding=(0, pad_h, pad_w)), - ) - elif up_sample: - self.conv1 = nn.Sequential( - nn.GroupNorm(norm_num_groups, in_dim), - nn.SiLU(), - nn.Conv3d(in_dim, out_dim * 2, (1, stride, stride), padding=(0, pad_h, pad_w)), - ) - else: - self.conv1 = nn.Sequential( - nn.GroupNorm(norm_num_groups, in_dim), - nn.SiLU(), - nn.Conv3d(in_dim, out_dim, (3, stride, stride), padding=(pad_t, pad_h, pad_w)), - ) - self.conv2 = nn.Sequential( - nn.GroupNorm(norm_num_groups, out_dim), - nn.SiLU(), - nn.Dropout(dropout), - nn.Conv3d(out_dim, in_dim, (3, stride, stride), padding=(pad_t, pad_h, pad_w)), - ) - self.conv3 = nn.Sequential( - nn.GroupNorm(norm_num_groups, out_dim), - nn.SiLU(), - nn.Dropout(dropout), - nn.Conv3d(out_dim, in_dim, (3, stride, stride), padding=(pad_t, pad_h, pad_h)), - ) - self.conv4 = nn.Sequential( - nn.GroupNorm(norm_num_groups, out_dim), - nn.SiLU(), - nn.Conv3d(out_dim, in_dim, (3, stride, stride), padding=(pad_t, pad_h, pad_h)), - ) - - @staticmethod - def _pad_temporal_dim(hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = torch.cat((hidden_states[:, :, 0:1], hidden_states), dim=2) - hidden_states = torch.cat((hidden_states, hidden_states[:, :, -1:]), dim=2) - return hidden_states - - def forward(self, hidden_states: torch.Tensor, batch_size: int) -> torch.Tensor: - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - - if self.down_sample: - identity = hidden_states[:, :, ::2] - elif self.up_sample: - identity = hidden_states.repeat_interleave(2, dim=2, output_size=hidden_states.shape[2] * 2) - else: - identity = hidden_states - - if self.down_sample or self.up_sample: - hidden_states = self.conv1(hidden_states) - else: - hidden_states = self._pad_temporal_dim(hidden_states) - hidden_states = self.conv1(hidden_states) - - if self.up_sample: - hidden_states = hidden_states.unflatten(1, (2, -1)).permute(0, 2, 3, 1, 4, 5).flatten(2, 3) - - hidden_states = self._pad_temporal_dim(hidden_states) - hidden_states = self.conv2(hidden_states) - - hidden_states = self._pad_temporal_dim(hidden_states) - hidden_states = self.conv3(hidden_states) - - hidden_states = self._pad_temporal_dim(hidden_states) - hidden_states = self.conv4(hidden_states) - - hidden_states = identity + hidden_states - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) - - return hidden_states - - -class AllegroDownBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - spatial_downsample: bool = True, - temporal_downsample: bool = False, - downsample_padding: int = 1, - ): - super().__init__() - - resnets = [] - temp_convs = [] - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=None, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - temp_convs.append( - AllegroTemporalConvLayer( - out_channels, - out_channels, - dropout=0.1, - norm_num_groups=resnet_groups, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.temp_convs = nn.ModuleList(temp_convs) - - if temporal_downsample: - self.temp_convs_down = AllegroTemporalConvLayer( - out_channels, out_channels, dropout=0.1, norm_num_groups=resnet_groups, down_sample=True, stride=3 - ) - self.add_temp_downsample = temporal_downsample - - if spatial_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op" - ) - ] - ) - else: - self.downsamplers = None - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size = hidden_states.shape[0] - - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) - - for resnet, temp_conv in zip(self.resnets, self.temp_convs): - hidden_states = resnet(hidden_states, temb=None) - hidden_states = temp_conv(hidden_states, batch_size=batch_size) - - if self.add_temp_downsample: - hidden_states = self.temp_convs_down(hidden_states, batch_size=batch_size) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - return hidden_states - - -class AllegroUpBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", # default, spatial - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - spatial_upsample: bool = True, - temporal_upsample: bool = False, - temb_channels: int | None = None, - ): - super().__init__() - - resnets = [] - temp_convs = [] - - for i in range(num_layers): - input_channels = in_channels if i == 0 else out_channels - - resnets.append( - ResnetBlock2D( - in_channels=input_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - temp_convs.append( - AllegroTemporalConvLayer( - out_channels, - out_channels, - dropout=0.1, - norm_num_groups=resnet_groups, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.temp_convs = nn.ModuleList(temp_convs) - - self.add_temp_upsample = temporal_upsample - if temporal_upsample: - self.temp_conv_up = AllegroTemporalConvLayer( - out_channels, out_channels, dropout=0.1, norm_num_groups=resnet_groups, up_sample=True, stride=3 - ) - - if spatial_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size = hidden_states.shape[0] - - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) - - for resnet, temp_conv in zip(self.resnets, self.temp_convs): - hidden_states = resnet(hidden_states, temb=None) - hidden_states = temp_conv(hidden_states, batch_size=batch_size) - - if self.add_temp_upsample: - hidden_states = self.temp_conv_up(hidden_states, batch_size=batch_size) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - return hidden_states - - -class AllegroMidBlock3DConv(nn.Module): - def __init__( - self, - in_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", # default, spatial - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - add_attention: bool = True, - attention_head_dim: int = 1, - output_scale_factor: float = 1.0, - ): - super().__init__() - - # there is always at least one resnet - resnets = [ - ResnetBlock2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ] - temp_convs = [ - AllegroTemporalConvLayer( - in_channels, - in_channels, - dropout=0.1, - norm_num_groups=resnet_groups, - ) - ] - attentions = [] - - if attention_head_dim is None: - attention_head_dim = in_channels - - for _ in range(num_layers): - if add_attention: - attentions.append( - Attention( - in_channels, - heads=in_channels // attention_head_dim, - dim_head=attention_head_dim, - rescale_output_factor=output_scale_factor, - eps=resnet_eps, - norm_num_groups=resnet_groups if resnet_time_scale_shift == "default" else None, - spatial_norm_dim=temb_channels if resnet_time_scale_shift == "spatial" else None, - residual_connection=True, - bias=True, - upcast_softmax=True, - _from_deprecated_attn_block=True, - ) - ) - else: - attentions.append(None) - - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - temp_convs.append( - AllegroTemporalConvLayer( - in_channels, - in_channels, - dropout=0.1, - norm_num_groups=resnet_groups, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.temp_convs = nn.ModuleList(temp_convs) - self.attentions = nn.ModuleList(attentions) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size = hidden_states.shape[0] - - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) - hidden_states = self.resnets[0](hidden_states, temb=None) - - hidden_states = self.temp_convs[0](hidden_states, batch_size=batch_size) - - for attn, resnet, temp_conv in zip(self.attentions, self.resnets[1:], self.temp_convs[1:]): - hidden_states = attn(hidden_states) - hidden_states = resnet(hidden_states, temb=None) - hidden_states = temp_conv(hidden_states, batch_size=batch_size) - - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - return hidden_states - - -class AllegroEncoder3D(nn.Module): - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - down_block_types: tuple[str, ...] = ( - "AllegroDownBlock3D", - "AllegroDownBlock3D", - "AllegroDownBlock3D", - "AllegroDownBlock3D", - ), - block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - temporal_downsample_blocks: tuple[bool, ...] = [True, True, False, False], - layers_per_block: int = 2, - norm_num_groups: int = 32, - act_fn: str = "silu", - double_z: bool = True, - ): - super().__init__() - - self.conv_in = nn.Conv2d( - in_channels, - block_out_channels[0], - kernel_size=3, - stride=1, - padding=1, - ) - - self.temp_conv_in = nn.Conv3d( - in_channels=block_out_channels[0], - out_channels=block_out_channels[0], - kernel_size=(3, 1, 1), - padding=(1, 0, 0), - ) - - self.down_blocks = nn.ModuleList([]) - - # down - output_channel = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - if down_block_type == "AllegroDownBlock3D": - down_block = AllegroDownBlock3D( - num_layers=layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - spatial_downsample=not is_final_block, - temporal_downsample=temporal_downsample_blocks[i], - resnet_eps=1e-6, - downsample_padding=0, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - ) - else: - raise ValueError("Invalid `down_block_type` encountered. Must be `AllegroDownBlock3D`") - - self.down_blocks.append(down_block) - - # mid - self.mid_block = AllegroMidBlock3DConv( - in_channels=block_out_channels[-1], - resnet_eps=1e-6, - resnet_act_fn=act_fn, - output_scale_factor=1, - resnet_time_scale_shift="default", - attention_head_dim=block_out_channels[-1], - resnet_groups=norm_num_groups, - temb_channels=None, - ) - - # out - self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6) - self.conv_act = nn.SiLU() - - conv_out_channels = 2 * out_channels if double_z else out_channels - - self.temp_conv_out = nn.Conv3d(block_out_channels[-1], block_out_channels[-1], (3, 1, 1), padding=(1, 0, 0)) - self.conv_out = nn.Conv2d(block_out_channels[-1], conv_out_channels, 3, padding=1) - - self.gradient_checkpointing = False - - def forward(self, sample: torch.Tensor) -> torch.Tensor: - batch_size = sample.shape[0] - - sample = sample.permute(0, 2, 1, 3, 4).flatten(0, 1) - sample = self.conv_in(sample) - - sample = sample.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - residual = sample - sample = self.temp_conv_in(sample) - sample = sample + residual - - if torch.is_grad_enabled() and self.gradient_checkpointing: - # Down blocks - for down_block in self.down_blocks: - sample = self._gradient_checkpointing_func(down_block, sample) - - # Mid block - sample = self._gradient_checkpointing_func(self.mid_block, sample) - else: - # Down blocks - for down_block in self.down_blocks: - sample = down_block(sample) - - # Mid block - sample = self.mid_block(sample) - - # Post process - sample = sample.permute(0, 2, 1, 3, 4).flatten(0, 1) - sample = self.conv_norm_out(sample) - sample = self.conv_act(sample) - - sample = sample.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - residual = sample - sample = self.temp_conv_out(sample) - sample = sample + residual - - sample = sample.permute(0, 2, 1, 3, 4).flatten(0, 1) - sample = self.conv_out(sample) - - sample = sample.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - return sample - - -class AllegroDecoder3D(nn.Module): - def __init__( - self, - in_channels: int = 4, - out_channels: int = 3, - up_block_types: tuple[str, ...] = ( - "AllegroUpBlock3D", - "AllegroUpBlock3D", - "AllegroUpBlock3D", - "AllegroUpBlock3D", - ), - temporal_upsample_blocks: tuple[bool, ...] = [False, True, True, False], - block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - layers_per_block: int = 2, - norm_num_groups: int = 32, - act_fn: str = "silu", - norm_type: str = "group", # group, spatial - ): - super().__init__() - - self.conv_in = nn.Conv2d( - in_channels, - block_out_channels[-1], - kernel_size=3, - stride=1, - padding=1, - ) - - self.temp_conv_in = nn.Conv3d(block_out_channels[-1], block_out_channels[-1], (3, 1, 1), padding=(1, 0, 0)) - - self.mid_block = None - self.up_blocks = nn.ModuleList([]) - - temb_channels = in_channels if norm_type == "spatial" else None - - # mid - self.mid_block = AllegroMidBlock3DConv( - in_channels=block_out_channels[-1], - resnet_eps=1e-6, - resnet_act_fn=act_fn, - output_scale_factor=1, - resnet_time_scale_shift="default" if norm_type == "group" else norm_type, - attention_head_dim=block_out_channels[-1], - resnet_groups=norm_num_groups, - temb_channels=temb_channels, - ) - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - output_channel = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - - is_final_block = i == len(block_out_channels) - 1 - - if up_block_type == "AllegroUpBlock3D": - up_block = AllegroUpBlock3D( - num_layers=layers_per_block + 1, - in_channels=prev_output_channel, - out_channels=output_channel, - spatial_upsample=not is_final_block, - temporal_upsample=temporal_upsample_blocks[i], - resnet_eps=1e-6, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - temb_channels=temb_channels, - resnet_time_scale_shift=norm_type, - ) - else: - raise ValueError("Invalid `UP_block_type` encountered. Must be `AllegroUpBlock3D`") - - self.up_blocks.append(up_block) - prev_output_channel = output_channel - - # out - if norm_type == "spatial": - self.conv_norm_out = SpatialNorm(block_out_channels[0], temb_channels) - else: - self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6) - - self.conv_act = nn.SiLU() - - self.temp_conv_out = nn.Conv3d(block_out_channels[0], block_out_channels[0], (3, 1, 1), padding=(1, 0, 0)) - self.conv_out = nn.Conv2d(block_out_channels[0], out_channels, 3, padding=1) - - self.gradient_checkpointing = False - - def forward(self, sample: torch.Tensor) -> torch.Tensor: - batch_size = sample.shape[0] - - sample = sample.permute(0, 2, 1, 3, 4).flatten(0, 1) - sample = self.conv_in(sample) - - sample = sample.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - residual = sample - sample = self.temp_conv_in(sample) - sample = sample + residual - - upscale_dtype = next(iter(self.up_blocks.parameters())).dtype - - if torch.is_grad_enabled() and self.gradient_checkpointing: - # Mid block - sample = self._gradient_checkpointing_func(self.mid_block, sample) - - # Up blocks - for up_block in self.up_blocks: - sample = self._gradient_checkpointing_func(up_block, sample) - - else: - # Mid block - sample = self.mid_block(sample) - sample = sample.to(upscale_dtype) - - # Up blocks - for up_block in self.up_blocks: - sample = up_block(sample) - - # Post process - sample = sample.permute(0, 2, 1, 3, 4).flatten(0, 1) - sample = self.conv_norm_out(sample) - sample = self.conv_act(sample) - - sample = sample.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - residual = sample - sample = self.temp_conv_out(sample) - sample = sample + residual - - sample = sample.permute(0, 2, 1, 3, 4).flatten(0, 1) - sample = self.conv_out(sample) - - sample = sample.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - return sample - - -class AutoencoderKLAllegro(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. Used in - [Allegro](https://github.com/rhymes-ai/Allegro). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - in_channels (int, defaults to `3`): - Number of channels in the input image. - out_channels (int, defaults to `3`): - Number of channels in the output. - down_block_types (`tuple[str, ...]`, defaults to `("AllegroDownBlock3D", "AllegroDownBlock3D", "AllegroDownBlock3D", "AllegroDownBlock3D")`): - tuple of strings denoting which types of down blocks to use. - up_block_types (`tuple[str, ...]`, defaults to `("AllegroUpBlock3D", "AllegroUpBlock3D", "AllegroUpBlock3D", "AllegroUpBlock3D")`): - tuple of strings denoting which types of up blocks to use. - block_out_channels (`tuple[int, ...]`, defaults to `(128, 256, 512, 512)`): - tuple of integers denoting number of output channels in each block. - temporal_downsample_blocks (`tuple[bool, ...]`, defaults to `(True, True, False, False)`): - tuple of booleans denoting which blocks to enable temporal downsampling in. - latent_channels (`int`, defaults to `4`): - Number of channels in latents. - layers_per_block (`int`, defaults to `2`): - Number of resnet or attention or temporal convolution layers per down/up block. - act_fn (`str`, defaults to `"silu"`): - The activation function to use. - norm_num_groups (`int`, defaults to `32`): - Number of groups to use in normalization layers. - temporal_compression_ratio (`int`, defaults to `4`): - Ratio by which temporal dimension of samples are compressed. - sample_size (`int`, defaults to `320`): - Default latent size. - scaling_factor (`float`, defaults to `0.13235`): - The component-wise standard deviation of the trained latent space computed using the first batch of the - training set. This is used to scale the latent space to have unit variance when training the diffusion - model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the - diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 - / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image - Synthesis with Latent Diffusion Models](https://huggingface.co/papers/2112.10752) paper. - force_upcast (`bool`, default to `True`): - If enabled it will force the VAE to run in float32 for high image resolution pipelines, such as SD-XL. VAE - can be fine-tuned / trained to a lower range without losing too much precision in which case `force_upcast` - can be set to `False` - see: https://huggingface.co/madebyollin/sdxl-vae-fp16-fix - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - down_block_types: tuple[str, ...] = ( - "AllegroDownBlock3D", - "AllegroDownBlock3D", - "AllegroDownBlock3D", - "AllegroDownBlock3D", - ), - up_block_types: tuple[str, ...] = ( - "AllegroUpBlock3D", - "AllegroUpBlock3D", - "AllegroUpBlock3D", - "AllegroUpBlock3D", - ), - block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - temporal_downsample_blocks: tuple[bool, ...] = (True, True, False, False), - temporal_upsample_blocks: tuple[bool, ...] = (False, True, True, False), - latent_channels: int = 4, - layers_per_block: int = 2, - act_fn: str = "silu", - norm_num_groups: int = 32, - temporal_compression_ratio: float = 4, - sample_size: int = 320, - scaling_factor: float = 0.13, - force_upcast: bool = True, - ) -> None: - super().__init__() - - self.encoder = AllegroEncoder3D( - in_channels=in_channels, - out_channels=latent_channels, - down_block_types=down_block_types, - temporal_downsample_blocks=temporal_downsample_blocks, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - double_z=True, - ) - self.decoder = AllegroDecoder3D( - in_channels=latent_channels, - out_channels=out_channels, - up_block_types=up_block_types, - temporal_upsample_blocks=temporal_upsample_blocks, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - norm_num_groups=norm_num_groups, - act_fn=act_fn, - ) - self.quant_conv = nn.Conv2d(2 * latent_channels, 2 * latent_channels, 1) - self.post_quant_conv = nn.Conv2d(latent_channels, latent_channels, 1) - - # TODO(aryan): For the 1.0.0 refactor, `temporal_compression_ratio` can be inferred directly and we don't need - # to use a specific parameter here or in other VAEs. - - self.use_slicing = False - self.use_tiling = False - - self.spatial_compression_ratio = 2 ** (len(block_out_channels) - 1) - self.tile_overlap_t = 8 - self.tile_overlap_h = 120 - self.tile_overlap_w = 80 - sample_frames = 24 - - self.kernel = (sample_frames, sample_size, sample_size) - self.stride = ( - sample_frames - self.tile_overlap_t, - sample_size - self.tile_overlap_h, - sample_size - self.tile_overlap_w, - ) - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - # TODO(aryan) - # if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): - if self.use_tiling: - return self.tiled_encode(x) - - raise NotImplementedError("Encoding without tiling has not been implemented yet.") - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - r""" - Encode a batch of videos into latents. - - Args: - x (`torch.Tensor`): - Input batch of videos. - return_dict (`bool`, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor) -> torch.Tensor: - # TODO(aryan): refactor tiling implementation - # if self.use_tiling and (width > self.tile_latent_min_width or height > self.tile_latent_min_height): - if self.use_tiling: - return self.tiled_decode(z) - - raise NotImplementedError("Decoding without tiling has not been implemented yet.") - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - """ - Decode a batch of videos. - - Args: - z (`torch.Tensor`): - Input batch of latent vectors. - return_dict (`bool`, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice) for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z) - - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - local_batch_size = 1 - rs = self.spatial_compression_ratio - rt = self.config.temporal_compression_ratio - - batch_size, num_channels, num_frames, height, width = x.shape - - output_num_frames = math.floor((num_frames - self.kernel[0]) / self.stride[0]) + 1 - output_height = math.floor((height - self.kernel[1]) / self.stride[1]) + 1 - output_width = math.floor((width - self.kernel[2]) / self.stride[2]) + 1 - - count = 0 - output_latent = x.new_zeros( - ( - output_num_frames * output_height * output_width, - 2 * self.config.latent_channels, - self.kernel[0] // rt, - self.kernel[1] // rs, - self.kernel[2] // rs, - ) - ) - vae_batch_input = x.new_zeros((local_batch_size, num_channels, self.kernel[0], self.kernel[1], self.kernel[2])) - - for i in range(output_num_frames): - for j in range(output_height): - for k in range(output_width): - n_start, n_end = i * self.stride[0], i * self.stride[0] + self.kernel[0] - h_start, h_end = j * self.stride[1], j * self.stride[1] + self.kernel[1] - w_start, w_end = k * self.stride[2], k * self.stride[2] + self.kernel[2] - - video_cube = x[:, :, n_start:n_end, h_start:h_end, w_start:w_end] - vae_batch_input[count % local_batch_size] = video_cube - - if ( - count % local_batch_size == local_batch_size - 1 - or count == output_num_frames * output_height * output_width - 1 - ): - latent = self.encoder(vae_batch_input) - - if ( - count == output_num_frames * output_height * output_width - 1 - and count % local_batch_size != local_batch_size - 1 - ): - output_latent[count - count % local_batch_size :] = latent[: count % local_batch_size + 1] - else: - output_latent[count - local_batch_size + 1 : count + 1] = latent - - vae_batch_input = x.new_zeros( - (local_batch_size, num_channels, self.kernel[0], self.kernel[1], self.kernel[2]) - ) - - count += 1 - - latent = x.new_zeros( - (batch_size, 2 * self.config.latent_channels, num_frames // rt, height // rs, width // rs) - ) - output_kernel = self.kernel[0] // rt, self.kernel[1] // rs, self.kernel[2] // rs - output_stride = self.stride[0] // rt, self.stride[1] // rs, self.stride[2] // rs - output_overlap = ( - output_kernel[0] - output_stride[0], - output_kernel[1] - output_stride[1], - output_kernel[2] - output_stride[2], - ) - - for i in range(output_num_frames): - n_start, n_end = i * output_stride[0], i * output_stride[0] + output_kernel[0] - for j in range(output_height): - h_start, h_end = j * output_stride[1], j * output_stride[1] + output_kernel[1] - for k in range(output_width): - w_start, w_end = k * output_stride[2], k * output_stride[2] + output_kernel[2] - latent_mean = _prepare_for_blend( - (i, output_num_frames, output_overlap[0]), - (j, output_height, output_overlap[1]), - (k, output_width, output_overlap[2]), - output_latent[i * output_height * output_width + j * output_width + k].unsqueeze(0), - ) - latent[:, :, n_start:n_end, h_start:h_end, w_start:w_end] += latent_mean - - latent = latent.permute(0, 2, 1, 3, 4).flatten(0, 1) - latent = self.quant_conv(latent) - latent = latent.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - return latent - - def tiled_decode(self, z: torch.Tensor) -> torch.Tensor: - local_batch_size = 1 - rs = self.spatial_compression_ratio - rt = self.config.temporal_compression_ratio - - latent_kernel = self.kernel[0] // rt, self.kernel[1] // rs, self.kernel[2] // rs - latent_stride = self.stride[0] // rt, self.stride[1] // rs, self.stride[2] // rs - - batch_size, num_channels, num_frames, height, width = z.shape - - ## post quant conv (a mapping) - z = z.permute(0, 2, 1, 3, 4).flatten(0, 1) - z = self.post_quant_conv(z) - z = z.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - - output_num_frames = math.floor((num_frames - latent_kernel[0]) / latent_stride[0]) + 1 - output_height = math.floor((height - latent_kernel[1]) / latent_stride[1]) + 1 - output_width = math.floor((width - latent_kernel[2]) / latent_stride[2]) + 1 - - count = 0 - decoded_videos = z.new_zeros( - ( - output_num_frames * output_height * output_width, - self.config.out_channels, - self.kernel[0], - self.kernel[1], - self.kernel[2], - ) - ) - vae_batch_input = z.new_zeros( - (local_batch_size, num_channels, latent_kernel[0], latent_kernel[1], latent_kernel[2]) - ) - - for i in range(output_num_frames): - for j in range(output_height): - for k in range(output_width): - n_start, n_end = i * latent_stride[0], i * latent_stride[0] + latent_kernel[0] - h_start, h_end = j * latent_stride[1], j * latent_stride[1] + latent_kernel[1] - w_start, w_end = k * latent_stride[2], k * latent_stride[2] + latent_kernel[2] - - current_latent = z[:, :, n_start:n_end, h_start:h_end, w_start:w_end] - vae_batch_input[count % local_batch_size] = current_latent - - if ( - count % local_batch_size == local_batch_size - 1 - or count == output_num_frames * output_height * output_width - 1 - ): - current_video = self.decoder(vae_batch_input) - - if ( - count == output_num_frames * output_height * output_width - 1 - and count % local_batch_size != local_batch_size - 1 - ): - decoded_videos[count - count % local_batch_size :] = current_video[ - : count % local_batch_size + 1 - ] - else: - decoded_videos[count - local_batch_size + 1 : count + 1] = current_video - - vae_batch_input = z.new_zeros( - (local_batch_size, num_channels, latent_kernel[0], latent_kernel[1], latent_kernel[2]) - ) - - count += 1 - - video = z.new_zeros((batch_size, self.config.out_channels, num_frames * rt, height * rs, width * rs)) - video_overlap = ( - self.kernel[0] - self.stride[0], - self.kernel[1] - self.stride[1], - self.kernel[2] - self.stride[2], - ) - - for i in range(output_num_frames): - n_start, n_end = i * self.stride[0], i * self.stride[0] + self.kernel[0] - for j in range(output_height): - h_start, h_end = j * self.stride[1], j * self.stride[1] + self.kernel[1] - for k in range(output_width): - w_start, w_end = k * self.stride[2], k * self.stride[2] + self.kernel[2] - out_video_blend = _prepare_for_blend( - (i, output_num_frames, video_overlap[0]), - (j, output_height, video_overlap[1]), - (k, output_width, video_overlap[2]), - decoded_videos[i * output_height * output_width + j * output_width + k].unsqueeze(0), - ) - video[:, :, n_start:n_end, h_start:h_end, w_start:w_end] += out_video_blend - - video = video.permute(0, 2, 1, 3, 4).contiguous() - return video - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - PyTorch random number generator. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z).sample - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - -def _prepare_for_blend(n_param, h_param, w_param, x): - # TODO(aryan): refactor - n, n_max, overlap_n = n_param - h, h_max, overlap_h = h_param - w, w_max, overlap_w = w_param - if overlap_n > 0: - if n > 0: # the head overlap part decays from 0 to 1 - x[:, :, 0:overlap_n, :, :] = x[:, :, 0:overlap_n, :, :] * ( - torch.arange(0, overlap_n).float().to(x.device) / overlap_n - ).reshape(overlap_n, 1, 1) - if n < n_max - 1: # the tail overlap part decays from 1 to 0 - x[:, :, -overlap_n:, :, :] = x[:, :, -overlap_n:, :, :] * ( - 1 - torch.arange(0, overlap_n).float().to(x.device) / overlap_n - ).reshape(overlap_n, 1, 1) - if h > 0: - x[:, :, :, 0:overlap_h, :] = x[:, :, :, 0:overlap_h, :] * ( - torch.arange(0, overlap_h).float().to(x.device) / overlap_h - ).reshape(overlap_h, 1) - if h < h_max - 1: - x[:, :, :, -overlap_h:, :] = x[:, :, :, -overlap_h:, :] * ( - 1 - torch.arange(0, overlap_h).float().to(x.device) / overlap_h - ).reshape(overlap_h, 1) - if w > 0: - x[:, :, :, :, 0:overlap_w] = x[:, :, :, :, 0:overlap_w] * ( - torch.arange(0, overlap_w).float().to(x.device) / overlap_w - ) - if w < w_max - 1: - x[:, :, :, :, -overlap_w:] = x[:, :, :, :, -overlap_w:] * ( - 1 - torch.arange(0, overlap_w).float().to(x.device) / overlap_w - ) - return x diff --git a/diffusers/models/autoencoders/autoencoder_kl_cogvideox.py b/diffusers/models/autoencoders/autoencoder_kl_cogvideox.py deleted file mode 100644 index ed624dc9e62e11a3c0c4d5ad3f809cb2ae775181..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_cogvideox.py +++ /dev/null @@ -1,1437 +0,0 @@ -# Copyright 2025 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import numpy as np -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders.single_file_model import FromOriginalModelMixin -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..downsampling import CogVideoXDownsample3D -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from ..upsampling import CogVideoXUpsample3D -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class CogVideoXSafeConv3d(nn.Conv3d): - r""" - A 3D convolution layer that splits the input tensor into smaller parts to avoid OOM in CogVideoX Model. - """ - - def forward(self, input: torch.Tensor) -> torch.Tensor: - memory_count = ( - (input.shape[0] * input.shape[1] * input.shape[2] * input.shape[3] * input.shape[4]) * 2 / 1024**3 - ) - - # Set to 2GB, suitable for CuDNN - if memory_count > 2: - kernel_size = self.kernel_size[0] - part_num = int(memory_count / 2) + 1 - input_chunks = torch.chunk(input, part_num, dim=2) - - if kernel_size > 1: - input_chunks = [input_chunks[0]] + [ - torch.cat((input_chunks[i - 1][:, :, -kernel_size + 1 :], input_chunks[i]), dim=2) - for i in range(1, len(input_chunks)) - ] - - output_chunks = [] - for input_chunk in input_chunks: - output_chunks.append(super().forward(input_chunk)) - output = torch.cat(output_chunks, dim=2) - return output - else: - return super().forward(input) - - -class CogVideoXCausalConv3d(nn.Module): - r"""A 3D causal convolution layer that pads the input tensor to ensure causality in CogVideoX Model. - - Args: - in_channels (`int`): Number of channels in the input tensor. - out_channels (`int`): Number of output channels produced by the convolution. - kernel_size (`int` or `tuple[int, int, int]`): Kernel size of the convolutional kernel. - stride (`int`, defaults to `1`): Stride of the convolution. - dilation (`int`, defaults to `1`): Dilation rate of the convolution. - pad_mode (`str`, defaults to `"constant"`): Padding mode. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int], - stride: int = 1, - dilation: int = 1, - pad_mode: str = "constant", - ): - super().__init__() - - if isinstance(kernel_size, int): - kernel_size = (kernel_size,) * 3 - - time_kernel_size, height_kernel_size, width_kernel_size = kernel_size - - # TODO(aryan): configure calculation based on stride and dilation in the future. - # Since CogVideoX does not use it, it is currently tailored to "just work" with Mochi - time_pad = time_kernel_size - 1 - height_pad = (height_kernel_size - 1) // 2 - width_pad = (width_kernel_size - 1) // 2 - - self.pad_mode = pad_mode - self.height_pad = height_pad - self.width_pad = width_pad - self.time_pad = time_pad - self.time_causal_padding = (width_pad, width_pad, height_pad, height_pad, time_pad, 0) - self.const_padding_conv3d = (0, self.width_pad, self.height_pad) - - self.temporal_dim = 2 - self.time_kernel_size = time_kernel_size - - stride = stride if isinstance(stride, tuple) else (stride, 1, 1) - dilation = (dilation, 1, 1) - self.conv = CogVideoXSafeConv3d( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=kernel_size, - stride=stride, - dilation=dilation, - padding=0 if self.pad_mode == "replicate" else self.const_padding_conv3d, - padding_mode="zeros", - ) - - def fake_context_parallel_forward( - self, inputs: torch.Tensor, conv_cache: torch.Tensor | None = None - ) -> torch.Tensor: - if self.pad_mode == "replicate": - inputs = F.pad(inputs, self.time_causal_padding, mode="replicate") - else: - kernel_size = self.time_kernel_size - if kernel_size > 1: - cached_inputs = [conv_cache] if conv_cache is not None else [inputs[:, :, :1]] * (kernel_size - 1) - inputs = torch.cat(cached_inputs + [inputs], dim=2) - return inputs - - def forward(self, inputs: torch.Tensor, conv_cache: torch.Tensor | None = None) -> torch.Tensor: - inputs = self.fake_context_parallel_forward(inputs, conv_cache) - - if self.pad_mode == "replicate": - conv_cache = None - else: - conv_cache = inputs[:, :, -self.time_kernel_size + 1 :].clone() - - output = self.conv(inputs) - return output, conv_cache - - -class CogVideoXSpatialNorm3D(nn.Module): - r""" - Spatially conditioned normalization as defined in https://huggingface.co/papers/2209.09002. This implementation is - specific to 3D-video like data. - - CogVideoXSafeConv3d is used instead of nn.Conv3d to avoid OOM in CogVideoX Model. - - Args: - f_channels (`int`): - The number of channels for input to group normalization layer, and output of the spatial norm layer. - zq_channels (`int`): - The number of channels for the quantized vector as described in the paper. - groups (`int`): - Number of groups to separate the channels into for group normalization. - """ - - def __init__( - self, - f_channels: int, - zq_channels: int, - groups: int = 32, - ): - super().__init__() - self.norm_layer = nn.GroupNorm(num_channels=f_channels, num_groups=groups, eps=1e-6, affine=True) - self.conv_y = CogVideoXCausalConv3d(zq_channels, f_channels, kernel_size=1, stride=1) - self.conv_b = CogVideoXCausalConv3d(zq_channels, f_channels, kernel_size=1, stride=1) - - def forward( - self, f: torch.Tensor, zq: torch.Tensor, conv_cache: dict[str, torch.Tensor] | None = None - ) -> torch.Tensor: - new_conv_cache = {} - conv_cache = conv_cache or {} - - if f.shape[2] > 1 and f.shape[2] % 2 == 1: - f_first, f_rest = f[:, :, :1], f[:, :, 1:] - f_first_size, f_rest_size = f_first.shape[-3:], f_rest.shape[-3:] - z_first, z_rest = zq[:, :, :1], zq[:, :, 1:] - z_first = F.interpolate(z_first, size=f_first_size) - z_rest = F.interpolate(z_rest, size=f_rest_size) - zq = torch.cat([z_first, z_rest], dim=2) - else: - zq = F.interpolate(zq, size=f.shape[-3:]) - - conv_y, new_conv_cache["conv_y"] = self.conv_y(zq, conv_cache=conv_cache.get("conv_y")) - conv_b, new_conv_cache["conv_b"] = self.conv_b(zq, conv_cache=conv_cache.get("conv_b")) - - norm_f = self.norm_layer(f) - new_f = norm_f * conv_y + conv_b - return new_f, new_conv_cache - - -class CogVideoXResnetBlock3D(nn.Module): - r""" - A 3D ResNet block used in the CogVideoX model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - dropout (`float`, defaults to `0.0`): - Dropout rate. - temb_channels (`int`, defaults to `512`): - Number of time embedding channels. - groups (`int`, defaults to `32`): - Number of groups to separate the channels into for group normalization. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - non_linearity (`str`, defaults to `"swish"`): - Activation function to use. - conv_shortcut (bool, defaults to `False`): - Whether or not to use a convolution shortcut. - spatial_norm_dim (`int`, *optional*): - The dimension to use for spatial norm if it is to be used instead of group norm. - pad_mode (str, defaults to `"first"`): - Padding mode. - """ - - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - dropout: float = 0.0, - temb_channels: int = 512, - groups: int = 32, - eps: float = 1e-6, - non_linearity: str = "swish", - conv_shortcut: bool = False, - spatial_norm_dim: int | None = None, - pad_mode: str = "first", - ): - super().__init__() - - out_channels = out_channels or in_channels - - self.in_channels = in_channels - self.out_channels = out_channels - self.nonlinearity = get_activation(non_linearity) - self.use_conv_shortcut = conv_shortcut - self.spatial_norm_dim = spatial_norm_dim - - if spatial_norm_dim is None: - self.norm1 = nn.GroupNorm(num_channels=in_channels, num_groups=groups, eps=eps) - self.norm2 = nn.GroupNorm(num_channels=out_channels, num_groups=groups, eps=eps) - else: - self.norm1 = CogVideoXSpatialNorm3D( - f_channels=in_channels, - zq_channels=spatial_norm_dim, - groups=groups, - ) - self.norm2 = CogVideoXSpatialNorm3D( - f_channels=out_channels, - zq_channels=spatial_norm_dim, - groups=groups, - ) - - self.conv1 = CogVideoXCausalConv3d( - in_channels=in_channels, out_channels=out_channels, kernel_size=3, pad_mode=pad_mode - ) - - if temb_channels > 0: - self.temb_proj = nn.Linear(in_features=temb_channels, out_features=out_channels) - - self.dropout = nn.Dropout(dropout) - self.conv2 = CogVideoXCausalConv3d( - in_channels=out_channels, out_channels=out_channels, kernel_size=3, pad_mode=pad_mode - ) - - if self.in_channels != self.out_channels: - if self.use_conv_shortcut: - self.conv_shortcut = CogVideoXCausalConv3d( - in_channels=in_channels, out_channels=out_channels, kernel_size=3, pad_mode=pad_mode - ) - else: - self.conv_shortcut = CogVideoXSafeConv3d( - in_channels=in_channels, out_channels=out_channels, kernel_size=1, stride=1, padding=0 - ) - - def forward( - self, - inputs: torch.Tensor, - temb: torch.Tensor | None = None, - zq: torch.Tensor | None = None, - conv_cache: dict[str, torch.Tensor] | None = None, - ) -> torch.Tensor: - new_conv_cache = {} - conv_cache = conv_cache or {} - - hidden_states = inputs - - if zq is not None: - hidden_states, new_conv_cache["norm1"] = self.norm1(hidden_states, zq, conv_cache=conv_cache.get("norm1")) - else: - hidden_states = self.norm1(hidden_states) - - hidden_states = self.nonlinearity(hidden_states) - hidden_states, new_conv_cache["conv1"] = self.conv1(hidden_states, conv_cache=conv_cache.get("conv1")) - - if temb is not None: - hidden_states = hidden_states + self.temb_proj(self.nonlinearity(temb))[:, :, None, None, None] - - if zq is not None: - hidden_states, new_conv_cache["norm2"] = self.norm2(hidden_states, zq, conv_cache=conv_cache.get("norm2")) - else: - hidden_states = self.norm2(hidden_states) - - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.dropout(hidden_states) - hidden_states, new_conv_cache["conv2"] = self.conv2(hidden_states, conv_cache=conv_cache.get("conv2")) - - if self.in_channels != self.out_channels: - if self.use_conv_shortcut: - inputs, new_conv_cache["conv_shortcut"] = self.conv_shortcut( - inputs, conv_cache=conv_cache.get("conv_shortcut") - ) - else: - inputs = self.conv_shortcut(inputs) - - hidden_states = hidden_states + inputs - return hidden_states, new_conv_cache - - -class CogVideoXDownBlock3D(nn.Module): - r""" - A downsampling block used in the CogVideoX model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - temb_channels (`int`, defaults to `512`): - Number of time embedding channels. - num_layers (`int`, defaults to `1`): - Number of resnet layers. - dropout (`float`, defaults to `0.0`): - Dropout rate. - resnet_eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - resnet_act_fn (`str`, defaults to `"swish"`): - Activation function to use. - resnet_groups (`int`, defaults to `32`): - Number of groups to separate the channels into for group normalization. - add_downsample (`bool`, defaults to `True`): - Whether or not to use a downsampling layer. If not used, output dimension would be same as input dimension. - compress_time (`bool`, defaults to `False`): - Whether or not to downsample across temporal dimension. - pad_mode (str, defaults to `"first"`): - Padding mode. - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - add_downsample: bool = True, - downsample_padding: int = 0, - compress_time: bool = False, - pad_mode: str = "first", - ): - super().__init__() - - resnets = [] - for i in range(num_layers): - in_channel = in_channels if i == 0 else out_channels - resnets.append( - CogVideoXResnetBlock3D( - in_channels=in_channel, - out_channels=out_channels, - dropout=dropout, - temb_channels=temb_channels, - groups=resnet_groups, - eps=resnet_eps, - non_linearity=resnet_act_fn, - pad_mode=pad_mode, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.downsamplers = None - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - CogVideoXDownsample3D( - out_channels, out_channels, padding=downsample_padding, compress_time=compress_time - ) - ] - ) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - zq: torch.Tensor | None = None, - conv_cache: dict[str, torch.Tensor] | None = None, - ) -> torch.Tensor: - r"""Forward method of the `CogVideoXDownBlock3D` class.""" - - new_conv_cache = {} - conv_cache = conv_cache or {} - - for i, resnet in enumerate(self.resnets): - conv_cache_key = f"resnet_{i}" - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, new_conv_cache[conv_cache_key] = self._gradient_checkpointing_func( - resnet, - hidden_states, - temb, - zq, - conv_cache.get(conv_cache_key), - ) - else: - hidden_states, new_conv_cache[conv_cache_key] = resnet( - hidden_states, temb, zq, conv_cache=conv_cache.get(conv_cache_key) - ) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - return hidden_states, new_conv_cache - - -class CogVideoXMidBlock3D(nn.Module): - r""" - A middle block used in the CogVideoX model. - - Args: - in_channels (`int`): - Number of input channels. - temb_channels (`int`, defaults to `512`): - Number of time embedding channels. - dropout (`float`, defaults to `0.0`): - Dropout rate. - num_layers (`int`, defaults to `1`): - Number of resnet layers. - resnet_eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - resnet_act_fn (`str`, defaults to `"swish"`): - Activation function to use. - resnet_groups (`int`, defaults to `32`): - Number of groups to separate the channels into for group normalization. - spatial_norm_dim (`int`, *optional*): - The dimension to use for spatial norm if it is to be used instead of group norm. - pad_mode (str, defaults to `"first"`): - Padding mode. - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - spatial_norm_dim: int | None = None, - pad_mode: str = "first", - ): - super().__init__() - - resnets = [] - for _ in range(num_layers): - resnets.append( - CogVideoXResnetBlock3D( - in_channels=in_channels, - out_channels=in_channels, - dropout=dropout, - temb_channels=temb_channels, - groups=resnet_groups, - eps=resnet_eps, - spatial_norm_dim=spatial_norm_dim, - non_linearity=resnet_act_fn, - pad_mode=pad_mode, - ) - ) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - zq: torch.Tensor | None = None, - conv_cache: dict[str, torch.Tensor] | None = None, - ) -> torch.Tensor: - r"""Forward method of the `CogVideoXMidBlock3D` class.""" - - new_conv_cache = {} - conv_cache = conv_cache or {} - - for i, resnet in enumerate(self.resnets): - conv_cache_key = f"resnet_{i}" - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, new_conv_cache[conv_cache_key] = self._gradient_checkpointing_func( - resnet, hidden_states, temb, zq, conv_cache.get(conv_cache_key) - ) - else: - hidden_states, new_conv_cache[conv_cache_key] = resnet( - hidden_states, temb, zq, conv_cache=conv_cache.get(conv_cache_key) - ) - - return hidden_states, new_conv_cache - - -class CogVideoXUpBlock3D(nn.Module): - r""" - An upsampling block used in the CogVideoX model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - temb_channels (`int`, defaults to `512`): - Number of time embedding channels. - dropout (`float`, defaults to `0.0`): - Dropout rate. - num_layers (`int`, defaults to `1`): - Number of resnet layers. - resnet_eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - resnet_act_fn (`str`, defaults to `"swish"`): - Activation function to use. - resnet_groups (`int`, defaults to `32`): - Number of groups to separate the channels into for group normalization. - spatial_norm_dim (`int`, defaults to `16`): - The dimension to use for spatial norm if it is to be used instead of group norm. - add_upsample (`bool`, defaults to `True`): - Whether or not to use a upsampling layer. If not used, output dimension would be same as input dimension. - compress_time (`bool`, defaults to `False`): - Whether or not to downsample across temporal dimension. - pad_mode (str, defaults to `"first"`): - Padding mode. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - spatial_norm_dim: int = 16, - add_upsample: bool = True, - upsample_padding: int = 1, - compress_time: bool = False, - pad_mode: str = "first", - ): - super().__init__() - - resnets = [] - for i in range(num_layers): - in_channel = in_channels if i == 0 else out_channels - resnets.append( - CogVideoXResnetBlock3D( - in_channels=in_channel, - out_channels=out_channels, - dropout=dropout, - temb_channels=temb_channels, - groups=resnet_groups, - eps=resnet_eps, - non_linearity=resnet_act_fn, - spatial_norm_dim=spatial_norm_dim, - pad_mode=pad_mode, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.upsamplers = None - - if add_upsample: - self.upsamplers = nn.ModuleList( - [ - CogVideoXUpsample3D( - out_channels, out_channels, padding=upsample_padding, compress_time=compress_time - ) - ] - ) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - zq: torch.Tensor | None = None, - conv_cache: dict[str, torch.Tensor] | None = None, - ) -> torch.Tensor: - r"""Forward method of the `CogVideoXUpBlock3D` class.""" - - new_conv_cache = {} - conv_cache = conv_cache or {} - - for i, resnet in enumerate(self.resnets): - conv_cache_key = f"resnet_{i}" - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, new_conv_cache[conv_cache_key] = self._gradient_checkpointing_func( - resnet, - hidden_states, - temb, - zq, - conv_cache.get(conv_cache_key), - ) - else: - hidden_states, new_conv_cache[conv_cache_key] = resnet( - hidden_states, temb, zq, conv_cache=conv_cache.get(conv_cache_key) - ) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - return hidden_states, new_conv_cache - - -class CogVideoXEncoder3D(nn.Module): - r""" - The `CogVideoXEncoder3D` layer of a variational autoencoder that encodes its input into a latent representation. - - Args: - in_channels (`int`, *optional*, defaults to 3): - The number of input channels. - out_channels (`int`, *optional*, defaults to 3): - The number of output channels. - down_block_types (`tuple[str, ...]`, *optional*, defaults to `("DownEncoderBlock2D",)`): - The types of down blocks to use. See `~diffusers.models.unet_2d_blocks.get_down_block` for available - options. - block_out_channels (`tuple[int, ...]`, *optional*, defaults to `(64,)`): - The number of output channels for each block. - act_fn (`str`, *optional*, defaults to `"silu"`): - The activation function to use. See `~diffusers.models.activations.get_activation` for available options. - layers_per_block (`int`, *optional*, defaults to 2): - The number of layers per block. - norm_num_groups (`int`, *optional*, defaults to 32): - The number of groups for normalization. - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 16, - down_block_types: tuple[str, ...] = ( - "CogVideoXDownBlock3D", - "CogVideoXDownBlock3D", - "CogVideoXDownBlock3D", - "CogVideoXDownBlock3D", - ), - block_out_channels: tuple[int, ...] = (128, 256, 256, 512), - layers_per_block: int = 3, - act_fn: str = "silu", - norm_eps: float = 1e-6, - norm_num_groups: int = 32, - dropout: float = 0.0, - pad_mode: str = "first", - temporal_compression_ratio: float = 4, - ): - super().__init__() - - # log2 of temporal_compress_times - temporal_compress_level = int(np.log2(temporal_compression_ratio)) - - self.conv_in = CogVideoXCausalConv3d(in_channels, block_out_channels[0], kernel_size=3, pad_mode=pad_mode) - self.down_blocks = nn.ModuleList([]) - - # down blocks - output_channel = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - compress_time = i < temporal_compress_level - - if down_block_type == "CogVideoXDownBlock3D": - down_block = CogVideoXDownBlock3D( - in_channels=input_channel, - out_channels=output_channel, - temb_channels=0, - dropout=dropout, - num_layers=layers_per_block, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - add_downsample=not is_final_block, - compress_time=compress_time, - ) - else: - raise ValueError("Invalid `down_block_type` encountered. Must be `CogVideoXDownBlock3D`") - - self.down_blocks.append(down_block) - - # mid block - self.mid_block = CogVideoXMidBlock3D( - in_channels=block_out_channels[-1], - temb_channels=0, - dropout=dropout, - num_layers=2, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - pad_mode=pad_mode, - ) - - self.norm_out = nn.GroupNorm(norm_num_groups, block_out_channels[-1], eps=1e-6) - self.conv_act = nn.SiLU() - self.conv_out = CogVideoXCausalConv3d( - block_out_channels[-1], 2 * out_channels, kernel_size=3, pad_mode=pad_mode - ) - - self.gradient_checkpointing = False - - def forward( - self, - sample: torch.Tensor, - temb: torch.Tensor | None = None, - conv_cache: dict[str, torch.Tensor] | None = None, - ) -> torch.Tensor: - r"""The forward method of the `CogVideoXEncoder3D` class.""" - - new_conv_cache = {} - conv_cache = conv_cache or {} - - hidden_states, new_conv_cache["conv_in"] = self.conv_in(sample, conv_cache=conv_cache.get("conv_in")) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - # 1. Down - for i, down_block in enumerate(self.down_blocks): - conv_cache_key = f"down_block_{i}" - hidden_states, new_conv_cache[conv_cache_key] = self._gradient_checkpointing_func( - down_block, - hidden_states, - temb, - None, - conv_cache.get(conv_cache_key), - ) - - # 2. Mid - hidden_states, new_conv_cache["mid_block"] = self._gradient_checkpointing_func( - self.mid_block, - hidden_states, - temb, - None, - conv_cache.get("mid_block"), - ) - else: - # 1. Down - for i, down_block in enumerate(self.down_blocks): - conv_cache_key = f"down_block_{i}" - hidden_states, new_conv_cache[conv_cache_key] = down_block( - hidden_states, temb, None, conv_cache.get(conv_cache_key) - ) - - # 2. Mid - hidden_states, new_conv_cache["mid_block"] = self.mid_block( - hidden_states, temb, None, conv_cache=conv_cache.get("mid_block") - ) - - # 3. Post-process - hidden_states = self.norm_out(hidden_states) - hidden_states = self.conv_act(hidden_states) - - hidden_states, new_conv_cache["conv_out"] = self.conv_out(hidden_states, conv_cache=conv_cache.get("conv_out")) - - return hidden_states, new_conv_cache - - -class CogVideoXDecoder3D(nn.Module): - r""" - The `CogVideoXDecoder3D` layer of a variational autoencoder that decodes its latent representation into an output - sample. - - Args: - in_channels (`int`, *optional*, defaults to 3): - The number of input channels. - out_channels (`int`, *optional*, defaults to 3): - The number of output channels. - up_block_types (`tuple[str, ...]`, *optional*, defaults to `("UpDecoderBlock2D",)`): - The types of up blocks to use. See `~diffusers.models.unet_2d_blocks.get_up_block` for available options. - block_out_channels (`tuple[int, ...]`, *optional*, defaults to `(64,)`): - The number of output channels for each block. - act_fn (`str`, *optional*, defaults to `"silu"`): - The activation function to use. See `~diffusers.models.activations.get_activation` for available options. - layers_per_block (`int`, *optional*, defaults to 2): - The number of layers per block. - norm_num_groups (`int`, *optional*, defaults to 32): - The number of groups for normalization. - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int = 16, - out_channels: int = 3, - up_block_types: tuple[str, ...] = ( - "CogVideoXUpBlock3D", - "CogVideoXUpBlock3D", - "CogVideoXUpBlock3D", - "CogVideoXUpBlock3D", - ), - block_out_channels: tuple[int, ...] = (128, 256, 256, 512), - layers_per_block: int = 3, - act_fn: str = "silu", - norm_eps: float = 1e-6, - norm_num_groups: int = 32, - dropout: float = 0.0, - pad_mode: str = "first", - temporal_compression_ratio: float = 4, - ): - super().__init__() - - reversed_block_out_channels = list(reversed(block_out_channels)) - - self.conv_in = CogVideoXCausalConv3d( - in_channels, reversed_block_out_channels[0], kernel_size=3, pad_mode=pad_mode - ) - - # mid block - self.mid_block = CogVideoXMidBlock3D( - in_channels=reversed_block_out_channels[0], - temb_channels=0, - num_layers=2, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - spatial_norm_dim=in_channels, - pad_mode=pad_mode, - ) - - # up blocks - self.up_blocks = nn.ModuleList([]) - - output_channel = reversed_block_out_channels[0] - temporal_compress_level = int(np.log2(temporal_compression_ratio)) - - for i, up_block_type in enumerate(up_block_types): - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - compress_time = i < temporal_compress_level - - if up_block_type == "CogVideoXUpBlock3D": - up_block = CogVideoXUpBlock3D( - in_channels=prev_output_channel, - out_channels=output_channel, - temb_channels=0, - dropout=dropout, - num_layers=layers_per_block + 1, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - spatial_norm_dim=in_channels, - add_upsample=not is_final_block, - compress_time=compress_time, - pad_mode=pad_mode, - ) - prev_output_channel = output_channel - else: - raise ValueError("Invalid `up_block_type` encountered. Must be `CogVideoXUpBlock3D`") - - self.up_blocks.append(up_block) - - self.norm_out = CogVideoXSpatialNorm3D(reversed_block_out_channels[-1], in_channels, groups=norm_num_groups) - self.conv_act = nn.SiLU() - self.conv_out = CogVideoXCausalConv3d( - reversed_block_out_channels[-1], out_channels, kernel_size=3, pad_mode=pad_mode - ) - - self.gradient_checkpointing = False - - def forward( - self, - sample: torch.Tensor, - temb: torch.Tensor | None = None, - conv_cache: dict[str, torch.Tensor] | None = None, - ) -> torch.Tensor: - r"""The forward method of the `CogVideoXDecoder3D` class.""" - - new_conv_cache = {} - conv_cache = conv_cache or {} - - hidden_states, new_conv_cache["conv_in"] = self.conv_in(sample, conv_cache=conv_cache.get("conv_in")) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - # 1. Mid - hidden_states, new_conv_cache["mid_block"] = self._gradient_checkpointing_func( - self.mid_block, - hidden_states, - temb, - sample, - conv_cache.get("mid_block"), - ) - - # 2. Up - for i, up_block in enumerate(self.up_blocks): - conv_cache_key = f"up_block_{i}" - hidden_states, new_conv_cache[conv_cache_key] = self._gradient_checkpointing_func( - up_block, - hidden_states, - temb, - sample, - conv_cache.get(conv_cache_key), - ) - else: - # 1. Mid - hidden_states, new_conv_cache["mid_block"] = self.mid_block( - hidden_states, temb, sample, conv_cache=conv_cache.get("mid_block") - ) - - # 2. Up - for i, up_block in enumerate(self.up_blocks): - conv_cache_key = f"up_block_{i}" - hidden_states, new_conv_cache[conv_cache_key] = up_block( - hidden_states, temb, sample, conv_cache=conv_cache.get(conv_cache_key) - ) - - # 3. Post-process - hidden_states, new_conv_cache["norm_out"] = self.norm_out( - hidden_states, sample, conv_cache=conv_cache.get("norm_out") - ) - hidden_states = self.conv_act(hidden_states) - hidden_states, new_conv_cache["conv_out"] = self.conv_out(hidden_states, conv_cache=conv_cache.get("conv_out")) - - return hidden_states, new_conv_cache - - -class AutoencoderKLCogVideoX(ModelMixin, AutoencoderMixin, ConfigMixin, FromOriginalModelMixin): - r""" - A VAE model with KL loss for encoding images into latents and decoding latent representations into images. Used in - [CogVideoX](https://github.com/THUDM/CogVideo). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - in_channels (int, *optional*, defaults to 3): Number of channels in the input image. - out_channels (int, *optional*, defaults to 3): Number of channels in the output. - down_block_types (`tuple[str]`, *optional*, defaults to `("DownEncoderBlock2D",)`): - tuple of downsample block types. - up_block_types (`tuple[str]`, *optional*, defaults to `("UpDecoderBlock2D",)`): - tuple of upsample block types. - block_out_channels (`tuple[int]`, *optional*, defaults to `(64,)`): - tuple of block output channels. - act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use. - sample_size (`int`, *optional*, defaults to `32`): Sample input size. - scaling_factor (`float`, *optional*, defaults to `1.15258426`): - The component-wise standard deviation of the trained latent space computed using the first batch of the - training set. This is used to scale the latent space to have unit variance when training the diffusion - model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the - diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 - / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image - Synthesis with Latent Diffusion Models](https://huggingface.co/papers/2112.10752) paper. - force_upcast (`bool`, *optional*, default to `True`): - If enabled it will force the VAE to run in float32 for high image resolution pipelines, such as SD-XL. VAE - can be fine-tuned / trained to a lower range without losing too much precision in which case `force_upcast` - can be set to `False` - see: https://huggingface.co/madebyollin/sdxl-vae-fp16-fix - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["CogVideoXResnetBlock3D"] - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - down_block_types: tuple[str] = ( - "CogVideoXDownBlock3D", - "CogVideoXDownBlock3D", - "CogVideoXDownBlock3D", - "CogVideoXDownBlock3D", - ), - up_block_types: tuple[str] = ( - "CogVideoXUpBlock3D", - "CogVideoXUpBlock3D", - "CogVideoXUpBlock3D", - "CogVideoXUpBlock3D", - ), - block_out_channels: tuple[int] = (128, 256, 256, 512), - latent_channels: int = 16, - layers_per_block: int = 3, - act_fn: str = "silu", - norm_eps: float = 1e-6, - norm_num_groups: int = 32, - temporal_compression_ratio: float = 4, - sample_height: int = 480, - sample_width: int = 720, - scaling_factor: float = 1.15258426, - shift_factor: float | None = None, - latents_mean: tuple[float] | None = None, - latents_std: tuple[float] | None = None, - force_upcast: float = True, - use_quant_conv: bool = False, - use_post_quant_conv: bool = False, - invert_scale_latents: bool = False, - ): - super().__init__() - - self.encoder = CogVideoXEncoder3D( - in_channels=in_channels, - out_channels=latent_channels, - down_block_types=down_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - act_fn=act_fn, - norm_eps=norm_eps, - norm_num_groups=norm_num_groups, - temporal_compression_ratio=temporal_compression_ratio, - ) - self.decoder = CogVideoXDecoder3D( - in_channels=latent_channels, - out_channels=out_channels, - up_block_types=up_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - act_fn=act_fn, - norm_eps=norm_eps, - norm_num_groups=norm_num_groups, - temporal_compression_ratio=temporal_compression_ratio, - ) - self.quant_conv = CogVideoXSafeConv3d(2 * out_channels, 2 * out_channels, 1) if use_quant_conv else None - self.post_quant_conv = CogVideoXSafeConv3d(out_channels, out_channels, 1) if use_post_quant_conv else None - - self.use_slicing = False - self.use_tiling = False - - # Can be increased to decode more latent frames at once, but comes at a reasonable memory cost and it is not - # recommended because the temporal parts of the VAE, here, are tricky to understand. - # If you decode X latent frames together, the number of output frames is: - # (X + (2 conv cache) + (2 time upscale_1) + (4 time upscale_2) - (2 causal conv downscale)) => X + 6 frames - # - # Example with num_latent_frames_batch_size = 2: - # - 12 latent frames: (0, 1), (2, 3), (4, 5), (6, 7), (8, 9), (10, 11) are processed together - # => (12 // 2 frame slices) * ((2 num_latent_frames_batch_size) + (2 conv cache) + (2 time upscale_1) + (4 time upscale_2) - (2 causal conv downscale)) - # => 6 * 8 = 48 frames - # - 13 latent frames: (0, 1, 2) (special case), (3, 4), (5, 6), (7, 8), (9, 10), (11, 12) are processed together - # => (1 frame slice) * ((3 num_latent_frames_batch_size) + (2 conv cache) + (2 time upscale_1) + (4 time upscale_2) - (2 causal conv downscale)) + - # ((13 - 3) // 2) * ((2 num_latent_frames_batch_size) + (2 conv cache) + (2 time upscale_1) + (4 time upscale_2) - (2 causal conv downscale)) - # => 1 * 9 + 5 * 8 = 49 frames - # It has been implemented this way so as to not have "magic values" in the code base that would be hard to explain. Note that - # setting it to anything other than 2 would give poor results because the VAE hasn't been trained to be adaptive with different - # number of temporal frames. - self.num_latent_frames_batch_size = 2 - self.num_sample_frames_batch_size = 8 - - # We make the minimum height and width of sample for tiling half that of the generally supported - self.tile_sample_min_height = sample_height // 2 - self.tile_sample_min_width = sample_width // 2 - self.tile_latent_min_height = int( - self.tile_sample_min_height / (2 ** (len(self.config.block_out_channels) - 1)) - ) - self.tile_latent_min_width = int(self.tile_sample_min_width / (2 ** (len(self.config.block_out_channels) - 1))) - - # These are experimental overlap factors that were chosen based on experimentation and seem to work best for - # 720x480 (WxH) resolution. The above resolution is the strongly recommended generation resolution in CogVideoX - # and so the tiling implementation has only been tested on those specific resolutions. - self.tile_overlap_factor_height = 1 / 6 - self.tile_overlap_factor_width = 1 / 5 - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_overlap_factor_height: float | None = None, - tile_overlap_factor_width: float | None = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_overlap_factor_height (`int`, *optional*): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. Must be between 0 and 1. Setting a higher - value might cause more tiles to be processed leading to slow down of the decoding process. - tile_overlap_factor_width (`int`, *optional*): - The minimum amount of overlap between two consecutive horizontal tiles. This is to ensure that there - are no tiling artifacts produced across the width dimension. Must be between 0 and 1. Setting a higher - value might cause more tiles to be processed leading to slow down of the decoding process. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_latent_min_height = int( - self.tile_sample_min_height / (2 ** (len(self.config.block_out_channels) - 1)) - ) - self.tile_latent_min_width = int(self.tile_sample_min_width / (2 ** (len(self.config.block_out_channels) - 1))) - self.tile_overlap_factor_height = tile_overlap_factor_height or self.tile_overlap_factor_height - self.tile_overlap_factor_width = tile_overlap_factor_width or self.tile_overlap_factor_width - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = x.shape - - if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): - return self.tiled_encode(x) - - frame_batch_size = self.num_sample_frames_batch_size - # Note: We expect the number of frames to be either `1` or `frame_batch_size * k` or `frame_batch_size * k + 1` for some k. - # As the extra single frame is handled inside the loop, it is not required to round up here. - num_batches = max(num_frames // frame_batch_size, 1) - conv_cache = None - enc = [] - - for i in range(num_batches): - remaining_frames = num_frames % frame_batch_size - start_frame = frame_batch_size * i + (0 if i == 0 else remaining_frames) - end_frame = frame_batch_size * (i + 1) + remaining_frames - x_intermediate = x[:, :, start_frame:end_frame] - x_intermediate, conv_cache = self.encoder(x_intermediate, conv_cache=conv_cache) - if self.quant_conv is not None: - x_intermediate = self.quant_conv(x_intermediate) - enc.append(x_intermediate) - - enc = torch.cat(enc, dim=2) - return enc - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - batch_size, num_channels, num_frames, height, width = z.shape - - if self.use_tiling and (width > self.tile_latent_min_width or height > self.tile_latent_min_height): - return self.tiled_decode(z, return_dict=return_dict) - - frame_batch_size = self.num_latent_frames_batch_size - num_batches = max(num_frames // frame_batch_size, 1) - conv_cache = None - dec = [] - - for i in range(num_batches): - remaining_frames = num_frames % frame_batch_size - start_frame = frame_batch_size * i + (0 if i == 0 else remaining_frames) - end_frame = frame_batch_size * (i + 1) + remaining_frames - z_intermediate = z[:, :, start_frame:end_frame] - if self.post_quant_conv is not None: - z_intermediate = self.post_quant_conv(z_intermediate) - z_intermediate, conv_cache = self.decoder(z_intermediate, conv_cache=conv_cache) - dec.append(z_intermediate) - - dec = torch.cat(dec, dim=2) - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - """ - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[3], b.shape[3], blend_extent) - for y in range(blend_extent): - b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[4], b.shape[4], blend_extent) - for x in range(blend_extent): - b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( - x / blend_extent - ) - return b - - def tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - r"""Encode a batch of images using a tiled encoder. - - When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several - steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is - different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the - tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the - output, but they should be much less noticeable. - - Args: - x (`torch.Tensor`): Input batch of videos. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - # For a rough memory estimate, take a look at the `tiled_decode` method. - batch_size, num_channels, num_frames, height, width = x.shape - - overlap_height = int(self.tile_sample_min_height * (1 - self.tile_overlap_factor_height)) - overlap_width = int(self.tile_sample_min_width * (1 - self.tile_overlap_factor_width)) - blend_extent_height = int(self.tile_latent_min_height * self.tile_overlap_factor_height) - blend_extent_width = int(self.tile_latent_min_width * self.tile_overlap_factor_width) - row_limit_height = self.tile_latent_min_height - blend_extent_height - row_limit_width = self.tile_latent_min_width - blend_extent_width - frame_batch_size = self.num_sample_frames_batch_size - - # Split x into overlapping tiles and encode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, overlap_height): - row = [] - for j in range(0, width, overlap_width): - # Note: We expect the number of frames to be either `1` or `frame_batch_size * k` or `frame_batch_size * k + 1` for some k. - # As the extra single frame is handled inside the loop, it is not required to round up here. - num_batches = max(num_frames // frame_batch_size, 1) - conv_cache = None - time = [] - - for k in range(num_batches): - remaining_frames = num_frames % frame_batch_size - start_frame = frame_batch_size * k + (0 if k == 0 else remaining_frames) - end_frame = frame_batch_size * (k + 1) + remaining_frames - tile = x[ - :, - :, - start_frame:end_frame, - i : i + self.tile_sample_min_height, - j : j + self.tile_sample_min_width, - ] - tile, conv_cache = self.encoder(tile, conv_cache=conv_cache) - if self.quant_conv is not None: - tile = self.quant_conv(tile) - time.append(tile) - - row.append(torch.cat(time, dim=2)) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent_width) - result_row.append(tile[:, :, :, :row_limit_height, :row_limit_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - enc = torch.cat(result_rows, dim=3) - return enc - - def tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - # Rough memory assessment: - # - In CogVideoX-2B, there are a total of 24 CausalConv3d layers. - # - The biggest intermediate dimensions are: [1, 128, 9, 480, 720]. - # - Assume fp16 (2 bytes per value). - # Memory required: 1 * 128 * 9 * 480 * 720 * 24 * 2 / 1024**3 = 17.8 GB - # - # Memory assessment when using tiling: - # - Assume everything as above but now HxW is 240x360 by tiling in half - # Memory required: 1 * 128 * 9 * 240 * 360 * 24 * 2 / 1024**3 = 4.5 GB - - batch_size, num_channels, num_frames, height, width = z.shape - - overlap_height = int(self.tile_latent_min_height * (1 - self.tile_overlap_factor_height)) - overlap_width = int(self.tile_latent_min_width * (1 - self.tile_overlap_factor_width)) - blend_extent_height = int(self.tile_sample_min_height * self.tile_overlap_factor_height) - blend_extent_width = int(self.tile_sample_min_width * self.tile_overlap_factor_width) - row_limit_height = self.tile_sample_min_height - blend_extent_height - row_limit_width = self.tile_sample_min_width - blend_extent_width - frame_batch_size = self.num_latent_frames_batch_size - - # Split z into overlapping tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, overlap_height): - row = [] - for j in range(0, width, overlap_width): - num_batches = max(num_frames // frame_batch_size, 1) - conv_cache = None - time = [] - - for k in range(num_batches): - remaining_frames = num_frames % frame_batch_size - start_frame = frame_batch_size * k + (0 if k == 0 else remaining_frames) - end_frame = frame_batch_size * (k + 1) + remaining_frames - tile = z[ - :, - :, - start_frame:end_frame, - i : i + self.tile_latent_min_height, - j : j + self.tile_latent_min_width, - ] - if self.post_quant_conv is not None: - tile = self.post_quant_conv(tile) - tile, conv_cache = self.decoder(tile, conv_cache=conv_cache) - time.append(tile) - - row.append(torch.cat(time, dim=2)) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent_width) - result_row.append(tile[:, :, :, :row_limit_height, :row_limit_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - dec = torch.cat(result_rows, dim=3) - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> torch.Tensor | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z).sample - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) diff --git a/diffusers/models/autoencoders/autoencoder_kl_cosmos.py b/diffusers/models/autoencoders/autoencoder_kl_cosmos.py deleted file mode 100644 index 362df0bd96a22a69baf9e893d4fa0f0e35c9ed7a..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_cosmos.py +++ /dev/null @@ -1,1106 +0,0 @@ -# Copyright 2025 The NVIDIA Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import math - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import get_logger -from ...utils.accelerate_utils import apply_forward_hook -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, IdentityDistribution - - -logger = get_logger(__name__) - - -# fmt: off -# These latents and means are from CV8x8x8-1.0. Each checkpoint has different values, but since this is the main VAE used, -# we will default to these values. -LATENTS_MEAN = [0.11362758, -0.0171717, 0.03071163, 0.02046862, 0.01931456, 0.02138567, 0.01999342, 0.02189187, 0.02011935, 0.01872694, 0.02168613, 0.02207148, 0.01986941, 0.01770413, 0.02067643, 0.02028245, 0.19125476, 0.04556972, 0.0595558, 0.05315534, 0.05496629, 0.05356264, 0.04856596, 0.05327453, 0.05410472, 0.05597149, 0.05524866, 0.05181874, 0.05071663, 0.05204537, 0.0564108, 0.05518042, 0.01306714, 0.03341161, 0.03847246, 0.02810185, 0.02790166, 0.02920026, 0.02823597, 0.02631033, 0.0278531, 0.02880507, 0.02977769, 0.03145441, 0.02888389, 0.03280773, 0.03484927, 0.03049198, -0.00197727, 0.07534957, 0.04963879, 0.05530893, 0.05410828, 0.05252541, 0.05029899, 0.05321025, 0.05149245, 0.0511921, 0.04643495, 0.04604527, 0.04631618, 0.04404101, 0.04403536, 0.04499495, -0.02994183, -0.04787003, -0.01064558, -0.01779824, -0.01490502, -0.02157517, -0.0204778, -0.02180816, -0.01945375, -0.02062863, -0.02192209, -0.02520639, -0.02246656, -0.02427533, -0.02683363, -0.02762006, 0.08019473, -0.13005368, -0.07568636, -0.06082374, -0.06036175, -0.05875364, -0.05921887, -0.05869788, -0.05273941, -0.052565, -0.05346428, -0.05456541, -0.053657, -0.05656897, -0.05728589, -0.05321847, 0.16718403, -0.00390146, 0.0379406, 0.0356561, 0.03554131, 0.03924074, 0.03873615, 0.04187329, 0.04226924, 0.04378717, 0.04684274, 0.05117614, 0.04547792, 0.05251586, 0.05048339, 0.04950784, 0.09564418, 0.0547128, 0.08183969, 0.07978633, 0.08076023, 0.08108605, 0.08011818, 0.07965573, 0.08187773, 0.08350263, 0.08101469, 0.0786941, 0.0774442, 0.07724521, 0.07830418, 0.07599796, -0.04987567, 0.05923908, -0.01058746, -0.01177603, -0.01116162, -0.01364149, -0.01546014, -0.0117213, -0.01780043, -0.01648314, -0.02100247, -0.02104417, -0.02482123, -0.02611689, -0.02561143, -0.02597336, -0.05364667, 0.08211684, 0.04686937, 0.04605641, 0.04304186, 0.0397355, 0.03686767, 0.04087112, 0.03704741, 0.03706401, 0.03120073, 0.03349091, 0.03319963, 0.03205781, 0.03195127, 0.03180481, 0.16427967, -0.11048453, -0.04595276, -0.04982893, -0.05213465, -0.04809378, -0.05080318, -0.04992863, -0.04493337, -0.0467619, -0.04884703, -0.04627892, -0.04913311, -0.04955709, -0.04533982, -0.04570218, -0.10612928, -0.05121198, -0.06761009, -0.07251801, -0.07265285, -0.07417855, -0.07202412, -0.07499027, -0.07625481, -0.07535747, -0.07638787, -0.07920305, -0.07596069, -0.07959418, -0.08265036, -0.07955471, -0.16888915, 0.0753242, 0.04062594, 0.03375093, 0.03337452, 0.03699376, 0.03651138, 0.03611023, 0.03555622, 0.03378554, 0.0300498, 0.03395559, 0.02941847, 0.03156432, 0.03431173, 0.03016853, -0.03415358, -0.01699573, -0.04029295, -0.04912157, -0.0498858, -0.04917918, -0.04918056, -0.0525189, -0.05325506, -0.05341973, -0.04983329, -0.04883146, -0.04985548, -0.04736718, -0.0462027, -0.04836091, 0.02055675, 0.03419799, -0.02907669, -0.04350509, -0.04156144, -0.04234421, -0.04446109, -0.04461774, -0.04882839, -0.04822346, -0.04502493, -0.0506244, -0.05146913, -0.04655267, -0.04862994, -0.04841615, 0.20312774, -0.07208502, -0.03635615, -0.03556088, -0.04246174, -0.04195838, -0.04293778, -0.04071276, -0.04240569, -0.04125213, -0.04395144, -0.03959096, -0.04044993, -0.04015875, -0.04088107, -0.03885176] -LATENTS_STD = [0.56700271, 0.65488982, 0.65589428, 0.66524369, 0.66619784, 0.6666382, 0.6720838, 0.66955978, 0.66928875, 0.67108786, 0.67092526, 0.67397463, 0.67894882, 0.67668313, 0.67769569, 0.67479557, 0.85245121, 0.8688373, 0.87348086, 0.88459337, 0.89135885, 0.8910504, 0.89714909, 0.89947474, 0.90201765, 0.90411824, 0.90692616, 0.90847772, 0.90648711, 0.91006982, 0.91033435, 0.90541548, 0.84960359, 0.85863352, 0.86895317, 0.88460612, 0.89245003, 0.89451706, 0.89931005, 0.90647358, 0.90338236, 0.90510076, 0.91008312, 0.90961218, 0.9123717, 0.91313171, 0.91435546, 0.91565102, 0.91877103, 0.85155135, 0.857804, 0.86998034, 0.87365264, 0.88161767, 0.88151032, 0.88758916, 0.89015514, 0.89245576, 0.89276224, 0.89450496, 0.90054202, 0.89994133, 0.90136105, 0.90114892, 0.77755755, 0.81456852, 0.81911844, 0.83137071, 0.83820474, 0.83890373, 0.84401101, 0.84425181, 0.84739357, 0.84798753, 0.85249585, 0.85114998, 0.85160935, 0.85626358, 0.85677862, 0.85641026, 0.69903517, 0.71697885, 0.71696913, 0.72583169, 0.72931731, 0.73254126, 0.73586977, 0.73734969, 0.73664582, 0.74084908, 0.74399322, 0.74471819, 0.74493188, 0.74824578, 0.75024873, 0.75274801, 0.8187142, 0.82251883, 0.82616025, 0.83164483, 0.84072375, 0.8396467, 0.84143305, 0.84880769, 0.8503468, 0.85196948, 0.85211051, 0.85386664, 0.85410017, 0.85439342, 0.85847849, 0.85385275, 0.67583984, 0.68259847, 0.69198853, 0.69928843, 0.70194328, 0.70467001, 0.70755547, 0.70917857, 0.71007699, 0.70963502, 0.71064079, 0.71027333, 0.71291167, 0.71537536, 0.71902508, 0.71604162, 0.72450989, 0.71979928, 0.72057378, 0.73035461, 0.73329622, 0.73660028, 0.73891461, 0.74279994, 0.74105692, 0.74002433, 0.74257588, 0.74416119, 0.74543899, 0.74694443, 0.74747062, 0.74586403, 0.90176988, 0.90990674, 0.91106802, 0.92163783, 0.92390233, 0.93056196, 0.93482202, 0.93642414, 0.93858379, 0.94064975, 0.94078934, 0.94325715, 0.94955301, 0.94814706, 0.95144123, 0.94923073, 0.49853548, 0.64968109, 0.6427654, 0.64966393, 0.6487664, 0.65203559, 0.6584242, 0.65351611, 0.65464371, 0.6574859, 0.65626335, 0.66123748, 0.66121179, 0.66077942, 0.66040152, 0.66474909, 0.61986589, 0.69138134, 0.6884557, 0.6955843, 0.69765401, 0.70015347, 0.70529598, 0.70468754, 0.70399523, 0.70479989, 0.70887572, 0.71126866, 0.7097227, 0.71249932, 0.71231949, 0.71175605, 0.35586974, 0.68723857, 0.68973219, 0.69958478, 0.6943453, 0.6995818, 0.70980215, 0.69899458, 0.70271689, 0.70095056, 0.69912851, 0.70522696, 0.70392174, 0.70916915, 0.70585734, 0.70373541, 0.98101336, 0.89024764, 0.89607251, 0.90678179, 0.91308665, 0.91812348, 0.91980827, 0.92480654, 0.92635667, 0.92887944, 0.93338072, 0.93468094, 0.93619436, 0.93906063, 0.94191772, 0.94471723, 0.83202779, 0.84106231, 0.84463632, 0.85829508, 0.86319661, 0.86751342, 0.86914337, 0.87085921, 0.87286359, 0.87537396, 0.87931138, 0.88054478, 0.8811838, 0.88872558, 0.88942474, 0.88934827, 0.44025335, 0.63061613, 0.63110614, 0.63601959, 0.6395812, 0.64104342, 0.65019929, 0.6502797, 0.64355946, 0.64657205, 0.64847094, 0.64728117, 0.64972943, 0.65162975, 0.65328044, 0.64914775] -_WAVELETS = { - "haar": torch.tensor([0.7071067811865476, 0.7071067811865476]), - "rearrange": torch.tensor([1.0, 1.0]), -} -# fmt: on - - -class CosmosCausalConv3d(nn.Conv3d): - def __init__( - self, - in_channels: int = 1, - out_channels: int = 1, - kernel_size: int | tuple[int, int, int] = (3, 3, 3), - dilation: int | tuple[int, int, int] = (1, 1, 1), - stride: int | tuple[int, int, int] = (1, 1, 1), - padding: int = 1, - pad_mode: str = "constant", - ) -> None: - kernel_size = (kernel_size, kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size - dilation = (dilation, dilation, dilation) if isinstance(dilation, int) else dilation - stride = (stride, stride, stride) if isinstance(stride, int) else stride - - _, height_kernel_size, width_kernel_size = kernel_size - assert height_kernel_size % 2 == 1 and width_kernel_size % 2 == 1 - - super().__init__( - in_channels, - out_channels, - kernel_size, - stride=stride, - dilation=dilation, - ) - - self.pad_mode = pad_mode - self.temporal_pad = dilation[0] * (kernel_size[0] - 1) + (1 - stride[0]) - self.spatial_pad = (padding, padding, padding, padding) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states_prev = hidden_states[:, :, :1, ...].repeat(1, 1, self.temporal_pad, 1, 1) - hidden_states = torch.cat([hidden_states_prev, hidden_states], dim=2) - hidden_states = F.pad(hidden_states, (*self.spatial_pad, 0, 0), mode=self.pad_mode, value=0.0) - return super().forward(hidden_states) - - -class CosmosCausalGroupNorm(torch.nn.Module): - def __init__(self, in_channels: int, num_groups: int = 1): - super().__init__() - self.norm = nn.GroupNorm( - num_groups=num_groups, - num_channels=in_channels, - eps=1e-6, - affine=True, - ) - self.num_groups = num_groups - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.num_groups == 1: - batch_size = hidden_states.size(0) - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) # [B, C, T, H, W] -> [B * T, C, H, W] - hidden_states = self.norm(hidden_states) - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute( - 0, 2, 1, 3, 4 - ) # [B * T, C, H, W] -> [B, C, T, H, W] - else: - hidden_states = self.norm(hidden_states) - return hidden_states - - -class CosmosPatchEmbed3d(nn.Module): - def __init__(self, patch_size: int = 1, patch_method: str = "haar") -> None: - super().__init__() - - self.patch_size = patch_size - self.patch_method = patch_method - - wavelets = _WAVELETS.get(patch_method).clone() - arange = torch.arange(wavelets.shape[0]) - - self.register_buffer("wavelets", wavelets, persistent=False) - self.register_buffer("_arange", arange, persistent=False) - - def _dwt(self, hidden_states: torch.Tensor, mode: str = "reflect", rescale=False) -> torch.Tensor: - dtype = hidden_states.dtype - wavelets = self.wavelets - - n = wavelets.shape[0] - g = hidden_states.shape[1] - hl = wavelets.flip(0).reshape(1, 1, -1).repeat(g, 1, 1) - hh = (wavelets * ((-1) ** self._arange)).reshape(1, 1, -1).repeat(g, 1, 1) - hh = hh.to(dtype=dtype) - hl = hl.to(dtype=dtype) - - # Handles temporal axis - hidden_states = F.pad(hidden_states, pad=(max(0, n - 2), n - 1, n - 2, n - 1, n - 2, n - 1), mode=mode).to( - dtype - ) - xl = F.conv3d(hidden_states, hl.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1)) - xh = F.conv3d(hidden_states, hh.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1)) - - # Handles spatial axes - xll = F.conv3d(xl, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)) - xlh = F.conv3d(xl, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)) - xhl = F.conv3d(xh, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)) - xhh = F.conv3d(xh, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)) - - xlll = F.conv3d(xll, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - xllh = F.conv3d(xll, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - xlhl = F.conv3d(xlh, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - xlhh = F.conv3d(xlh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - xhll = F.conv3d(xhl, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - xhlh = F.conv3d(xhl, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - xhhl = F.conv3d(xhh, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - xhhh = F.conv3d(xhh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - - hidden_states = torch.cat([xlll, xllh, xlhl, xlhh, xhll, xhlh, xhhl, xhhh], dim=1) - if rescale: - hidden_states = hidden_states / 8**0.5 - return hidden_states - - def _haar(self, hidden_states: torch.Tensor) -> torch.Tensor: - xi, xv = torch.split(hidden_states, [1, hidden_states.shape[2] - 1], dim=2) - hidden_states = torch.cat([xi.repeat_interleave(self.patch_size, dim=2), xv], dim=2) - for _ in range(int(math.log2(self.patch_size))): - hidden_states = self._dwt(hidden_states, rescale=True) - return hidden_states - - def _arrange(self, hidden_states: torch.Tensor) -> torch.Tensor: - xi, xv = torch.split(hidden_states, [1, hidden_states.shape[2] - 1], dim=2) - hidden_states = torch.cat([xi.repeat_interleave(self.patch_size, dim=2), xv], dim=2) - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p = self.patch_size - - hidden_states = hidden_states.reshape( - batch_size, num_channels, num_frames // p, p, height // p, p, width // p, p - ) - hidden_states = hidden_states.permute(0, 1, 3, 5, 7, 2, 4, 6).flatten(1, 4).contiguous() - return hidden_states - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.patch_method == "haar": - return self._haar(hidden_states) - elif self.patch_method == "rearrange": - return self._arrange(hidden_states) - else: - raise ValueError(f"Unsupported patch method: {self.patch_method}") - - -class CosmosUnpatcher3d(nn.Module): - def __init__(self, patch_size: int = 1, patch_method: str = "haar"): - super().__init__() - - self.patch_size = patch_size - self.patch_method = patch_method - - wavelets = _WAVELETS.get(patch_method).clone() - arange = torch.arange(wavelets.shape[0]) - - self.register_buffer("wavelets", wavelets, persistent=False) - self.register_buffer("_arange", arange, persistent=False) - - def _idwt(self, hidden_states: torch.Tensor, rescale: bool = False) -> torch.Tensor: - device = hidden_states.device - dtype = hidden_states.dtype - h = self.wavelets.to(device) - - g = hidden_states.shape[1] // 8 # split into 8 spatio-temporal filtered tesnors. - hl = h.flip([0]).reshape(1, 1, -1).repeat([g, 1, 1]) - hh = (h * ((-1) ** self._arange.to(device))).reshape(1, 1, -1).repeat(g, 1, 1) - hl = hl.to(dtype=dtype) - hh = hh.to(dtype=dtype) - - xlll, xllh, xlhl, xlhh, xhll, xhlh, xhhl, xhhh = torch.chunk(hidden_states, 8, dim=1) - - # Handle height transposed convolutions - xll = F.conv_transpose3d(xlll, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - xll = F.conv_transpose3d(xllh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) + xll - - xlh = F.conv_transpose3d(xlhl, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - xlh = F.conv_transpose3d(xlhh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) + xlh - - xhl = F.conv_transpose3d(xhll, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - xhl = F.conv_transpose3d(xhlh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) + xhl - - xhh = F.conv_transpose3d(xhhl, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) - xhh = F.conv_transpose3d(xhhh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)) + xhh - - # Handles width transposed convolutions - xl = F.conv_transpose3d(xll, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)) - xl = F.conv_transpose3d(xlh, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)) + xl - xh = F.conv_transpose3d(xhl, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)) - xh = F.conv_transpose3d(xhh, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)) + xh - - # Handles time axis transposed convolutions - hidden_states = F.conv_transpose3d(xl, hl.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1)) - hidden_states = ( - F.conv_transpose3d(xh, hh.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1)) + hidden_states - ) - - if rescale: - hidden_states = hidden_states * 8**0.5 - - return hidden_states - - def _ihaar(self, hidden_states: torch.Tensor) -> torch.Tensor: - for _ in range(int(math.log2(self.patch_size))): - hidden_states = self._idwt(hidden_states, rescale=True) - hidden_states = hidden_states[:, :, self.patch_size - 1 :, ...] - return hidden_states - - def _irearrange(self, hidden_states: torch.Tensor) -> torch.Tensor: - p = self.patch_size - hidden_states = hidden_states.unflatten(1, (-1, p, p, p)) - hidden_states = hidden_states.permute(0, 1, 5, 2, 6, 3, 7, 4) - hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - hidden_states = hidden_states[:, :, p - 1 :, ...] - return hidden_states - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.patch_method == "haar": - return self._ihaar(hidden_states) - elif self.patch_method == "rearrange": - return self._irearrange(hidden_states) - else: - raise ValueError("Unknown patch method: " + self.patch_method) - - -class CosmosConvProjection3d(nn.Module): - def __init__(self, in_channels: int, out_channels: int) -> None: - super().__init__() - - self.conv_s = CosmosCausalConv3d(in_channels, out_channels, kernel_size=(1, 3, 3), stride=1, padding=1) - self.conv_t = CosmosCausalConv3d(out_channels, out_channels, kernel_size=(3, 1, 1), stride=1, padding=0) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv_s(hidden_states) - hidden_states = self.conv_t(hidden_states) - return hidden_states - - -class CosmosResnetBlock3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - dropout: float = 0.0, - num_groups: int = 1, - ) -> None: - super().__init__() - out_channels = out_channels or in_channels - - self.norm1 = CosmosCausalGroupNorm(in_channels, num_groups) - self.conv1 = CosmosConvProjection3d(in_channels, out_channels) - - self.norm2 = CosmosCausalGroupNorm(out_channels, num_groups) - self.dropout = nn.Dropout(dropout) - self.conv2 = CosmosConvProjection3d(out_channels, out_channels) - - if in_channels != out_channels: - self.conv_shortcut = CosmosCausalConv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0) - else: - self.conv_shortcut = nn.Identity() - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - residual = hidden_states - residual = self.conv_shortcut(residual) - - hidden_states = self.norm1(hidden_states) - hidden_states = F.silu(hidden_states) - hidden_states = self.conv1(hidden_states) - - hidden_states = self.norm2(hidden_states) - hidden_states = F.silu(hidden_states) - hidden_states = self.dropout(hidden_states) - hidden_states = self.conv2(hidden_states) - - return hidden_states + residual - - -class CosmosDownsample3d(nn.Module): - def __init__( - self, - in_channels: int, - spatial_downsample: bool = True, - temporal_downsample: bool = True, - ) -> None: - super().__init__() - - self.spatial_downsample = spatial_downsample - self.temporal_downsample = temporal_downsample - - self.conv1 = nn.Identity() - self.conv2 = nn.Identity() - self.conv3 = nn.Identity() - - if spatial_downsample: - self.conv1 = CosmosCausalConv3d( - in_channels, in_channels, kernel_size=(1, 3, 3), stride=(1, 2, 2), padding=0 - ) - if temporal_downsample: - self.conv2 = CosmosCausalConv3d( - in_channels, in_channels, kernel_size=(3, 1, 1), stride=(2, 1, 1), padding=0 - ) - if spatial_downsample or temporal_downsample: - self.conv3 = CosmosCausalConv3d( - in_channels, in_channels, kernel_size=(1, 1, 1), stride=(1, 1, 1), padding=0 - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if not self.spatial_downsample and not self.temporal_downsample: - return hidden_states - - if self.spatial_downsample: - pad = (0, 1, 0, 1, 0, 0) - hidden_states = F.pad(hidden_states, pad, mode="constant", value=0) - conv_out = self.conv1(hidden_states) - pool_out = F.avg_pool3d(hidden_states, kernel_size=(1, 2, 2), stride=(1, 2, 2)) - hidden_states = conv_out + pool_out - - if self.temporal_downsample: - hidden_states = torch.cat([hidden_states[:, :, :1, ...], hidden_states], dim=2) - conv_out = self.conv2(hidden_states) - pool_out = F.avg_pool3d(hidden_states, kernel_size=(2, 1, 1), stride=(2, 1, 1)) - hidden_states = conv_out + pool_out - - hidden_states = self.conv3(hidden_states) - return hidden_states - - -class CosmosUpsample3d(nn.Module): - def __init__( - self, - in_channels: int, - spatial_upsample: bool = True, - temporal_upsample: bool = True, - ) -> None: - super().__init__() - - self.spatial_upsample = spatial_upsample - self.temporal_upsample = temporal_upsample - - self.conv1 = nn.Identity() - self.conv2 = nn.Identity() - self.conv3 = nn.Identity() - - if temporal_upsample: - self.conv1 = CosmosCausalConv3d( - in_channels, in_channels, kernel_size=(3, 1, 1), stride=(1, 1, 1), padding=0 - ) - if spatial_upsample: - self.conv2 = CosmosCausalConv3d( - in_channels, in_channels, kernel_size=(1, 3, 3), stride=(1, 1, 1), padding=1 - ) - if spatial_upsample or temporal_upsample: - self.conv3 = CosmosCausalConv3d( - in_channels, in_channels, kernel_size=(1, 1, 1), stride=(1, 1, 1), padding=0 - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if not self.spatial_upsample and not self.temporal_upsample: - return hidden_states - - if self.temporal_upsample: - num_frames = hidden_states.size(2) - time_factor = int(1.0 + 1.0 * (num_frames > 1)) - hidden_states = hidden_states.repeat_interleave(int(time_factor), dim=2) - hidden_states = hidden_states[..., time_factor - 1 :, :, :] - hidden_states = self.conv1(hidden_states) + hidden_states - - if self.spatial_upsample: - hidden_states = hidden_states.repeat_interleave(2, dim=3).repeat_interleave(2, dim=4) - hidden_states = self.conv2(hidden_states) + hidden_states - - hidden_states = self.conv3(hidden_states) - return hidden_states - - -class CosmosCausalAttention(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - num_groups: int = 1, - dropout: float = 0.0, - processor: "CosmosSpatialAttentionProcessor2_0" | "CosmosTemporalAttentionProcessor2_0" = None, - ) -> None: - super().__init__() - self.num_attention_heads = num_attention_heads - - self.norm = CosmosCausalGroupNorm(attention_head_dim, num_groups=num_groups) - self.to_q = CosmosCausalConv3d(attention_head_dim, attention_head_dim, kernel_size=1, stride=1, padding=0) - self.to_k = CosmosCausalConv3d(attention_head_dim, attention_head_dim, kernel_size=1, stride=1, padding=0) - self.to_v = CosmosCausalConv3d(attention_head_dim, attention_head_dim, kernel_size=1, stride=1, padding=0) - self.to_out = nn.ModuleList([]) - self.to_out.append( - CosmosCausalConv3d(attention_head_dim, attention_head_dim, kernel_size=1, stride=1, padding=0) - ) - self.to_out.append(nn.Dropout(dropout)) - - self.processor = processor - if self.processor is None: - raise ValueError("CosmosCausalAttention requires a processor.") - - def forward(self, hidden_states: torch.Tensor, attention_mask: torch.Tensor | None = None) -> torch.Tensor: - return self.processor(self, hidden_states=hidden_states, attention_mask=attention_mask) - - -class CosmosSpatialAttentionProcessor2_0: - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "CosmosSpatialAttentionProcessor2_0 requires PyTorch 2.0 or higher. To use it, please upgrade PyTorch." - ) - - def __call__( - self, attn: CosmosCausalAttention, hidden_states: torch.Tensor, attention_mask: torch.Tensor | None = None - ) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - residual = hidden_states - - hidden_states = attn.norm(hidden_states) - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - # [B, C, T, H, W] -> [B * T, H * W, C] - query = query.permute(0, 2, 3, 4, 1).flatten(2, 3).flatten(0, 1) - key = key.permute(0, 2, 3, 4, 1).flatten(2, 3).flatten(0, 1) - value = value.permute(0, 2, 3, 4, 1).flatten(2, 3).flatten(0, 1) - - # [B * T, H * W, C] -> [B * T, N, H * W, C // N] - query = query.unflatten(2, (attn.num_attention_heads, -1)).transpose(1, 2) - key = key.unflatten(2, (attn.num_attention_heads, -1)).transpose(1, 2) - value = value.unflatten(2, (attn.num_attention_heads, -1)).transpose(1, 2) - - hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask) - hidden_states = hidden_states.transpose(1, 2).flatten(2, 3).type_as(query) - hidden_states = hidden_states.unflatten(1, (height, width)).unflatten(0, (batch_size, num_frames)) - hidden_states = hidden_states.permute(0, 4, 1, 2, 3) - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - return hidden_states + residual - - -class CosmosTemporalAttentionProcessor2_0: - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "CosmosSpatialAttentionProcessor2_0 requires PyTorch 2.0 or higher. To use it, please upgrade PyTorch." - ) - - def __call__( - self, attn: CosmosCausalAttention, hidden_states: torch.Tensor, attention_mask: torch.Tensor | None = None - ) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - residual = hidden_states - - hidden_states = attn.norm(hidden_states) - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - # [B, C, T, H, W] -> [B * T, H * W, C] - query = query.permute(0, 3, 4, 2, 1).flatten(0, 2) - key = key.permute(0, 3, 4, 2, 1).flatten(0, 2) - value = value.permute(0, 3, 4, 2, 1).flatten(0, 2) - - # [B * T, H * W, C] -> [B * T, N, H * W, C // N] - query = query.unflatten(2, (attn.num_attention_heads, -1)).transpose(1, 2) - key = key.unflatten(2, (attn.num_attention_heads, -1)).transpose(1, 2) - value = value.unflatten(2, (attn.num_attention_heads, -1)).transpose(1, 2) - - hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask) - hidden_states = hidden_states.transpose(1, 2).flatten(2, 3).type_as(query) - hidden_states = hidden_states.unflatten(0, (batch_size, height, width)) - hidden_states = hidden_states.permute(0, 4, 3, 1, 2) - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - return hidden_states + residual - - -class CosmosDownBlock3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int, - dropout: float, - use_attention: bool, - use_downsample: bool, - spatial_downsample: bool, - temporal_downsample: bool, - ) -> None: - super().__init__() - - resnets, attentions, temp_attentions = [], [], [] - in_channel, out_channel = in_channels, out_channels - - for _ in range(num_layers): - resnets.append(CosmosResnetBlock3d(in_channel, out_channel, dropout, num_groups=1)) - in_channel = out_channel - - if use_attention: - attentions.append( - CosmosCausalAttention( - num_attention_heads=1, - attention_head_dim=out_channel, - num_groups=1, - dropout=dropout, - processor=CosmosSpatialAttentionProcessor2_0(), - ) - ) - temp_attentions.append( - CosmosCausalAttention( - num_attention_heads=1, - attention_head_dim=out_channel, - num_groups=1, - dropout=dropout, - processor=CosmosTemporalAttentionProcessor2_0(), - ) - ) - else: - attentions.append(None) - temp_attentions.append(None) - - self.resnets = nn.ModuleList(resnets) - self.attentions = nn.ModuleList(attentions) - self.temp_attentions = nn.ModuleList(temp_attentions) - - self.downsamplers = None - if use_downsample: - self.downsamplers = nn.ModuleList([]) - self.downsamplers.append(CosmosDownsample3d(out_channel, spatial_downsample, temporal_downsample)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - for resnet, attention, temp_attention in zip(self.resnets, self.attentions, self.temp_attentions): - hidden_states = resnet(hidden_states) - if attention is not None: - hidden_states = attention(hidden_states) - if temp_attention is not None: - num_frames = hidden_states.size(2) - attention_mask = torch.tril(hidden_states.new_ones(num_frames, num_frames)).bool() - hidden_states = temp_attention(hidden_states, attention_mask) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - return hidden_states - - -class CosmosMidBlock3d(nn.Module): - def __init__(self, in_channels: int, num_layers: int, dropout: float, num_groups: int = 1) -> None: - super().__init__() - - resnets, attentions, temp_attentions = [], [], [] - - resnets.append(CosmosResnetBlock3d(in_channels, in_channels, dropout, num_groups)) - for _ in range(num_layers): - attentions.append( - CosmosCausalAttention( - num_attention_heads=1, - attention_head_dim=in_channels, - num_groups=num_groups, - dropout=dropout, - processor=CosmosSpatialAttentionProcessor2_0(), - ) - ) - temp_attentions.append( - CosmosCausalAttention( - num_attention_heads=1, - attention_head_dim=in_channels, - num_groups=num_groups, - dropout=dropout, - processor=CosmosTemporalAttentionProcessor2_0(), - ) - ) - resnets.append(CosmosResnetBlock3d(in_channels, in_channels, dropout, num_groups)) - - self.resnets = nn.ModuleList(resnets) - self.attentions = nn.ModuleList(attentions) - self.temp_attentions = nn.ModuleList(temp_attentions) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.resnets[0](hidden_states) - - for attention, temp_attention, resnet in zip(self.attentions, self.temp_attentions, self.resnets[1:]): - num_frames = hidden_states.size(2) - attention_mask = torch.tril(hidden_states.new_ones(num_frames, num_frames)).bool() - - hidden_states = attention(hidden_states) - hidden_states = temp_attention(hidden_states, attention_mask) - hidden_states = resnet(hidden_states) - - return hidden_states - - -class CosmosUpBlock3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int, - dropout: float, - use_attention: bool, - use_upsample: bool, - spatial_upsample: bool, - temporal_upsample: bool, - ) -> None: - super().__init__() - - resnets, attention, temp_attentions = [], [], [] - in_channel, out_channel = in_channels, out_channels - - for _ in range(num_layers): - resnets.append(CosmosResnetBlock3d(in_channel, out_channel, dropout, num_groups=1)) - in_channel = out_channel - - if use_attention: - attention.append( - CosmosCausalAttention( - num_attention_heads=1, - attention_head_dim=out_channel, - num_groups=1, - dropout=dropout, - processor=CosmosSpatialAttentionProcessor2_0(), - ) - ) - temp_attentions.append( - CosmosCausalAttention( - num_attention_heads=1, - attention_head_dim=out_channel, - num_groups=1, - dropout=dropout, - processor=CosmosTemporalAttentionProcessor2_0(), - ) - ) - else: - attention.append(None) - temp_attentions.append(None) - - self.resnets = nn.ModuleList(resnets) - self.attentions = nn.ModuleList(attention) - self.temp_attentions = nn.ModuleList(temp_attentions) - - self.upsamplers = None - if use_upsample: - self.upsamplers = nn.ModuleList([]) - self.upsamplers.append(CosmosUpsample3d(out_channel, spatial_upsample, temporal_upsample)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - for resnet, attention, temp_attention in zip(self.resnets, self.attentions, self.temp_attentions): - hidden_states = resnet(hidden_states) - if attention is not None: - hidden_states = attention(hidden_states) - if temp_attention is not None: - num_frames = hidden_states.size(2) - attention_mask = torch.tril(hidden_states.new_ones(num_frames, num_frames)).bool() - hidden_states = temp_attention(hidden_states, attention_mask) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - return hidden_states - - -class CosmosEncoder3d(nn.Module): - def __init__( - self, - in_channels: int = 3, - out_channels: int = 16, - block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - num_resnet_blocks: int = 2, - attention_resolutions: tuple[int, ...] = (32,), - resolution: int = 1024, - patch_size: int = 4, - patch_type: str = "haar", - dropout: float = 0.0, - spatial_compression_ratio: int = 8, - temporal_compression_ratio: int = 8, - ) -> None: - super().__init__() - inner_dim = in_channels * patch_size**3 - num_spatial_layers = int(math.log2(spatial_compression_ratio)) - int(math.log2(patch_size)) - num_temporal_layers = int(math.log2(temporal_compression_ratio)) - int(math.log2(patch_size)) - - # 1. Input patching & projection - self.patch_embed = CosmosPatchEmbed3d(patch_size, patch_type) - - self.conv_in = CosmosConvProjection3d(inner_dim, block_out_channels[0]) - - # 2. Down blocks - current_resolution = resolution // patch_size - down_blocks = [] - for i in range(len(block_out_channels) - 1): - in_channel = block_out_channels[i] - out_channel = block_out_channels[i + 1] - - use_attention = current_resolution in attention_resolutions - spatial_downsample = temporal_downsample = False - if i < len(block_out_channels) - 2: - use_downsample = True - spatial_downsample = i < num_spatial_layers - temporal_downsample = i < num_temporal_layers - current_resolution = current_resolution // 2 - else: - use_downsample = False - - down_blocks.append( - CosmosDownBlock3d( - in_channel, - out_channel, - num_resnet_blocks, - dropout, - use_attention, - use_downsample, - spatial_downsample, - temporal_downsample, - ) - ) - self.down_blocks = nn.ModuleList(down_blocks) - - # 3. Mid block - self.mid_block = CosmosMidBlock3d(block_out_channels[-1], num_layers=1, dropout=dropout, num_groups=1) - - # 4. Output norm & projection - self.norm_out = CosmosCausalGroupNorm(block_out_channels[-1], num_groups=1) - self.conv_out = CosmosConvProjection3d(block_out_channels[-1], out_channels) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.patch_embed(hidden_states) - hidden_states = self.conv_in(hidden_states) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.down_blocks: - hidden_states = self._gradient_checkpointing_func(block, hidden_states) - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states) - else: - for block in self.down_blocks: - hidden_states = block(hidden_states) - hidden_states = self.mid_block(hidden_states) - - hidden_states = self.norm_out(hidden_states) - hidden_states = F.silu(hidden_states) - hidden_states = self.conv_out(hidden_states) - return hidden_states - - -class CosmosDecoder3d(nn.Module): - def __init__( - self, - in_channels: int = 16, - out_channels: int = 3, - block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - num_resnet_blocks: int = 2, - attention_resolutions: tuple[int, ...] = (32,), - resolution: int = 1024, - patch_size: int = 4, - patch_type: str = "haar", - dropout: float = 0.0, - spatial_compression_ratio: int = 8, - temporal_compression_ratio: int = 8, - ) -> None: - super().__init__() - inner_dim = out_channels * patch_size**3 - num_spatial_layers = int(math.log2(spatial_compression_ratio)) - int(math.log2(patch_size)) - num_temporal_layers = int(math.log2(temporal_compression_ratio)) - int(math.log2(patch_size)) - reversed_block_out_channels = list(reversed(block_out_channels)) - - # 1. Input projection - self.conv_in = CosmosConvProjection3d(in_channels, reversed_block_out_channels[0]) - - # 2. Mid block - self.mid_block = CosmosMidBlock3d(reversed_block_out_channels[0], num_layers=1, dropout=dropout, num_groups=1) - - # 3. Up blocks - current_resolution = (resolution // patch_size) // 2 ** (len(block_out_channels) - 2) - up_blocks = [] - for i in range(len(block_out_channels) - 1): - in_channel = reversed_block_out_channels[i] - out_channel = reversed_block_out_channels[i + 1] - - use_attention = current_resolution in attention_resolutions - spatial_upsample = temporal_upsample = False - if i < len(block_out_channels) - 2: - use_upsample = True - temporal_upsample = 0 < i < num_temporal_layers + 1 - spatial_upsample = temporal_upsample or ( - i < num_spatial_layers and num_spatial_layers > num_temporal_layers - ) - current_resolution = current_resolution * 2 - else: - use_upsample = False - - up_blocks.append( - CosmosUpBlock3d( - in_channel, - out_channel, - num_resnet_blocks + 1, - dropout, - use_attention, - use_upsample, - spatial_upsample, - temporal_upsample, - ) - ) - self.up_blocks = nn.ModuleList(up_blocks) - - # 4. Output norm & projection & unpatching - self.norm_out = CosmosCausalGroupNorm(reversed_block_out_channels[-1], num_groups=1) - self.conv_out = CosmosConvProjection3d(reversed_block_out_channels[-1], inner_dim) - - self.unpatch_embed = CosmosUnpatcher3d(patch_size, patch_type) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv_in(hidden_states) - hidden_states = self.mid_block(hidden_states) - - for block in self.up_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(block, hidden_states) - else: - hidden_states = block(hidden_states) - - hidden_states = self.norm_out(hidden_states) - hidden_states = F.silu(hidden_states) - hidden_states = self.conv_out(hidden_states) - hidden_states = self.unpatch_embed(hidden_states) - return hidden_states - - -class AutoencoderKLCosmos(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - Autoencoder used in [Cosmos](https://huggingface.co/papers/2501.03575). - - Args: - in_channels (`int`, defaults to `3`): - Number of input channels. - out_channels (`int`, defaults to `3`): - Number of output channels. - latent_channels (`int`, defaults to `16`): - Number of latent channels. - encoder_block_out_channels (`tuple[int, ...]`, defaults to `(128, 256, 512, 512)`): - Number of output channels for each encoder down block. - decode_block_out_channels (`tuple[int, ...]`, defaults to `(256, 512, 512, 512)`): - Number of output channels for each decoder up block. - attention_resolutions (`tuple[int, ...]`, defaults to `(32,)`): - list of image/video resolutions at which to apply attention. - resolution (`int`, defaults to `1024`): - Base image/video resolution used for computing whether a block should have attention layers. - num_layers (`int`, defaults to `2`): - Number of resnet blocks in each encoder/decoder block. - patch_size (`int`, defaults to `4`): - Patch size used for patching the input image/video. - patch_type (`str`, defaults to `haar`): - Patch type used for patching the input image/video. Can be either `haar` or `rearrange`. - scaling_factor (`float`, defaults to `1.0`): - The component-wise standard deviation of the trained latent space computed using the first batch of the - training set. This is used to scale the latent space to have unit variance when training the diffusion - model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the - diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 - / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image - Synthesis with Latent Diffusion Models](https://huggingface.co/papers/2112.10752) paper. Not applicable in - Cosmos, but we default to 1.0 for consistency. - spatial_compression_ratio (`int`, defaults to `8`): - The spatial compression ratio to apply in the VAE. The number of downsample blocks is determined using - this. - temporal_compression_ratio (`int`, defaults to `8`): - The temporal compression ratio to apply in the VAE. The number of downsample blocks is determined using - this. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - latent_channels: int = 16, - encoder_block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - decode_block_out_channels: tuple[int, ...] = (256, 512, 512, 512), - attention_resolutions: tuple[int, ...] = (32,), - resolution: int = 1024, - num_layers: int = 2, - patch_size: int = 4, - patch_type: str = "haar", - scaling_factor: float = 1.0, - spatial_compression_ratio: int = 8, - temporal_compression_ratio: int = 8, - latents_mean: list[float] | None = LATENTS_MEAN, - latents_std: list[float] | None = LATENTS_STD, - ) -> None: - super().__init__() - - self.encoder = CosmosEncoder3d( - in_channels=in_channels, - out_channels=latent_channels, - block_out_channels=encoder_block_out_channels, - num_resnet_blocks=num_layers, - attention_resolutions=attention_resolutions, - resolution=resolution, - patch_size=patch_size, - patch_type=patch_type, - spatial_compression_ratio=spatial_compression_ratio, - temporal_compression_ratio=temporal_compression_ratio, - ) - self.decoder = CosmosDecoder3d( - in_channels=latent_channels, - out_channels=out_channels, - block_out_channels=decode_block_out_channels, - num_resnet_blocks=num_layers, - attention_resolutions=attention_resolutions, - resolution=resolution, - patch_size=patch_size, - patch_type=patch_type, - spatial_compression_ratio=spatial_compression_ratio, - temporal_compression_ratio=temporal_compression_ratio, - ) - - self.quant_conv = CosmosCausalConv3d(latent_channels, latent_channels, kernel_size=1, padding=0) - self.post_quant_conv = CosmosCausalConv3d(latent_channels, latent_channels, kernel_size=1, padding=0) - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent - # frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the - # intermediate tiles together, the memory requirement can be lowered. - self.use_tiling = False - - # When decoding temporally long video latents, the memory requirement is very high. By decoding latent frames - # at a fixed frame batch size (based on `self.num_latent_frames_batch_sizes`), the memory requirement can be lowered. - self.use_framewise_encoding = False - self.use_framewise_decoding = False - - # This can be configured based on the amount of GPU memory available. - # `16` for sample frames and `2` for latent frames are sensible defaults for consumer GPUs. - # Setting it to higher values results in higher memory usage. - self.num_sample_frames_batch_size = 16 - self.num_latent_frames_batch_size = 2 - - # The minimal tile height and width for spatial tiling to be used - self.tile_sample_min_height = 512 - self.tile_sample_min_width = 512 - self.tile_sample_min_num_frames = 16 - - # The minimal distance between two spatial tiles - self.tile_sample_stride_height = 448 - self.tile_sample_stride_width = 448 - self.tile_sample_stride_num_frames = 8 - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_sample_min_num_frames: int | None = None, - tile_sample_stride_height: float | None = None, - tile_sample_stride_width: float | None = None, - tile_sample_stride_num_frames: float | None = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_sample_stride_height (`int`, *optional*): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. - tile_sample_stride_width (`int`, *optional*): - The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling - artifacts produced across the width dimension. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_sample_min_num_frames = tile_sample_min_num_frames or self.tile_sample_min_num_frames - self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height - self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width - self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - x = self.encoder(x) - enc = self.quant_conv(x) - return enc - - @apply_forward_hook - def encode(self, x: torch.Tensor, return_dict: bool = True) -> torch.Tensor: - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = IdentityDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | tuple[torch.Tensor]: - z = self.post_quant_conv(z) - dec = self.decoder(z) - - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | tuple[torch.Tensor]: - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> tuple[torch.Tensor] | DecoderOutput: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z).sample - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) diff --git a/diffusers/models/autoencoders/autoencoder_kl_flux2.py b/diffusers/models/autoencoders/autoencoder_kl_flux2.py deleted file mode 100644 index 24a8b024c52450963303f989a3dc10b3977cbc52..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_flux2.py +++ /dev/null @@ -1,496 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import math - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...loaders.single_file_model import FromOriginalModelMixin -from ...utils import deprecate -from ...utils.accelerate_utils import apply_forward_hook -from ..attention import AttentionMixin -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - Attention, - AttnAddedKVProcessor, - AttnProcessor, - FusedAttnProcessor2_0, -) -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, Decoder, DecoderOutput, DiagonalGaussianDistribution, Encoder - - -class AutoencoderKLFlux2( - ModelMixin, AutoencoderMixin, AttentionMixin, ConfigMixin, FromOriginalModelMixin, PeftAdapterMixin -): - r""" - A VAE model with KL loss for encoding images into latents and decoding latent representations into images. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - in_channels (int, *optional*, defaults to 3): Number of channels in the input image. - out_channels (int, *optional*, defaults to 3): Number of channels in the output. - down_block_types (`tuple[str]`, *optional*, defaults to `("DownEncoderBlock2D",)`): - Tuple of downsample block types. - up_block_types (`tuple[str]`, *optional*, defaults to `("UpDecoderBlock2D",)`): - Tuple of upsample block types. - block_out_channels (`tuple[int]`, *optional*, defaults to `(64,)`): - Tuple of block output channels. - act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use. - latent_channels (`int`, *optional*, defaults to 4): Number of channels in the latent space. - sample_size (`int`, *optional*, defaults to `32`): Sample input size. - force_upcast (`bool`, *optional*, default to `True`): - If enabled it will force the VAE to run in float32 for high image resolution pipelines, such as SD-XL. VAE - can be fine-tuned / trained to a lower range without losing too much precision in which case `force_upcast` - can be set to `False` - see: https://huggingface.co/madebyollin/sdxl-vae-fp16-fix - mid_block_add_attention (`bool`, *optional*, default to `True`): - If enabled, the mid_block of the Encoder and Decoder will have attention blocks. If set to false, the - mid_block will only have resnet blocks - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["BasicTransformerBlock", "ResnetBlock2D"] - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - down_block_types: tuple[str, ...] = ( - "DownEncoderBlock2D", - "DownEncoderBlock2D", - "DownEncoderBlock2D", - "DownEncoderBlock2D", - ), - up_block_types: tuple[str, ...] = ( - "UpDecoderBlock2D", - "UpDecoderBlock2D", - "UpDecoderBlock2D", - "UpDecoderBlock2D", - ), - block_out_channels: tuple[int, ...] = ( - 128, - 256, - 512, - 512, - ), - decoder_block_out_channels: tuple[int, ...] | None = None, - layers_per_block: int = 2, - act_fn: str = "silu", - latent_channels: int = 32, - norm_num_groups: int = 32, - sample_size: int = 1024, # YiYi notes: not sure - force_upcast: bool = True, - use_quant_conv: bool = True, - use_post_quant_conv: bool = True, - mid_block_add_attention: bool = True, - batch_norm_eps: float = 1e-4, - batch_norm_momentum: float = 0.1, - patch_size: tuple[int, int] = (2, 2), - ): - super().__init__() - - # pass init params to Encoder - self.encoder = Encoder( - in_channels=in_channels, - out_channels=latent_channels, - down_block_types=down_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - double_z=True, - mid_block_add_attention=mid_block_add_attention, - ) - - # pass init params to Decoder - self.decoder = Decoder( - in_channels=latent_channels, - out_channels=out_channels, - up_block_types=up_block_types, - block_out_channels=decoder_block_out_channels or block_out_channels, - layers_per_block=layers_per_block, - norm_num_groups=norm_num_groups, - act_fn=act_fn, - mid_block_add_attention=mid_block_add_attention, - ) - - self.quant_conv = nn.Conv2d(2 * latent_channels, 2 * latent_channels, 1) if use_quant_conv else None - self.post_quant_conv = nn.Conv2d(latent_channels, latent_channels, 1) if use_post_quant_conv else None - - self.bn = nn.BatchNorm2d( - math.prod(patch_size) * latent_channels, - eps=batch_norm_eps, - momentum=batch_norm_momentum, - affine=False, - track_running_stats=True, - ) - - self.use_slicing = False - self.use_tiling = False - - # only relevant if vae tiling is enabled - self.tile_sample_min_size = self.config.sample_size - sample_size = ( - self.config.sample_size[0] - if isinstance(self.config.sample_size, (list, tuple)) - else self.config.sample_size - ) - self.tile_latent_min_size = int(sample_size / (2 ** (len(self.config.block_out_channels) - 1))) - self.tile_overlap_factor = 0.25 - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, height, width = x.shape - - if self.use_tiling and (width > self.tile_sample_min_size or height > self.tile_sample_min_size): - return self._tiled_encode(x) - - enc = self.encoder(x) - if self.quant_conv is not None: - enc = self.quant_conv(enc) - - return enc - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded images. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - if self.use_tiling and (z.shape[-1] > self.tile_latent_min_size or z.shape[-2] > self.tile_latent_min_size): - return self.tiled_decode(z, return_dict=return_dict) - - if self.post_quant_conv is not None: - z = self.post_quant_conv(z) - - dec = self.decoder(z) - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode( - self, z: torch.FloatTensor, return_dict: bool = True, generator=None - ) -> DecoderOutput | torch.FloatTensor: - """ - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[2], b.shape[2], blend_extent) - for y in range(blend_extent): - b[:, :, y, :] = a[:, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, y, :] * (y / blend_extent) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[3], b.shape[3], blend_extent) - for x in range(blend_extent): - b[:, :, :, x] = a[:, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, x] * (x / blend_extent) - return b - - def _tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - r"""Encode a batch of images using a tiled encoder. - - When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several - steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is - different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the - tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the - output, but they should be much less noticeable. - - Args: - x (`torch.Tensor`): Input batch of images. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - - overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor)) - blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor) - row_limit = self.tile_latent_min_size - blend_extent - - # Split the image into 512x512 tiles and encode them separately. - rows = [] - for i in range(0, x.shape[2], overlap_size): - row = [] - for j in range(0, x.shape[3], overlap_size): - tile = x[:, :, i : i + self.tile_sample_min_size, j : j + self.tile_sample_min_size] - tile = self.encoder(tile) - if self.config.use_quant_conv: - tile = self.quant_conv(tile) - row.append(tile) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=3)) - - enc = torch.cat(result_rows, dim=2) - return enc - - def tiled_encode(self, x: torch.Tensor, return_dict: bool = True) -> AutoencoderKLOutput: - r"""Encode a batch of images using a tiled encoder. - - When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several - steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is - different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the - tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the - output, but they should be much less noticeable. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - [`~models.autoencoder_kl.AutoencoderKLOutput`] or `tuple`: - If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain - `tuple` is returned. - """ - deprecation_message = ( - "The tiled_encode implementation supporting the `return_dict` parameter is deprecated. In the future, the " - "implementation of this method will be replaced with that of `_tiled_encode` and you will no longer be able " - "to pass `return_dict`. You will also have to create a `DiagonalGaussianDistribution()` from the returned value." - ) - deprecate("tiled_encode", "1.0.0", deprecation_message, standard_warn=False) - - overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor)) - blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor) - row_limit = self.tile_latent_min_size - blend_extent - - # Split the image into 512x512 tiles and encode them separately. - rows = [] - for i in range(0, x.shape[2], overlap_size): - row = [] - for j in range(0, x.shape[3], overlap_size): - tile = x[:, :, i : i + self.tile_sample_min_size, j : j + self.tile_sample_min_size] - tile = self.encoder(tile) - if self.config.use_quant_conv: - tile = self.quant_conv(tile) - row.append(tile) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=3)) - - moments = torch.cat(result_rows, dim=2) - posterior = DiagonalGaussianDistribution(moments) - - if not return_dict: - return (posterior,) - - return AutoencoderKLOutput(latent_dist=posterior) - - def tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor)) - blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor) - row_limit = self.tile_sample_min_size - blend_extent - - # Split z into overlapping 64x64 tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, z.shape[2], overlap_size): - row = [] - for j in range(0, z.shape[3], overlap_size): - tile = z[:, :, i : i + self.tile_latent_min_size, j : j + self.tile_latent_min_size] - if self.config.use_post_quant_conv: - tile = self.post_quant_conv(tile) - decoded = self.decoder(tile) - row.append(decoded) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=3)) - - dec = torch.cat(result_rows, dim=2) - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z).sample - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) diff --git a/diffusers/models/autoencoders/autoencoder_kl_hunyuan_video.py b/diffusers/models/autoencoders/autoencoder_kl_hunyuan_video.py deleted file mode 100644 index fece756ebec65441ed4a7b61c6cfa910694571d8..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_hunyuan_video.py +++ /dev/null @@ -1,1080 +0,0 @@ -# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import numpy as np -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..attention_processor import Attention -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def prepare_causal_attention_mask( - num_frames: int, height_width: int, dtype: torch.dtype, device: torch.device, batch_size: int = None -) -> torch.Tensor: - indices = torch.arange(1, num_frames + 1, dtype=torch.int32, device=device) - indices_blocks = indices.repeat_interleave(height_width) - x, y = torch.meshgrid(indices_blocks, indices_blocks, indexing="xy") - mask = torch.where(x <= y, 0, -float("inf")).to(dtype=dtype) - - if batch_size is not None: - mask = mask.unsqueeze(0).expand(batch_size, -1, -1) - return mask - - -class HunyuanVideoCausalConv3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int] = 3, - stride: int | tuple[int, int, int] = 1, - padding: int | tuple[int, int, int] = 0, - dilation: int | tuple[int, int, int] = 1, - bias: bool = True, - pad_mode: str = "replicate", - ) -> None: - super().__init__() - - kernel_size = (kernel_size, kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size - - self.pad_mode = pad_mode - self.time_causal_padding = ( - kernel_size[0] // 2, - kernel_size[0] // 2, - kernel_size[1] // 2, - kernel_size[1] // 2, - kernel_size[2] - 1, - 0, - ) - - self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, dilation, bias=bias) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = F.pad(hidden_states, self.time_causal_padding, mode=self.pad_mode) - return self.conv(hidden_states) - - -class HunyuanVideoUpsampleCausal3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - kernel_size: int = 3, - stride: int = 1, - bias: bool = True, - upsample_factor: tuple[float, float, float] = (2, 2, 2), - ) -> None: - super().__init__() - - out_channels = out_channels or in_channels - self.upsample_factor = upsample_factor - - self.conv = HunyuanVideoCausalConv3d(in_channels, out_channels, kernel_size, stride, bias=bias) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - num_frames = hidden_states.size(2) - - first_frame, other_frames = hidden_states.split((1, num_frames - 1), dim=2) - first_frame = F.interpolate( - first_frame.squeeze(2), scale_factor=self.upsample_factor[1:], mode="nearest" - ).unsqueeze(2) - - if num_frames > 1: - # See: https://github.com/pytorch/pytorch/issues/81665 - # Unless you have a version of pytorch where non-contiguous implementation of F.interpolate - # is fixed, this will raise either a runtime error, or fail silently with bad outputs. - # If you are encountering an error here, make sure to try running encoding/decoding with - # `vae.enable_tiling()` first. If that doesn't work, open an issue at: - # https://github.com/huggingface/diffusers/issues - other_frames = other_frames.contiguous() - other_frames = F.interpolate(other_frames, scale_factor=self.upsample_factor, mode="nearest") - hidden_states = torch.cat((first_frame, other_frames), dim=2) - else: - hidden_states = first_frame - - hidden_states = self.conv(hidden_states) - return hidden_states - - -class HunyuanVideoDownsampleCausal3D(nn.Module): - def __init__( - self, - channels: int, - out_channels: int | None = None, - padding: int = 1, - kernel_size: int = 3, - bias: bool = True, - stride=2, - ) -> None: - super().__init__() - out_channels = out_channels or channels - - self.conv = HunyuanVideoCausalConv3d(channels, out_channels, kernel_size, stride, padding, bias=bias) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv(hidden_states) - return hidden_states - - -class HunyuanVideoResnetBlockCausal3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - dropout: float = 0.0, - groups: int = 32, - eps: float = 1e-6, - non_linearity: str = "swish", - ) -> None: - super().__init__() - out_channels = out_channels or in_channels - - self.nonlinearity = get_activation(non_linearity) - - self.norm1 = nn.GroupNorm(groups, in_channels, eps=eps, affine=True) - self.conv1 = HunyuanVideoCausalConv3d(in_channels, out_channels, 3, 1, 0) - - self.norm2 = nn.GroupNorm(groups, out_channels, eps=eps, affine=True) - self.dropout = nn.Dropout(dropout) - self.conv2 = HunyuanVideoCausalConv3d(out_channels, out_channels, 3, 1, 0) - - self.conv_shortcut = None - if in_channels != out_channels: - self.conv_shortcut = HunyuanVideoCausalConv3d(in_channels, out_channels, 1, 1, 0) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = hidden_states.contiguous() - residual = hidden_states - - hidden_states = self.norm1(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv1(hidden_states) - - hidden_states = self.norm2(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.dropout(hidden_states) - hidden_states = self.conv2(hidden_states) - - if self.conv_shortcut is not None: - residual = self.conv_shortcut(residual) - - hidden_states = hidden_states + residual - return hidden_states - - -class HunyuanVideoMidBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - add_attention: bool = True, - attention_head_dim: int = 1, - ) -> None: - super().__init__() - resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32) - self.add_attention = add_attention - - # There is always at least one resnet - resnets = [ - HunyuanVideoResnetBlockCausal3D( - in_channels=in_channels, - out_channels=in_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - non_linearity=resnet_act_fn, - ) - ] - attentions = [] - - for _ in range(num_layers): - if self.add_attention: - attentions.append( - Attention( - in_channels, - heads=in_channels // attention_head_dim, - dim_head=attention_head_dim, - eps=resnet_eps, - norm_num_groups=resnet_groups, - residual_connection=True, - bias=True, - upcast_softmax=True, - _from_deprecated_attn_block=True, - ) - ) - else: - attentions.append(None) - - resnets.append( - HunyuanVideoResnetBlockCausal3D( - in_channels=in_channels, - out_channels=in_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - non_linearity=resnet_act_fn, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(self.resnets[0], hidden_states) - - for attn, resnet in zip(self.attentions, self.resnets[1:]): - if attn is not None: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - hidden_states = hidden_states.permute(0, 2, 3, 4, 1).flatten(1, 3) - attention_mask = prepare_causal_attention_mask( - num_frames, height * width, hidden_states.dtype, hidden_states.device, batch_size=batch_size - ) - hidden_states = attn(hidden_states, attention_mask=attention_mask) - hidden_states = hidden_states.unflatten(1, (num_frames, height, width)).permute(0, 4, 1, 2, 3) - - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states) - - else: - hidden_states = self.resnets[0](hidden_states) - - for attn, resnet in zip(self.attentions, self.resnets[1:]): - if attn is not None: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - hidden_states = hidden_states.permute(0, 2, 3, 4, 1).flatten(1, 3) - attention_mask = prepare_causal_attention_mask( - num_frames, height * width, hidden_states.dtype, hidden_states.device, batch_size=batch_size - ) - hidden_states = attn(hidden_states, attention_mask=attention_mask) - hidden_states = hidden_states.unflatten(1, (num_frames, height, width)).permute(0, 4, 1, 2, 3) - - hidden_states = resnet(hidden_states) - - return hidden_states - - -class HunyuanVideoDownBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - add_downsample: bool = True, - downsample_stride: int = 2, - downsample_padding: int = 1, - ) -> None: - super().__init__() - resnets = [] - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - HunyuanVideoResnetBlockCausal3D( - in_channels=in_channels, - out_channels=out_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - non_linearity=resnet_act_fn, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - HunyuanVideoDownsampleCausal3D( - out_channels, - out_channels=out_channels, - padding=downsample_padding, - stride=downsample_stride, - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if torch.is_grad_enabled() and self.gradient_checkpointing: - for resnet in self.resnets: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states) - else: - for resnet in self.resnets: - hidden_states = resnet(hidden_states) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - return hidden_states - - -class HunyuanVideoUpBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - add_upsample: bool = True, - upsample_scale_factor: tuple[int, int, int] = (2, 2, 2), - ) -> None: - super().__init__() - resnets = [] - - for i in range(num_layers): - input_channels = in_channels if i == 0 else out_channels - - resnets.append( - HunyuanVideoResnetBlockCausal3D( - in_channels=input_channels, - out_channels=out_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - non_linearity=resnet_act_fn, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList( - [ - HunyuanVideoUpsampleCausal3D( - out_channels, - out_channels=out_channels, - upsample_factor=upsample_scale_factor, - ) - ] - ) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if torch.is_grad_enabled() and self.gradient_checkpointing: - for resnet in self.resnets: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states) - - else: - for resnet in self.resnets: - hidden_states = resnet(hidden_states) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - return hidden_states - - -class HunyuanVideoEncoder3D(nn.Module): - r""" - Causal encoder for 3D video-like data introduced in [Hunyuan Video](https://huggingface.co/papers/2412.03603). - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - down_block_types: tuple[str, ...] = ( - "HunyuanVideoDownBlock3D", - "HunyuanVideoDownBlock3D", - "HunyuanVideoDownBlock3D", - "HunyuanVideoDownBlock3D", - ), - block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - layers_per_block: int = 2, - norm_num_groups: int = 32, - act_fn: str = "silu", - double_z: bool = True, - mid_block_add_attention=True, - temporal_compression_ratio: int = 4, - spatial_compression_ratio: int = 8, - ) -> None: - super().__init__() - - self.conv_in = HunyuanVideoCausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1) - self.mid_block = None - self.down_blocks = nn.ModuleList([]) - - output_channel = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - if down_block_type != "HunyuanVideoDownBlock3D": - raise ValueError(f"Unsupported down_block_type: {down_block_type}") - - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - num_spatial_downsample_layers = int(np.log2(spatial_compression_ratio)) - num_time_downsample_layers = int(np.log2(temporal_compression_ratio)) - - if temporal_compression_ratio == 4: - add_spatial_downsample = bool(i < num_spatial_downsample_layers) - add_time_downsample = bool( - i >= (len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block - ) - elif temporal_compression_ratio == 8: - add_spatial_downsample = bool(i < num_spatial_downsample_layers) - add_time_downsample = bool(i < num_time_downsample_layers) - else: - raise ValueError(f"Unsupported time_compression_ratio: {temporal_compression_ratio}") - - downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1) - downsample_stride_T = (2,) if add_time_downsample else (1,) - downsample_stride = tuple(downsample_stride_T + downsample_stride_HW) - - down_block = HunyuanVideoDownBlock3D( - num_layers=layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - add_downsample=bool(add_spatial_downsample or add_time_downsample), - resnet_eps=1e-6, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - downsample_stride=downsample_stride, - downsample_padding=0, - ) - - self.down_blocks.append(down_block) - - self.mid_block = HunyuanVideoMidBlock3D( - in_channels=block_out_channels[-1], - resnet_eps=1e-6, - resnet_act_fn=act_fn, - attention_head_dim=block_out_channels[-1], - resnet_groups=norm_num_groups, - add_attention=mid_block_add_attention, - ) - - self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6) - self.conv_act = nn.SiLU() - - conv_out_channels = 2 * out_channels if double_z else out_channels - self.conv_out = HunyuanVideoCausalConv3d(block_out_channels[-1], conv_out_channels, kernel_size=3) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv_in(hidden_states) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for down_block in self.down_blocks: - hidden_states = self._gradient_checkpointing_func(down_block, hidden_states) - - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states) - else: - for down_block in self.down_blocks: - hidden_states = down_block(hidden_states) - - hidden_states = self.mid_block(hidden_states) - - hidden_states = self.conv_norm_out(hidden_states) - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states) - - return hidden_states - - -class HunyuanVideoDecoder3D(nn.Module): - r""" - Causal decoder for 3D video-like data introduced in [Hunyuan Video](https://huggingface.co/papers/2412.03603). - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - up_block_types: tuple[str, ...] = ( - "HunyuanVideoUpBlock3D", - "HunyuanVideoUpBlock3D", - "HunyuanVideoUpBlock3D", - "HunyuanVideoUpBlock3D", - ), - block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - layers_per_block: int = 2, - norm_num_groups: int = 32, - act_fn: str = "silu", - mid_block_add_attention=True, - time_compression_ratio: int = 4, - spatial_compression_ratio: int = 8, - ): - super().__init__() - self.layers_per_block = layers_per_block - - self.conv_in = HunyuanVideoCausalConv3d(in_channels, block_out_channels[-1], kernel_size=3, stride=1) - self.up_blocks = nn.ModuleList([]) - - # mid - self.mid_block = HunyuanVideoMidBlock3D( - in_channels=block_out_channels[-1], - resnet_eps=1e-6, - resnet_act_fn=act_fn, - attention_head_dim=block_out_channels[-1], - resnet_groups=norm_num_groups, - add_attention=mid_block_add_attention, - ) - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - output_channel = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - if up_block_type != "HunyuanVideoUpBlock3D": - raise ValueError(f"Unsupported up_block_type: {up_block_type}") - - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - num_spatial_upsample_layers = int(np.log2(spatial_compression_ratio)) - num_time_upsample_layers = int(np.log2(time_compression_ratio)) - - if time_compression_ratio == 4: - add_spatial_upsample = bool(i < num_spatial_upsample_layers) - add_time_upsample = bool( - i >= len(block_out_channels) - 1 - num_time_upsample_layers and not is_final_block - ) - else: - raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}") - - upsample_scale_factor_HW = (2, 2) if add_spatial_upsample else (1, 1) - upsample_scale_factor_T = (2,) if add_time_upsample else (1,) - upsample_scale_factor = tuple(upsample_scale_factor_T + upsample_scale_factor_HW) - - up_block = HunyuanVideoUpBlock3D( - num_layers=self.layers_per_block + 1, - in_channels=prev_output_channel, - out_channels=output_channel, - add_upsample=bool(add_spatial_upsample or add_time_upsample), - upsample_scale_factor=upsample_scale_factor, - resnet_eps=1e-6, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - ) - - self.up_blocks.append(up_block) - prev_output_channel = output_channel - - # out - self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6) - self.conv_act = nn.SiLU() - self.conv_out = HunyuanVideoCausalConv3d(block_out_channels[0], out_channels, kernel_size=3) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv_in(hidden_states) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states) - - for up_block in self.up_blocks: - hidden_states = self._gradient_checkpointing_func(up_block, hidden_states) - else: - hidden_states = self.mid_block(hidden_states) - - for up_block in self.up_blocks: - hidden_states = up_block(hidden_states) - - # post-process - hidden_states = self.conv_norm_out(hidden_states) - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states) - - return hidden_states - - -class AutoencoderKLHunyuanVideo(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. - Introduced in [HunyuanVideo](https://huggingface.co/papers/2412.03603). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - latent_channels: int = 16, - down_block_types: tuple[str, ...] = ( - "HunyuanVideoDownBlock3D", - "HunyuanVideoDownBlock3D", - "HunyuanVideoDownBlock3D", - "HunyuanVideoDownBlock3D", - ), - up_block_types: tuple[str, ...] = ( - "HunyuanVideoUpBlock3D", - "HunyuanVideoUpBlock3D", - "HunyuanVideoUpBlock3D", - "HunyuanVideoUpBlock3D", - ), - block_out_channels: tuple[int] = (128, 256, 512, 512), - layers_per_block: int = 2, - act_fn: str = "silu", - norm_num_groups: int = 32, - scaling_factor: float = 0.476986, - spatial_compression_ratio: int = 8, - temporal_compression_ratio: int = 4, - mid_block_add_attention: bool = True, - ) -> None: - super().__init__() - - self.time_compression_ratio = temporal_compression_ratio - - self.encoder = HunyuanVideoEncoder3D( - in_channels=in_channels, - out_channels=latent_channels, - down_block_types=down_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - norm_num_groups=norm_num_groups, - act_fn=act_fn, - double_z=True, - mid_block_add_attention=mid_block_add_attention, - temporal_compression_ratio=temporal_compression_ratio, - spatial_compression_ratio=spatial_compression_ratio, - ) - - self.decoder = HunyuanVideoDecoder3D( - in_channels=latent_channels, - out_channels=out_channels, - up_block_types=up_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - norm_num_groups=norm_num_groups, - act_fn=act_fn, - time_compression_ratio=temporal_compression_ratio, - spatial_compression_ratio=spatial_compression_ratio, - mid_block_add_attention=mid_block_add_attention, - ) - - self.quant_conv = nn.Conv3d(2 * latent_channels, 2 * latent_channels, kernel_size=1) - self.post_quant_conv = nn.Conv3d(latent_channels, latent_channels, kernel_size=1) - - self.spatial_compression_ratio = spatial_compression_ratio - self.temporal_compression_ratio = temporal_compression_ratio - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent - # frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the - # intermediate tiles together, the memory requirement can be lowered. - self.use_tiling = False - - # When decoding temporally long video latents, the memory requirement is very high. By decoding latent frames - # at a fixed frame batch size (based on `self.tile_sample_min_num_frames`), the memory requirement can be lowered. - self.use_framewise_encoding = True - self.use_framewise_decoding = True - - # The minimal tile height and width for spatial tiling to be used - self.tile_sample_min_height = 256 - self.tile_sample_min_width = 256 - self.tile_sample_min_num_frames = 16 - - # The minimal distance between two spatial tiles - self.tile_sample_stride_height = 192 - self.tile_sample_stride_width = 192 - self.tile_sample_stride_num_frames = 12 - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_sample_min_num_frames: int | None = None, - tile_sample_stride_height: float | None = None, - tile_sample_stride_width: float | None = None, - tile_sample_stride_num_frames: float | None = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_sample_min_num_frames (`int`, *optional*): - The minimum number of frames required for a sample to be separated into tiles across the frame - dimension. - tile_sample_stride_height (`int`, *optional*): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. - tile_sample_stride_width (`int`, *optional*): - The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling - artifacts produced across the width dimension. - tile_sample_stride_num_frames (`int`, *optional*): - The stride between two consecutive frame tiles. This is to ensure that there are no tiling artifacts - produced across the frame dimension. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_sample_min_num_frames = tile_sample_min_num_frames or self.tile_sample_min_num_frames - self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height - self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width - self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = x.shape - - if self.use_framewise_encoding and num_frames > self.tile_sample_min_num_frames: - return self._temporal_tiled_encode(x) - - if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): - return self.tiled_encode(x) - - x = self.encoder(x) - enc = self.quant_conv(x) - return enc - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - r""" - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - batch_size, num_channels, num_frames, height, width = z.shape - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio - - if self.use_framewise_decoding and num_frames > tile_latent_min_num_frames: - return self._temporal_tiled_decode(z, return_dict=return_dict) - - if self.use_tiling and (width > tile_latent_min_width or height > tile_latent_min_height): - return self.tiled_decode(z, return_dict=return_dict) - - z = self.post_quant_conv(z) - dec = self.decoder(z) - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-2], b.shape[-2], blend_extent) - for y in range(blend_extent): - b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-1], b.shape[-1], blend_extent) - for x in range(blend_extent): - b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( - x / blend_extent - ) - return b - - def blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-3], b.shape[-3], blend_extent) - for x in range(blend_extent): - b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (1 - x / blend_extent) + b[:, :, x, :, :] * ( - x / blend_extent - ) - return b - - def tiled_encode(self, x: torch.Tensor) -> AutoencoderKLOutput: - r"""Encode a batch of images using a tiled encoder. - - Args: - x (`torch.Tensor`): Input batch of videos. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - batch_size, num_channels, num_frames, height, width = x.shape - latent_height = height // self.spatial_compression_ratio - latent_width = width // self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = tile_latent_min_height - tile_latent_stride_height - blend_width = tile_latent_min_width - tile_latent_stride_width - - # Split x into overlapping tiles and encode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, self.tile_sample_stride_height): - row = [] - for j in range(0, width, self.tile_sample_stride_width): - tile = x[:, :, :, i : i + self.tile_sample_min_height, j : j + self.tile_sample_min_width] - tile = self.encoder(tile) - tile = self.quant_conv(tile) - row.append(tile) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :tile_latent_stride_height, :tile_latent_stride_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - enc = torch.cat(result_rows, dim=3)[:, :, :, :latent_height, :latent_width] - return enc - - def tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - - batch_size, num_channels, num_frames, height, width = z.shape - sample_height = height * self.spatial_compression_ratio - sample_width = width * self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = self.tile_sample_min_height - self.tile_sample_stride_height - blend_width = self.tile_sample_min_width - self.tile_sample_stride_width - - # Split z into overlapping tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, tile_latent_stride_height): - row = [] - for j in range(0, width, tile_latent_stride_width): - tile = z[:, :, :, i : i + tile_latent_min_height, j : j + tile_latent_min_width] - tile = self.post_quant_conv(tile) - decoded = self.decoder(tile) - row.append(decoded) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, : self.tile_sample_stride_height, : self.tile_sample_stride_width]) - result_rows.append(torch.cat(result_row, dim=-1)) - - dec = torch.cat(result_rows, dim=3)[:, :, :, :sample_height, :sample_width] - - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) - - def _temporal_tiled_encode(self, x: torch.Tensor) -> AutoencoderKLOutput: - batch_size, num_channels, num_frames, height, width = x.shape - latent_num_frames = (num_frames - 1) // self.temporal_compression_ratio + 1 - - tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio - tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio - blend_num_frames = tile_latent_min_num_frames - tile_latent_stride_num_frames - - row = [] - for i in range(0, num_frames, self.tile_sample_stride_num_frames): - tile = x[:, :, i : i + self.tile_sample_min_num_frames + 1, :, :] - if self.use_tiling and (height > self.tile_sample_min_height or width > self.tile_sample_min_width): - tile = self.tiled_encode(tile) - else: - tile = self.encoder(tile) - tile = self.quant_conv(tile) - if i > 0: - tile = tile[:, :, 1:, :, :] - row.append(tile) - - result_row = [] - for i, tile in enumerate(row): - if i > 0: - tile = self.blend_t(row[i - 1], tile, blend_num_frames) - result_row.append(tile[:, :, :tile_latent_stride_num_frames, :, :]) - else: - result_row.append(tile[:, :, : tile_latent_stride_num_frames + 1, :, :]) - - enc = torch.cat(result_row, dim=2)[:, :, :latent_num_frames] - return enc - - def _temporal_tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - batch_size, num_channels, num_frames, height, width = z.shape - num_sample_frames = (num_frames - 1) * self.temporal_compression_ratio + 1 - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio - tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio - blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames - - row = [] - for i in range(0, num_frames, tile_latent_stride_num_frames): - tile = z[:, :, i : i + tile_latent_min_num_frames + 1, :, :] - if self.use_tiling and (tile.shape[-1] > tile_latent_min_width or tile.shape[-2] > tile_latent_min_height): - decoded = self.tiled_decode(tile, return_dict=True).sample - else: - tile = self.post_quant_conv(tile) - decoded = self.decoder(tile) - if i > 0: - decoded = decoded[:, :, 1:, :, :] - row.append(decoded) - - result_row = [] - for i, tile in enumerate(row): - if i > 0: - tile = self.blend_t(row[i - 1], tile, blend_num_frames) - result_row.append(tile[:, :, : self.tile_sample_stride_num_frames, :, :]) - else: - result_row.append(tile[:, :, : self.tile_sample_stride_num_frames + 1, :, :]) - - dec = torch.cat(result_row, dim=2)[:, :, :num_sample_frames] - - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z, return_dict=return_dict) - return dec diff --git a/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage.py b/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage.py deleted file mode 100644 index c1d975ae6bb7373146443b60244f77a1fd1407a3..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage.py +++ /dev/null @@ -1,697 +0,0 @@ -# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import numpy as np -import torch -import torch.nn as nn -import torch.nn.functional as F -import torch.utils.checkpoint - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class HunyuanImageResnetBlock(nn.Module): - r""" - Residual block with two convolutions and optional channel change. - - Args: - in_channels (int): Number of input channels. - out_channels (int): Number of output channels. - non_linearity (str, optional): Type of non-linearity to use. Default is "silu". - """ - - def __init__(self, in_channels: int, out_channels: int, non_linearity: str = "silu") -> None: - super().__init__() - self.in_channels = in_channels - self.out_channels = out_channels - self.nonlinearity = get_activation(non_linearity) - - # layers - self.norm1 = nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) - self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) - self.norm2 = nn.GroupNorm(num_groups=32, num_channels=out_channels, eps=1e-6, affine=True) - self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1) - if in_channels != out_channels: - self.conv_shortcut = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0) - else: - self.conv_shortcut = None - - def forward(self, x): - # Apply shortcut connection - residual = x - - # First normalization and activation - x = self.norm1(x) - x = self.nonlinearity(x) - - x = self.conv1(x) - x = self.norm2(x) - x = self.nonlinearity(x) - x = self.conv2(x) - - if self.conv_shortcut is not None: - x = self.conv_shortcut(x) - # Add residual connection - return x + residual - - -class HunyuanImageAttentionBlock(nn.Module): - r""" - Self-attention with a single head. - - Args: - in_channels (int): The number of channels in the input tensor. - """ - - def __init__(self, in_channels: int): - super().__init__() - - # layers - self.norm = nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) - self.to_q = nn.Conv2d(in_channels, in_channels, 1) - self.to_k = nn.Conv2d(in_channels, in_channels, 1) - self.to_v = nn.Conv2d(in_channels, in_channels, 1) - self.proj = nn.Conv2d(in_channels, in_channels, 1) - - def forward(self, x): - identity = x - x = self.norm(x) - - # compute query, key, value - query = self.to_q(x) - key = self.to_k(x) - value = self.to_v(x) - - batch_size, channels, height, width = query.shape - query = query.permute(0, 2, 3, 1).reshape(batch_size, height * width, channels).contiguous() - key = key.permute(0, 2, 3, 1).reshape(batch_size, height * width, channels).contiguous() - value = value.permute(0, 2, 3, 1).reshape(batch_size, height * width, channels).contiguous() - - # apply attention - x = F.scaled_dot_product_attention(query, key, value) - - x = x.reshape(batch_size, height, width, channels).permute(0, 3, 1, 2) - # output projection - x = self.proj(x) - - return x + identity - - -class HunyuanImageDownsample(nn.Module): - """ - Downsampling block for spatial reduction. - - Args: - in_channels (int): Number of input channels. - out_channels (int): Number of output channels. - """ - - def __init__(self, in_channels: int, out_channels: int): - super().__init__() - factor = 4 - if out_channels % factor != 0: - raise ValueError(f"out_channels % factor != 0: {out_channels % factor}") - - self.conv = nn.Conv2d(in_channels, out_channels // factor, kernel_size=3, stride=1, padding=1) - self.group_size = factor * in_channels // out_channels - - def forward(self, x: torch.Tensor) -> torch.Tensor: - h = self.conv(x) - - B, C, H, W = h.shape - h = h.reshape(B, C, H // 2, 2, W // 2, 2) - h = h.permute(0, 3, 5, 1, 2, 4) # b, r1, r2, c, h, w - h = h.reshape(B, 4 * C, H // 2, W // 2) - - B, C, H, W = x.shape - shortcut = x.reshape(B, C, H // 2, 2, W // 2, 2) - shortcut = shortcut.permute(0, 3, 5, 1, 2, 4) # b, r1, r2, c, h, w - shortcut = shortcut.reshape(B, 4 * C, H // 2, W // 2) - - B, C, H, W = shortcut.shape - shortcut = shortcut.view(B, h.shape[1], self.group_size, H, W).mean(dim=2) - return h + shortcut - - -class HunyuanImageUpsample(nn.Module): - """ - Upsampling block for spatial expansion. - - Args: - in_channels (int): Number of input channels. - out_channels (int): Number of output channels. - """ - - def __init__(self, in_channels: int, out_channels: int): - super().__init__() - factor = 4 - self.conv = nn.Conv2d(in_channels, out_channels * factor, kernel_size=3, stride=1, padding=1) - self.repeats = factor * out_channels // in_channels - - def forward(self, x: torch.Tensor) -> torch.Tensor: - h = self.conv(x) - - B, C, H, W = h.shape - h = h.reshape(B, 2, 2, C // 4, H, W) # b, r1, r2, c, h, w - h = h.permute(0, 3, 4, 1, 5, 2) # b, c, h, r1, w, r2 - h = h.reshape(B, C // 4, H * 2, W * 2) - - shortcut = x.repeat_interleave(repeats=self.repeats, dim=1) - - B, C, H, W = shortcut.shape - shortcut = shortcut.reshape(B, 2, 2, C // 4, H, W) # b, r1, r2, c, h, w - shortcut = shortcut.permute(0, 3, 4, 1, 5, 2) # b, c, h, r1, w, r2 - shortcut = shortcut.reshape(B, C // 4, H * 2, W * 2) - return h + shortcut - - -class HunyuanImageMidBlock(nn.Module): - """ - Middle block for HunyuanImageVAE encoder and decoder. - - Args: - in_channels (int): Number of input channels. - num_layers (int): Number of layers. - """ - - def __init__(self, in_channels: int, num_layers: int = 1): - super().__init__() - - resnets = [HunyuanImageResnetBlock(in_channels=in_channels, out_channels=in_channels)] - - attentions = [] - for _ in range(num_layers): - attentions.append(HunyuanImageAttentionBlock(in_channels)) - resnets.append(HunyuanImageResnetBlock(in_channels=in_channels, out_channels=in_channels)) - - self.resnets = nn.ModuleList(resnets) - self.attentions = nn.ModuleList(attentions) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = self.resnets[0](x) - - for attn, resnet in zip(self.attentions, self.resnets[1:]): - x = attn(x) - x = resnet(x) - - return x - - -class HunyuanImageEncoder2D(nn.Module): - r""" - Encoder network that compresses input to latent representation. - - Args: - in_channels (int): Number of input channels. - z_channels (int): Number of latent channels. - block_out_channels (list of int): Output channels for each block. - num_res_blocks (int): Number of residual blocks per block. - spatial_compression_ratio (int): Spatial downsampling factor. - non_linearity (str): Type of non-linearity to use. Default is "silu". - downsample_match_channel (bool): Whether to match channels during downsampling. - """ - - def __init__( - self, - in_channels: int, - z_channels: int, - block_out_channels: tuple[int, ...], - num_res_blocks: int, - spatial_compression_ratio: int, - non_linearity: str = "silu", - downsample_match_channel: bool = True, - ): - super().__init__() - if block_out_channels[-1] % (2 * z_channels) != 0: - raise ValueError( - f"block_out_channels[-1 has to be divisible by 2 * out_channels, you have block_out_channels = {block_out_channels[-1]} and out_channels = {z_channels}" - ) - - self.in_channels = in_channels - self.z_channels = z_channels - self.block_out_channels = block_out_channels - self.num_res_blocks = num_res_blocks - self.spatial_compression_ratio = spatial_compression_ratio - - self.group_size = block_out_channels[-1] // (2 * z_channels) - self.nonlinearity = get_activation(non_linearity) - - # init block - self.conv_in = nn.Conv2d(in_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1) - - # downsample blocks - self.down_blocks = nn.ModuleList([]) - - block_in_channel = block_out_channels[0] - for i in range(len(block_out_channels)): - block_out_channel = block_out_channels[i] - # residual blocks - for _ in range(num_res_blocks): - self.down_blocks.append( - HunyuanImageResnetBlock(in_channels=block_in_channel, out_channels=block_out_channel) - ) - block_in_channel = block_out_channel - - # downsample block - if i < np.log2(spatial_compression_ratio) and i != len(block_out_channels) - 1: - if downsample_match_channel: - block_out_channel = block_out_channels[i + 1] - self.down_blocks.append( - HunyuanImageDownsample(in_channels=block_in_channel, out_channels=block_out_channel) - ) - block_in_channel = block_out_channel - - # middle blocks - self.mid_block = HunyuanImageMidBlock(in_channels=block_out_channels[-1], num_layers=1) - - # output blocks - # Output layers - self.norm_out = nn.GroupNorm(num_groups=32, num_channels=block_out_channels[-1], eps=1e-6, affine=True) - self.conv_out = nn.Conv2d(block_out_channels[-1], 2 * z_channels, kernel_size=3, stride=1, padding=1) - - self.gradient_checkpointing = False - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = self.conv_in(x) - - ## downsamples - for down_block in self.down_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - x = self._gradient_checkpointing_func(down_block, x) - else: - x = down_block(x) - - ## middle - if torch.is_grad_enabled() and self.gradient_checkpointing: - x = self._gradient_checkpointing_func(self.mid_block, x) - else: - x = self.mid_block(x) - - ## head - B, C, H, W = x.shape - residual = x.view(B, C // self.group_size, self.group_size, H, W).mean(dim=2) - - x = self.norm_out(x) - x = self.nonlinearity(x) - x = self.conv_out(x) - return x + residual - - -class HunyuanImageDecoder2D(nn.Module): - r""" - Decoder network that reconstructs output from latent representation. - - Args: - z_channels : int - Number of latent channels. - out_channels : int - Number of output channels. - block_out_channels : tuple[int, ...] - Output channels for each block. - num_res_blocks : int - Number of residual blocks per block. - spatial_compression_ratio : int - Spatial upsampling factor. - upsample_match_channel : bool - Whether to match channels during upsampling. - non_linearity (str): Type of non-linearity to use. Default is "silu". - """ - - def __init__( - self, - z_channels: int, - out_channels: int, - block_out_channels: tuple[int, ...], - num_res_blocks: int, - spatial_compression_ratio: int, - upsample_match_channel: bool = True, - non_linearity: str = "silu", - ): - super().__init__() - if block_out_channels[0] % z_channels != 0: - raise ValueError( - f"block_out_channels[0] should be divisible by z_channels but has block_out_channels[0] = {block_out_channels[0]} and z_channels = {z_channels}" - ) - - self.z_channels = z_channels - self.block_out_channels = block_out_channels - self.num_res_blocks = num_res_blocks - self.repeat = block_out_channels[0] // z_channels - self.spatial_compression_ratio = spatial_compression_ratio - self.nonlinearity = get_activation(non_linearity) - - self.conv_in = nn.Conv2d(z_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1) - - # Middle blocks with attention - self.mid_block = HunyuanImageMidBlock(in_channels=block_out_channels[0], num_layers=1) - - # Upsampling blocks - block_in_channel = block_out_channels[0] - self.up_blocks = nn.ModuleList() - for i in range(len(block_out_channels)): - block_out_channel = block_out_channels[i] - for _ in range(self.num_res_blocks + 1): - self.up_blocks.append( - HunyuanImageResnetBlock(in_channels=block_in_channel, out_channels=block_out_channel) - ) - block_in_channel = block_out_channel - - if i < np.log2(spatial_compression_ratio) and i != len(block_out_channels) - 1: - if upsample_match_channel: - block_out_channel = block_out_channels[i + 1] - self.up_blocks.append(HunyuanImageUpsample(block_in_channel, block_out_channel)) - block_in_channel = block_out_channel - - # Output layers - self.norm_out = nn.GroupNorm(num_groups=32, num_channels=block_out_channels[-1], eps=1e-6, affine=True) - self.conv_out = nn.Conv2d(block_out_channels[-1], out_channels, kernel_size=3, stride=1, padding=1) - - self.gradient_checkpointing = False - - def forward(self, x: torch.Tensor) -> torch.Tensor: - h = self.conv_in(x) + x.repeat_interleave(repeats=self.repeat, dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - h = self._gradient_checkpointing_func(self.mid_block, h) - else: - h = self.mid_block(h) - - for up_block in self.up_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - h = self._gradient_checkpointing_func(up_block, h) - else: - h = up_block(h) - h = self.norm_out(h) - h = self.nonlinearity(h) - h = self.conv_out(h) - return h - - -class AutoencoderKLHunyuanImage(ModelMixin, AutoencoderMixin, ConfigMixin, FromOriginalModelMixin): - r""" - A VAE model for 2D images with spatial tiling support. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - """ - - _supports_gradient_checkpointing = False - - # fmt: off - @register_to_config - def __init__( - self, - in_channels: int, - out_channels: int, - latent_channels: int, - block_out_channels: tuple[int, ...], - layers_per_block: int, - spatial_compression_ratio: int, - sample_size: int, - scaling_factor: float = None, - downsample_match_channel: bool = True, - upsample_match_channel: bool = True, - ) -> None: - # fmt: on - super().__init__() - - self.encoder = HunyuanImageEncoder2D( - in_channels=in_channels, - z_channels=latent_channels, - block_out_channels=block_out_channels, - num_res_blocks=layers_per_block, - spatial_compression_ratio=spatial_compression_ratio, - downsample_match_channel=downsample_match_channel, - ) - - self.decoder = HunyuanImageDecoder2D( - z_channels=latent_channels, - out_channels=out_channels, - block_out_channels=list(reversed(block_out_channels)), - num_res_blocks=layers_per_block, - spatial_compression_ratio=spatial_compression_ratio, - upsample_match_channel=upsample_match_channel, - ) - - # Tiling and slicing configuration - self.use_slicing = False - self.use_tiling = False - - # Tiling parameters - self.tile_sample_min_size = sample_size - self.tile_latent_min_size = sample_size // spatial_compression_ratio - self.tile_overlap_factor = 0.25 - - def enable_tiling( - self, - tile_sample_min_size: int | None = None, - tile_overlap_factor: float | None = None, - ) -> None: - r""" - Enable spatial tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles - to compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to - allow processing larger images. - - Args: - tile_sample_min_size (`int`, *optional*): - The minimum size required for a sample to be separated into tiles across the spatial dimension. - tile_overlap_factor (`float`, *optional*): - The overlap factor required for a latent to be separated into tiles across the spatial dimension. - """ - self.use_tiling = True - self.tile_sample_min_size = tile_sample_min_size or self.tile_sample_min_size - self.tile_overlap_factor = tile_overlap_factor or self.tile_overlap_factor - self.tile_latent_min_size = self.tile_sample_min_size // self.config.spatial_compression_ratio - - def _encode(self, x: torch.Tensor): - - batch_size, num_channels, height, width = x.shape - - if self.use_tiling and (width > self.tile_sample_min_size or height > self.tile_sample_min_size): - return self.tiled_encode(x) - - enc = self.encoder(x) - - return enc - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - r""" - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, return_dict: bool = True): - - batch_size, num_channels, height, width = z.shape - - if self.use_tiling and (width > self.tile_latent_min_size or height > self.tile_latent_min_size): - return self.tiled_decode(z, return_dict=return_dict) - - dec = self.decoder(z) - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-2], b.shape[-2], blend_extent) - for y in range(blend_extent): - b[:, :, y, :] = a[:, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-1], b.shape[-1], blend_extent) - for x in range(blend_extent): - b[:, :, :, x] = a[:, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, x] * ( - x / blend_extent - ) - return b - - def tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - """ - Encode input using spatial tiling strategy. - - Args: - x (`torch.Tensor`): Input tensor of shape (B, C, T, H, W). - - Returns: - `torch.Tensor`: - The latent representation of the encoded images. - """ - _, _, _, height, width = x.shape - overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor)) - blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor) - row_limit = self.tile_latent_min_size - blend_extent - - rows = [] - for i in range(0, height, overlap_size): - row = [] - for j in range(0, width, overlap_size): - tile = x[:, :, :, i : i + self.tile_sample_min_size, j : j + self.tile_sample_min_size] - tile = self.encoder(tile) - row.append(tile) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=-1)) - - moments = torch.cat(result_rows, dim=-2) - - return moments - - def tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - """ - Decode latent using spatial tiling strategy. - - Args: - z (`torch.Tensor`): Latent tensor of shape (B, C, H, W). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - _, _, height, width = z.shape - overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor)) - blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor) - row_limit = self.tile_sample_min_size - blend_extent - - rows = [] - for i in range(0, height, overlap_size): - row = [] - for j in range(0, width, overlap_size): - tile = z[:, :, i : i + self.tile_latent_min_size, j : j + self.tile_latent_min_size] - decoded = self.decoder(tile) - row.append(decoded) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=-1)) - - dec = torch.cat(result_rows, dim=-2) - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) - - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - """ - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - posterior = self.encode(sample).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z, return_dict=return_dict) - - return dec diff --git a/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage_refiner.py b/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage_refiner.py deleted file mode 100644 index 5297e3c850bacc890134dbaf0aed8ca8ce91c7c3..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage_refiner.py +++ /dev/null @@ -1,927 +0,0 @@ -# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import numpy as np -import torch -import torch.nn as nn -import torch.nn.functional as F -import torch.utils.checkpoint - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class HunyuanImageRefinerCausalConv3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int] = 3, - stride: int | tuple[int, int, int] = 1, - padding: int | tuple[int, int, int] = 0, - dilation: int | tuple[int, int, int] = 1, - bias: bool = True, - pad_mode: str = "replicate", - ) -> None: - super().__init__() - - kernel_size = (kernel_size, kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size - - self.pad_mode = pad_mode - self.time_causal_padding = ( - kernel_size[0] // 2, - kernel_size[0] // 2, - kernel_size[1] // 2, - kernel_size[1] // 2, - kernel_size[2] - 1, - 0, - ) - - self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, dilation, bias=bias) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = F.pad(hidden_states, self.time_causal_padding, mode=self.pad_mode) - return self.conv(hidden_states) - - -class HunyuanImageRefinerRMS_norm(nn.Module): - r""" - A custom RMS normalization layer. - - Args: - dim (int): The number of dimensions to normalize over. - channel_first (bool, optional): Whether the input tensor has channels as the first dimension. - Default is True. - images (bool, optional): Whether the input represents image data. Default is True. - bias (bool, optional): Whether to include a learnable bias term. Default is False. - """ - - def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bias: bool = False) -> None: - super().__init__() - broadcastable_dims = (1, 1, 1) if not images else (1, 1) - shape = (dim, *broadcastable_dims) if channel_first else (dim,) - - self.channel_first = channel_first - self.scale = dim**0.5 - self.gamma = nn.Parameter(torch.ones(shape)) - self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - - def forward(self, x): - needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( - t in str(x.dtype) for t in ("float4_", "float8_") - ) - normalized = F.normalize(x.float() if needs_fp32_normalize else x, dim=(1 if self.channel_first else -1)).to( - x.dtype - ) - - return normalized * self.scale * self.gamma + self.bias - - -class HunyuanImageRefinerAttnBlock(nn.Module): - def __init__(self, in_channels: int): - super().__init__() - self.in_channels = in_channels - - self.norm = HunyuanImageRefinerRMS_norm(in_channels, images=False) - - self.to_q = nn.Conv3d(in_channels, in_channels, kernel_size=1) - self.to_k = nn.Conv3d(in_channels, in_channels, kernel_size=1) - self.to_v = nn.Conv3d(in_channels, in_channels, kernel_size=1) - self.proj_out = nn.Conv3d(in_channels, in_channels, kernel_size=1) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - identity = x - - x = self.norm(x) - - query = self.to_q(x) - key = self.to_k(x) - value = self.to_v(x) - - batch_size, channels, frames, height, width = query.shape - - query = query.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous() - key = key.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous() - value = value.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous() - - x = nn.functional.scaled_dot_product_attention(query, key, value, attn_mask=None) - - # batch_size, 1, frames * height * width, channels - - x = x.squeeze(1).reshape(batch_size, frames, height, width, channels).permute(0, 4, 1, 2, 3) - x = self.proj_out(x) - - return x + identity - - -class HunyuanImageRefinerUpsampleDCAE(nn.Module): - def __init__(self, in_channels: int, out_channels: int, add_temporal_upsample: bool = True): - super().__init__() - factor = 2 * 2 * 2 if add_temporal_upsample else 1 * 2 * 2 - self.conv = HunyuanImageRefinerCausalConv3d(in_channels, out_channels * factor, kernel_size=3) - - self.add_temporal_upsample = add_temporal_upsample - self.repeats = factor * out_channels // in_channels - - @staticmethod - def _dcae_upsample_rearrange(tensor, r1=1, r2=2, r3=2): - """ - Convert (b, r1*r2*r3*c, f, h, w) -> (b, c, r1*f, r2*h, r3*w) - - Args: - tensor: Input tensor of shape (b, r1*r2*r3*c, f, h, w) - r1: temporal upsampling factor - r2: height upsampling factor - r3: width upsampling factor - """ - b, packed_c, f, h, w = tensor.shape - factor = r1 * r2 * r3 - c = packed_c // factor - - tensor = tensor.view(b, r1, r2, r3, c, f, h, w) - tensor = tensor.permute(0, 4, 5, 1, 6, 2, 7, 3) - return tensor.reshape(b, c, f * r1, h * r2, w * r3) - - def forward(self, x: torch.Tensor): - r1 = 2 if self.add_temporal_upsample else 1 - h = self.conv(x) - if self.add_temporal_upsample: - h = self._dcae_upsample_rearrange(h, r1=1, r2=2, r3=2) - h = h[:, : h.shape[1] // 2] - - # shortcut computation - shortcut = self._dcae_upsample_rearrange(x, r1=1, r2=2, r3=2) - shortcut = shortcut.repeat_interleave(repeats=self.repeats // 2, dim=1) - - else: - h = self._dcae_upsample_rearrange(h, r1=r1, r2=2, r3=2) - shortcut = x.repeat_interleave(repeats=self.repeats, dim=1) - shortcut = self._dcae_upsample_rearrange(shortcut, r1=r1, r2=2, r3=2) - return h + shortcut - - -class HunyuanImageRefinerDownsampleDCAE(nn.Module): - def __init__(self, in_channels: int, out_channels: int, add_temporal_downsample: bool = True): - super().__init__() - factor = 2 * 2 * 2 if add_temporal_downsample else 1 * 2 * 2 - assert out_channels % factor == 0 - # self.conv = Conv3d(in_channels, out_channels // factor, kernel_size=3, stride=1, padding=1) - self.conv = HunyuanImageRefinerCausalConv3d(in_channels, out_channels // factor, kernel_size=3) - - self.add_temporal_downsample = add_temporal_downsample - self.group_size = factor * in_channels // out_channels - - @staticmethod - def _dcae_downsample_rearrange(tensor, r1=1, r2=2, r3=2): - """ - Convert (b, c, r1*f, r2*h, r3*w) -> (b, r1*r2*r3*c, f, h, w) - - This packs spatial/temporal dimensions into channels (opposite of upsample) - """ - b, c, packed_f, packed_h, packed_w = tensor.shape - f, h, w = packed_f // r1, packed_h // r2, packed_w // r3 - - tensor = tensor.view(b, c, f, r1, h, r2, w, r3) - tensor = tensor.permute(0, 3, 5, 7, 1, 2, 4, 6) - return tensor.reshape(b, r1 * r2 * r3 * c, f, h, w) - - def forward(self, x: torch.Tensor): - r1 = 2 if self.add_temporal_downsample else 1 - h = self.conv(x) - if self.add_temporal_downsample: - # h = rearrange(h, "b c f (h r2) (w r3) -> b (r2 r3 c) f h w", r2=2, r3=2) - h = self._dcae_downsample_rearrange(h, r1=1, r2=2, r3=2) - h = torch.cat([h, h], dim=1) - # shortcut computation - # shortcut = rearrange(x, "b c f (h r2) (w r3) -> b (r2 r3 c) f h w", r2=2, r3=2) - shortcut = self._dcae_downsample_rearrange(x, r1=1, r2=2, r3=2) - B, C, T, H, W = shortcut.shape - shortcut = shortcut.view(B, h.shape[1], self.group_size // 2, T, H, W).mean(dim=2) - else: - # h = rearrange(h, "b c (f r1) (h r2) (w r3) -> b (r1 r2 r3 c) f h w", r1=r1, r2=2, r3=2) - h = self._dcae_downsample_rearrange(h, r1=r1, r2=2, r3=2) - # shortcut = rearrange(x, "b c (f r1) (h r2) (w r3) -> b (r1 r2 r3 c) f h w", r1=r1, r2=2, r3=2) - shortcut = self._dcae_downsample_rearrange(x, r1=r1, r2=2, r3=2) - B, C, T, H, W = shortcut.shape - shortcut = shortcut.view(B, h.shape[1], self.group_size, T, H, W).mean(dim=2) - - return h + shortcut - - -class HunyuanImageRefinerResnetBlock(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - non_linearity: str = "swish", - ) -> None: - super().__init__() - out_channels = out_channels or in_channels - - self.nonlinearity = get_activation(non_linearity) - - self.norm1 = HunyuanImageRefinerRMS_norm(in_channels, images=False) - self.conv1 = HunyuanImageRefinerCausalConv3d(in_channels, out_channels, kernel_size=3) - - self.norm2 = HunyuanImageRefinerRMS_norm(out_channels, images=False) - self.conv2 = HunyuanImageRefinerCausalConv3d(out_channels, out_channels, kernel_size=3) - - self.conv_shortcut = None - if in_channels != out_channels: - self.conv_shortcut = nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - residual = hidden_states - - hidden_states = self.norm1(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv1(hidden_states) - - hidden_states = self.norm2(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv2(hidden_states) - - if self.conv_shortcut is not None: - residual = self.conv_shortcut(residual) - - return hidden_states + residual - - -class HunyuanImageRefinerMidBlock(nn.Module): - def __init__( - self, - in_channels: int, - num_layers: int = 1, - add_attention: bool = True, - ) -> None: - super().__init__() - self.add_attention = add_attention - - # There is always at least one resnet - resnets = [ - HunyuanImageRefinerResnetBlock( - in_channels=in_channels, - out_channels=in_channels, - ) - ] - attentions = [] - - for _ in range(num_layers): - if self.add_attention: - attentions.append(HunyuanImageRefinerAttnBlock(in_channels)) - else: - attentions.append(None) - - resnets.append( - HunyuanImageRefinerResnetBlock( - in_channels=in_channels, - out_channels=in_channels, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.resnets[0](hidden_states) - - for attn, resnet in zip(self.attentions, self.resnets[1:]): - if attn is not None: - hidden_states = attn(hidden_states) - hidden_states = resnet(hidden_states) - - return hidden_states - - -class HunyuanImageRefinerDownBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int = 1, - downsample_out_channels: int | None = None, - add_temporal_downsample: int = True, - ) -> None: - super().__init__() - resnets = [] - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - HunyuanImageRefinerResnetBlock( - in_channels=in_channels, - out_channels=out_channels, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if downsample_out_channels is not None: - self.downsamplers = nn.ModuleList( - [ - HunyuanImageRefinerDownsampleDCAE( - out_channels, - out_channels=downsample_out_channels, - add_temporal_downsample=add_temporal_downsample, - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - for resnet in self.resnets: - hidden_states = resnet(hidden_states) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - return hidden_states - - -class HunyuanImageRefinerUpBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int = 1, - upsample_out_channels: int | None = None, - add_temporal_upsample: bool = True, - ) -> None: - super().__init__() - resnets = [] - - for i in range(num_layers): - input_channels = in_channels if i == 0 else out_channels - - resnets.append( - HunyuanImageRefinerResnetBlock( - in_channels=input_channels, - out_channels=out_channels, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if upsample_out_channels is not None: - self.upsamplers = nn.ModuleList( - [ - HunyuanImageRefinerUpsampleDCAE( - out_channels, - out_channels=upsample_out_channels, - add_temporal_upsample=add_temporal_upsample, - ) - ] - ) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if torch.is_grad_enabled() and self.gradient_checkpointing: - for resnet in self.resnets: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states) - - else: - for resnet in self.resnets: - hidden_states = resnet(hidden_states) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - return hidden_states - - -class HunyuanImageRefinerEncoder3D(nn.Module): - r""" - 3D vae encoder for HunyuanImageRefiner. - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 64, - block_out_channels: tuple[int, ...] = (128, 256, 512, 1024, 1024), - layers_per_block: int = 2, - temporal_compression_ratio: int = 4, - spatial_compression_ratio: int = 16, - downsample_match_channel: bool = True, - ) -> None: - super().__init__() - - self.in_channels = in_channels - self.out_channels = out_channels - self.group_size = block_out_channels[-1] // self.out_channels - - self.conv_in = HunyuanImageRefinerCausalConv3d(in_channels, block_out_channels[0], kernel_size=3) - self.mid_block = None - self.down_blocks = nn.ModuleList([]) - - input_channel = block_out_channels[0] - for i in range(len(block_out_channels)): - add_spatial_downsample = i < np.log2(spatial_compression_ratio) - output_channel = block_out_channels[i] - if not add_spatial_downsample: - down_block = HunyuanImageRefinerDownBlock3D( - num_layers=layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - downsample_out_channels=None, - add_temporal_downsample=False, - ) - input_channel = output_channel - else: - add_temporal_downsample = i >= np.log2(spatial_compression_ratio // temporal_compression_ratio) - downsample_out_channels = block_out_channels[i + 1] if downsample_match_channel else output_channel - down_block = HunyuanImageRefinerDownBlock3D( - num_layers=layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - downsample_out_channels=downsample_out_channels, - add_temporal_downsample=add_temporal_downsample, - ) - input_channel = downsample_out_channels - - self.down_blocks.append(down_block) - - self.mid_block = HunyuanImageRefinerMidBlock(in_channels=block_out_channels[-1]) - - self.norm_out = HunyuanImageRefinerRMS_norm(block_out_channels[-1], images=False) - self.conv_act = nn.SiLU() - self.conv_out = HunyuanImageRefinerCausalConv3d(block_out_channels[-1], out_channels, kernel_size=3) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv_in(hidden_states) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for down_block in self.down_blocks: - hidden_states = self._gradient_checkpointing_func(down_block, hidden_states) - - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states) - else: - for down_block in self.down_blocks: - hidden_states = down_block(hidden_states) - - hidden_states = self.mid_block(hidden_states) - - # short_cut = rearrange(hidden_states, "b (c r) f h w -> b c r f h w", r=self.group_size).mean(dim=2) - batch_size, _, frame, height, width = hidden_states.shape - short_cut = hidden_states.view(batch_size, -1, self.group_size, frame, height, width).mean(dim=2) - - hidden_states = self.norm_out(hidden_states) - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states) - - hidden_states += short_cut - - return hidden_states - - -class HunyuanImageRefinerDecoder3D(nn.Module): - r""" - Causal decoder for 3D video-like data used for HunyuanImage-2.1 Refiner. - """ - - def __init__( - self, - in_channels: int = 32, - out_channels: int = 3, - block_out_channels: tuple[int, ...] = (1024, 1024, 512, 256, 128), - layers_per_block: int = 2, - spatial_compression_ratio: int = 16, - temporal_compression_ratio: int = 4, - upsample_match_channel: bool = True, - ): - super().__init__() - self.layers_per_block = layers_per_block - self.in_channels = in_channels - self.out_channels = out_channels - self.repeat = block_out_channels[0] // self.in_channels - - self.conv_in = HunyuanImageRefinerCausalConv3d(self.in_channels, block_out_channels[0], kernel_size=3) - self.up_blocks = nn.ModuleList([]) - - # mid - self.mid_block = HunyuanImageRefinerMidBlock(in_channels=block_out_channels[0]) - - # up - input_channel = block_out_channels[0] - for i in range(len(block_out_channels)): - output_channel = block_out_channels[i] - - add_spatial_upsample = i < np.log2(spatial_compression_ratio) - add_temporal_upsample = i < np.log2(temporal_compression_ratio) - if add_spatial_upsample or add_temporal_upsample: - upsample_out_channels = block_out_channels[i + 1] if upsample_match_channel else output_channel - up_block = HunyuanImageRefinerUpBlock3D( - num_layers=self.layers_per_block + 1, - in_channels=input_channel, - out_channels=output_channel, - upsample_out_channels=upsample_out_channels, - add_temporal_upsample=add_temporal_upsample, - ) - input_channel = upsample_out_channels - else: - up_block = HunyuanImageRefinerUpBlock3D( - num_layers=self.layers_per_block + 1, - in_channels=input_channel, - out_channels=output_channel, - upsample_out_channels=None, - add_temporal_upsample=False, - ) - input_channel = output_channel - - self.up_blocks.append(up_block) - - # out - self.norm_out = HunyuanImageRefinerRMS_norm(block_out_channels[-1], images=False) - self.conv_act = nn.SiLU() - self.conv_out = HunyuanImageRefinerCausalConv3d(block_out_channels[-1], out_channels, kernel_size=3) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv_in(hidden_states) + hidden_states.repeat_interleave(repeats=self.repeat, dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states) - - for up_block in self.up_blocks: - hidden_states = self._gradient_checkpointing_func(up_block, hidden_states) - else: - hidden_states = self.mid_block(hidden_states) - - for up_block in self.up_blocks: - hidden_states = up_block(hidden_states) - - # post-process - hidden_states = self.norm_out(hidden_states) - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states) - return hidden_states - - -class AutoencoderKLHunyuanImageRefiner(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. Used for - HunyuanImage-2.1 Refiner. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - latent_channels: int = 32, - block_out_channels: tuple[int, ...] = (128, 256, 512, 1024, 1024), - layers_per_block: int = 2, - spatial_compression_ratio: int = 16, - temporal_compression_ratio: int = 4, - downsample_match_channel: bool = True, - upsample_match_channel: bool = True, - scaling_factor: float = 1.03682, - ) -> None: - super().__init__() - - self.encoder = HunyuanImageRefinerEncoder3D( - in_channels=in_channels, - out_channels=latent_channels * 2, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - temporal_compression_ratio=temporal_compression_ratio, - spatial_compression_ratio=spatial_compression_ratio, - downsample_match_channel=downsample_match_channel, - ) - - self.decoder = HunyuanImageRefinerDecoder3D( - in_channels=latent_channels, - out_channels=out_channels, - block_out_channels=list(reversed(block_out_channels)), - layers_per_block=layers_per_block, - temporal_compression_ratio=temporal_compression_ratio, - spatial_compression_ratio=spatial_compression_ratio, - upsample_match_channel=upsample_match_channel, - ) - - self.spatial_compression_ratio = spatial_compression_ratio - self.temporal_compression_ratio = temporal_compression_ratio - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent - # frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the - # intermediate tiles together, the memory requirement can be lowered. - self.use_tiling = False - - # The minimal tile height and width for spatial tiling to be used - self.tile_sample_min_height = 256 - self.tile_sample_min_width = 256 - - # The minimal distance between two spatial tiles - self.tile_sample_stride_height = 192 - self.tile_sample_stride_width = 192 - - self.tile_overlap_factor = 0.25 - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_sample_stride_height: float | None = None, - tile_sample_stride_width: float | None = None, - tile_overlap_factor: float | None = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_sample_stride_height (`int`, *optional*): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. - tile_sample_stride_width (`int`, *optional*): - The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling - artifacts produced across the width dimension. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height - self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width - self.tile_overlap_factor = tile_overlap_factor or self.tile_overlap_factor - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - _, _, _, height, width = x.shape - - if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): - return self.tiled_encode(x) - - x = self.encoder(x) - return x - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - r""" - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor) -> torch.Tensor: - _, _, _, height, width = z.shape - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - - if self.use_tiling and (width > tile_latent_min_width or height > tile_latent_min_height): - return self.tiled_decode(z) - - dec = self.decoder(z) - - return dec - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice) for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z) - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-2], b.shape[-2], blend_extent) - for y in range(blend_extent): - b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-1], b.shape[-1], blend_extent) - for x in range(blend_extent): - b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( - x / blend_extent - ) - return b - - def blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-3], b.shape[-3], blend_extent) - for x in range(blend_extent): - b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (1 - x / blend_extent) + b[:, :, x, :, :] * ( - x / blend_extent - ) - return b - - def tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - r"""Encode a batch of images using a tiled encoder. - - Args: - x (`torch.Tensor`): Input batch of videos. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - _, _, _, height, width = x.shape - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - overlap_height = int(tile_latent_min_height * (1 - self.tile_overlap_factor)) # 256 * (1 - 0.25) = 192 - overlap_width = int(tile_latent_min_width * (1 - self.tile_overlap_factor)) # 256 * (1 - 0.25) = 192 - blend_height = int(tile_latent_min_height * self.tile_overlap_factor) # 8 * 0.25 = 2 - blend_width = int(tile_latent_min_width * self.tile_overlap_factor) # 8 * 0.25 = 2 - row_limit_height = tile_latent_min_height - blend_height # 8 - 2 = 6 - row_limit_width = tile_latent_min_width - blend_width # 8 - 2 = 6 - - rows = [] - for i in range(0, height, overlap_height): - row = [] - for j in range(0, width, overlap_width): - tile = x[ - :, - :, - :, - i : i + self.tile_sample_min_height, - j : j + self.tile_sample_min_width, - ] - tile = self.encoder(tile) - row.append(tile) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :row_limit_height, :row_limit_width]) - result_rows.append(torch.cat(result_row, dim=-1)) - moments = torch.cat(result_rows, dim=-2) - - return moments - - def tiled_decode(self, z: torch.Tensor) -> torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - - _, _, _, height, width = z.shape - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - overlap_height = int(tile_latent_min_height * (1 - self.tile_overlap_factor)) # 8 * (1 - 0.25) = 6 - overlap_width = int(tile_latent_min_width * (1 - self.tile_overlap_factor)) # 8 * (1 - 0.25) = 6 - blend_height = int(tile_latent_min_height * self.tile_overlap_factor) # 256 * 0.25 = 64 - blend_width = int(tile_latent_min_width * self.tile_overlap_factor) # 256 * 0.25 = 64 - row_limit_height = tile_latent_min_height - blend_height # 256 - 64 = 192 - row_limit_width = tile_latent_min_width - blend_width # 256 - 64 = 192 - - rows = [] - for i in range(0, height, overlap_height): - row = [] - for j in range(0, width, overlap_width): - tile = z[ - :, - :, - :, - i : i + tile_latent_min_height, - j : j + tile_latent_min_width, - ] - decoded = self.decoder(tile) - row.append(decoded) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :row_limit_height, :row_limit_width]) - result_rows.append(torch.cat(result_row, dim=-1)) - dec = torch.cat(result_rows, dim=-2) - - return dec - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z, return_dict=return_dict) - return dec diff --git a/diffusers/models/autoencoders/autoencoder_kl_hunyuanvideo15.py b/diffusers/models/autoencoders/autoencoder_kl_hunyuanvideo15.py deleted file mode 100644 index dec20aacb7d513b37e34a078e052b7621347f877..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_hunyuanvideo15.py +++ /dev/null @@ -1,960 +0,0 @@ -# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import numpy as np -import torch -import torch.nn as nn -import torch.nn.functional as F -import torch.utils.checkpoint - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class HunyuanVideo15CausalConv3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int] = 3, - stride: int | tuple[int, int, int] = 1, - padding: int | tuple[int, int, int] = 0, - dilation: int | tuple[int, int, int] = 1, - bias: bool = True, - pad_mode: str = "replicate", - ) -> None: - super().__init__() - - kernel_size = (kernel_size, kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size - - self.pad_mode = pad_mode - self.time_causal_padding = ( - kernel_size[0] // 2, - kernel_size[0] // 2, - kernel_size[1] // 2, - kernel_size[1] // 2, - kernel_size[2] - 1, - 0, - ) - - self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, dilation, bias=bias) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = F.pad(hidden_states, self.time_causal_padding, mode=self.pad_mode) - return self.conv(hidden_states) - - -class HunyuanVideo15RMS_norm(nn.Module): - r""" - A custom RMS normalization layer. - - Args: - dim (int): The number of dimensions to normalize over. - channel_first (bool, optional): Whether the input tensor has channels as the first dimension. - Default is True. - images (bool, optional): Whether the input represents image data. Default is True. - bias (bool, optional): Whether to include a learnable bias term. Default is False. - """ - - def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bias: bool = False) -> None: - super().__init__() - broadcastable_dims = (1, 1, 1) if not images else (1, 1) - shape = (dim, *broadcastable_dims) if channel_first else (dim,) - - self.channel_first = channel_first - self.scale = dim**0.5 - self.gamma = nn.Parameter(torch.ones(shape)) - self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - - def forward(self, x): - needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( - t in str(x.dtype) for t in ("float4_", "float8_") - ) - normalized = F.normalize(x.float() if needs_fp32_normalize else x, dim=(1 if self.channel_first else -1)).to( - x.dtype - ) - - return normalized * self.scale * self.gamma + self.bias - - -class HunyuanVideo15AttnBlock(nn.Module): - def __init__(self, in_channels: int): - super().__init__() - self.in_channels = in_channels - - self.norm = HunyuanVideo15RMS_norm(in_channels, images=False) - - self.to_q = nn.Conv3d(in_channels, in_channels, kernel_size=1) - self.to_k = nn.Conv3d(in_channels, in_channels, kernel_size=1) - self.to_v = nn.Conv3d(in_channels, in_channels, kernel_size=1) - self.proj_out = nn.Conv3d(in_channels, in_channels, kernel_size=1) - - @staticmethod - def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None): - """Prepare a causal attention mask for 3D videos. - - Args: - n_frame (int): Number of frames (temporal length). - n_hw (int): Product of height and width. - dtype: Desired mask dtype. - device: Device for the mask. - batch_size (int, optional): If set, expands for batch. - - Returns: - torch.Tensor: Causal attention mask. - """ - seq_len = n_frame * n_hw - mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device) - for i in range(seq_len): - i_frame = i // n_hw - mask[i, : (i_frame + 1) * n_hw] = 0 - if batch_size is not None: - mask = mask.unsqueeze(0).expand(batch_size, -1, -1) - return mask - - def forward(self, x: torch.Tensor) -> torch.Tensor: - identity = x - - x = self.norm(x) - - query = self.to_q(x) - key = self.to_k(x) - value = self.to_v(x) - - batch_size, channels, frames, height, width = query.shape - - query = query.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous() - key = key.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous() - value = value.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous() - - attention_mask = self.prepare_causal_attention_mask( - frames, height * width, query.dtype, query.device, batch_size=batch_size - ) - - x = nn.functional.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask) - - # batch_size, 1, frames * height * width, channels - - x = x.squeeze(1).reshape(batch_size, frames, height, width, channels).permute(0, 4, 1, 2, 3) - x = self.proj_out(x) - - return x + identity - - -class HunyuanVideo15Upsample(nn.Module): - def __init__(self, in_channels: int, out_channels: int, add_temporal_upsample: bool = True): - super().__init__() - factor = 2 * 2 * 2 if add_temporal_upsample else 1 * 2 * 2 - self.conv = HunyuanVideo15CausalConv3d(in_channels, out_channels * factor, kernel_size=3) - - self.add_temporal_upsample = add_temporal_upsample - self.repeats = factor * out_channels // in_channels - - @staticmethod - def _dcae_upsample_rearrange(tensor, r1=1, r2=2, r3=2): - """ - Convert (b, r1*r2*r3*c, f, h, w) -> (b, c, r1*f, r2*h, r3*w) - - Args: - tensor: Input tensor of shape (b, r1*r2*r3*c, f, h, w) - r1: temporal upsampling factor - r2: height upsampling factor - r3: width upsampling factor - """ - b, packed_c, f, h, w = tensor.shape - factor = r1 * r2 * r3 - c = packed_c // factor - - tensor = tensor.view(b, r1, r2, r3, c, f, h, w) - tensor = tensor.permute(0, 4, 5, 1, 6, 2, 7, 3) - return tensor.reshape(b, c, f * r1, h * r2, w * r3) - - def forward(self, x: torch.Tensor): - r1 = 2 if self.add_temporal_upsample else 1 - h = self.conv(x) - if self.add_temporal_upsample: - h_first = h[:, :, :1, :, :] - h_first = self._dcae_upsample_rearrange(h_first, r1=1, r2=2, r3=2) - h_first = h_first[:, : h_first.shape[1] // 2] - h_next = h[:, :, 1:, :, :] - h_next = self._dcae_upsample_rearrange(h_next, r1=r1, r2=2, r3=2) - h = torch.cat([h_first, h_next], dim=2) - - # shortcut computation - x_first = x[:, :, :1, :, :] - x_first = self._dcae_upsample_rearrange(x_first, r1=1, r2=2, r3=2) - x_first = x_first.repeat_interleave(repeats=self.repeats // 2, dim=1) - - x_next = x[:, :, 1:, :, :] - x_next = self._dcae_upsample_rearrange(x_next, r1=r1, r2=2, r3=2) - x_next = x_next.repeat_interleave(repeats=self.repeats, dim=1) - shortcut = torch.cat([x_first, x_next], dim=2) - - else: - h = self._dcae_upsample_rearrange(h, r1=r1, r2=2, r3=2) - shortcut = x.repeat_interleave(repeats=self.repeats, dim=1) - shortcut = self._dcae_upsample_rearrange(shortcut, r1=r1, r2=2, r3=2) - return h + shortcut - - -class HunyuanVideo15Downsample(nn.Module): - def __init__(self, in_channels: int, out_channels: int, add_temporal_downsample: bool = True): - super().__init__() - factor = 2 * 2 * 2 if add_temporal_downsample else 1 * 2 * 2 - self.conv = HunyuanVideo15CausalConv3d(in_channels, out_channels // factor, kernel_size=3) - - self.add_temporal_downsample = add_temporal_downsample - self.group_size = factor * in_channels // out_channels - - @staticmethod - def _dcae_downsample_rearrange(tensor, r1=1, r2=2, r3=2): - """ - Convert (b, c, r1*f, r2*h, r3*w) -> (b, r1*r2*r3*c, f, h, w) - - This packs spatial/temporal dimensions into channels (opposite of upsample) - """ - b, c, packed_f, packed_h, packed_w = tensor.shape - f, h, w = packed_f // r1, packed_h // r2, packed_w // r3 - - tensor = tensor.view(b, c, f, r1, h, r2, w, r3) - tensor = tensor.permute(0, 3, 5, 7, 1, 2, 4, 6) - return tensor.reshape(b, r1 * r2 * r3 * c, f, h, w) - - def forward(self, x: torch.Tensor): - r1 = 2 if self.add_temporal_downsample else 1 - h = self.conv(x) - if self.add_temporal_downsample: - h_first = h[:, :, :1, :, :] - h_first = self._dcae_downsample_rearrange(h_first, r1=1, r2=2, r3=2) - h_first = torch.cat([h_first, h_first], dim=1) - h_next = h[:, :, 1:, :, :] - h_next = self._dcae_downsample_rearrange(h_next, r1=r1, r2=2, r3=2) - h = torch.cat([h_first, h_next], dim=2) - - # shortcut computation - x_first = x[:, :, :1, :, :] - x_first = self._dcae_downsample_rearrange(x_first, r1=1, r2=2, r3=2) - B, C, T, H, W = x_first.shape - x_first = x_first.view(B, h.shape[1], self.group_size // 2, T, H, W).mean(dim=2) - x_next = x[:, :, 1:, :, :] - x_next = self._dcae_downsample_rearrange(x_next, r1=r1, r2=2, r3=2) - B, C, T, H, W = x_next.shape - x_next = x_next.view(B, h.shape[1], self.group_size, T, H, W).mean(dim=2) - shortcut = torch.cat([x_first, x_next], dim=2) - else: - h = self._dcae_downsample_rearrange(h, r1=r1, r2=2, r3=2) - shortcut = self._dcae_downsample_rearrange(x, r1=r1, r2=2, r3=2) - B, C, T, H, W = shortcut.shape - shortcut = shortcut.view(B, h.shape[1], self.group_size, T, H, W).mean(dim=2) - - return h + shortcut - - -class HunyuanVideo15ResnetBlock(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - non_linearity: str = "swish", - ) -> None: - super().__init__() - out_channels = out_channels or in_channels - - self.nonlinearity = get_activation(non_linearity) - - self.norm1 = HunyuanVideo15RMS_norm(in_channels, images=False) - self.conv1 = HunyuanVideo15CausalConv3d(in_channels, out_channels, kernel_size=3) - - self.norm2 = HunyuanVideo15RMS_norm(out_channels, images=False) - self.conv2 = HunyuanVideo15CausalConv3d(out_channels, out_channels, kernel_size=3) - - self.conv_shortcut = None - if in_channels != out_channels: - self.conv_shortcut = nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - residual = hidden_states - - hidden_states = self.norm1(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv1(hidden_states) - - hidden_states = self.norm2(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv2(hidden_states) - - if self.conv_shortcut is not None: - residual = self.conv_shortcut(residual) - - return hidden_states + residual - - -class HunyuanVideo15MidBlock(nn.Module): - def __init__( - self, - in_channels: int, - num_layers: int = 1, - add_attention: bool = True, - ) -> None: - super().__init__() - self.add_attention = add_attention - - # There is always at least one resnet - resnets = [ - HunyuanVideo15ResnetBlock( - in_channels=in_channels, - out_channels=in_channels, - ) - ] - attentions = [] - - for _ in range(num_layers): - if self.add_attention: - attentions.append(HunyuanVideo15AttnBlock(in_channels)) - else: - attentions.append(None) - - resnets.append( - HunyuanVideo15ResnetBlock( - in_channels=in_channels, - out_channels=in_channels, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.resnets[0](hidden_states) - - for attn, resnet in zip(self.attentions, self.resnets[1:]): - if attn is not None: - hidden_states = attn(hidden_states) - hidden_states = resnet(hidden_states) - - return hidden_states - - -class HunyuanVideo15DownBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int = 1, - downsample_out_channels: int | None = None, - add_temporal_downsample: int = True, - ) -> None: - super().__init__() - resnets = [] - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - HunyuanVideo15ResnetBlock( - in_channels=in_channels, - out_channels=out_channels, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if downsample_out_channels is not None: - self.downsamplers = nn.ModuleList( - [ - HunyuanVideo15Downsample( - out_channels, - out_channels=downsample_out_channels, - add_temporal_downsample=add_temporal_downsample, - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - for resnet in self.resnets: - hidden_states = resnet(hidden_states) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - return hidden_states - - -class HunyuanVideo15UpBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int = 1, - upsample_out_channels: int | None = None, - add_temporal_upsample: bool = True, - ) -> None: - super().__init__() - resnets = [] - - for i in range(num_layers): - input_channels = in_channels if i == 0 else out_channels - - resnets.append( - HunyuanVideo15ResnetBlock( - in_channels=input_channels, - out_channels=out_channels, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if upsample_out_channels is not None: - self.upsamplers = nn.ModuleList( - [ - HunyuanVideo15Upsample( - out_channels, - out_channels=upsample_out_channels, - add_temporal_upsample=add_temporal_upsample, - ) - ] - ) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if torch.is_grad_enabled() and self.gradient_checkpointing: - for resnet in self.resnets: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states) - - else: - for resnet in self.resnets: - hidden_states = resnet(hidden_states) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - return hidden_states - - -class HunyuanVideo15Encoder3D(nn.Module): - r""" - 3D vae encoder for HunyuanImageRefiner. - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 64, - block_out_channels: tuple[int, ...] = (128, 256, 512, 1024, 1024), - layers_per_block: int = 2, - temporal_compression_ratio: int = 4, - spatial_compression_ratio: int = 16, - downsample_match_channel: bool = True, - ) -> None: - super().__init__() - - self.in_channels = in_channels - self.out_channels = out_channels - self.group_size = block_out_channels[-1] // self.out_channels - - self.conv_in = HunyuanVideo15CausalConv3d(in_channels, block_out_channels[0], kernel_size=3) - self.mid_block = None - self.down_blocks = nn.ModuleList([]) - - input_channel = block_out_channels[0] - for i in range(len(block_out_channels)): - add_spatial_downsample = i < np.log2(spatial_compression_ratio) - output_channel = block_out_channels[i] - if not add_spatial_downsample: - down_block = HunyuanVideo15DownBlock3D( - num_layers=layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - downsample_out_channels=None, - add_temporal_downsample=False, - ) - input_channel = output_channel - else: - add_temporal_downsample = i >= np.log2(spatial_compression_ratio // temporal_compression_ratio) - downsample_out_channels = block_out_channels[i + 1] if downsample_match_channel else output_channel - down_block = HunyuanVideo15DownBlock3D( - num_layers=layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - downsample_out_channels=downsample_out_channels, - add_temporal_downsample=add_temporal_downsample, - ) - input_channel = downsample_out_channels - - self.down_blocks.append(down_block) - - self.mid_block = HunyuanVideo15MidBlock(in_channels=block_out_channels[-1]) - - self.norm_out = HunyuanVideo15RMS_norm(block_out_channels[-1], images=False) - self.conv_act = nn.SiLU() - self.conv_out = HunyuanVideo15CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv_in(hidden_states) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for down_block in self.down_blocks: - hidden_states = self._gradient_checkpointing_func(down_block, hidden_states) - - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states) - else: - for down_block in self.down_blocks: - hidden_states = down_block(hidden_states) - - hidden_states = self.mid_block(hidden_states) - - batch_size, _, frame, height, width = hidden_states.shape - short_cut = hidden_states.view(batch_size, -1, self.group_size, frame, height, width).mean(dim=2) - - hidden_states = self.norm_out(hidden_states) - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states) - - hidden_states += short_cut - - return hidden_states - - -class HunyuanVideo15Decoder3D(nn.Module): - r""" - Causal decoder for 3D video-like data used for HunyuanImage-1.5 Refiner. - """ - - def __init__( - self, - in_channels: int = 32, - out_channels: int = 3, - block_out_channels: tuple[int, ...] = (1024, 1024, 512, 256, 128), - layers_per_block: int = 2, - spatial_compression_ratio: int = 16, - temporal_compression_ratio: int = 4, - upsample_match_channel: bool = True, - ): - super().__init__() - self.layers_per_block = layers_per_block - self.in_channels = in_channels - self.out_channels = out_channels - self.repeat = block_out_channels[0] // self.in_channels - - self.conv_in = HunyuanVideo15CausalConv3d(self.in_channels, block_out_channels[0], kernel_size=3) - self.up_blocks = nn.ModuleList([]) - - # mid - self.mid_block = HunyuanVideo15MidBlock(in_channels=block_out_channels[0]) - - # up - input_channel = block_out_channels[0] - for i in range(len(block_out_channels)): - output_channel = block_out_channels[i] - - add_spatial_upsample = i < np.log2(spatial_compression_ratio) - add_temporal_upsample = i < np.log2(temporal_compression_ratio) - if add_spatial_upsample or add_temporal_upsample: - upsample_out_channels = block_out_channels[i + 1] if upsample_match_channel else output_channel - up_block = HunyuanVideo15UpBlock3D( - num_layers=self.layers_per_block + 1, - in_channels=input_channel, - out_channels=output_channel, - upsample_out_channels=upsample_out_channels, - add_temporal_upsample=add_temporal_upsample, - ) - input_channel = upsample_out_channels - else: - up_block = HunyuanVideo15UpBlock3D( - num_layers=self.layers_per_block + 1, - in_channels=input_channel, - out_channels=output_channel, - upsample_out_channels=None, - add_temporal_upsample=False, - ) - input_channel = output_channel - - self.up_blocks.append(up_block) - - # out - self.norm_out = HunyuanVideo15RMS_norm(block_out_channels[-1], images=False) - self.conv_act = nn.SiLU() - self.conv_out = HunyuanVideo15CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv_in(hidden_states) + hidden_states.repeat_interleave(repeats=self.repeat, dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states) - - for up_block in self.up_blocks: - hidden_states = self._gradient_checkpointing_func(up_block, hidden_states) - else: - hidden_states = self.mid_block(hidden_states) - - for up_block in self.up_blocks: - hidden_states = up_block(hidden_states) - - # post-process - hidden_states = self.norm_out(hidden_states) - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states) - return hidden_states - - -class AutoencoderKLHunyuanVideo15(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. Used for - HunyuanVideo-1.5. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - latent_channels: int = 32, - block_out_channels: tuple[int] = (128, 256, 512, 1024, 1024), - layers_per_block: int = 2, - spatial_compression_ratio: int = 16, - temporal_compression_ratio: int = 4, - downsample_match_channel: bool = True, - upsample_match_channel: bool = True, - scaling_factor: float = 1.03682, - ) -> None: - super().__init__() - - self.encoder = HunyuanVideo15Encoder3D( - in_channels=in_channels, - out_channels=latent_channels * 2, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - temporal_compression_ratio=temporal_compression_ratio, - spatial_compression_ratio=spatial_compression_ratio, - downsample_match_channel=downsample_match_channel, - ) - - self.decoder = HunyuanVideo15Decoder3D( - in_channels=latent_channels, - out_channels=out_channels, - block_out_channels=list(reversed(block_out_channels)), - layers_per_block=layers_per_block, - temporal_compression_ratio=temporal_compression_ratio, - spatial_compression_ratio=spatial_compression_ratio, - upsample_match_channel=upsample_match_channel, - ) - - self.spatial_compression_ratio = spatial_compression_ratio - self.temporal_compression_ratio = temporal_compression_ratio - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent - # frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the - # intermediate tiles together, the memory requirement can be lowered. - self.use_tiling = False - - # The minimal tile height and width for spatial tiling to be used - self.tile_sample_min_height = 256 - self.tile_sample_min_width = 256 - - # The minimal tile height and width in latent space - self.tile_latent_min_height = self.tile_sample_min_height // spatial_compression_ratio - self.tile_latent_min_width = self.tile_sample_min_width // spatial_compression_ratio - self.tile_overlap_factor = 0.25 - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_latent_min_height: int | None = None, - tile_latent_min_width: int | None = None, - tile_overlap_factor: float | None = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_latent_min_height (`int`, *optional*): - The minimum height required for a latent to be separated into tiles across the height dimension. - tile_latent_min_width (`int`, *optional*): - The minimum width required for a latent to be separated into tiles across the width dimension. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_latent_min_height = tile_latent_min_height or self.tile_latent_min_height - self.tile_latent_min_width = tile_latent_min_width or self.tile_latent_min_width - self.tile_overlap_factor = tile_overlap_factor or self.tile_overlap_factor - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - _, _, _, height, width = x.shape - - if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): - return self.tiled_encode(x) - - x = self.encoder(x) - return x - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - r""" - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor) -> torch.Tensor: - _, _, _, height, width = z.shape - - if self.use_tiling and (width > self.tile_latent_min_width or height > self.tile_latent_min_height): - return self.tiled_decode(z) - - dec = self.decoder(z) - - return dec - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice) for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z) - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-2], b.shape[-2], blend_extent) - for y in range(blend_extent): - b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-1], b.shape[-1], blend_extent) - for x in range(blend_extent): - b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( - x / blend_extent - ) - return b - - def blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-3], b.shape[-3], blend_extent) - for x in range(blend_extent): - b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (1 - x / blend_extent) + b[:, :, x, :, :] * ( - x / blend_extent - ) - return b - - def tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - r"""Encode a batch of images using a tiled encoder. - - Args: - x (`torch.Tensor`): Input batch of videos. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - _, _, _, height, width = x.shape - - overlap_height = int(self.tile_sample_min_height * (1 - self.tile_overlap_factor)) # 256 * (1 - 0.25) = 192 - overlap_width = int(self.tile_sample_min_width * (1 - self.tile_overlap_factor)) # 256 * (1 - 0.25) = 192 - blend_height = int(self.tile_latent_min_height * self.tile_overlap_factor) # 8 * 0.25 = 2 - blend_width = int(self.tile_latent_min_width * self.tile_overlap_factor) # 8 * 0.25 = 2 - row_limit_height = self.tile_latent_min_height - blend_height # 8 - 2 = 6 - row_limit_width = self.tile_latent_min_width - blend_width # 8 - 2 = 6 - - rows = [] - for i in range(0, height, overlap_height): - row = [] - for j in range(0, width, overlap_width): - tile = x[ - :, - :, - :, - i : i + self.tile_sample_min_height, - j : j + self.tile_sample_min_width, - ] - tile = self.encoder(tile) - row.append(tile) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :row_limit_height, :row_limit_width]) - result_rows.append(torch.cat(result_row, dim=-1)) - moments = torch.cat(result_rows, dim=-2) - - return moments - - def tiled_decode(self, z: torch.Tensor) -> torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - - _, _, _, height, width = z.shape - - overlap_height = int(self.tile_latent_min_height * (1 - self.tile_overlap_factor)) # 8 * (1 - 0.25) = 6 - overlap_width = int(self.tile_latent_min_width * (1 - self.tile_overlap_factor)) # 8 * (1 - 0.25) = 6 - blend_height = int(self.tile_sample_min_height * self.tile_overlap_factor) # 256 * 0.25 = 64 - blend_width = int(self.tile_sample_min_width * self.tile_overlap_factor) # 256 * 0.25 = 64 - row_limit_height = self.tile_sample_min_height - blend_height # 256 - 64 = 192 - row_limit_width = self.tile_sample_min_width - blend_width # 256 - 64 = 192 - - rows = [] - for i in range(0, height, overlap_height): - row = [] - for j in range(0, width, overlap_width): - tile = z[ - :, - :, - :, - i : i + self.tile_latent_min_height, - j : j + self.tile_latent_min_width, - ] - decoded = self.decoder(tile) - row.append(decoded) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :row_limit_height, :row_limit_width]) - result_rows.append(torch.cat(result_row, dim=-1)) - dec = torch.cat(result_rows, dim=-2) - - return dec - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z, return_dict=return_dict) - return dec diff --git a/diffusers/models/autoencoders/autoencoder_kl_kvae.py b/diffusers/models/autoencoders/autoencoder_kl_kvae.py deleted file mode 100644 index dc8b9e4c36e7e2775335d0610faf20a88fe15228..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_kvae.py +++ /dev/null @@ -1,810 +0,0 @@ -# Copyright 2025 The Kandinsky Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from typing import Optional, Tuple, Union - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -class KVAEResnetBlock2D(nn.Module): - r""" - A Resnet block with optional guidance. - - Parameters: - in_channels (`int`): The number of channels in the input. - out_channels (`int`, *optional*, default to `None`): - The number of output channels for the first conv2d layer. If None, same as `in_channels`. - conv_shortcut (`bool`, *optional*, default to `False`): - If `True` and `in_channels` not equal to `out_channels`, add a 3x3 nn.conv2d layer for skip-connection. - temb_channels (`int`, *optional*, default to `512`): The number of channels in timestep embedding. - zq_ch (`int`, *optional*, default to `None`): Guidance channels for normalization. - add_conv (`bool`, *optional*, default to `False`): - If `True` add conv2d layer for normalization. - normalization (`nn.Module`, *optional*, default to `None`): The normalization layer. - act_fn (`str`, *optional*, default to `"swish"`): The activation function to use. - """ - - def __init__( - self, - *, - in_channels: int, - out_channels: Optional[int] = None, - conv_shortcut: bool = False, - temb_channels: int = 512, - zq_ch: Optional[int] = None, - add_conv: bool = False, - act_fn: str = "swish", - ): - super().__init__() - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - self.use_conv_shortcut = conv_shortcut - self.nonlinearity = get_activation(act_fn) - - if zq_ch is None: - self.norm1 = nn.GroupNorm(num_channels=in_channels, num_groups=32, eps=1e-6, affine=True) - else: - self.norm1 = KVAEDecoderSpatialNorm2D(in_channels, zq_channels=zq_ch, add_conv=add_conv) - - self.conv1 = nn.Conv2d( - in_channels=in_channels, out_channels=out_channels, kernel_size=3, padding=(1, 1), padding_mode="replicate" - ) - if temb_channels > 0: - self.temb_proj = torch.nn.Linear(temb_channels, out_channels) - if zq_ch is None: - self.norm2 = nn.GroupNorm(num_channels=out_channels, num_groups=32, eps=1e-6, affine=True) - else: - self.norm2 = KVAEDecoderSpatialNorm2D(out_channels, zq_channels=zq_ch, add_conv=add_conv) - self.conv2 = nn.Conv2d( - in_channels=out_channels, - out_channels=out_channels, - kernel_size=3, - padding=(1, 1), - padding_mode="replicate", - ) - if self.in_channels != self.out_channels: - if self.use_conv_shortcut: - self.conv_shortcut = nn.Conv2d( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=3, - padding=(1, 1), - padding_mode="replicate", - ) - else: - self.nin_shortcut = nn.Conv2d( - in_channels, - out_channels, - kernel_size=1, - stride=1, - padding=0, - ) - - def forward(self, x: torch.Tensor, temb: torch.Tensor, zq: torch.Tensor = None) -> torch.Tensor: - h = x - - if zq is None: - h = self.norm1(h) - else: - h = self.norm1(h, zq) - - h = self.nonlinearity(h) - h = self.conv1(h) - - if temb is not None: - h = h + self.temb_proj(self.nonlinearity(temb))[:, :, None, None, None] - - if zq is None: - h = self.norm2(h) - else: - h = self.norm2(h, zq) - - h = self.nonlinearity(h) - - h = self.conv2(h) - - if self.in_channels != self.out_channels: - if self.use_conv_shortcut: - x = self.conv_shortcut(x) - else: - x = self.nin_shortcut(x) - - return x + h - - -class KVAEPXSDownsample(nn.Module): - def __init__(self, in_channels: int, factor: int = 2): - r""" - A Downsampling module. - - Args: - in_channels (`int`): The number of channels in the input. - factor (`int`, *optional*, default to `2`): The downsampling factor. - """ - super().__init__() - self.factor = factor - self.unshuffle = nn.PixelUnshuffle(self.factor) - self.spatial_conv = nn.Conv2d( - in_channels, in_channels, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), padding_mode="reflect" - ) - self.linear = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - # x: (bchw) - pxs_interm = self.unshuffle(x) - b, c, h, w = pxs_interm.shape - pxs_interm_view = pxs_interm.view(b, c // self.factor**2, self.factor**2, h, w) - pxs_out = torch.mean(pxs_interm_view, dim=2) - - conv_out = self.spatial_conv(x) - - # adding it all together - out = conv_out + pxs_out - return self.linear(out) - - -class KVAEPXSUpsample(nn.Module): - def __init__(self, in_channels: int, factor: int = 2): - r""" - An Upsampling module. - - Args: - in_channels (`int`): The number of channels in the input. - factor (`int`, *optional*, default to `2`): The upsampling factor. - """ - super().__init__() - self.factor = factor - self.shuffle = nn.PixelShuffle(self.factor) - self.spatial_conv = nn.Conv2d( - in_channels, in_channels, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), padding_mode="reflect" - ) - - self.linear = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - repeated = x.repeat_interleave(self.factor**2, dim=1) - pxs_interm = self.shuffle(repeated) - - image_like_ups = F.interpolate(x, scale_factor=2, mode="nearest") - conv_out = self.spatial_conv(image_like_ups) - - # adding it all together - out = conv_out + pxs_interm - return self.linear(out) - - -class KVAEDecoderSpatialNorm2D(nn.Module): - r""" - A 2D normalization module for decoder. - - Args: - in_channels (`int`): The number of channels in the input. - zq_channels (`int`): The number of channels in the guidance. - add_conv (`bool`, *optional*, default to `false`): - If `True` add conv2d 3x3 layer for guidance in the beginning. - """ - - def __init__( - self, - in_channels: int, - zq_channels: int, - add_conv: bool = False, - ): - super().__init__() - self.norm_layer = nn.GroupNorm(num_channels=in_channels, num_groups=32, eps=1e-6, affine=True) - - self.add_conv = add_conv - if add_conv: - self.conv = nn.Conv2d( - in_channels=zq_channels, - out_channels=zq_channels, - kernel_size=3, - padding=(1, 1), - padding_mode="replicate", - ) - - self.conv_y = nn.Conv2d( - in_channels=zq_channels, - out_channels=in_channels, - kernel_size=1, - ) - self.conv_b = nn.Conv2d( - in_channels=zq_channels, - out_channels=in_channels, - kernel_size=1, - ) - - def forward(self, f: torch.Tensor, zq: torch.Tensor) -> torch.Tensor: - f_first = f - f_first_size = f_first.shape[2:] - zq = F.interpolate(zq, size=f_first_size, mode="nearest") - - if self.add_conv: - zq = self.conv(zq) - - norm_f = self.norm_layer(f) - new_f = norm_f * self.conv_y(zq) + self.conv_b(zq) - return new_f - - -class KVAEEncoder2D(nn.Module): - r""" - A 2D encoder module. - - Args: - ch (`int`): The base number of channels in multiresolution blocks. - ch_mult (`Tuple[int, ...]`, *optional*, default to `(1, 2, 4, 8)`): - The channel multipliers in multiresolution blocks. - num_res_blocks (`int`): The number of Resnet blocks. - in_channels (`int`): The number of channels in the input. - z_channels (`int`): The number of output channels. - double_z (`bool`, *optional*, defaults to `True`): - Whether to double the number of output channels for the last block. - act_fn (`str`, *optional*, default to `"swish"`): The activation function to use. - """ - - def __init__( - self, - *, - ch: int, - ch_mult: Tuple[int, ...] = (1, 2, 4, 8), - num_res_blocks: int, - in_channels: int, - z_channels: int, - double_z: bool = True, - act_fn: str = "swish", - ): - super().__init__() - self.ch = ch - self.temb_ch = 0 - self.num_resolutions = len(ch_mult) - if isinstance(num_res_blocks, int): - self.num_res_blocks = [num_res_blocks] * self.num_resolutions - else: - self.num_res_blocks = num_res_blocks - self.nonlinearity = get_activation(act_fn) - - self.in_channels = in_channels - - self.conv_in = nn.Conv2d( - in_channels=in_channels, - out_channels=self.ch, - kernel_size=3, - padding=(1, 1), - ) - - in_ch_mult = (1,) + tuple(ch_mult) - self.down = nn.ModuleList() - for i_level in range(self.num_resolutions): - block = nn.ModuleList() - attn = nn.ModuleList() - block_in = ch * in_ch_mult[i_level] - block_out = ch * ch_mult[i_level] - for i_block in range(self.num_res_blocks[i_level]): - block.append( - KVAEResnetBlock2D( - in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - ) - ) - block_in = block_out - down = nn.Module() - down.block = block - down.attn = attn - if i_level < self.num_resolutions - 1: - down.downsample = KVAEPXSDownsample(in_channels=block_in) # mb: bad out channels - self.down.append(down) - - # middle - self.mid = nn.Module() - self.mid.block_1 = KVAEResnetBlock2D( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - ) - - self.mid.block_2 = KVAEResnetBlock2D( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - ) - - # end - self.norm_out = nn.GroupNorm(num_channels=block_in, num_groups=32, eps=1e-6, affine=True) - - self.conv_out = nn.Conv2d( - in_channels=block_in, - out_channels=2 * z_channels if double_z else z_channels, - kernel_size=3, - padding=(1, 1), - ) - - self.gradient_checkpointing = False - - def forward(self, x: torch.Tensor) -> torch.Tensor: - # timestep embedding - temb = None - - # downsampling - h = self.conv_in(x) - for i_level in range(self.num_resolutions): - for i_block in range(self.num_res_blocks[i_level]): - if torch.is_grad_enabled() and self.gradient_checkpointing: - h = self._gradient_checkpointing_func(self.down[i_level].block[i_block], h, temb) - else: - h = self.down[i_level].block[i_block](h, temb) - if len(self.down[i_level].attn) > 0: - h = self.down[i_level].attn[i_block](h) - if i_level != self.num_resolutions - 1: - h = self.down[i_level].downsample(h) - - # middle - if torch.is_grad_enabled() and self.gradient_checkpointing: - h = self._gradient_checkpointing_func(self.mid.block_1, h, temb) - h = self._gradient_checkpointing_func(self.mid.block_2, h, temb) - else: - h = self.mid.block_1(h, temb) - h = self.mid.block_2(h, temb) - - # end - h = self.norm_out(h) - h = self.nonlinearity(h) - h = self.conv_out(h) - - return h - - -class KVAEDecoder2D(nn.Module): - r""" - A 2D decoder module. - - Args: - ch (`int`): The base number of channels in multiresolution blocks. - out_ch (`int`): The number of output channels. - ch_mult (`Tuple[int, ...]`, *optional*, default to `(1, 2, 4, 8)`): - The channel multipliers in multiresolution blocks. - num_res_blocks (`int`): The number of Resnet blocks. - in_channels (`int`): The number of channels in the input. - z_channels (`int`): The number of input channels. - give_pre_end (`bool`, *optional*, default to `false`): - If `True` exit the forward pass early and return the penultimate feature map. - zq_ch (`bool`, *optional*, default to `None`): The number of channels in the guidance. - add_conv (`bool`, *optional*, default to `false`): If `True` add conv2d layer for Resnet normalization layer. - act_fn (`str`, *optional*, default to `"swish"`): The activation function to use. - """ - - def __init__( - self, - *, - ch: int, - out_ch: int, - ch_mult: Tuple[int, ...] = (1, 2, 4, 8), - num_res_blocks: int, - in_channels: int, - z_channels: int, - give_pre_end: bool = False, - zq_ch: Optional[int] = None, - add_conv: bool = False, - act_fn: str = "swish", - ): - super().__init__() - self.ch = ch - self.temb_ch = 0 - self.num_resolutions = len(ch_mult) - self.num_res_blocks = num_res_blocks - self.in_channels = in_channels - self.give_pre_end = give_pre_end - self.nonlinearity = get_activation(act_fn) - - if zq_ch is None: - zq_ch = z_channels - - # compute in_ch_mult, block_in and curr_res at lowest res - block_in = ch * ch_mult[self.num_resolutions - 1] - - self.conv_in = nn.Conv2d( - in_channels=z_channels, out_channels=block_in, kernel_size=3, padding=(1, 1), padding_mode="replicate" - ) - - # middle - self.mid = nn.Module() - self.mid.block_1 = KVAEResnetBlock2D( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - zq_ch=zq_ch, - add_conv=add_conv, - ) - - self.mid.block_2 = KVAEResnetBlock2D( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - zq_ch=zq_ch, - add_conv=add_conv, - ) - - # upsampling - self.up = nn.ModuleList() - for i_level in reversed(range(self.num_resolutions)): - block = nn.ModuleList() - attn = nn.ModuleList() - block_out = ch * ch_mult[i_level] - for i_block in range(self.num_res_blocks + 1): - block.append( - KVAEResnetBlock2D( - in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - zq_ch=zq_ch, - add_conv=add_conv, - ) - ) - block_in = block_out - up = nn.Module() - up.block = block - up.attn = attn - if i_level != 0: - up.upsample = KVAEPXSUpsample(in_channels=block_in) - self.up.insert(0, up) - - self.norm_out = KVAEDecoderSpatialNorm2D(block_in, zq_ch, add_conv=add_conv) # , gather=gather_norm) - - self.conv_out = nn.Conv2d( - in_channels=block_in, out_channels=out_ch, kernel_size=3, padding=(1, 1), padding_mode="replicate" - ) - - self.gradient_checkpointing = False - - def forward(self, z: torch.Tensor) -> torch.Tensor: - self.last_z_shape = z.shape - - # timestep embedding - temb = None - - # z to block_in - zq = z - h = self.conv_in(z) - - # middle - if torch.is_grad_enabled() and self.gradient_checkpointing: - h = self._gradient_checkpointing_func(self.mid.block_1, h, temb, zq) - h = self._gradient_checkpointing_func(self.mid.block_2, h, temb, zq) - else: - h = self.mid.block_1(h, temb, zq) - h = self.mid.block_2(h, temb, zq) - - # upsampling - for i_level in reversed(range(self.num_resolutions)): - for i_block in range(self.num_res_blocks + 1): - if torch.is_grad_enabled() and self.gradient_checkpointing: - h = self._gradient_checkpointing_func(self.up[i_level].block[i_block], h, temb, zq) - else: - h = self.up[i_level].block[i_block](h, temb, zq) - if len(self.up[i_level].attn) > 0: - h = self.up[i_level].attn[i_block](h, zq) - if i_level != 0: - h = self.up[i_level].upsample(h) - - # end - if self.give_pre_end: - return h - - h = self.norm_out(h, zq) - h = self.nonlinearity(h) - h = self.conv_out(h) - - return h - - -class AutoencoderKLKVAE(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - A VAE model with KL loss for encoding images into latents and decoding latent representations into images. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for its generic methods implemented for - all models (such as downloading or saving). - - Parameters: - in_channels (int, *optional*, defaults to 3): Number of channels in the input image. - channels (int, *optional*, defaults to 128): The base number of channels in multiresolution blocks. - num_enc_blocks (int, *optional*, defaults to 2): - The number of Resnet blocks in encoder multiresolution layers. - num_dec_blocks (int, *optional*, defaults to 2): - The number of Resnet blocks in decoder multiresolution layers. - z_channels (int, *optional*, defaults to 16): Number of channels in the latent space. - double_z (`bool`, *optional*, defaults to `True`): - Whether to double the number of output channels of encoder. - ch_mult (`Tuple[int, ...]`, *optional*, default to `(1, 2, 4, 8)`): - The channel multipliers in multiresolution blocks. - sample_size (`int`, *optional*, defaults to `1024`): Sample input size. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - channels: int = 128, - num_enc_blocks: int = 2, - num_dec_blocks: int = 2, - z_channels: int = 16, - double_z: bool = True, - ch_mult: Tuple[int, ...] = (1, 2, 4, 8), - sample_size: int = 1024, - ): - super().__init__() - - # pass init params to Encoder - self.encoder = KVAEEncoder2D( - in_channels=in_channels, - ch=channels, - ch_mult=ch_mult, - num_res_blocks=num_enc_blocks, - z_channels=z_channels, - double_z=double_z, - ) - - # pass init params to Decoder - self.decoder = KVAEDecoder2D( - out_ch=in_channels, - ch=channels, - ch_mult=ch_mult, - num_res_blocks=num_dec_blocks, - in_channels=None, - z_channels=z_channels, - ) - - self.use_slicing = False - self.use_tiling = False - - # only relevant if vae tiling is enabled - self.tile_sample_min_size = self.config.sample_size - sample_size = ( - self.config.sample_size[0] - if isinstance(self.config.sample_size, (list, tuple)) - else self.config.sample_size - ) - self.tile_latent_min_size = int(sample_size / (2 ** (len(self.config.ch_mult) - 1))) - self.tile_overlap_factor = 0.25 - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, height, width = x.shape - - if self.use_tiling and (width > self.tile_sample_min_size or height > self.tile_sample_min_size): - return self._tiled_encode(x) - - enc = self.encoder(x) - - return enc - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded images. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, return_dict: bool = True) -> Union[DecoderOutput, torch.Tensor]: - if self.use_tiling and (z.shape[-1] > self.tile_latent_min_size or z.shape[-2] > self.tile_latent_min_size): - return self.tiled_decode(z, return_dict=return_dict) - - dec = self.decoder(z) - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode( - self, z: torch.FloatTensor, return_dict: bool = True, generator=None - ) -> Union[DecoderOutput, torch.FloatTensor]: - """ - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[2], b.shape[2], blend_extent) - for y in range(blend_extent): - b[:, :, y, :] = a[:, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, y, :] * (y / blend_extent) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[3], b.shape[3], blend_extent) - for x in range(blend_extent): - b[:, :, :, x] = a[:, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, x] * (x / blend_extent) - return b - - def _tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - r"""Encode a batch of images using a tiled encoder. - - When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several - steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is - different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the - tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the - output, but they should be much less noticeable. - - Args: - x (`torch.Tensor`): Input batch of images. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - - overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor)) - blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor) - row_limit = self.tile_latent_min_size - blend_extent - - # Split the image into 512x512 tiles and encode them separately. - rows = [] - for i in range(0, x.shape[2], overlap_size): - row = [] - for j in range(0, x.shape[3], overlap_size): - tile = x[:, :, i : i + self.tile_sample_min_size, j : j + self.tile_sample_min_size] - tile = self.encoder(tile) - row.append(tile) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=3)) - - enc = torch.cat(result_rows, dim=2) - return enc - - def tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> Union[DecoderOutput, torch.Tensor]: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor)) - blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor) - row_limit = self.tile_sample_min_size - blend_extent - - # Split z into overlapping 64x64 tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, z.shape[2], overlap_size): - row = [] - for j in range(0, z.shape[3], overlap_size): - tile = z[:, :, i : i + self.tile_latent_min_size, j : j + self.tile_latent_min_size] - decoded = self.decoder(tile) - row.append(decoded) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=3)) - - dec = torch.cat(result_rows, dim=2) - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: Optional[torch.Generator] = None, - ) -> Union[DecoderOutput, torch.Tensor]: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z).sample - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) diff --git a/diffusers/models/autoencoders/autoencoder_kl_kvae_video.py b/diffusers/models/autoencoders/autoencoder_kl_kvae_video.py deleted file mode 100644 index 26a7d5b2ef1c0e0becd4bbb126e981b0ca996f40..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_kvae_video.py +++ /dev/null @@ -1,970 +0,0 @@ -# Copyright 2025 The Kandinsky Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import math -from typing import Dict, Optional, Tuple, Union - -import numpy as np -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders.single_file_model import FromOriginalModelMixin -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def nonlinearity(x: torch.Tensor) -> torch.Tensor: - return F.silu(x) - - -# ============================================================================= -# Base layers -# ============================================================================= - - -class KVAESafeConv3d(nn.Conv3d): - r""" - A 3D convolution layer that splits the input tensor into smaller parts to avoid OOM. - """ - - def forward(self, input: torch.Tensor, write_to: torch.Tensor = None) -> torch.Tensor: - memory_count = input.numel() * input.element_size() / (10**9) - - if memory_count > 3: - kernel_size = self.kernel_size[0] - part_num = math.ceil(memory_count / 2) - input_chunks = torch.chunk(input, part_num, dim=2) - - if write_to is None: - output = [] - for i, chunk in enumerate(input_chunks): - if i == 0 or kernel_size == 1: - z = torch.clone(chunk) - else: - z = torch.cat([z[:, :, -kernel_size + 1 :], chunk], dim=2) - output.append(super().forward(z)) - return torch.cat(output, dim=2) - else: - time_offset = 0 - for i, chunk in enumerate(input_chunks): - if i == 0 or kernel_size == 1: - z = torch.clone(chunk) - else: - z = torch.cat([z[:, :, -kernel_size + 1 :], chunk], dim=2) - z_time = z.size(2) - (kernel_size - 1) - write_to[:, :, time_offset : time_offset + z_time] = super().forward(z) - time_offset += z_time - return write_to - else: - if write_to is None: - return super().forward(input) - else: - write_to[...] = super().forward(input) - return write_to - - -class KVAECausalConv3d(nn.Module): - r""" - A 3D causal convolution layer. - """ - - def __init__( - self, - chan_in: int, - chan_out: int, - kernel_size: Union[int, Tuple[int, int, int]], - stride: Tuple[int, int, int] = (1, 1, 1), - dilation: Tuple[int, int, int] = (1, 1, 1), - **kwargs, - ): - super().__init__() - if isinstance(kernel_size, int): - kernel_size = (kernel_size, kernel_size, kernel_size) - - time_kernel_size, height_kernel_size, width_kernel_size = kernel_size - - self.height_pad = height_kernel_size // 2 - self.width_pad = width_kernel_size // 2 - self.time_pad = time_kernel_size - 1 - self.time_kernel_size = time_kernel_size - self.stride = stride - - self.conv = KVAESafeConv3d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs) - - def forward(self, input: torch.Tensor) -> torch.Tensor: - padding_3d = (self.width_pad, self.width_pad, self.height_pad, self.height_pad, self.time_pad, 0) - input_padded = F.pad(input, padding_3d, mode="replicate") - return self.conv(input_padded) - - -class KVAECachedCausalConv3d(nn.Module): - r""" - A 3D causal convolution layer with caching for temporal processing. - """ - - def __init__( - self, - chan_in: int, - chan_out: int, - kernel_size: Union[int, Tuple[int, int, int]], - stride: Tuple[int, int, int] = (1, 1, 1), - dilation: Tuple[int, int, int] = (1, 1, 1), - **kwargs, - ): - super().__init__() - if isinstance(kernel_size, int): - kernel_size = (kernel_size, kernel_size, kernel_size) - - time_kernel_size, height_kernel_size, width_kernel_size = kernel_size - - self.height_pad = height_kernel_size // 2 - self.width_pad = width_kernel_size // 2 - self.time_pad = time_kernel_size - 1 - self.time_kernel_size = time_kernel_size - self.stride = stride - - self.conv = KVAESafeConv3d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs) - - def forward(self, input: torch.Tensor, cache: Dict) -> torch.Tensor: - t_stride = self.stride[0] - padding_3d = (self.height_pad, self.height_pad, self.width_pad, self.width_pad, 0, 0) - input_parallel = F.pad(input, padding_3d, mode="replicate") - - if cache["padding"] is None: - first_frame = input_parallel[:, :, :1] - time_pad_shape = list(first_frame.shape) - time_pad_shape[2] = self.time_pad - padding = first_frame.expand(time_pad_shape) - else: - padding = cache["padding"] - - out_size = list(input.shape) - out_size[1] = self.conv.out_channels - if t_stride == 2: - out_size[2] = (input.size(2) + 1) // 2 - output = torch.empty(tuple(out_size), dtype=input.dtype, device=input.device) - - offset_out = math.ceil(padding.size(2) / t_stride) - offset_in = offset_out * t_stride - padding.size(2) - - if offset_out > 0: - padding_poisoned = torch.cat( - [padding, input_parallel[:, :, : offset_in + self.time_kernel_size - t_stride]], dim=2 - ) - output[:, :, :offset_out] = self.conv(padding_poisoned) - - if offset_out < output.size(2): - output[:, :, offset_out:] = self.conv(input_parallel[:, :, offset_in:]) - - pad_offset = ( - offset_in - + t_stride * math.trunc((input_parallel.size(2) - offset_in - self.time_kernel_size) / t_stride) - + t_stride - ) - cache["padding"] = torch.clone(input_parallel[:, :, pad_offset:]) - - return output - - -class KVAECachedGroupNorm(nn.Module): - r""" - GroupNorm with caching support for temporal processing. - """ - - def __init__(self, in_channels: int): - super().__init__() - self.norm_layer = nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) - - def forward(self, x: torch.Tensor, cache: Dict = None) -> torch.Tensor: - out = self.norm_layer(x) - if cache is not None and cache.get("mean") is None and cache.get("var") is None: - cache["mean"] = 1 - cache["var"] = 1 - return out - - -# ============================================================================= -# Cached layers -# ============================================================================= - - -class KVAECachedSpatialNorm3D(nn.Module): - r""" - Spatially conditioned normalization for decoder with caching. - """ - - def __init__( - self, - f_channels: int, - zq_channels: int, - add_conv: bool = False, - ): - super().__init__() - self.norm_layer = KVAECachedGroupNorm(f_channels) - self.add_conv = add_conv - - if add_conv: - self.conv = KVAECachedCausalConv3d(chan_in=zq_channels, chan_out=zq_channels, kernel_size=3) - - self.conv_y = KVAESafeConv3d(zq_channels, f_channels, kernel_size=1) - self.conv_b = KVAESafeConv3d(zq_channels, f_channels, kernel_size=1) - - def forward(self, f: torch.Tensor, zq: torch.Tensor, cache: Dict) -> torch.Tensor: - if cache["norm"].get("mean") is None and cache["norm"].get("var") is None: - f_first, f_rest = f[:, :, :1], f[:, :, 1:] - f_first_size, f_rest_size = f_first.shape[-3:], f_rest.shape[-3:] - zq_first, zq_rest = zq[:, :, :1], zq[:, :, 1:] - - zq_first = F.interpolate(zq_first, size=f_first_size, mode="nearest") - - if zq.size(2) > 1: - zq_rest_splits = torch.split(zq_rest, 32, dim=1) - interpolated_splits = [ - F.interpolate(split, size=f_rest_size, mode="nearest") for split in zq_rest_splits - ] - zq_rest = torch.cat(interpolated_splits, dim=1) - zq = torch.cat([zq_first, zq_rest], dim=2) - else: - zq = zq_first - else: - f_size = f.shape[-3:] - zq_splits = torch.split(zq, 32, dim=1) - interpolated_splits = [F.interpolate(split, size=f_size, mode="nearest") for split in zq_splits] - zq = torch.cat(interpolated_splits, dim=1) - - if self.add_conv: - zq = self.conv(zq, cache["add_conv"]) - - norm_f = self.norm_layer(f, cache["norm"]) - norm_f = norm_f * self.conv_y(zq) - norm_f = norm_f + self.conv_b(zq) - - return norm_f - - -class KVAECachedResnetBlock3D(nn.Module): - r""" - A 3D ResNet block with caching. - """ - - def __init__( - self, - in_channels: int, - out_channels: Optional[int] = None, - conv_shortcut: bool = False, - dropout: float = 0.0, - temb_channels: int = 0, - zq_ch: Optional[int] = None, - add_conv: bool = False, - gather_norm: bool = False, - ): - super().__init__() - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - self.use_conv_shortcut = conv_shortcut - - if zq_ch is None: - self.norm1 = KVAECachedGroupNorm(in_channels) - else: - self.norm1 = KVAECachedSpatialNorm3D(in_channels, zq_ch, add_conv=add_conv) - - self.conv1 = KVAECachedCausalConv3d(chan_in=in_channels, chan_out=out_channels, kernel_size=3) - - if temb_channels > 0: - self.temb_proj = nn.Linear(temb_channels, out_channels) - - if zq_ch is None: - self.norm2 = KVAECachedGroupNorm(out_channels) - else: - self.norm2 = KVAECachedSpatialNorm3D(out_channels, zq_ch, add_conv=add_conv) - - self.conv2 = KVAECachedCausalConv3d(chan_in=out_channels, chan_out=out_channels, kernel_size=3) - - if self.in_channels != self.out_channels: - if self.use_conv_shortcut: - self.conv_shortcut = KVAECachedCausalConv3d(chan_in=in_channels, chan_out=out_channels, kernel_size=3) - else: - self.nin_shortcut = KVAESafeConv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0) - - def forward(self, x: torch.Tensor, temb: torch.Tensor, layer_cache: Dict, zq: torch.Tensor = None) -> torch.Tensor: - h = x - - if zq is None: - # Encoder path - norm takes cache - h = self.norm1(h, cache=layer_cache["norm1"]) - else: - # Decoder path - spatial norm takes zq and cache - h = self.norm1(h, zq, cache=layer_cache["norm1"]) - - h = F.silu(h) - h = self.conv1(h, cache=layer_cache["conv1"]) - - if temb is not None: - h = h + self.temb_proj(nonlinearity(temb))[:, :, None, None, None] - - if zq is None: - h = self.norm2(h, cache=layer_cache["norm2"]) - else: - h = self.norm2(h, zq, cache=layer_cache["norm2"]) - - h = F.silu(h) - h = self.conv2(h, cache=layer_cache["conv2"]) - - if self.in_channels != self.out_channels: - if self.use_conv_shortcut: - x = self.conv_shortcut(x, cache=layer_cache["conv_shortcut"]) - else: - x = self.nin_shortcut(x) - - return x + h - - -class KVAECachedPXSDownsample(nn.Module): - r""" - A 3D downsampling layer using PixelUnshuffle with caching. - """ - - def __init__(self, in_channels: int, compress_time: bool, factor: int = 2): - super().__init__() - self.temporal_compress = compress_time - self.factor = factor - self.unshuffle = nn.PixelUnshuffle(self.factor) - self.s_pool = nn.AvgPool3d((1, 2, 2), (1, 2, 2)) - - self.spatial_conv = KVAESafeConv3d( - in_channels, - in_channels, - kernel_size=(1, 3, 3), - stride=(1, 2, 2), - padding=(0, 1, 1), - padding_mode="reflect", - ) - - if self.temporal_compress: - self.temporal_conv = KVAECachedCausalConv3d( - in_channels, in_channels, kernel_size=(3, 1, 1), stride=(2, 1, 1), dilation=(1, 1, 1) - ) - - self.linear = nn.Conv3d(in_channels, in_channels, kernel_size=1, stride=1) - - def spatial_downsample(self, input: torch.Tensor) -> torch.Tensor: - b, c, t, h, w = input.shape - pxs_input = input.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w) - # pxs_input = rearrange(input, 'b c t h w -> (b t) c h w') - pxs_interm = self.unshuffle(pxs_input) - b_it, c_it, h_it, w_it = pxs_interm.shape - pxs_interm_view = pxs_interm.view(b_it, c_it // self.factor**2, self.factor**2, h_it, w_it) - pxs_out = torch.mean(pxs_interm_view, dim=2) - pxs_out = pxs_out.view(b, t, -1, h_it, w_it).permute(0, 2, 1, 3, 4) - # pxs_out = rearrange(pxs_out, '(b t) c h w -> b c t h w', t=input.size(2)) - conv_out = self.spatial_conv(input) - return conv_out + pxs_out - - def temporal_downsample(self, input: torch.Tensor, cache: list) -> torch.Tensor: - b, c, t, h, w = input.shape - - permuted = input.permute(0, 3, 4, 1, 2).reshape(b * h * w, c, t) - - if cache[0]["padding"] is None: - first, rest = permuted[..., :1], permuted[..., 1:] - if rest.size(-1) > 0: - rest_interp = F.avg_pool1d(rest, kernel_size=2, stride=2) - full_interp = torch.cat([first, rest_interp], dim=-1) - else: - full_interp = first - else: - rest = permuted - if rest.size(-1) > 0: - full_interp = F.avg_pool1d(rest, kernel_size=2, stride=2) - - t_new = full_interp.size(-1) - full_interp = full_interp.view(b, h, w, c, t_new).permute(0, 3, 4, 1, 2) - conv_out = self.temporal_conv(input, cache[0]) - return conv_out + full_interp - - def forward(self, x: torch.Tensor, cache: list) -> torch.Tensor: - out = self.spatial_downsample(x) - - if self.temporal_compress: - out = self.temporal_downsample(out, cache=cache) - - return self.linear(out) - - -class KVAECachedPXSUpsample(nn.Module): - r""" - A 3D upsampling layer using PixelShuffle with caching. - """ - - def __init__(self, in_channels: int, compress_time: bool, factor: int = 2): - super().__init__() - self.temporal_compress = compress_time - self.factor = factor - self.shuffle = nn.PixelShuffle(self.factor) - - self.spatial_conv = KVAESafeConv3d( - in_channels, - in_channels, - kernel_size=(1, 3, 3), - stride=(1, 1, 1), - padding=(0, 1, 1), - padding_mode="reflect", - ) - - if self.temporal_compress: - self.temporal_conv = KVAECachedCausalConv3d( - in_channels, in_channels, kernel_size=(3, 1, 1), stride=(1, 1, 1), dilation=(1, 1, 1) - ) - - self.linear = KVAESafeConv3d(in_channels, in_channels, kernel_size=1, stride=1) - - def spatial_upsample(self, input: torch.Tensor) -> torch.Tensor: - b, c, t, h, w = input.shape - input_view = input.permute(0, 2, 1, 3, 4).reshape(b, t * c, h, w) - input_interp = F.interpolate(input_view, scale_factor=2, mode="nearest") - input_interp = input_interp.view(b, t, c, 2 * h, 2 * w).permute(0, 2, 1, 3, 4) - - out = self.spatial_conv(input_interp) - return input_interp + out - - def temporal_upsample(self, input: torch.Tensor, cache: Dict) -> torch.Tensor: - time_factor = 1.0 + 1.0 * (input.size(2) > 1) - if isinstance(time_factor, torch.Tensor): - time_factor = time_factor.item() - - repeated = input.repeat_interleave(int(time_factor), dim=2) - - if cache["padding"] is None: - tail = repeated[..., int(time_factor - 1) :, :, :] - else: - tail = repeated - - conv_out = self.temporal_conv(tail, cache) - return conv_out + tail - - def forward(self, x: torch.Tensor, cache: Dict) -> torch.Tensor: - if self.temporal_compress: - x = self.temporal_upsample(x, cache) - - s_out = self.spatial_upsample(x) - to = torch.empty_like(s_out) - lin_out = self.linear(s_out, write_to=to) - return lin_out - - -# ============================================================================= -# Cached Encoder/Decoder -# ============================================================================= - - -class KVAECachedEncoder3D(nn.Module): - r""" - Cached 3D Encoder for KVAE. - """ - - def __init__( - self, - ch: int = 128, - ch_mult: Tuple[int, ...] = (1, 2, 4, 8), - num_res_blocks: int = 2, - dropout: float = 0.0, - in_channels: int = 3, - z_channels: int = 16, - double_z: bool = True, - temporal_compress_times: int = 4, - ): - super().__init__() - self.ch = ch - self.temb_ch = 0 - self.num_resolutions = len(ch_mult) - self.num_res_blocks = num_res_blocks - self.in_channels = in_channels - self.temporal_compress_level = int(np.log2(temporal_compress_times)) - - self.conv_in = KVAECachedCausalConv3d(chan_in=in_channels, chan_out=self.ch, kernel_size=3) - - in_ch_mult = (1,) + tuple(ch_mult) - self.down = nn.ModuleList() - block_in = ch - - for i_level in range(self.num_resolutions): - block = nn.ModuleList() - attn = nn.ModuleList() - - block_in = ch * in_ch_mult[i_level] - block_out = ch * ch_mult[i_level] - - for i_block in range(self.num_res_blocks): - block.append( - KVAECachedResnetBlock3D( - in_channels=block_in, - out_channels=block_out, - dropout=dropout, - temb_channels=self.temb_ch, - ) - ) - block_in = block_out - - down = nn.Module() - down.block = block - down.attn = attn - - if i_level != self.num_resolutions - 1: - if i_level < self.temporal_compress_level: - down.downsample = KVAECachedPXSDownsample(block_in, compress_time=True) - else: - down.downsample = KVAECachedPXSDownsample(block_in, compress_time=False) - self.down.append(down) - - self.mid = nn.Module() - self.mid.block_1 = KVAECachedResnetBlock3D( - in_channels=block_in, out_channels=block_in, temb_channels=self.temb_ch, dropout=dropout - ) - self.mid.block_2 = KVAECachedResnetBlock3D( - in_channels=block_in, out_channels=block_in, temb_channels=self.temb_ch, dropout=dropout - ) - - self.norm_out = KVAECachedGroupNorm(block_in) - self.conv_out = KVAECachedCausalConv3d( - chan_in=block_in, chan_out=2 * z_channels if double_z else z_channels, kernel_size=3 - ) - - self.gradient_checkpointing = False - - def forward(self, x: torch.Tensor, cache_dict: Dict) -> torch.Tensor: - temb = None - - h = self.conv_in(x, cache=cache_dict["conv_in"]) - - for i_level in range(self.num_resolutions): - for i_block in range(self.num_res_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - h = self._gradient_checkpointing_func( - self.down[i_level].block[i_block], h, temb, cache_dict[i_level][i_block] - ) - else: - h = self.down[i_level].block[i_block](h, temb, layer_cache=cache_dict[i_level][i_block]) - if len(self.down[i_level].attn) > 0: - h = self.down[i_level].attn[i_block](h) - if i_level != self.num_resolutions - 1: - h = self.down[i_level].downsample(h, cache=cache_dict[i_level]["down"]) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - h = self._gradient_checkpointing_func(self.mid.block_1, h, temb, cache_dict["mid_1"]) - h = self._gradient_checkpointing_func(self.mid.block_2, h, temb, cache_dict["mid_2"]) - else: - h = self.mid.block_1(h, temb, layer_cache=cache_dict["mid_1"]) - h = self.mid.block_2(h, temb, layer_cache=cache_dict["mid_2"]) - - h = self.norm_out(h, cache=cache_dict["norm_out"]) - h = nonlinearity(h) - h = self.conv_out(h, cache=cache_dict["conv_out"]) - - return h - - -class KVAECachedDecoder3D(nn.Module): - r""" - Cached 3D Decoder for KVAE. - """ - - def __init__( - self, - ch: int = 128, - out_ch: int = 3, - ch_mult: Tuple[int, ...] = (1, 2, 4, 8), - num_res_blocks: int = 2, - dropout: float = 0.0, - z_channels: int = 16, - zq_ch: Optional[int] = None, - add_conv: bool = False, - temporal_compress_times: int = 4, - ): - super().__init__() - self.ch = ch - self.temb_ch = 0 - self.num_resolutions = len(ch_mult) - self.num_res_blocks = num_res_blocks - self.temporal_compress_level = int(np.log2(temporal_compress_times)) - - if zq_ch is None: - zq_ch = z_channels - - block_in = ch * ch_mult[self.num_resolutions - 1] - - self.conv_in = KVAECachedCausalConv3d(chan_in=z_channels, chan_out=block_in, kernel_size=3) - - self.mid = nn.Module() - self.mid.block_1 = KVAECachedResnetBlock3D( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - zq_ch=zq_ch, - add_conv=add_conv, - ) - self.mid.block_2 = KVAECachedResnetBlock3D( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - zq_ch=zq_ch, - add_conv=add_conv, - ) - - self.up = nn.ModuleList() - for i_level in reversed(range(self.num_resolutions)): - block = nn.ModuleList() - attn = nn.ModuleList() - block_out = ch * ch_mult[i_level] - - for i_block in range(self.num_res_blocks + 1): - block.append( - KVAECachedResnetBlock3D( - in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - zq_ch=zq_ch, - add_conv=add_conv, - ) - ) - block_in = block_out - - up = nn.Module() - up.block = block - up.attn = attn - - if i_level != 0: - if i_level < self.num_resolutions - self.temporal_compress_level: - up.upsample = KVAECachedPXSUpsample(block_in, compress_time=False) - else: - up.upsample = KVAECachedPXSUpsample(block_in, compress_time=True) - self.up.insert(0, up) - - self.norm_out = KVAECachedSpatialNorm3D(block_in, zq_ch, add_conv=add_conv) - self.conv_out = KVAECachedCausalConv3d(chan_in=block_in, chan_out=out_ch, kernel_size=3) - - self.gradient_checkpointing = False - - def forward(self, z: torch.Tensor, cache_dict: Dict) -> torch.Tensor: - temb = None - zq = z - - h = self.conv_in(z, cache_dict["conv_in"]) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - h = self._gradient_checkpointing_func(self.mid.block_1, h, temb, cache_dict["mid_1"], zq) - h = self._gradient_checkpointing_func(self.mid.block_2, h, temb, cache_dict["mid_2"], zq) - else: - h = self.mid.block_1(h, temb, layer_cache=cache_dict["mid_1"], zq=zq) - h = self.mid.block_2(h, temb, layer_cache=cache_dict["mid_2"], zq=zq) - - for i_level in reversed(range(self.num_resolutions)): - for i_block in range(self.num_res_blocks + 1): - if torch.is_grad_enabled() and self.gradient_checkpointing: - h = self._gradient_checkpointing_func( - self.up[i_level].block[i_block], h, temb, cache_dict[i_level][i_block], zq - ) - else: - h = self.up[i_level].block[i_block](h, temb, layer_cache=cache_dict[i_level][i_block], zq=zq) - if len(self.up[i_level].attn) > 0: - h = self.up[i_level].attn[i_block](h, zq) - if i_level != 0: - h = self.up[i_level].upsample(h, cache_dict[i_level]["up"]) - - h = self.norm_out(h, zq, cache_dict["norm_out"]) - h = nonlinearity(h) - h = self.conv_out(h, cache_dict["conv_out"]) - - return h - - -# ============================================================================= -# Main AutoencoderKL class -# ============================================================================= - - -class AutoencoderKLKVAEVideo(ModelMixin, AutoencoderMixin, ConfigMixin, FromOriginalModelMixin): - r""" - A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. Used in - [KVAE](https://github.com/kandinskylab/kvae-1). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for its generic methods implemented for - all models (such as downloading or saving). - - Parameters: - ch (`int`, *optional*, defaults to 128): Base channel count. - ch_mult (`Tuple[int]`, *optional*, defaults to `(1, 2, 4, 8)`): Channel multipliers per level. - num_res_blocks (`int`, *optional*, defaults to 2): Number of residual blocks per level. - in_channels (`int`, *optional*, defaults to 3): Number of input channels. - out_ch (`int`, *optional*, defaults to 3): Number of output channels. - z_channels (`int`, *optional*, defaults to 16): Number of latent channels. - temporal_compress_times (`int`, *optional*, defaults to 4): Temporal compression factor. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["KVAECachedResnetBlock3D"] - - @register_to_config - def __init__( - self, - ch: int = 128, - ch_mult: Tuple[int, ...] = (1, 2, 4, 8), - num_res_blocks: int = 2, - in_channels: int = 3, - out_ch: int = 3, - z_channels: int = 16, - temporal_compress_times: int = 4, - ): - super().__init__() - - self.encoder = KVAECachedEncoder3D( - ch=ch, - ch_mult=ch_mult, - num_res_blocks=num_res_blocks, - in_channels=in_channels, - z_channels=z_channels, - double_z=True, - temporal_compress_times=temporal_compress_times, - ) - - self.decoder = KVAECachedDecoder3D( - ch=ch, - ch_mult=ch_mult, - num_res_blocks=num_res_blocks, - out_ch=out_ch, - z_channels=z_channels, - temporal_compress_times=temporal_compress_times, - ) - - self.use_slicing = False - self.use_tiling = False - - def _make_encoder_cache(self) -> Dict: - """Create empty cache for cached encoder.""" - - def make_dict(name, p=None): - if name == "conv": - return {"padding": None} - - layer, module = name.split("_") - if layer == "norm": - if module == "enc": - return {"mean": None, "var": None} - else: - return {"norm": make_dict("norm_enc"), "add_conv": make_dict("conv")} - elif layer == "resblock": - return { - "norm1": make_dict(f"norm_{module}"), - "norm2": make_dict(f"norm_{module}"), - "conv1": make_dict("conv"), - "conv2": make_dict("conv"), - "conv_shortcut": make_dict("conv"), - } - elif layer.isdigit(): - out_dict = {"down": [make_dict("conv"), make_dict("conv")], "up": make_dict("conv")} - for i in range(p): - out_dict[i] = make_dict(f"resblock_{module}") - return out_dict - - cache = { - "conv_in": make_dict("conv"), - "mid_1": make_dict("resblock_enc"), - "mid_2": make_dict("resblock_enc"), - "norm_out": make_dict("norm_enc"), - "conv_out": make_dict("conv"), - } - # Encoder uses num_res_blocks per level - for i in range(len(self.config.ch_mult)): - cache[i] = make_dict(f"{i}_enc", p=self.config.num_res_blocks) - return cache - - def _make_decoder_cache(self) -> Dict: - """Create empty cache for decoder.""" - - def make_dict(name, p=None): - if name == "conv": - return {"padding": None} - - layer, module = name.split("_") - if layer == "norm": - if module == "enc": - return {"mean": None, "var": None} - else: - return {"norm": make_dict("norm_enc"), "add_conv": make_dict("conv")} - elif layer == "resblock": - return { - "norm1": make_dict(f"norm_{module}"), - "norm2": make_dict(f"norm_{module}"), - "conv1": make_dict("conv"), - "conv2": make_dict("conv"), - "conv_shortcut": make_dict("conv"), - } - elif layer.isdigit(): - out_dict = {"down": [make_dict("conv"), make_dict("conv")], "up": make_dict("conv")} - for i in range(p): - out_dict[i] = make_dict(f"resblock_{module}") - return out_dict - - cache = { - "conv_in": make_dict("conv"), - "mid_1": make_dict("resblock_dec"), - "mid_2": make_dict("resblock_dec"), - "norm_out": make_dict("norm_dec"), - "conv_out": make_dict("conv"), - } - for i in range(len(self.config.ch_mult)): - cache[i] = make_dict(f"{i}_dec", p=self.config.num_res_blocks + 1) - return cache - - def enable_slicing(self) -> None: - r"""Enable sliced VAE decoding.""" - self.use_slicing = True - - def disable_slicing(self) -> None: - r"""Disable sliced VAE decoding.""" - self.use_slicing = False - - def _encode(self, x: torch.Tensor, seg_len: int = 16) -> torch.Tensor: - # Cached encoder processes by segments - cache = self._make_encoder_cache() - - split_list = [seg_len + 1] - n_frames = x.size(2) - (seg_len + 1) - while n_frames > 0: - split_list.append(seg_len) - n_frames -= seg_len - split_list[-1] += n_frames - - latent = [] - for chunk in torch.split(x, split_list, dim=2): - l = self.encoder(chunk, cache) - sample, _ = torch.chunk(l, 2, dim=1) - latent.append(sample) - - return torch.cat(latent, dim=2) - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]: - """ - Encode a batch of videos into latents. - - Args: - x (`torch.Tensor`): Input batch of videos with shape (B, C, T, H, W). - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - # For cached encoder, we already did the split in _encode - h_double = torch.cat([h, torch.zeros_like(h)], dim=1) - posterior = DiagonalGaussianDistribution(h_double) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, seg_len: int = 16) -> torch.Tensor: - cache = self._make_decoder_cache() - temporal_compress = self.config.temporal_compress_times - - split_list = [seg_len + 1] - n_frames = temporal_compress * (z.size(2) - 1) - seg_len - while n_frames > 0: - split_list.append(seg_len) - n_frames -= seg_len - split_list[-1] += n_frames - split_list = [math.ceil(size / temporal_compress) for size in split_list] - - recs = [] - for chunk in torch.split(z, split_list, dim=2): - out = self.decoder(chunk, cache) - recs.append(out) - - return torch.cat(recs, dim=2) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> Union[DecoderOutput, torch.Tensor]: - """ - Decode a batch of videos. - - Args: - z (`torch.Tensor`): Input batch of latent vectors with shape (B, C, T, H, W). - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: Decoded video. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice) for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z) - - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: Optional[torch.Generator] = None, - ) -> Union[DecoderOutput, torch.Tensor]: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z).sample - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) diff --git a/diffusers/models/autoencoders/autoencoder_kl_ltx.py b/diffusers/models/autoencoders/autoencoder_kl_ltx.py deleted file mode 100644 index 8cb646e8b5db14b8496e857307da60610207874a..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_ltx.py +++ /dev/null @@ -1,1552 +0,0 @@ -# Copyright 2025 The Lightricks team and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..embeddings import PixArtAlphaCombinedTimestepSizeEmbeddings -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from ..normalization import RMSNorm -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -class LTXVideoCausalConv3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int] = 3, - stride: int | tuple[int, int, int] = 1, - dilation: int | tuple[int, int, int] = 1, - groups: int = 1, - padding_mode: str = "zeros", - is_causal: bool = True, - ): - super().__init__() - - self.in_channels = in_channels - self.out_channels = out_channels - self.is_causal = is_causal - self.kernel_size = kernel_size if isinstance(kernel_size, tuple) else (kernel_size, kernel_size, kernel_size) - - dilation = dilation if isinstance(dilation, tuple) else (dilation, 1, 1) - stride = stride if isinstance(stride, tuple) else (stride, stride, stride) - height_pad = self.kernel_size[1] // 2 - width_pad = self.kernel_size[2] // 2 - padding = (0, height_pad, width_pad) - - self.conv = nn.Conv3d( - in_channels, - out_channels, - self.kernel_size, - stride=stride, - dilation=dilation, - groups=groups, - padding=padding, - padding_mode=padding_mode, - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - time_kernel_size = self.kernel_size[0] - - if self.is_causal: - pad_left = hidden_states[:, :, :1, :, :].repeat((1, 1, time_kernel_size - 1, 1, 1)) - hidden_states = torch.concatenate([pad_left, hidden_states], dim=2) - else: - pad_left = hidden_states[:, :, :1, :, :].repeat((1, 1, (time_kernel_size - 1) // 2, 1, 1)) - pad_right = hidden_states[:, :, -1:, :, :].repeat((1, 1, (time_kernel_size - 1) // 2, 1, 1)) - hidden_states = torch.concatenate([pad_left, hidden_states, pad_right], dim=2) - - hidden_states = self.conv(hidden_states) - return hidden_states - - -class LTXVideoResnetBlock3d(nn.Module): - r""" - A 3D ResNet block used in the LTXVideo model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - dropout (`float`, defaults to `0.0`): - Dropout rate. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - elementwise_affine (`bool`, defaults to `False`): - Whether to enable elementwise affinity in the normalization layers. - non_linearity (`str`, defaults to `"swish"`): - Activation function to use. - conv_shortcut (bool, defaults to `False`): - Whether or not to use a convolution shortcut. - """ - - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - dropout: float = 0.0, - eps: float = 1e-6, - elementwise_affine: bool = False, - non_linearity: str = "swish", - is_causal: bool = True, - inject_noise: bool = False, - timestep_conditioning: bool = False, - ) -> None: - super().__init__() - - out_channels = out_channels or in_channels - - self.nonlinearity = get_activation(non_linearity) - - self.norm1 = RMSNorm(in_channels, eps=1e-8, elementwise_affine=elementwise_affine) - self.conv1 = LTXVideoCausalConv3d( - in_channels=in_channels, out_channels=out_channels, kernel_size=3, is_causal=is_causal - ) - - self.norm2 = RMSNorm(out_channels, eps=1e-8, elementwise_affine=elementwise_affine) - self.dropout = nn.Dropout(dropout) - self.conv2 = LTXVideoCausalConv3d( - in_channels=out_channels, out_channels=out_channels, kernel_size=3, is_causal=is_causal - ) - - self.norm3 = None - self.conv_shortcut = None - if in_channels != out_channels: - self.norm3 = nn.LayerNorm(in_channels, eps=eps, elementwise_affine=True, bias=True) - self.conv_shortcut = LTXVideoCausalConv3d( - in_channels=in_channels, out_channels=out_channels, kernel_size=1, stride=1, is_causal=is_causal - ) - - self.per_channel_scale1 = None - self.per_channel_scale2 = None - if inject_noise: - self.per_channel_scale1 = nn.Parameter(torch.zeros(in_channels, 1, 1)) - self.per_channel_scale2 = nn.Parameter(torch.zeros(in_channels, 1, 1)) - - self.scale_shift_table = None - if timestep_conditioning: - self.scale_shift_table = nn.Parameter(torch.randn(4, in_channels) / in_channels**0.5) - - def forward( - self, inputs: torch.Tensor, temb: torch.Tensor | None = None, generator: torch.Generator | None = None - ) -> torch.Tensor: - hidden_states = inputs - - hidden_states = self.norm1(hidden_states.movedim(1, -1)).movedim(-1, 1) - - if self.scale_shift_table is not None: - temb = temb.unflatten(1, (4, -1)) + self.scale_shift_table[None, ..., None, None, None] - shift_1, scale_1, shift_2, scale_2 = temb.unbind(dim=1) - hidden_states = hidden_states * (1 + scale_1) + shift_1 - - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv1(hidden_states) - - if self.per_channel_scale1 is not None: - spatial_shape = hidden_states.shape[-2:] - spatial_noise = torch.randn( - spatial_shape, generator=generator, device=hidden_states.device, dtype=hidden_states.dtype - )[None] - hidden_states = hidden_states + (spatial_noise * self.per_channel_scale1)[None, :, None, ...] - - hidden_states = self.norm2(hidden_states.movedim(1, -1)).movedim(-1, 1) - - if self.scale_shift_table is not None: - hidden_states = hidden_states * (1 + scale_2) + shift_2 - - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.dropout(hidden_states) - hidden_states = self.conv2(hidden_states) - - if self.per_channel_scale2 is not None: - spatial_shape = hidden_states.shape[-2:] - spatial_noise = torch.randn( - spatial_shape, generator=generator, device=hidden_states.device, dtype=hidden_states.dtype - )[None] - hidden_states = hidden_states + (spatial_noise * self.per_channel_scale2)[None, :, None, ...] - - if self.norm3 is not None: - inputs = self.norm3(inputs.movedim(1, -1)).movedim(-1, 1) - - if self.conv_shortcut is not None: - inputs = self.conv_shortcut(inputs) - - hidden_states = hidden_states + inputs - return hidden_states - - -class LTXVideoDownsampler3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - stride: int | tuple[int, int, int] = 1, - is_causal: bool = True, - padding_mode: str = "zeros", - ) -> None: - super().__init__() - - self.stride = stride if isinstance(stride, tuple) else (stride, stride, stride) - self.group_size = (in_channels * stride[0] * stride[1] * stride[2]) // out_channels - - out_channels = out_channels // (self.stride[0] * self.stride[1] * self.stride[2]) - - self.conv = LTXVideoCausalConv3d( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=3, - stride=1, - is_causal=is_causal, - padding_mode=padding_mode, - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = torch.cat([hidden_states[:, :, : self.stride[0] - 1], hidden_states], dim=2) - - residual = ( - hidden_states.unflatten(4, (-1, self.stride[2])) - .unflatten(3, (-1, self.stride[1])) - .unflatten(2, (-1, self.stride[0])) - ) - residual = residual.permute(0, 1, 3, 5, 7, 2, 4, 6).flatten(1, 4) - residual = residual.unflatten(1, (-1, self.group_size)) - residual = residual.mean(dim=2) - - hidden_states = self.conv(hidden_states) - hidden_states = ( - hidden_states.unflatten(4, (-1, self.stride[2])) - .unflatten(3, (-1, self.stride[1])) - .unflatten(2, (-1, self.stride[0])) - ) - hidden_states = hidden_states.permute(0, 1, 3, 5, 7, 2, 4, 6).flatten(1, 4) - hidden_states = hidden_states + residual - - return hidden_states - - -class LTXVideoUpsampler3d(nn.Module): - def __init__( - self, - in_channels: int, - stride: int | tuple[int, int, int] = 1, - is_causal: bool = True, - residual: bool = False, - upscale_factor: int = 1, - padding_mode: str = "zeros", - ) -> None: - super().__init__() - - self.stride = stride if isinstance(stride, tuple) else (stride, stride, stride) - self.residual = residual - self.upscale_factor = upscale_factor - - out_channels = (in_channels * stride[0] * stride[1] * stride[2]) // upscale_factor - - self.conv = LTXVideoCausalConv3d( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=3, - stride=1, - is_causal=is_causal, - padding_mode=padding_mode, - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - - if self.residual: - residual = hidden_states.reshape( - batch_size, -1, self.stride[0], self.stride[1], self.stride[2], num_frames, height, width - ) - residual = residual.permute(0, 1, 5, 2, 6, 3, 7, 4).flatten(6, 7).flatten(4, 5).flatten(2, 3) - repeats = (self.stride[0] * self.stride[1] * self.stride[2]) // self.upscale_factor - residual = residual.repeat(1, repeats, 1, 1, 1) - residual = residual[:, :, self.stride[0] - 1 :] - - hidden_states = self.conv(hidden_states) - hidden_states = hidden_states.reshape( - batch_size, -1, self.stride[0], self.stride[1], self.stride[2], num_frames, height, width - ) - hidden_states = hidden_states.permute(0, 1, 5, 2, 6, 3, 7, 4).flatten(6, 7).flatten(4, 5).flatten(2, 3) - hidden_states = hidden_states[:, :, self.stride[0] - 1 :] - - if self.residual: - hidden_states = hidden_states + residual - - return hidden_states - - -class LTXVideoDownBlock3D(nn.Module): - r""" - Down block used in the LTXVideo model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - num_layers (`int`, defaults to `1`): - Number of resnet layers. - dropout (`float`, defaults to `0.0`): - Dropout rate. - resnet_eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - resnet_act_fn (`str`, defaults to `"swish"`): - Activation function to use. - spatio_temporal_scale (`bool`, defaults to `True`): - Whether or not to use a downsampling layer. If not used, output dimension would be same as input dimension. - Whether or not to downsample across temporal dimension. - is_causal (`bool`, defaults to `True`): - Whether this layer behaves causally (future frames depend only on past frames) or not. - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - num_layers: int = 1, - dropout: float = 0.0, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - spatio_temporal_scale: bool = True, - is_causal: bool = True, - ): - super().__init__() - - out_channels = out_channels or in_channels - - resnets = [] - for _ in range(num_layers): - resnets.append( - LTXVideoResnetBlock3d( - in_channels=in_channels, - out_channels=in_channels, - dropout=dropout, - eps=resnet_eps, - non_linearity=resnet_act_fn, - is_causal=is_causal, - ) - ) - self.resnets = nn.ModuleList(resnets) - - self.downsamplers = None - if spatio_temporal_scale: - self.downsamplers = nn.ModuleList( - [ - LTXVideoCausalConv3d( - in_channels=in_channels, - out_channels=in_channels, - kernel_size=3, - stride=(2, 2, 2), - is_causal=is_causal, - ) - ] - ) - - self.conv_out = None - if in_channels != out_channels: - self.conv_out = LTXVideoResnetBlock3d( - in_channels=in_channels, - out_channels=out_channels, - dropout=dropout, - eps=resnet_eps, - non_linearity=resnet_act_fn, - is_causal=is_causal, - ) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - generator: torch.Generator | None = None, - ) -> torch.Tensor: - r"""Forward method of the `LTXDownBlock3D` class.""" - - for i, resnet in enumerate(self.resnets): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, generator) - else: - hidden_states = resnet(hidden_states, temb, generator) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - if self.conv_out is not None: - hidden_states = self.conv_out(hidden_states, temb, generator) - - return hidden_states - - -class LTXVideo095DownBlock3D(nn.Module): - r""" - Down block used in the LTXVideo model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - num_layers (`int`, defaults to `1`): - Number of resnet layers. - dropout (`float`, defaults to `0.0`): - Dropout rate. - resnet_eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - resnet_act_fn (`str`, defaults to `"swish"`): - Activation function to use. - spatio_temporal_scale (`bool`, defaults to `True`): - Whether or not to use a downsampling layer. If not used, output dimension would be same as input dimension. - Whether or not to downsample across temporal dimension. - is_causal (`bool`, defaults to `True`): - Whether this layer behaves causally (future frames depend only on past frames) or not. - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - num_layers: int = 1, - dropout: float = 0.0, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - spatio_temporal_scale: bool = True, - is_causal: bool = True, - downsample_type: str = "conv", - ): - super().__init__() - - out_channels = out_channels or in_channels - - resnets = [] - for _ in range(num_layers): - resnets.append( - LTXVideoResnetBlock3d( - in_channels=in_channels, - out_channels=in_channels, - dropout=dropout, - eps=resnet_eps, - non_linearity=resnet_act_fn, - is_causal=is_causal, - ) - ) - self.resnets = nn.ModuleList(resnets) - - self.downsamplers = None - if spatio_temporal_scale: - self.downsamplers = nn.ModuleList() - - if downsample_type == "conv": - self.downsamplers.append( - LTXVideoCausalConv3d( - in_channels=in_channels, - out_channels=in_channels, - kernel_size=3, - stride=(2, 2, 2), - is_causal=is_causal, - ) - ) - elif downsample_type == "spatial": - self.downsamplers.append( - LTXVideoDownsampler3d( - in_channels=in_channels, out_channels=out_channels, stride=(1, 2, 2), is_causal=is_causal - ) - ) - elif downsample_type == "temporal": - self.downsamplers.append( - LTXVideoDownsampler3d( - in_channels=in_channels, out_channels=out_channels, stride=(2, 1, 1), is_causal=is_causal - ) - ) - elif downsample_type == "spatiotemporal": - self.downsamplers.append( - LTXVideoDownsampler3d( - in_channels=in_channels, out_channels=out_channels, stride=(2, 2, 2), is_causal=is_causal - ) - ) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - generator: torch.Generator | None = None, - ) -> torch.Tensor: - r"""Forward method of the `LTXDownBlock3D` class.""" - - for i, resnet in enumerate(self.resnets): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, generator) - else: - hidden_states = resnet(hidden_states, temb, generator) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - return hidden_states - - -# Adapted from diffusers.models.autoencoders.autoencoder_kl_cogvideox.CogVideoMidBlock3d -class LTXVideoMidBlock3d(nn.Module): - r""" - A middle block used in the LTXVideo model. - - Args: - in_channels (`int`): - Number of input channels. - num_layers (`int`, defaults to `1`): - Number of resnet layers. - dropout (`float`, defaults to `0.0`): - Dropout rate. - resnet_eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - resnet_act_fn (`str`, defaults to `"swish"`): - Activation function to use. - is_causal (`bool`, defaults to `True`): - Whether this layer behaves causally (future frames depend only on past frames) or not. - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int, - num_layers: int = 1, - dropout: float = 0.0, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - is_causal: bool = True, - inject_noise: bool = False, - timestep_conditioning: bool = False, - ) -> None: - super().__init__() - - self.time_embedder = None - if timestep_conditioning: - self.time_embedder = PixArtAlphaCombinedTimestepSizeEmbeddings(in_channels * 4, 0) - - resnets = [] - for _ in range(num_layers): - resnets.append( - LTXVideoResnetBlock3d( - in_channels=in_channels, - out_channels=in_channels, - dropout=dropout, - eps=resnet_eps, - non_linearity=resnet_act_fn, - is_causal=is_causal, - inject_noise=inject_noise, - timestep_conditioning=timestep_conditioning, - ) - ) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - generator: torch.Generator | None = None, - ) -> torch.Tensor: - r"""Forward method of the `LTXMidBlock3D` class.""" - - if self.time_embedder is not None: - temb = self.time_embedder( - timestep=temb.flatten(), - resolution=None, - aspect_ratio=None, - batch_size=hidden_states.size(0), - hidden_dtype=hidden_states.dtype, - ) - temb = temb.view(hidden_states.size(0), -1, 1, 1, 1) - - for i, resnet in enumerate(self.resnets): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, generator) - else: - hidden_states = resnet(hidden_states, temb, generator) - - return hidden_states - - -class LTXVideoUpBlock3d(nn.Module): - r""" - Up block used in the LTXVideo model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - num_layers (`int`, defaults to `1`): - Number of resnet layers. - dropout (`float`, defaults to `0.0`): - Dropout rate. - resnet_eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - resnet_act_fn (`str`, defaults to `"swish"`): - Activation function to use. - spatio_temporal_scale (`bool`, defaults to `True`): - Whether or not to use a downsampling layer. If not used, output dimension would be same as input dimension. - Whether or not to downsample across temporal dimension. - is_causal (`bool`, defaults to `True`): - Whether this layer behaves causally (future frames depend only on past frames) or not. - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - num_layers: int = 1, - dropout: float = 0.0, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - spatio_temporal_scale: bool = True, - is_causal: bool = True, - inject_noise: bool = False, - timestep_conditioning: bool = False, - upsample_residual: bool = False, - upscale_factor: int = 1, - ): - super().__init__() - - out_channels = out_channels or in_channels - - self.time_embedder = None - if timestep_conditioning: - self.time_embedder = PixArtAlphaCombinedTimestepSizeEmbeddings(in_channels * 4, 0) - - self.conv_in = None - if in_channels != out_channels: - self.conv_in = LTXVideoResnetBlock3d( - in_channels=in_channels, - out_channels=out_channels, - dropout=dropout, - eps=resnet_eps, - non_linearity=resnet_act_fn, - is_causal=is_causal, - inject_noise=inject_noise, - timestep_conditioning=timestep_conditioning, - ) - - self.upsamplers = None - if spatio_temporal_scale: - self.upsamplers = nn.ModuleList( - [ - LTXVideoUpsampler3d( - out_channels * upscale_factor, - stride=(2, 2, 2), - is_causal=is_causal, - residual=upsample_residual, - upscale_factor=upscale_factor, - ) - ] - ) - - resnets = [] - for _ in range(num_layers): - resnets.append( - LTXVideoResnetBlock3d( - in_channels=out_channels, - out_channels=out_channels, - dropout=dropout, - eps=resnet_eps, - non_linearity=resnet_act_fn, - is_causal=is_causal, - inject_noise=inject_noise, - timestep_conditioning=timestep_conditioning, - ) - ) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - generator: torch.Generator | None = None, - ) -> torch.Tensor: - if self.conv_in is not None: - hidden_states = self.conv_in(hidden_states, temb, generator) - - if self.time_embedder is not None: - temb = self.time_embedder( - timestep=temb.flatten(), - resolution=None, - aspect_ratio=None, - batch_size=hidden_states.size(0), - hidden_dtype=hidden_states.dtype, - ) - temb = temb.view(hidden_states.size(0), -1, 1, 1, 1) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - for i, resnet in enumerate(self.resnets): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, generator) - else: - hidden_states = resnet(hidden_states, temb, generator) - - return hidden_states - - -class LTXVideoEncoder3d(nn.Module): - r""" - The `LTXVideoEncoder3d` layer of a variational autoencoder that encodes input video samples to its latent - representation. - - Args: - in_channels (`int`, defaults to 3): - Number of input channels. - out_channels (`int`, defaults to 128): - Number of latent channels. - block_out_channels (`tuple[int, ...]`, defaults to `(128, 256, 512, 512)`): - The number of output channels for each block. - spatio_temporal_scaling (`tuple[bool, ...], defaults to `(True, True, True, False)`: - Whether a block should contain spatio-temporal downscaling layers or not. - layers_per_block (`tuple[int, ...]`, defaults to `(4, 3, 3, 3, 4)`): - The number of layers per block. - patch_size (`int`, defaults to `4`): - The size of spatial patches. - patch_size_t (`int`, defaults to `1`): - The size of temporal patches. - resnet_norm_eps (`float`, defaults to `1e-6`): - Epsilon value for ResNet normalization layers. - is_causal (`bool`, defaults to `True`): - Whether this layer behaves causally (future frames depend only on past frames) or not. - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 128, - block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - down_block_types: tuple[str, ...] = ( - "LTXVideoDownBlock3D", - "LTXVideoDownBlock3D", - "LTXVideoDownBlock3D", - "LTXVideoDownBlock3D", - ), - spatio_temporal_scaling: tuple[bool, ...] = (True, True, True, False), - layers_per_block: tuple[int, ...] = (4, 3, 3, 3, 4), - downsample_type: tuple[str, ...] = ("conv", "conv", "conv", "conv"), - patch_size: int = 4, - patch_size_t: int = 1, - resnet_norm_eps: float = 1e-6, - is_causal: bool = True, - ): - super().__init__() - - self.patch_size = patch_size - self.patch_size_t = patch_size_t - self.in_channels = in_channels * patch_size**2 - - output_channel = block_out_channels[0] - - self.conv_in = LTXVideoCausalConv3d( - in_channels=self.in_channels, - out_channels=output_channel, - kernel_size=3, - stride=1, - is_causal=is_causal, - ) - - # down blocks - is_ltx_095 = down_block_types[-1] == "LTXVideo095DownBlock3D" - num_block_out_channels = len(block_out_channels) - (1 if is_ltx_095 else 0) - self.down_blocks = nn.ModuleList([]) - for i in range(num_block_out_channels): - input_channel = output_channel - if not is_ltx_095: - output_channel = block_out_channels[i + 1] if i + 1 < num_block_out_channels else block_out_channels[i] - else: - output_channel = block_out_channels[i + 1] - - if down_block_types[i] == "LTXVideoDownBlock3D": - down_block = LTXVideoDownBlock3D( - in_channels=input_channel, - out_channels=output_channel, - num_layers=layers_per_block[i], - resnet_eps=resnet_norm_eps, - spatio_temporal_scale=spatio_temporal_scaling[i], - is_causal=is_causal, - ) - elif down_block_types[i] == "LTXVideo095DownBlock3D": - down_block = LTXVideo095DownBlock3D( - in_channels=input_channel, - out_channels=output_channel, - num_layers=layers_per_block[i], - resnet_eps=resnet_norm_eps, - spatio_temporal_scale=spatio_temporal_scaling[i], - is_causal=is_causal, - downsample_type=downsample_type[i], - ) - else: - raise ValueError(f"Unknown down block type: {down_block_types[i]}") - - self.down_blocks.append(down_block) - - # mid block - self.mid_block = LTXVideoMidBlock3d( - in_channels=output_channel, - num_layers=layers_per_block[-1], - resnet_eps=resnet_norm_eps, - is_causal=is_causal, - ) - - # out - self.norm_out = RMSNorm(out_channels, eps=1e-8, elementwise_affine=False) - self.conv_act = nn.SiLU() - self.conv_out = LTXVideoCausalConv3d( - in_channels=output_channel, out_channels=out_channels + 1, kernel_size=3, stride=1, is_causal=is_causal - ) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - r"""The forward method of the `LTXVideoEncoder3d` class.""" - - p = self.patch_size - p_t = self.patch_size_t - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p - post_patch_width = width // p - - hidden_states = hidden_states.reshape( - batch_size, num_channels, post_patch_num_frames, p_t, post_patch_height, p, post_patch_width, p - ) - # Thanks for driving me insane with the weird patching order :( - hidden_states = hidden_states.permute(0, 1, 3, 7, 5, 2, 4, 6).flatten(1, 4) - hidden_states = self.conv_in(hidden_states) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for down_block in self.down_blocks: - hidden_states = self._gradient_checkpointing_func(down_block, hidden_states) - - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states) - else: - for down_block in self.down_blocks: - hidden_states = down_block(hidden_states) - - hidden_states = self.mid_block(hidden_states) - - hidden_states = self.norm_out(hidden_states.movedim(1, -1)).movedim(-1, 1) - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states) - - last_channel = hidden_states[:, -1:] - last_channel = last_channel.repeat(1, hidden_states.size(1) - 2, 1, 1, 1) - hidden_states = torch.cat([hidden_states, last_channel], dim=1) - - return hidden_states - - -class LTXVideoDecoder3d(nn.Module): - r""" - The `LTXVideoDecoder3d` layer of a variational autoencoder that decodes its latent representation into an output - sample. - - Args: - in_channels (`int`, defaults to 128): - Number of latent channels. - out_channels (`int`, defaults to 3): - Number of output channels. - block_out_channels (`tuple[int, ...]`, defaults to `(128, 256, 512, 512)`): - The number of output channels for each block. - spatio_temporal_scaling (`tuple[bool, ...], defaults to `(True, True, True, False)`: - Whether a block should contain spatio-temporal upscaling layers or not. - layers_per_block (`tuple[int, ...]`, defaults to `(4, 3, 3, 3, 4)`): - The number of layers per block. - patch_size (`int`, defaults to `4`): - The size of spatial patches. - patch_size_t (`int`, defaults to `1`): - The size of temporal patches. - resnet_norm_eps (`float`, defaults to `1e-6`): - Epsilon value for ResNet normalization layers. - is_causal (`bool`, defaults to `False`): - Whether this layer behaves causally (future frames depend only on past frames) or not. - timestep_conditioning (`bool`, defaults to `False`): - Whether to condition the model on timesteps. - """ - - def __init__( - self, - in_channels: int = 128, - out_channels: int = 3, - block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - spatio_temporal_scaling: tuple[bool, ...] = (True, True, True, False), - layers_per_block: tuple[int, ...] = (4, 3, 3, 3, 4), - patch_size: int = 4, - patch_size_t: int = 1, - resnet_norm_eps: float = 1e-6, - is_causal: bool = False, - inject_noise: tuple[bool, ...] = (False, False, False, False), - timestep_conditioning: bool = False, - upsample_residual: tuple[bool, ...] = (False, False, False, False), - upsample_factor: tuple[bool, ...] = (1, 1, 1, 1), - ) -> None: - super().__init__() - - self.patch_size = patch_size - self.patch_size_t = patch_size_t - self.out_channels = out_channels * patch_size**2 - - block_out_channels = tuple(reversed(block_out_channels)) - spatio_temporal_scaling = tuple(reversed(spatio_temporal_scaling)) - layers_per_block = tuple(reversed(layers_per_block)) - inject_noise = tuple(reversed(inject_noise)) - upsample_residual = tuple(reversed(upsample_residual)) - upsample_factor = tuple(reversed(upsample_factor)) - output_channel = block_out_channels[0] - - self.conv_in = LTXVideoCausalConv3d( - in_channels=in_channels, out_channels=output_channel, kernel_size=3, stride=1, is_causal=is_causal - ) - - self.mid_block = LTXVideoMidBlock3d( - in_channels=output_channel, - num_layers=layers_per_block[0], - resnet_eps=resnet_norm_eps, - is_causal=is_causal, - inject_noise=inject_noise[0], - timestep_conditioning=timestep_conditioning, - ) - - # up blocks - num_block_out_channels = len(block_out_channels) - self.up_blocks = nn.ModuleList([]) - for i in range(num_block_out_channels): - input_channel = output_channel // upsample_factor[i] - output_channel = block_out_channels[i] // upsample_factor[i] - - up_block = LTXVideoUpBlock3d( - in_channels=input_channel, - out_channels=output_channel, - num_layers=layers_per_block[i + 1], - resnet_eps=resnet_norm_eps, - spatio_temporal_scale=spatio_temporal_scaling[i], - is_causal=is_causal, - inject_noise=inject_noise[i + 1], - timestep_conditioning=timestep_conditioning, - upsample_residual=upsample_residual[i], - upscale_factor=upsample_factor[i], - ) - - self.up_blocks.append(up_block) - - # out - self.norm_out = RMSNorm(out_channels, eps=1e-8, elementwise_affine=False) - self.conv_act = nn.SiLU() - self.conv_out = LTXVideoCausalConv3d( - in_channels=output_channel, out_channels=self.out_channels, kernel_size=3, stride=1, is_causal=is_causal - ) - - # timestep embedding - self.time_embedder = None - self.scale_shift_table = None - self.timestep_scale_multiplier = None - if timestep_conditioning: - self.timestep_scale_multiplier = nn.Parameter(torch.tensor(1000.0, dtype=torch.float32)) - self.time_embedder = PixArtAlphaCombinedTimestepSizeEmbeddings(output_channel * 2, 0) - self.scale_shift_table = nn.Parameter(torch.randn(2, output_channel) / output_channel**0.5) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - hidden_states = self.conv_in(hidden_states) - - if self.timestep_scale_multiplier is not None: - temb = temb * self.timestep_scale_multiplier - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states, temb) - - for up_block in self.up_blocks: - hidden_states = self._gradient_checkpointing_func(up_block, hidden_states, temb) - else: - hidden_states = self.mid_block(hidden_states, temb) - - for up_block in self.up_blocks: - hidden_states = up_block(hidden_states, temb) - - hidden_states = self.norm_out(hidden_states.movedim(1, -1)).movedim(-1, 1) - - if self.time_embedder is not None: - temb = self.time_embedder( - timestep=temb.flatten(), - resolution=None, - aspect_ratio=None, - batch_size=hidden_states.size(0), - hidden_dtype=hidden_states.dtype, - ) - temb = temb.view(hidden_states.size(0), -1, 1, 1, 1).unflatten(1, (2, -1)) - temb = temb + self.scale_shift_table[None, ..., None, None, None] - shift, scale = temb.unbind(dim=1) - hidden_states = hidden_states * (1 + scale) + shift - - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states) - - p = self.patch_size - p_t = self.patch_size_t - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - hidden_states = hidden_states.reshape(batch_size, -1, p_t, p, p, num_frames, height, width) - hidden_states = hidden_states.permute(0, 1, 5, 2, 6, 4, 7, 3).flatten(6, 7).flatten(4, 5).flatten(2, 3) - - return hidden_states - - -class AutoencoderKLLTXVideo(ModelMixin, AutoencoderMixin, ConfigMixin, FromOriginalModelMixin): - r""" - A VAE model with KL loss for encoding images into latents and decoding latent representations into images. Used in - [LTX](https://huggingface.co/Lightricks/LTX-Video). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Args: - in_channels (`int`, defaults to `3`): - Number of input channels. - out_channels (`int`, defaults to `3`): - Number of output channels. - latent_channels (`int`, defaults to `128`): - Number of latent channels. - block_out_channels (`tuple[int, ...]`, defaults to `(128, 256, 512, 512)`): - The number of output channels for each block. - spatio_temporal_scaling (`tuple[bool, ...], defaults to `(True, True, True, False)`: - Whether a block should contain spatio-temporal downscaling or not. - layers_per_block (`tuple[int, ...]`, defaults to `(4, 3, 3, 3, 4)`): - The number of layers per block. - patch_size (`int`, defaults to `4`): - The size of spatial patches. - patch_size_t (`int`, defaults to `1`): - The size of temporal patches. - resnet_norm_eps (`float`, defaults to `1e-6`): - Epsilon value for ResNet normalization layers. - scaling_factor (`float`, *optional*, defaults to `1.0`): - The component-wise standard deviation of the trained latent space computed using the first batch of the - training set. This is used to scale the latent space to have unit variance when training the diffusion - model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the - diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 - / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image - Synthesis with Latent Diffusion Models](https://huggingface.co/papers/2112.10752) paper. - encoder_causal (`bool`, defaults to `True`): - Whether the encoder should behave causally (future frames depend only on past frames) or not. - decoder_causal (`bool`, defaults to `False`): - Whether the decoder should behave causally (future frames depend only on past frames) or not. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - latent_channels: int = 128, - block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - down_block_types: tuple[str, ...] = ( - "LTXVideoDownBlock3D", - "LTXVideoDownBlock3D", - "LTXVideoDownBlock3D", - "LTXVideoDownBlock3D", - ), - decoder_block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - layers_per_block: tuple[int, ...] = (4, 3, 3, 3, 4), - decoder_layers_per_block: tuple[int, ...] = (4, 3, 3, 3, 4), - spatio_temporal_scaling: tuple[bool, ...] = (True, True, True, False), - decoder_spatio_temporal_scaling: tuple[bool, ...] = (True, True, True, False), - decoder_inject_noise: tuple[bool, ...] = (False, False, False, False, False), - downsample_type: tuple[str, ...] = ("conv", "conv", "conv", "conv"), - upsample_residual: tuple[bool, ...] = (False, False, False, False), - upsample_factor: tuple[int, ...] = (1, 1, 1, 1), - timestep_conditioning: bool = False, - patch_size: int = 4, - patch_size_t: int = 1, - resnet_norm_eps: float = 1e-6, - scaling_factor: float = 1.0, - encoder_causal: bool = True, - decoder_causal: bool = False, - spatial_compression_ratio: int = None, - temporal_compression_ratio: int = None, - ) -> None: - super().__init__() - - self.encoder = LTXVideoEncoder3d( - in_channels=in_channels, - out_channels=latent_channels, - block_out_channels=block_out_channels, - down_block_types=down_block_types, - spatio_temporal_scaling=spatio_temporal_scaling, - layers_per_block=layers_per_block, - downsample_type=downsample_type, - patch_size=patch_size, - patch_size_t=patch_size_t, - resnet_norm_eps=resnet_norm_eps, - is_causal=encoder_causal, - ) - self.decoder = LTXVideoDecoder3d( - in_channels=latent_channels, - out_channels=out_channels, - block_out_channels=decoder_block_out_channels, - spatio_temporal_scaling=decoder_spatio_temporal_scaling, - layers_per_block=decoder_layers_per_block, - patch_size=patch_size, - patch_size_t=patch_size_t, - resnet_norm_eps=resnet_norm_eps, - is_causal=decoder_causal, - timestep_conditioning=timestep_conditioning, - inject_noise=decoder_inject_noise, - upsample_residual=upsample_residual, - upsample_factor=upsample_factor, - ) - - latents_mean = torch.zeros((latent_channels,), requires_grad=False) - latents_std = torch.ones((latent_channels,), requires_grad=False) - self.register_buffer("latents_mean", latents_mean, persistent=True) - self.register_buffer("latents_std", latents_std, persistent=True) - - self.spatial_compression_ratio = ( - patch_size * 2 ** sum(spatio_temporal_scaling) - if spatial_compression_ratio is None - else spatial_compression_ratio - ) - self.temporal_compression_ratio = ( - patch_size_t * 2 ** sum(spatio_temporal_scaling) - if temporal_compression_ratio is None - else temporal_compression_ratio - ) - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent - # frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the - # intermediate tiles together, the memory requirement can be lowered. - self.use_tiling = False - - # When decoding temporally long video latents, the memory requirement is very high. By decoding latent frames - # at a fixed frame batch size (based on `self.num_latent_frames_batch_sizes`), the memory requirement can be lowered. - self.use_framewise_encoding = False - self.use_framewise_decoding = False - - # This can be configured based on the amount of GPU memory available. - # `16` for sample frames and `2` for latent frames are sensible defaults for consumer GPUs. - # Setting it to higher values results in higher memory usage. - self.num_sample_frames_batch_size = 16 - self.num_latent_frames_batch_size = 2 - - # The minimal tile height and width for spatial tiling to be used - self.tile_sample_min_height = 512 - self.tile_sample_min_width = 512 - self.tile_sample_min_num_frames = 16 - - # The minimal distance between two spatial tiles - self.tile_sample_stride_height = 448 - self.tile_sample_stride_width = 448 - self.tile_sample_stride_num_frames = 8 - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_sample_min_num_frames: int | None = None, - tile_sample_stride_height: float | None = None, - tile_sample_stride_width: float | None = None, - tile_sample_stride_num_frames: float | None = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_sample_stride_height (`int`, *optional*): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. - tile_sample_stride_width (`int`, *optional*): - The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling - artifacts produced across the width dimension. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_sample_min_num_frames = tile_sample_min_num_frames or self.tile_sample_min_num_frames - self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height - self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width - self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = x.shape - - if self.use_framewise_decoding and num_frames > self.tile_sample_min_num_frames: - return self._temporal_tiled_encode(x) - - if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): - return self.tiled_encode(x) - - enc = self.encoder(x) - - return enc - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode( - self, z: torch.Tensor, temb: torch.Tensor | None = None, return_dict: bool = True - ) -> DecoderOutput | torch.Tensor: - batch_size, num_channels, num_frames, height, width = z.shape - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio - - if self.use_framewise_decoding and num_frames > tile_latent_min_num_frames: - return self._temporal_tiled_decode(z, temb, return_dict=return_dict) - - if self.use_tiling and (width > tile_latent_min_width or height > tile_latent_min_height): - return self.tiled_decode(z, temb, return_dict=return_dict) - - dec = self.decoder(z, temb) - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode( - self, z: torch.Tensor, temb: torch.Tensor | None = None, return_dict: bool = True - ) -> DecoderOutput | torch.Tensor: - """ - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - if temb is not None: - decoded_slices = [ - self._decode(z_slice, t_slice).sample for z_slice, t_slice in (z.split(1), temb.split(1)) - ] - else: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z, temb).sample - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[3], b.shape[3], blend_extent) - for y in range(blend_extent): - b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[4], b.shape[4], blend_extent) - for x in range(blend_extent): - b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( - x / blend_extent - ) - return b - - def blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-3], b.shape[-3], blend_extent) - for x in range(blend_extent): - b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (1 - x / blend_extent) + b[:, :, x, :, :] * ( - x / blend_extent - ) - return b - - def tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - r"""Encode a batch of images using a tiled encoder. - - Args: - x (`torch.Tensor`): Input batch of videos. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - batch_size, num_channels, num_frames, height, width = x.shape - latent_height = height // self.spatial_compression_ratio - latent_width = width // self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = tile_latent_min_height - tile_latent_stride_height - blend_width = tile_latent_min_width - tile_latent_stride_width - - # Split x into overlapping tiles and encode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, self.tile_sample_stride_height): - row = [] - for j in range(0, width, self.tile_sample_stride_width): - time = self.encoder( - x[:, :, :, i : i + self.tile_sample_min_height, j : j + self.tile_sample_min_width] - ) - - row.append(time) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :tile_latent_stride_height, :tile_latent_stride_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - enc = torch.cat(result_rows, dim=3)[:, :, :, :latent_height, :latent_width] - return enc - - def tiled_decode( - self, z: torch.Tensor, temb: torch.Tensor | None, return_dict: bool = True - ) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - - batch_size, num_channels, num_frames, height, width = z.shape - sample_height = height * self.spatial_compression_ratio - sample_width = width * self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = self.tile_sample_min_height - self.tile_sample_stride_height - blend_width = self.tile_sample_min_width - self.tile_sample_stride_width - - # Split z into overlapping tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, tile_latent_stride_height): - row = [] - for j in range(0, width, tile_latent_stride_width): - time = self.decoder(z[:, :, :, i : i + tile_latent_min_height, j : j + tile_latent_min_width], temb) - - row.append(time) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, : self.tile_sample_stride_height, : self.tile_sample_stride_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - dec = torch.cat(result_rows, dim=3)[:, :, :, :sample_height, :sample_width] - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - def _temporal_tiled_encode(self, x: torch.Tensor) -> AutoencoderKLOutput: - batch_size, num_channels, num_frames, height, width = x.shape - latent_num_frames = (num_frames - 1) // self.temporal_compression_ratio + 1 - - tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio - tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio - blend_num_frames = tile_latent_min_num_frames - tile_latent_stride_num_frames - - row = [] - for i in range(0, num_frames, self.tile_sample_stride_num_frames): - tile = x[:, :, i : i + self.tile_sample_min_num_frames + 1, :, :] - if self.use_tiling and (height > self.tile_sample_min_height or width > self.tile_sample_min_width): - tile = self.tiled_encode(tile) - else: - tile = self.encoder(tile) - if i > 0: - tile = tile[:, :, 1:, :, :] - row.append(tile) - - result_row = [] - for i, tile in enumerate(row): - if i > 0: - tile = self.blend_t(row[i - 1], tile, blend_num_frames) - result_row.append(tile[:, :, :tile_latent_stride_num_frames, :, :]) - else: - result_row.append(tile[:, :, : tile_latent_stride_num_frames + 1, :, :]) - - enc = torch.cat(result_row, dim=2)[:, :, :latent_num_frames] - return enc - - def _temporal_tiled_decode( - self, z: torch.Tensor, temb: torch.Tensor | None, return_dict: bool = True - ) -> DecoderOutput | torch.Tensor: - batch_size, num_channels, num_frames, height, width = z.shape - num_sample_frames = (num_frames - 1) * self.temporal_compression_ratio + 1 - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio - tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio - blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames - - row = [] - for i in range(0, num_frames, tile_latent_stride_num_frames): - tile = z[:, :, i : i + tile_latent_min_num_frames + 1, :, :] - if self.use_tiling and (tile.shape[-1] > tile_latent_min_width or tile.shape[-2] > tile_latent_min_height): - decoded = self.tiled_decode(tile, temb, return_dict=True).sample - else: - decoded = self.decoder(tile, temb) - if i > 0: - decoded = decoded[:, :, :-1, :, :] - row.append(decoded) - - result_row = [] - for i, tile in enumerate(row): - if i > 0: - tile = self.blend_t(row[i - 1], tile, blend_num_frames) - tile = tile[:, :, : self.tile_sample_stride_num_frames, :, :] - result_row.append(tile) - else: - result_row.append(tile[:, :, : self.tile_sample_stride_num_frames + 1, :, :]) - - dec = torch.cat(result_row, dim=2)[:, :, :num_sample_frames] - - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) - - def forward( - self, - sample: torch.Tensor, - temb: torch.Tensor | None = None, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> torch.Tensor | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - temb (`torch.Tensor`, *optional*): - Optional timestep embedding tensor used to condition the decoder. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z, temb) - if not return_dict: - return (dec.sample,) - return dec diff --git a/diffusers/models/autoencoders/autoencoder_kl_ltx2.py b/diffusers/models/autoencoders/autoencoder_kl_ltx2.py deleted file mode 100644 index 959a9fdb9e11093a8ff214359cd82e62a499b18b..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_ltx2.py +++ /dev/null @@ -1,1576 +0,0 @@ -# Copyright 2025 The Lightricks team and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..embeddings import PixArtAlphaCombinedTimestepSizeEmbeddings -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -class PerChannelRMSNorm(nn.Module): - """ - Per-pixel (per-location) RMS normalization layer. - - For each element along the chosen dimension, this layer normalizes the tensor by the root-mean-square of its values - across that dimension: - - y = x / sqrt(mean(x^2, dim=dim, keepdim=True) + eps) - """ - - def __init__(self, channel_dim: int = 1, eps: float = 1e-8) -> None: - """ - Args: - dim: Dimension along which to compute the RMS (typically channels). - eps: Small constant added for numerical stability. - """ - super().__init__() - self.channel_dim = channel_dim - self.eps = eps - - def forward(self, x: torch.Tensor, channel_dim: int | None = None) -> torch.Tensor: - """ - Apply RMS normalization along the configured dimension. - """ - channel_dim = channel_dim or self.channel_dim - # Compute mean of squared values along `dim`, keep dimensions for broadcasting. - mean_sq = torch.mean(x**2, dim=self.channel_dim, keepdim=True) - # Normalize by the root-mean-square (RMS). - rms = torch.sqrt(mean_sq + self.eps) - return x / rms - - -# Like LTXCausalConv3d, but whether causal inference is performed can be specified at runtime -class LTX2VideoCausalConv3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int] = 3, - stride: int | tuple[int, int, int] = 1, - dilation: int | tuple[int, int, int] = 1, - groups: int = 1, - spatial_padding_mode: str = "zeros", - ): - super().__init__() - - self.in_channels = in_channels - self.out_channels = out_channels - self.kernel_size = kernel_size if isinstance(kernel_size, tuple) else (kernel_size, kernel_size, kernel_size) - - dilation = dilation if isinstance(dilation, tuple) else (dilation, 1, 1) - stride = stride if isinstance(stride, tuple) else (stride, stride, stride) - height_pad = self.kernel_size[1] // 2 - width_pad = self.kernel_size[2] // 2 - padding = (0, height_pad, width_pad) - - self.conv = nn.Conv3d( - in_channels, - out_channels, - self.kernel_size, - stride=stride, - dilation=dilation, - groups=groups, - padding=padding, - padding_mode=spatial_padding_mode, - ) - - def forward(self, hidden_states: torch.Tensor, causal: bool = True) -> torch.Tensor: - time_kernel_size = self.kernel_size[0] - - if causal: - pad_left = hidden_states[:, :, :1, :, :].repeat((1, 1, time_kernel_size - 1, 1, 1)) - hidden_states = torch.concatenate([pad_left, hidden_states], dim=2) - else: - pad_left = hidden_states[:, :, :1, :, :].repeat((1, 1, (time_kernel_size - 1) // 2, 1, 1)) - pad_right = hidden_states[:, :, -1:, :, :].repeat((1, 1, (time_kernel_size - 1) // 2, 1, 1)) - hidden_states = torch.concatenate([pad_left, hidden_states, pad_right], dim=2) - - hidden_states = self.conv(hidden_states) - return hidden_states - - -# Like LTXVideoResnetBlock3d, but uses new causal Conv3d, normal Conv3d for the conv_shortcut, and the spatial padding -# mode is configurable -class LTX2VideoResnetBlock3d(nn.Module): - r""" - A 3D ResNet block used in the LTX 2.0 audiovisual model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - dropout (`float`, defaults to `0.0`): - Dropout rate. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - elementwise_affine (`bool`, defaults to `False`): - Whether to enable elementwise affinity in the normalization layers. - non_linearity (`str`, defaults to `"swish"`): - Activation function to use. - conv_shortcut (bool, defaults to `False`): - Whether or not to use a convolution shortcut. - """ - - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - dropout: float = 0.0, - eps: float = 1e-6, - elementwise_affine: bool = False, - non_linearity: str = "swish", - inject_noise: bool = False, - timestep_conditioning: bool = False, - spatial_padding_mode: str = "zeros", - ) -> None: - super().__init__() - - out_channels = out_channels or in_channels - - self.nonlinearity = get_activation(non_linearity) - - self.norm1 = PerChannelRMSNorm() - self.conv1 = LTX2VideoCausalConv3d( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=3, - spatial_padding_mode=spatial_padding_mode, - ) - - self.norm2 = PerChannelRMSNorm() - self.dropout = nn.Dropout(dropout) - self.conv2 = LTX2VideoCausalConv3d( - in_channels=out_channels, - out_channels=out_channels, - kernel_size=3, - spatial_padding_mode=spatial_padding_mode, - ) - - self.norm3 = None - self.conv_shortcut = None - if in_channels != out_channels: - self.norm3 = nn.LayerNorm(in_channels, eps=eps, elementwise_affine=True, bias=True) - # LTX 2.0 uses a normal nn.Conv3d here rather than LTXVideoCausalConv3d - self.conv_shortcut = nn.Conv3d(in_channels=in_channels, out_channels=out_channels, kernel_size=1, stride=1) - - self.per_channel_scale1 = None - self.per_channel_scale2 = None - if inject_noise: - self.per_channel_scale1 = nn.Parameter(torch.zeros(in_channels, 1, 1)) - self.per_channel_scale2 = nn.Parameter(torch.zeros(in_channels, 1, 1)) - - self.scale_shift_table = None - if timestep_conditioning: - self.scale_shift_table = nn.Parameter(torch.randn(4, in_channels) / in_channels**0.5) - - def forward( - self, - inputs: torch.Tensor, - temb: torch.Tensor | None = None, - generator: torch.Generator | None = None, - causal: bool = True, - ) -> torch.Tensor: - hidden_states = inputs - - hidden_states = self.norm1(hidden_states) - - if self.scale_shift_table is not None: - temb = temb.unflatten(1, (4, -1)) + self.scale_shift_table[None, ..., None, None, None] - shift_1, scale_1, shift_2, scale_2 = temb.unbind(dim=1) - hidden_states = hidden_states * (1 + scale_1) + shift_1 - - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv1(hidden_states, causal=causal) - - if self.per_channel_scale1 is not None: - spatial_shape = hidden_states.shape[-2:] - spatial_noise = torch.randn( - spatial_shape, generator=generator, device=hidden_states.device, dtype=hidden_states.dtype - )[None] - hidden_states = hidden_states + (spatial_noise * self.per_channel_scale1)[None, :, None, ...] - - hidden_states = self.norm2(hidden_states) - - if self.scale_shift_table is not None: - hidden_states = hidden_states * (1 + scale_2) + shift_2 - - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.dropout(hidden_states) - hidden_states = self.conv2(hidden_states, causal=causal) - - if self.per_channel_scale2 is not None: - spatial_shape = hidden_states.shape[-2:] - spatial_noise = torch.randn( - spatial_shape, generator=generator, device=hidden_states.device, dtype=hidden_states.dtype - )[None] - hidden_states = hidden_states + (spatial_noise * self.per_channel_scale2)[None, :, None, ...] - - if self.norm3 is not None: - inputs = self.norm3(inputs.movedim(1, -1)).movedim(-1, 1) - - if self.conv_shortcut is not None: - inputs = self.conv_shortcut(inputs) - - hidden_states = hidden_states + inputs - return hidden_states - - -# Like LTX 1.0 LTXVideoDownsampler3d, but uses new causal Conv3d -class LTX2VideoDownsampler3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - stride: int | tuple[int, int, int] = 1, - spatial_padding_mode: str = "zeros", - ) -> None: - super().__init__() - - self.stride = stride if isinstance(stride, tuple) else (stride, stride, stride) - self.group_size = (in_channels * stride[0] * stride[1] * stride[2]) // out_channels - - out_channels = out_channels // (self.stride[0] * self.stride[1] * self.stride[2]) - - self.conv = LTX2VideoCausalConv3d( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=3, - stride=1, - spatial_padding_mode=spatial_padding_mode, - ) - - def forward(self, hidden_states: torch.Tensor, causal: bool = True) -> torch.Tensor: - hidden_states = torch.cat([hidden_states[:, :, : self.stride[0] - 1], hidden_states], dim=2) - - residual = ( - hidden_states.unflatten(4, (-1, self.stride[2])) - .unflatten(3, (-1, self.stride[1])) - .unflatten(2, (-1, self.stride[0])) - ) - residual = residual.permute(0, 1, 3, 5, 7, 2, 4, 6).flatten(1, 4) - residual = residual.unflatten(1, (-1, self.group_size)) - residual = residual.mean(dim=2) - - hidden_states = self.conv(hidden_states, causal=causal) - hidden_states = ( - hidden_states.unflatten(4, (-1, self.stride[2])) - .unflatten(3, (-1, self.stride[1])) - .unflatten(2, (-1, self.stride[0])) - ) - hidden_states = hidden_states.permute(0, 1, 3, 5, 7, 2, 4, 6).flatten(1, 4) - hidden_states = hidden_states + residual - - return hidden_states - - -# Like LTX 1.0 LTXVideoUpsampler3d, but uses new causal Conv3d -class LTX2VideoUpsampler3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - stride: int | tuple[int, int, int] = 1, - residual: bool = False, - upscale_factor: int = 1, - spatial_padding_mode: str = "zeros", - ) -> None: - super().__init__() - - self.stride = stride if isinstance(stride, tuple) else (stride, stride, stride) - self.residual = residual - self.upscale_factor = upscale_factor - - out_channels = out_channels or in_channels - out_channels = (out_channels * stride[0] * stride[1] * stride[2]) // upscale_factor - - self.conv = LTX2VideoCausalConv3d( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=3, - stride=1, - spatial_padding_mode=spatial_padding_mode, - ) - - def forward(self, hidden_states: torch.Tensor, causal: bool = True) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - - if self.residual: - residual = hidden_states.reshape( - batch_size, -1, self.stride[0], self.stride[1], self.stride[2], num_frames, height, width - ) - residual = residual.permute(0, 1, 5, 2, 6, 3, 7, 4).flatten(6, 7).flatten(4, 5).flatten(2, 3) - repeats = (self.stride[0] * self.stride[1] * self.stride[2]) // self.upscale_factor - residual = residual.repeat(1, repeats, 1, 1, 1) - residual = residual[:, :, self.stride[0] - 1 :] - - hidden_states = self.conv(hidden_states, causal=causal) - hidden_states = hidden_states.reshape( - batch_size, -1, self.stride[0], self.stride[1], self.stride[2], num_frames, height, width - ) - hidden_states = hidden_states.permute(0, 1, 5, 2, 6, 3, 7, 4).flatten(6, 7).flatten(4, 5).flatten(2, 3) - hidden_states = hidden_states[:, :, self.stride[0] - 1 :] - - if self.residual: - hidden_states = hidden_states + residual - - return hidden_states - - -# Like LTX 1.0 LTXVideo095DownBlock3D, but with the updated LTX2VideoResnetBlock3d -class LTX2VideoDownBlock3D(nn.Module): - r""" - Down block used in the LTXVideo model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - num_layers (`int`, defaults to `1`): - Number of resnet layers. - dropout (`float`, defaults to `0.0`): - Dropout rate. - resnet_eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - resnet_act_fn (`str`, defaults to `"swish"`): - Activation function to use. - spatio_temporal_scale (`bool`, defaults to `True`): - Whether or not to use a downsampling layer. If not used, output dimension would be same as input dimension. - Whether or not to downsample across temporal dimension. - is_causal (`bool`, defaults to `True`): - Whether this layer behaves causally (future frames depend only on past frames) or not. - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - num_layers: int = 1, - dropout: float = 0.0, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - spatio_temporal_scale: bool = True, - downsample_type: str = "conv", - spatial_padding_mode: str = "zeros", - ): - super().__init__() - - out_channels = out_channels or in_channels - - resnets = [] - for _ in range(num_layers): - resnets.append( - LTX2VideoResnetBlock3d( - in_channels=in_channels, - out_channels=in_channels, - dropout=dropout, - eps=resnet_eps, - non_linearity=resnet_act_fn, - spatial_padding_mode=spatial_padding_mode, - ) - ) - self.resnets = nn.ModuleList(resnets) - - self.downsamplers = None - if spatio_temporal_scale: - self.downsamplers = nn.ModuleList() - - if downsample_type == "conv": - self.downsamplers.append( - LTX2VideoCausalConv3d( - in_channels=in_channels, - out_channels=in_channels, - kernel_size=3, - stride=(2, 2, 2), - spatial_padding_mode=spatial_padding_mode, - ) - ) - elif downsample_type == "spatial": - self.downsamplers.append( - LTX2VideoDownsampler3d( - in_channels=in_channels, - out_channels=out_channels, - stride=(1, 2, 2), - spatial_padding_mode=spatial_padding_mode, - ) - ) - elif downsample_type == "temporal": - self.downsamplers.append( - LTX2VideoDownsampler3d( - in_channels=in_channels, - out_channels=out_channels, - stride=(2, 1, 1), - spatial_padding_mode=spatial_padding_mode, - ) - ) - elif downsample_type == "spatiotemporal": - self.downsamplers.append( - LTX2VideoDownsampler3d( - in_channels=in_channels, - out_channels=out_channels, - stride=(2, 2, 2), - spatial_padding_mode=spatial_padding_mode, - ) - ) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - generator: torch.Generator | None = None, - causal: bool = True, - ) -> torch.Tensor: - r"""Forward method of the `LTXDownBlock3D` class.""" - - for i, resnet in enumerate(self.resnets): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, generator, causal) - else: - hidden_states = resnet(hidden_states, temb, generator, causal=causal) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states, causal=causal) - - return hidden_states - - -# Adapted from diffusers.models.autoencoders.autoencoder_kl_cogvideox.CogVideoMidBlock3d -# Like LTX 1.0 LTXVideoMidBlock3d, but with the updated LTX2VideoResnetBlock3d -class LTX2VideoMidBlock3d(nn.Module): - r""" - A middle block used in the LTXVideo model. - - Args: - in_channels (`int`): - Number of input channels. - num_layers (`int`, defaults to `1`): - Number of resnet layers. - dropout (`float`, defaults to `0.0`): - Dropout rate. - resnet_eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - resnet_act_fn (`str`, defaults to `"swish"`): - Activation function to use. - is_causal (`bool`, defaults to `True`): - Whether this layer behaves causally (future frames depend only on past frames) or not. - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int, - num_layers: int = 1, - dropout: float = 0.0, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - inject_noise: bool = False, - timestep_conditioning: bool = False, - spatial_padding_mode: str = "zeros", - ) -> None: - super().__init__() - - self.time_embedder = None - if timestep_conditioning: - self.time_embedder = PixArtAlphaCombinedTimestepSizeEmbeddings(in_channels * 4, 0) - - resnets = [] - for _ in range(num_layers): - resnets.append( - LTX2VideoResnetBlock3d( - in_channels=in_channels, - out_channels=in_channels, - dropout=dropout, - eps=resnet_eps, - non_linearity=resnet_act_fn, - inject_noise=inject_noise, - timestep_conditioning=timestep_conditioning, - spatial_padding_mode=spatial_padding_mode, - ) - ) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - generator: torch.Generator | None = None, - causal: bool = True, - ) -> torch.Tensor: - r"""Forward method of the `LTXMidBlock3D` class.""" - - if self.time_embedder is not None: - temb = self.time_embedder( - timestep=temb.flatten(), - resolution=None, - aspect_ratio=None, - batch_size=hidden_states.size(0), - hidden_dtype=hidden_states.dtype, - ) - temb = temb.view(hidden_states.size(0), -1, 1, 1, 1) - - for i, resnet in enumerate(self.resnets): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, generator, causal) - else: - hidden_states = resnet(hidden_states, temb, generator, causal=causal) - - return hidden_states - - -# Like LTXVideoUpBlock3d but with no conv_in and the updated LTX2VideoResnetBlock3d -class LTX2VideoUpBlock3d(nn.Module): - r""" - Up block used in the LTXVideo model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - num_layers (`int`, defaults to `1`): - Number of resnet layers. - dropout (`float`, defaults to `0.0`): - Dropout rate. - resnet_eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - resnet_act_fn (`str`, defaults to `"swish"`): - Activation function to use. - spatio_temporal_scale (`bool`, defaults to `True`): - Whether or not to use a downsampling layer. If not used, output dimension would be same as input dimension. - Whether or not to downsample across temporal dimension. - is_causal (`bool`, defaults to `True`): - Whether this layer behaves causally (future frames depend only on past frames) or not. - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - num_layers: int = 1, - dropout: float = 0.0, - resnet_eps: float = 1e-6, - resnet_act_fn: str = "swish", - spatio_temporal_scale: bool = True, - upsample_type: str = "spatiotemporal", - inject_noise: bool = False, - timestep_conditioning: bool = False, - upsample_residual: bool = False, - upscale_factor: int = 1, - spatial_padding_mode: str = "zeros", - ): - super().__init__() - - out_channels = out_channels or in_channels - - self.time_embedder = None - if timestep_conditioning: - self.time_embedder = PixArtAlphaCombinedTimestepSizeEmbeddings(in_channels * 4, 0) - - self.conv_in = None - if in_channels != out_channels: - self.conv_in = LTX2VideoResnetBlock3d( - in_channels=in_channels, - out_channels=out_channels, - dropout=dropout, - eps=resnet_eps, - non_linearity=resnet_act_fn, - inject_noise=inject_noise, - timestep_conditioning=timestep_conditioning, - spatial_padding_mode=spatial_padding_mode, - ) - - self.upsamplers = None - if spatio_temporal_scale: - self.upsamplers = nn.ModuleList() - - if upsample_type == "spatial": - upsample_stride = (1, 2, 2) - elif upsample_type == "temporal": - upsample_stride = (2, 1, 1) - elif upsample_type == "spatiotemporal": - upsample_stride = (2, 2, 2) - - self.upsamplers.append( - LTX2VideoUpsampler3d( - in_channels=out_channels * upscale_factor, - stride=upsample_stride, - residual=upsample_residual, - upscale_factor=upscale_factor, - spatial_padding_mode=spatial_padding_mode, - ) - ) - - resnets = [] - for _ in range(num_layers): - resnets.append( - LTX2VideoResnetBlock3d( - in_channels=out_channels, - out_channels=out_channels, - dropout=dropout, - eps=resnet_eps, - non_linearity=resnet_act_fn, - inject_noise=inject_noise, - timestep_conditioning=timestep_conditioning, - spatial_padding_mode=spatial_padding_mode, - ) - ) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - generator: torch.Generator | None = None, - causal: bool = True, - ) -> torch.Tensor: - if self.conv_in is not None: - hidden_states = self.conv_in(hidden_states, temb, generator, causal=causal) - - if self.time_embedder is not None: - temb = self.time_embedder( - timestep=temb.flatten(), - resolution=None, - aspect_ratio=None, - batch_size=hidden_states.size(0), - hidden_dtype=hidden_states.dtype, - ) - temb = temb.view(hidden_states.size(0), -1, 1, 1, 1) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, causal=causal) - - for i, resnet in enumerate(self.resnets): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, generator, causal) - else: - hidden_states = resnet(hidden_states, temb, generator, causal=causal) - - return hidden_states - - -# Like LTX 1.0 LTXVideoEncoder3d but with different default args - the spatiotemporal downsampling pattern is -# different, as is the layers_per_block (the 2.0 VAE is bigger) -class LTX2VideoEncoder3d(nn.Module): - r""" - The `LTXVideoEncoder3d` layer of a variational autoencoder that encodes input video samples to its latent - representation. - - Args: - in_channels (`int`, defaults to 3): - Number of input channels. - out_channels (`int`, defaults to 128): - Number of latent channels. - block_out_channels (`tuple[int, ...]`, defaults to `(256, 512, 1024, 2048)`): - The number of output channels for each block. - spatio_temporal_scaling (`tuple[bool, ...], defaults to `(True, True, True, True)`: - Whether a block should contain spatio-temporal downscaling layers or not. - layers_per_block (`tuple[int, ...]`, defaults to `(4, 6, 6, 2, 2)`): - The number of layers per block. - downsample_type (`tuple[str, ...]`, defaults to `("spatial", "temporal", "spatiotemporal", "spatiotemporal")`): - The spatiotemporal downsampling pattern per block. Per-layer values can be - - `"spatial"` (downsample spatial dims by 2x) - - `"temporal"` (downsample temporal dim by 2x) - - `"spatiotemporal"` (downsample both spatial and temporal dims by 2x) - patch_size (`int`, defaults to `4`): - The size of spatial patches. - patch_size_t (`int`, defaults to `1`): - The size of temporal patches. - resnet_norm_eps (`float`, defaults to `1e-6`): - Epsilon value for ResNet normalization layers. - is_causal (`bool`, defaults to `True`): - Whether this layer behaves causally (future frames depend only on past frames) or not. - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 128, - block_out_channels: tuple[int, ...] = (256, 512, 1024, 2048), - down_block_types: tuple[str, ...] = ( - "LTX2VideoDownBlock3D", - "LTX2VideoDownBlock3D", - "LTX2VideoDownBlock3D", - "LTX2VideoDownBlock3D", - ), - spatio_temporal_scaling: bool | tuple[bool, ...] = (True, True, True, True), - layers_per_block: tuple[int, ...] = (4, 6, 6, 2, 2), - downsample_type: tuple[str, ...] = ("spatial", "temporal", "spatiotemporal", "spatiotemporal"), - patch_size: int = 4, - patch_size_t: int = 1, - resnet_norm_eps: float = 1e-6, - is_causal: bool = True, - spatial_padding_mode: str = "zeros", - ): - super().__init__() - num_encoder_blocks = len(layers_per_block) - if isinstance(spatio_temporal_scaling, bool): - spatio_temporal_scaling = (spatio_temporal_scaling,) * (num_encoder_blocks - 1) - - self.patch_size = patch_size - self.patch_size_t = patch_size_t - self.in_channels = in_channels * patch_size**2 - self.is_causal = is_causal - - output_channel = out_channels - - self.conv_in = LTX2VideoCausalConv3d( - in_channels=self.in_channels, - out_channels=output_channel, - kernel_size=3, - stride=1, - spatial_padding_mode=spatial_padding_mode, - ) - - # down blocks - num_block_out_channels = len(block_out_channels) - self.down_blocks = nn.ModuleList([]) - for i in range(num_block_out_channels): - input_channel = output_channel - output_channel = block_out_channels[i] - - if down_block_types[i] == "LTX2VideoDownBlock3D": - down_block = LTX2VideoDownBlock3D( - in_channels=input_channel, - out_channels=output_channel, - num_layers=layers_per_block[i], - resnet_eps=resnet_norm_eps, - spatio_temporal_scale=spatio_temporal_scaling[i], - downsample_type=downsample_type[i], - spatial_padding_mode=spatial_padding_mode, - ) - else: - raise ValueError(f"Unknown down block type: {down_block_types[i]}") - - self.down_blocks.append(down_block) - - # mid block - self.mid_block = LTX2VideoMidBlock3d( - in_channels=output_channel, - num_layers=layers_per_block[-1], - resnet_eps=resnet_norm_eps, - spatial_padding_mode=spatial_padding_mode, - ) - - # out - self.norm_out = PerChannelRMSNorm() - self.conv_act = nn.SiLU() - self.conv_out = LTX2VideoCausalConv3d( - in_channels=output_channel, - out_channels=out_channels + 1, - kernel_size=3, - stride=1, - spatial_padding_mode=spatial_padding_mode, - ) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor, causal: bool | None = None) -> torch.Tensor: - r"""The forward method of the `LTXVideoEncoder3d` class.""" - - p = self.patch_size - p_t = self.patch_size_t - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p - post_patch_width = width // p - causal = causal or self.is_causal - - hidden_states = hidden_states.reshape( - batch_size, num_channels, post_patch_num_frames, p_t, post_patch_height, p, post_patch_width, p - ) - # Thanks for driving me insane with the weird patching order :( - hidden_states = hidden_states.permute(0, 1, 3, 7, 5, 2, 4, 6).flatten(1, 4) - hidden_states = self.conv_in(hidden_states, causal=causal) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for down_block in self.down_blocks: - hidden_states = self._gradient_checkpointing_func(down_block, hidden_states, None, None, causal) - - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states, None, None, causal) - else: - for down_block in self.down_blocks: - hidden_states = down_block(hidden_states, causal=causal) - - hidden_states = self.mid_block(hidden_states, causal=causal) - - hidden_states = self.norm_out(hidden_states) - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states, causal=causal) - - last_channel = hidden_states[:, -1:] - last_channel = last_channel.repeat(1, hidden_states.size(1) - 2, 1, 1, 1) - hidden_states = torch.cat([hidden_states, last_channel], dim=1) - - return hidden_states - - -# Like LTX 1.0 LTXVideoDecoder3d, but has only 3 symmetric up blocks which are causal and residual with upsample_factor 2 -class LTX2VideoDecoder3d(nn.Module): - r""" - The `LTXVideoDecoder3d` layer of a variational autoencoder that decodes its latent representation into an output - sample. - - Args: - in_channels (`int`, defaults to 128): - Number of latent channels. - out_channels (`int`, defaults to 3): - Number of output channels. - block_out_channels (`tuple[int, ...]`, defaults to `(128, 256, 512, 512)`): - The number of output channels for each block. - spatio_temporal_scaling (`tuple[bool, ...], defaults to `(True, True, True, False)`: - Whether a block should contain spatio-temporal upscaling layers or not. - layers_per_block (`tuple[int, ...]`, defaults to `(4, 3, 3, 3, 4)`): - The number of layers per block. - patch_size (`int`, defaults to `4`): - The size of spatial patches. - patch_size_t (`int`, defaults to `1`): - The size of temporal patches. - resnet_norm_eps (`float`, defaults to `1e-6`): - Epsilon value for ResNet normalization layers. - is_causal (`bool`, defaults to `False`): - Whether this layer behaves causally (future frames depend only on past frames) or not. - timestep_conditioning (`bool`, defaults to `False`): - Whether to condition the model on timesteps. - """ - - def __init__( - self, - in_channels: int = 128, - out_channels: int = 3, - block_out_channels: tuple[int, ...] = (256, 512, 1024), - spatio_temporal_scaling: bool | tuple[bool, ...] = (True, True, True), - layers_per_block: tuple[int, ...] = (5, 5, 5, 5), - upsample_type: tuple[str, ...] = ("spatiotemporal", "spatiotemporal", "spatiotemporal"), - patch_size: int = 4, - patch_size_t: int = 1, - resnet_norm_eps: float = 1e-6, - is_causal: bool = False, - inject_noise: bool | tuple[bool, ...] = (False, False, False), - timestep_conditioning: bool = False, - upsample_residual: bool | tuple[bool, ...] = (True, True, True), - upsample_factor: tuple[bool, ...] = (2, 2, 2), - spatial_padding_mode: str = "reflect", - ) -> None: - super().__init__() - num_decoder_blocks = len(layers_per_block) - if isinstance(spatio_temporal_scaling, bool): - spatio_temporal_scaling = (spatio_temporal_scaling,) * (num_decoder_blocks - 1) - if isinstance(inject_noise, bool): - inject_noise = (inject_noise,) * num_decoder_blocks - if isinstance(upsample_residual, bool): - upsample_residual = (upsample_residual,) * (num_decoder_blocks - 1) - - self.patch_size = patch_size - self.patch_size_t = patch_size_t - self.out_channels = out_channels * patch_size**2 - self.is_causal = is_causal - - block_out_channels = tuple(reversed(block_out_channels)) - spatio_temporal_scaling = tuple(reversed(spatio_temporal_scaling)) - layers_per_block = tuple(reversed(layers_per_block)) - inject_noise = tuple(reversed(inject_noise)) - upsample_residual = tuple(reversed(upsample_residual)) - upsample_factor = tuple(reversed(upsample_factor)) - output_channel = block_out_channels[0] - - self.conv_in = LTX2VideoCausalConv3d( - in_channels=in_channels, - out_channels=output_channel, - kernel_size=3, - stride=1, - spatial_padding_mode=spatial_padding_mode, - ) - - self.mid_block = LTX2VideoMidBlock3d( - in_channels=output_channel, - num_layers=layers_per_block[0], - resnet_eps=resnet_norm_eps, - inject_noise=inject_noise[0], - timestep_conditioning=timestep_conditioning, - spatial_padding_mode=spatial_padding_mode, - ) - - # up blocks - num_block_out_channels = len(block_out_channels) - self.up_blocks = nn.ModuleList([]) - for i in range(num_block_out_channels): - input_channel = output_channel // upsample_factor[i] - output_channel = block_out_channels[i] // upsample_factor[i] - - up_block = LTX2VideoUpBlock3d( - in_channels=input_channel, - out_channels=output_channel, - num_layers=layers_per_block[i + 1], - resnet_eps=resnet_norm_eps, - spatio_temporal_scale=spatio_temporal_scaling[i], - upsample_type=upsample_type[i], - inject_noise=inject_noise[i + 1], - timestep_conditioning=timestep_conditioning, - upsample_residual=upsample_residual[i], - upscale_factor=upsample_factor[i], - spatial_padding_mode=spatial_padding_mode, - ) - - self.up_blocks.append(up_block) - - # out - self.norm_out = PerChannelRMSNorm() - self.conv_act = nn.SiLU() - self.conv_out = LTX2VideoCausalConv3d( - in_channels=output_channel, - out_channels=self.out_channels, - kernel_size=3, - stride=1, - spatial_padding_mode=spatial_padding_mode, - ) - - # timestep embedding - self.time_embedder = None - self.scale_shift_table = None - self.timestep_scale_multiplier = None - if timestep_conditioning: - self.timestep_scale_multiplier = nn.Parameter(torch.tensor(1000.0, dtype=torch.float32)) - self.time_embedder = PixArtAlphaCombinedTimestepSizeEmbeddings(output_channel * 2, 0) - self.scale_shift_table = nn.Parameter(torch.randn(2, output_channel) / output_channel**0.5) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - causal: bool | None = None, - ) -> torch.Tensor: - causal = causal or self.is_causal - - hidden_states = self.conv_in(hidden_states, causal=causal) - - if self.timestep_scale_multiplier is not None: - temb = temb * self.timestep_scale_multiplier - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states, temb, None, causal) - - for up_block in self.up_blocks: - hidden_states = self._gradient_checkpointing_func(up_block, hidden_states, temb, None, causal) - else: - hidden_states = self.mid_block(hidden_states, temb, causal=causal) - - for up_block in self.up_blocks: - hidden_states = up_block(hidden_states, temb, causal=causal) - - hidden_states = self.norm_out(hidden_states) - - if self.time_embedder is not None: - temb = self.time_embedder( - timestep=temb.flatten(), - resolution=None, - aspect_ratio=None, - batch_size=hidden_states.size(0), - hidden_dtype=hidden_states.dtype, - ) - temb = temb.view(hidden_states.size(0), -1, 1, 1, 1).unflatten(1, (2, -1)) - temb = temb + self.scale_shift_table[None, ..., None, None, None] - shift, scale = temb.unbind(dim=1) - hidden_states = hidden_states * (1 + scale) + shift - - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states, causal=causal) - - p = self.patch_size - p_t = self.patch_size_t - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - hidden_states = hidden_states.reshape(batch_size, -1, p_t, p, p, num_frames, height, width) - hidden_states = hidden_states.permute(0, 1, 5, 2, 6, 4, 7, 3).flatten(6, 7).flatten(4, 5).flatten(2, 3) - - return hidden_states - - -class AutoencoderKLLTX2Video(ModelMixin, AutoencoderMixin, ConfigMixin, FromOriginalModelMixin): - r""" - A VAE model with KL loss for encoding images into latents and decoding latent representations into images. Used in - [LTX-2](https://huggingface.co/Lightricks/LTX-2). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Args: - in_channels (`int`, defaults to `3`): - Number of input channels. - out_channels (`int`, defaults to `3`): - Number of output channels. - latent_channels (`int`, defaults to `128`): - Number of latent channels. - block_out_channels (`tuple[int, ...]`, defaults to `(128, 256, 512, 512)`): - The number of output channels for each block. - spatio_temporal_scaling (`tuple[bool, ...], defaults to `(True, True, True, False)`: - Whether a block should contain spatio-temporal downscaling or not. - layers_per_block (`tuple[int, ...]`, defaults to `(4, 3, 3, 3, 4)`): - The number of layers per block. - patch_size (`int`, defaults to `4`): - The size of spatial patches. - patch_size_t (`int`, defaults to `1`): - The size of temporal patches. - resnet_norm_eps (`float`, defaults to `1e-6`): - Epsilon value for ResNet normalization layers. - scaling_factor (`float`, *optional*, defaults to `1.0`): - The component-wise standard deviation of the trained latent space computed using the first batch of the - training set. This is used to scale the latent space to have unit variance when training the diffusion - model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the - diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 - / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image - Synthesis with Latent Diffusion Models](https://huggingface.co/papers/2112.10752) paper. - encoder_causal (`bool`, defaults to `True`): - Whether the encoder should behave causally (future frames depend only on past frames) or not. - decoder_causal (`bool`, defaults to `False`): - Whether the decoder should behave causally (future frames depend only on past frames) or not. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - latent_channels: int = 128, - block_out_channels: tuple[int, ...] = (256, 512, 1024, 2048), - down_block_types: tuple[str, ...] = ( - "LTX2VideoDownBlock3D", - "LTX2VideoDownBlock3D", - "LTX2VideoDownBlock3D", - "LTX2VideoDownBlock3D", - ), - decoder_block_out_channels: tuple[int, ...] = (256, 512, 1024), - layers_per_block: tuple[int, ...] = (4, 6, 6, 2, 2), - decoder_layers_per_block: tuple[int, ...] = (5, 5, 5, 5), - spatio_temporal_scaling: bool | tuple[bool, ...] = (True, True, True, True), - decoder_spatio_temporal_scaling: bool | tuple[bool, ...] = (True, True, True), - decoder_inject_noise: bool | tuple[bool, ...] = (False, False, False, False), - downsample_type: tuple[str, ...] = ("spatial", "temporal", "spatiotemporal", "spatiotemporal"), - upsample_type: tuple[str, ...] = ("spatiotemporal", "spatiotemporal", "spatiotemporal"), - upsample_residual: bool | tuple[bool, ...] = (True, True, True), - upsample_factor: tuple[int, ...] = (2, 2, 2), - timestep_conditioning: bool = False, - patch_size: int = 4, - patch_size_t: int = 1, - resnet_norm_eps: float = 1e-6, - scaling_factor: float = 1.0, - encoder_causal: bool = True, - decoder_causal: bool = True, - encoder_spatial_padding_mode: str = "zeros", - decoder_spatial_padding_mode: str = "reflect", - spatial_compression_ratio: int = None, - temporal_compression_ratio: int = None, - ) -> None: - super().__init__() - num_encoder_blocks = len(layers_per_block) - num_decoder_blocks = len(decoder_layers_per_block) - if isinstance(spatio_temporal_scaling, bool): - spatio_temporal_scaling = (spatio_temporal_scaling,) * (num_encoder_blocks - 1) - if isinstance(decoder_spatio_temporal_scaling, bool): - decoder_spatio_temporal_scaling = (decoder_spatio_temporal_scaling,) * (num_decoder_blocks - 1) - if isinstance(decoder_inject_noise, bool): - decoder_inject_noise = (decoder_inject_noise,) * num_decoder_blocks - if isinstance(upsample_residual, bool): - upsample_residual = (upsample_residual,) * (num_decoder_blocks - 1) - - self.encoder = LTX2VideoEncoder3d( - in_channels=in_channels, - out_channels=latent_channels, - block_out_channels=block_out_channels, - down_block_types=down_block_types, - spatio_temporal_scaling=spatio_temporal_scaling, - layers_per_block=layers_per_block, - downsample_type=downsample_type, - patch_size=patch_size, - patch_size_t=patch_size_t, - resnet_norm_eps=resnet_norm_eps, - is_causal=encoder_causal, - spatial_padding_mode=encoder_spatial_padding_mode, - ) - self.decoder = LTX2VideoDecoder3d( - in_channels=latent_channels, - out_channels=out_channels, - block_out_channels=decoder_block_out_channels, - spatio_temporal_scaling=decoder_spatio_temporal_scaling, - layers_per_block=decoder_layers_per_block, - upsample_type=upsample_type, - patch_size=patch_size, - patch_size_t=patch_size_t, - resnet_norm_eps=resnet_norm_eps, - is_causal=decoder_causal, - timestep_conditioning=timestep_conditioning, - inject_noise=decoder_inject_noise, - upsample_residual=upsample_residual, - upsample_factor=upsample_factor, - spatial_padding_mode=decoder_spatial_padding_mode, - ) - - latents_mean = torch.zeros((latent_channels,), requires_grad=False) - latents_std = torch.ones((latent_channels,), requires_grad=False) - self.register_buffer("latents_mean", latents_mean, persistent=True) - self.register_buffer("latents_std", latents_std, persistent=True) - - self.spatial_compression_ratio = ( - patch_size * 2 ** sum(spatio_temporal_scaling) - if spatial_compression_ratio is None - else spatial_compression_ratio - ) - self.temporal_compression_ratio = ( - patch_size_t * 2 ** sum(spatio_temporal_scaling) - if temporal_compression_ratio is None - else temporal_compression_ratio - ) - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent - # frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the - # intermediate tiles together, the memory requirement can be lowered. - self.use_tiling = False - - # When decoding temporally long video latents, the memory requirement is very high. By decoding latent frames - # at a fixed frame batch size (based on `self.num_latent_frames_batch_sizes`), the memory requirement can be lowered. - self.use_framewise_encoding = False - self.use_framewise_decoding = False - - # This can be configured based on the amount of GPU memory available. - # `16` for sample frames and `2` for latent frames are sensible defaults for consumer GPUs. - # Setting it to higher values results in higher memory usage. - self.num_sample_frames_batch_size = 16 - self.num_latent_frames_batch_size = 2 - - # The minimal tile height and width for spatial tiling to be used - self.tile_sample_min_height = 512 - self.tile_sample_min_width = 512 - self.tile_sample_min_num_frames = 16 - - # The minimal distance between two spatial tiles - self.tile_sample_stride_height = 448 - self.tile_sample_stride_width = 448 - self.tile_sample_stride_num_frames = 8 - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_sample_min_num_frames: int | None = None, - tile_sample_stride_height: float | None = None, - tile_sample_stride_width: float | None = None, - tile_sample_stride_num_frames: float | None = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_sample_stride_height (`int`, *optional*): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. - tile_sample_stride_width (`int`, *optional*): - The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling - artifacts produced across the width dimension. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_sample_min_num_frames = tile_sample_min_num_frames or self.tile_sample_min_num_frames - self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height - self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width - self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames - - def _encode(self, x: torch.Tensor, causal: bool | None = None) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = x.shape - - if self.use_framewise_decoding and num_frames > self.tile_sample_min_num_frames: - return self._temporal_tiled_encode(x, causal=causal) - - if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): - return self.tiled_encode(x, causal=causal) - - enc = self.encoder(x, causal=causal) - - return enc - - @apply_forward_hook - def encode( - self, x: torch.Tensor, causal: bool | None = None, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice, causal=causal) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x, causal=causal) - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode( - self, - z: torch.Tensor, - temb: torch.Tensor | None = None, - causal: bool | None = None, - return_dict: bool = True, - ) -> DecoderOutput | torch.Tensor: - batch_size, num_channels, num_frames, height, width = z.shape - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio - - if self.use_framewise_decoding and num_frames > tile_latent_min_num_frames: - return self._temporal_tiled_decode(z, temb, causal=causal, return_dict=return_dict) - - if self.use_tiling and (width > tile_latent_min_width or height > tile_latent_min_height): - return self.tiled_decode(z, temb, causal=causal, return_dict=return_dict) - - dec = self.decoder(z, temb, causal=causal) - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode( - self, - z: torch.Tensor, - temb: torch.Tensor | None = None, - causal: bool | None = None, - return_dict: bool = True, - ) -> DecoderOutput | torch.Tensor: - """ - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - if temb is not None: - decoded_slices = [ - self._decode(z_slice, t_slice, causal=causal).sample - for z_slice, t_slice in (z.split(1), temb.split(1)) - ] - else: - decoded_slices = [self._decode(z_slice, causal=causal).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z, temb, causal=causal).sample - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[3], b.shape[3], blend_extent) - for y in range(blend_extent): - b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[4], b.shape[4], blend_extent) - for x in range(blend_extent): - b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( - x / blend_extent - ) - return b - - def blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-3], b.shape[-3], blend_extent) - for x in range(blend_extent): - b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (1 - x / blend_extent) + b[:, :, x, :, :] * ( - x / blend_extent - ) - return b - - def tiled_encode(self, x: torch.Tensor, causal: bool | None = None) -> torch.Tensor: - r"""Encode a batch of images using a tiled encoder. - - Args: - x (`torch.Tensor`): Input batch of videos. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - batch_size, num_channels, num_frames, height, width = x.shape - latent_height = height // self.spatial_compression_ratio - latent_width = width // self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = tile_latent_min_height - tile_latent_stride_height - blend_width = tile_latent_min_width - tile_latent_stride_width - - # Split x into overlapping tiles and encode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, self.tile_sample_stride_height): - row = [] - for j in range(0, width, self.tile_sample_stride_width): - time = self.encoder( - x[:, :, :, i : i + self.tile_sample_min_height, j : j + self.tile_sample_min_width], - causal=causal, - ) - - row.append(time) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :tile_latent_stride_height, :tile_latent_stride_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - enc = torch.cat(result_rows, dim=3)[:, :, :, :latent_height, :latent_width] - return enc - - def tiled_decode( - self, z: torch.Tensor, temb: torch.Tensor | None, causal: bool | None = None, return_dict: bool = True - ) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - - batch_size, num_channels, num_frames, height, width = z.shape - sample_height = height * self.spatial_compression_ratio - sample_width = width * self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = self.tile_sample_min_height - self.tile_sample_stride_height - blend_width = self.tile_sample_min_width - self.tile_sample_stride_width - - # Split z into overlapping tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, tile_latent_stride_height): - row = [] - for j in range(0, width, tile_latent_stride_width): - time = self.decoder( - z[:, :, :, i : i + tile_latent_min_height, j : j + tile_latent_min_width], temb, causal=causal - ) - - row.append(time) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, : self.tile_sample_stride_height, : self.tile_sample_stride_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - dec = torch.cat(result_rows, dim=3)[:, :, :, :sample_height, :sample_width] - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - def _temporal_tiled_encode(self, x: torch.Tensor, causal: bool | None = None) -> AutoencoderKLOutput: - batch_size, num_channels, num_frames, height, width = x.shape - latent_num_frames = (num_frames - 1) // self.temporal_compression_ratio + 1 - - tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio - tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio - blend_num_frames = tile_latent_min_num_frames - tile_latent_stride_num_frames - - row = [] - for i in range(0, num_frames, self.tile_sample_stride_num_frames): - tile = x[:, :, i : i + self.tile_sample_min_num_frames + 1, :, :] - if self.use_tiling and (height > self.tile_sample_min_height or width > self.tile_sample_min_width): - tile = self.tiled_encode(tile, causal=causal) - else: - tile = self.encoder(tile, causal=causal) - if i > 0: - tile = tile[:, :, 1:, :, :] - row.append(tile) - - result_row = [] - for i, tile in enumerate(row): - if i > 0: - tile = self.blend_t(row[i - 1], tile, blend_num_frames) - result_row.append(tile[:, :, :tile_latent_stride_num_frames, :, :]) - else: - result_row.append(tile[:, :, : tile_latent_stride_num_frames + 1, :, :]) - - enc = torch.cat(result_row, dim=2)[:, :, :latent_num_frames] - return enc - - def _temporal_tiled_decode( - self, z: torch.Tensor, temb: torch.Tensor | None, causal: bool | None = None, return_dict: bool = True - ) -> DecoderOutput | torch.Tensor: - batch_size, num_channels, num_frames, height, width = z.shape - num_sample_frames = (num_frames - 1) * self.temporal_compression_ratio + 1 - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio - tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio - blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames - - row = [] - for i in range(0, num_frames, tile_latent_stride_num_frames): - tile = z[:, :, i : i + tile_latent_min_num_frames + 1, :, :] - if self.use_tiling and (tile.shape[-1] > tile_latent_min_width or tile.shape[-2] > tile_latent_min_height): - decoded = self.tiled_decode(tile, temb, causal=causal, return_dict=True).sample - else: - decoded = self.decoder(tile, temb, causal=causal) - if i > 0: - decoded = decoded[:, :, :-1, :, :] - row.append(decoded) - - result_row = [] - for i, tile in enumerate(row): - if i > 0: - tile = self.blend_t(row[i - 1], tile, blend_num_frames) - tile = tile[:, :, : self.tile_sample_stride_num_frames, :, :] - result_row.append(tile) - else: - result_row.append(tile[:, :, : self.tile_sample_stride_num_frames + 1, :, :]) - - dec = torch.cat(result_row, dim=2)[:, :, :num_sample_frames] - - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) - - def forward( - self, - sample: torch.Tensor, - temb: torch.Tensor | None = None, - sample_posterior: bool = False, - encoder_causal: bool | None = None, - decoder_causal: bool | None = None, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> torch.Tensor | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - temb (`torch.Tensor`, *optional*): - Optional timestep embedding tensor used to condition the decoder. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - encoder_causal (`bool`, *optional*): - Whether the encoder should use causal convolutions. If `None`, falls back to the model default. - decoder_causal (`bool`, *optional*): - Whether the decoder should use causal convolutions. If `None`, falls back to the model default. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x, causal=encoder_causal).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z, temb, causal=decoder_causal) - if not return_dict: - return (dec.sample,) - return dec diff --git a/diffusers/models/autoencoders/autoencoder_kl_ltx2_audio.py b/diffusers/models/autoencoders/autoencoder_kl_ltx2_audio.py deleted file mode 100644 index fb773dbdc01edfc3159be0b798ec80d2b1b885a2..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_ltx2_audio.py +++ /dev/null @@ -1,818 +0,0 @@ -# Copyright 2025 The Lightricks team and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils.accelerate_utils import apply_forward_hook -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -LATENT_DOWNSAMPLE_FACTOR = 4 - - -class LTX2AudioCausalConv2d(nn.Module): - """ - A causal 2D convolution that pads asymmetrically along the causal axis. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int], - stride: int = 1, - dilation: int | tuple[int, int] = 1, - groups: int = 1, - bias: bool = True, - causality_axis: str = "height", - ) -> None: - super().__init__() - - self.causality_axis = causality_axis - kernel_size = (kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size - dilation = (dilation, dilation) if isinstance(dilation, int) else dilation - - pad_h = (kernel_size[0] - 1) * dilation[0] - pad_w = (kernel_size[1] - 1) * dilation[1] - - if self.causality_axis == "none": - padding = (pad_w // 2, pad_w - pad_w // 2, pad_h // 2, pad_h - pad_h // 2) - elif self.causality_axis in {"width", "width-compatibility"}: - padding = (pad_w, 0, pad_h // 2, pad_h - pad_h // 2) - elif self.causality_axis == "height": - padding = (pad_w // 2, pad_w - pad_w // 2, pad_h, 0) - else: - raise ValueError(f"Invalid causality_axis: {causality_axis}") - - self.padding = padding - self.conv = nn.Conv2d( - in_channels, - out_channels, - kernel_size, - stride=stride, - padding=0, - dilation=dilation, - groups=groups, - bias=bias, - ) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = F.pad(x, self.padding) - return self.conv(x) - - -class LTX2AudioPixelNorm(nn.Module): - """ - Per-pixel (per-location) RMS normalization layer. - """ - - def __init__(self, dim: int = 1, eps: float = 1e-8) -> None: - super().__init__() - self.dim = dim - self.eps = eps - - def forward(self, x: torch.Tensor) -> torch.Tensor: - mean_sq = torch.mean(x**2, dim=self.dim, keepdim=True) - rms = torch.sqrt(mean_sq + self.eps) - return x / rms - - -class LTX2AudioAttnBlock(nn.Module): - def __init__( - self, - in_channels: int, - norm_type: str = "group", - ) -> None: - super().__init__() - self.in_channels = in_channels - - if norm_type == "group": - self.norm = nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) - elif norm_type == "pixel": - self.norm = LTX2AudioPixelNorm(dim=1, eps=1e-6) - else: - raise ValueError(f"Invalid normalization type: {norm_type}") - self.q = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.k = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.v = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.proj_out = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - h_ = self.norm(x) - q = self.q(h_) - k = self.k(h_) - v = self.v(h_) - - batch, channels, height, width = q.shape - q = q.reshape(batch, channels, height * width).permute(0, 2, 1).contiguous() - k = k.reshape(batch, channels, height * width).contiguous() - attn = torch.bmm(q, k) * (int(channels) ** (-0.5)) - attn = torch.nn.functional.softmax(attn, dim=2) - - v = v.reshape(batch, channels, height * width) - attn = attn.permute(0, 2, 1).contiguous() - h_ = torch.bmm(v, attn).reshape(batch, channels, height, width) - - h_ = self.proj_out(h_) - return x + h_ - - -class LTX2AudioResnetBlock(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - conv_shortcut: bool = False, - dropout: float = 0.0, - temb_channels: int = 512, - norm_type: str = "group", - causality_axis: str = "height", - ) -> None: - super().__init__() - self.causality_axis = causality_axis - - if self.causality_axis is not None and self.causality_axis != "none" and norm_type == "group": - raise ValueError("Causal ResnetBlock with GroupNorm is not supported.") - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - self.use_conv_shortcut = conv_shortcut - - if norm_type == "group": - self.norm1 = nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) - elif norm_type == "pixel": - self.norm1 = LTX2AudioPixelNorm(dim=1, eps=1e-6) - else: - raise ValueError(f"Invalid normalization type: {norm_type}") - self.non_linearity = nn.SiLU() - if causality_axis is not None: - self.conv1 = LTX2AudioCausalConv2d( - in_channels, out_channels, kernel_size=3, stride=1, causality_axis=causality_axis - ) - else: - self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) - if temb_channels > 0: - self.temb_proj = nn.Linear(temb_channels, out_channels) - if norm_type == "group": - self.norm2 = nn.GroupNorm(num_groups=32, num_channels=out_channels, eps=1e-6, affine=True) - elif norm_type == "pixel": - self.norm2 = LTX2AudioPixelNorm(dim=1, eps=1e-6) - else: - raise ValueError(f"Invalid normalization type: {norm_type}") - self.dropout = nn.Dropout(dropout) - if causality_axis is not None: - self.conv2 = LTX2AudioCausalConv2d( - out_channels, out_channels, kernel_size=3, stride=1, causality_axis=causality_axis - ) - else: - self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1) - if self.in_channels != self.out_channels: - if self.use_conv_shortcut: - if causality_axis is not None: - self.conv_shortcut = LTX2AudioCausalConv2d( - in_channels, out_channels, kernel_size=3, stride=1, causality_axis=causality_axis - ) - else: - self.conv_shortcut = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) - else: - if causality_axis is not None: - self.nin_shortcut = LTX2AudioCausalConv2d( - in_channels, out_channels, kernel_size=1, stride=1, causality_axis=causality_axis - ) - else: - self.nin_shortcut = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0) - - def forward(self, x: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - h = self.norm1(x) - h = self.non_linearity(h) - h = self.conv1(h) - - if temb is not None: - h = h + self.temb_proj(self.non_linearity(temb))[:, :, None, None] - - h = self.norm2(h) - h = self.non_linearity(h) - h = self.dropout(h) - h = self.conv2(h) - - if self.in_channels != self.out_channels: - x = self.conv_shortcut(x) if self.use_conv_shortcut else self.nin_shortcut(x) - - return x + h - - -class LTX2AudioDownsample(nn.Module): - def __init__(self, in_channels: int, with_conv: bool, causality_axis: str | None = "height") -> None: - super().__init__() - self.with_conv = with_conv - self.causality_axis = causality_axis - - if self.with_conv: - self.conv = torch.nn.Conv2d(in_channels, in_channels, kernel_size=3, stride=2, padding=0) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - if self.with_conv: - # Padding tuple is in the order: (left, right, top, bottom). - if self.causality_axis == "none": - pad = (0, 1, 0, 1) - elif self.causality_axis == "width": - pad = (2, 0, 0, 1) - elif self.causality_axis == "height": - pad = (0, 1, 2, 0) - elif self.causality_axis == "width-compatibility": - pad = (1, 0, 0, 1) - else: - raise ValueError( - f"Invalid `causality_axis` {self.causality_axis}; supported values are `none`, `width`, `height`," - f" and `width-compatibility`." - ) - - x = F.pad(x, pad, mode="constant", value=0) - x = self.conv(x) - else: - # with_conv=False implies that causality_axis is "none" - x = F.avg_pool2d(x, kernel_size=2, stride=2) - return x - - -class LTX2AudioUpsample(nn.Module): - def __init__(self, in_channels: int, with_conv: bool, causality_axis: str | None = "height") -> None: - super().__init__() - self.with_conv = with_conv - self.causality_axis = causality_axis - if self.with_conv: - if causality_axis is not None: - self.conv = LTX2AudioCausalConv2d( - in_channels, in_channels, kernel_size=3, stride=1, causality_axis=causality_axis - ) - else: - self.conv = nn.Conv2d(in_channels, in_channels, kernel_size=3, stride=1, padding=1) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest") - if self.with_conv: - x = self.conv(x) - if self.causality_axis is None or self.causality_axis == "none": - pass - elif self.causality_axis == "height": - x = x[:, :, 1:, :] - elif self.causality_axis == "width": - x = x[:, :, :, 1:] - elif self.causality_axis == "width-compatibility": - pass - else: - raise ValueError(f"Invalid causality_axis: {self.causality_axis}") - - return x - - -class LTX2AudioAudioPatchifier: - """ - Patchifier for spectrogram/audio latents. - """ - - def __init__( - self, - patch_size: int, - sample_rate: int = 16000, - hop_length: int = 160, - audio_latent_downsample_factor: int = 4, - is_causal: bool = True, - ): - self.hop_length = hop_length - self.sample_rate = sample_rate - self.audio_latent_downsample_factor = audio_latent_downsample_factor - self.is_causal = is_causal - self._patch_size = (1, patch_size, patch_size) - - def patchify(self, audio_latents: torch.Tensor) -> torch.Tensor: - batch, channels, time, freq = audio_latents.shape - return audio_latents.permute(0, 2, 1, 3).reshape(batch, time, channels * freq) - - def unpatchify(self, audio_latents: torch.Tensor, channels: int, mel_bins: int) -> torch.Tensor: - batch, time, _ = audio_latents.shape - return audio_latents.view(batch, time, channels, mel_bins).permute(0, 2, 1, 3) - - @property - def patch_size(self) -> tuple[int, int, int]: - return self._patch_size - - -class LTX2AudioEncoder(nn.Module): - def __init__( - self, - base_channels: int = 128, - output_channels: int = 1, - num_res_blocks: int = 2, - attn_resolutions: tuple[int, ...] | None = None, - in_channels: int = 2, - resolution: int = 256, - latent_channels: int = 8, - ch_mult: tuple[int, ...] = (1, 2, 4), - norm_type: str = "group", - causality_axis: str | None = "width", - dropout: float = 0.0, - mid_block_add_attention: bool = False, - sample_rate: int = 16000, - mel_hop_length: int = 160, - is_causal: bool = True, - mel_bins: int | None = 64, - double_z: bool = True, - ): - super().__init__() - - self.sample_rate = sample_rate - self.mel_hop_length = mel_hop_length - self.is_causal = is_causal - self.mel_bins = mel_bins - - self.base_channels = base_channels - self.temb_ch = 0 - self.num_resolutions = len(ch_mult) - self.num_res_blocks = num_res_blocks - self.resolution = resolution - self.in_channels = in_channels - self.out_ch = output_channels - self.give_pre_end = False - self.tanh_out = False - self.norm_type = norm_type - self.latent_channels = latent_channels - self.channel_multipliers = ch_mult - self.attn_resolutions = attn_resolutions - self.causality_axis = causality_axis - - base_block_channels = base_channels - base_resolution = resolution - self.z_shape = (1, latent_channels, base_resolution, base_resolution) - - if self.causality_axis is not None: - self.conv_in = LTX2AudioCausalConv2d( - in_channels, base_block_channels, kernel_size=3, stride=1, causality_axis=self.causality_axis - ) - else: - self.conv_in = nn.Conv2d(in_channels, base_block_channels, kernel_size=3, stride=1, padding=1) - - self.down = nn.ModuleList() - block_in = base_block_channels - curr_res = self.resolution - - for level in range(self.num_resolutions): - stage = nn.Module() - stage.block = nn.ModuleList() - stage.attn = nn.ModuleList() - block_out = self.base_channels * self.channel_multipliers[level] - - for _ in range(self.num_res_blocks): - stage.block.append( - LTX2AudioResnetBlock( - in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - norm_type=self.norm_type, - causality_axis=self.causality_axis, - ) - ) - block_in = block_out - if self.attn_resolutions: - if curr_res in self.attn_resolutions: - stage.attn.append(LTX2AudioAttnBlock(block_in, norm_type=self.norm_type)) - - if level != self.num_resolutions - 1: - stage.downsample = LTX2AudioDownsample(block_in, True, causality_axis=self.causality_axis) - curr_res = curr_res // 2 - - self.down.append(stage) - - self.mid = nn.Module() - self.mid.block_1 = LTX2AudioResnetBlock( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - norm_type=self.norm_type, - causality_axis=self.causality_axis, - ) - if mid_block_add_attention: - self.mid.attn_1 = LTX2AudioAttnBlock(block_in, norm_type=self.norm_type) - else: - self.mid.attn_1 = nn.Identity() - self.mid.block_2 = LTX2AudioResnetBlock( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - norm_type=self.norm_type, - causality_axis=self.causality_axis, - ) - - final_block_channels = block_in - z_channels = 2 * latent_channels if double_z else latent_channels - if self.norm_type == "group": - self.norm_out = nn.GroupNorm(num_groups=32, num_channels=final_block_channels, eps=1e-6, affine=True) - elif self.norm_type == "pixel": - self.norm_out = LTX2AudioPixelNorm(dim=1, eps=1e-6) - else: - raise ValueError(f"Invalid normalization type: {self.norm_type}") - self.non_linearity = nn.SiLU() - - if self.causality_axis is not None: - self.conv_out = LTX2AudioCausalConv2d( - final_block_channels, z_channels, kernel_size=3, stride=1, causality_axis=self.causality_axis - ) - else: - self.conv_out = nn.Conv2d(final_block_channels, z_channels, kernel_size=3, stride=1, padding=1) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - # hidden_states expected shape: (batch_size, channels, time, num_mel_bins) - hidden_states = self.conv_in(hidden_states) - - for level in range(self.num_resolutions): - stage = self.down[level] - for block_idx, block in enumerate(stage.block): - hidden_states = block(hidden_states, temb=None) - if stage.attn: - hidden_states = stage.attn[block_idx](hidden_states) - - if level != self.num_resolutions - 1 and hasattr(stage, "downsample"): - hidden_states = stage.downsample(hidden_states) - - hidden_states = self.mid.block_1(hidden_states, temb=None) - hidden_states = self.mid.attn_1(hidden_states) - hidden_states = self.mid.block_2(hidden_states, temb=None) - - hidden_states = self.norm_out(hidden_states) - hidden_states = self.non_linearity(hidden_states) - hidden_states = self.conv_out(hidden_states) - - return hidden_states - - -class LTX2AudioDecoder(nn.Module): - """ - Symmetric decoder that reconstructs audio spectrograms from latent features. - - The decoder mirrors the encoder structure with configurable channel multipliers, attention resolutions, and causal - convolutions. - """ - - def __init__( - self, - base_channels: int = 128, - output_channels: int = 1, - num_res_blocks: int = 2, - attn_resolutions: tuple[int, ...] | None = None, - in_channels: int = 2, - resolution: int = 256, - latent_channels: int = 8, - ch_mult: tuple[int, ...] = (1, 2, 4), - norm_type: str = "group", - causality_axis: str | None = "width", - dropout: float = 0.0, - mid_block_add_attention: bool = False, - sample_rate: int = 16000, - mel_hop_length: int = 160, - is_causal: bool = True, - mel_bins: int | None = 64, - ) -> None: - super().__init__() - - self.sample_rate = sample_rate - self.mel_hop_length = mel_hop_length - self.is_causal = is_causal - self.mel_bins = mel_bins - self.patchifier = LTX2AudioAudioPatchifier( - patch_size=1, - audio_latent_downsample_factor=LATENT_DOWNSAMPLE_FACTOR, - sample_rate=sample_rate, - hop_length=mel_hop_length, - is_causal=is_causal, - ) - - self.base_channels = base_channels - self.temb_ch = 0 - self.num_resolutions = len(ch_mult) - self.num_res_blocks = num_res_blocks - self.resolution = resolution - self.in_channels = in_channels - self.out_ch = output_channels - self.give_pre_end = False - self.tanh_out = False - self.norm_type = norm_type - self.latent_channels = latent_channels - self.channel_multipliers = ch_mult - self.attn_resolutions = attn_resolutions - self.causality_axis = causality_axis - - base_block_channels = base_channels * self.channel_multipliers[-1] - base_resolution = resolution // (2 ** (self.num_resolutions - 1)) - self.z_shape = (1, latent_channels, base_resolution, base_resolution) - - if self.causality_axis is not None: - self.conv_in = LTX2AudioCausalConv2d( - latent_channels, base_block_channels, kernel_size=3, stride=1, causality_axis=self.causality_axis - ) - else: - self.conv_in = nn.Conv2d(latent_channels, base_block_channels, kernel_size=3, stride=1, padding=1) - self.non_linearity = nn.SiLU() - self.mid = nn.Module() - self.mid.block_1 = LTX2AudioResnetBlock( - in_channels=base_block_channels, - out_channels=base_block_channels, - temb_channels=self.temb_ch, - dropout=dropout, - norm_type=self.norm_type, - causality_axis=self.causality_axis, - ) - if mid_block_add_attention: - self.mid.attn_1 = LTX2AudioAttnBlock(base_block_channels, norm_type=self.norm_type) - else: - self.mid.attn_1 = nn.Identity() - self.mid.block_2 = LTX2AudioResnetBlock( - in_channels=base_block_channels, - out_channels=base_block_channels, - temb_channels=self.temb_ch, - dropout=dropout, - norm_type=self.norm_type, - causality_axis=self.causality_axis, - ) - - self.up = nn.ModuleList() - block_in = base_block_channels - curr_res = self.resolution // (2 ** (self.num_resolutions - 1)) - - for level in reversed(range(self.num_resolutions)): - stage = nn.Module() - stage.block = nn.ModuleList() - stage.attn = nn.ModuleList() - block_out = self.base_channels * self.channel_multipliers[level] - - for _ in range(self.num_res_blocks + 1): - stage.block.append( - LTX2AudioResnetBlock( - in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - norm_type=self.norm_type, - causality_axis=self.causality_axis, - ) - ) - block_in = block_out - if self.attn_resolutions: - if curr_res in self.attn_resolutions: - stage.attn.append(LTX2AudioAttnBlock(block_in, norm_type=self.norm_type)) - - if level != 0: - stage.upsample = LTX2AudioUpsample(block_in, True, causality_axis=self.causality_axis) - curr_res *= 2 - - self.up.insert(0, stage) - - final_block_channels = block_in - - if self.norm_type == "group": - self.norm_out = nn.GroupNorm(num_groups=32, num_channels=final_block_channels, eps=1e-6, affine=True) - elif self.norm_type == "pixel": - self.norm_out = LTX2AudioPixelNorm(dim=1, eps=1e-6) - else: - raise ValueError(f"Invalid normalization type: {self.norm_type}") - - if self.causality_axis is not None: - self.conv_out = LTX2AudioCausalConv2d( - final_block_channels, output_channels, kernel_size=3, stride=1, causality_axis=self.causality_axis - ) - else: - self.conv_out = nn.Conv2d(final_block_channels, output_channels, kernel_size=3, stride=1, padding=1) - - def forward( - self, - sample: torch.Tensor, - ) -> torch.Tensor: - _, _, frames, mel_bins = sample.shape - - target_frames = frames * LATENT_DOWNSAMPLE_FACTOR - - if self.causality_axis is not None: - target_frames = max(target_frames - (LATENT_DOWNSAMPLE_FACTOR - 1), 1) - - target_channels = self.out_ch - target_mel_bins = self.mel_bins if self.mel_bins is not None else mel_bins - - hidden_features = self.conv_in(sample) - hidden_features = self.mid.block_1(hidden_features, temb=None) - hidden_features = self.mid.attn_1(hidden_features) - hidden_features = self.mid.block_2(hidden_features, temb=None) - - for level in reversed(range(self.num_resolutions)): - stage = self.up[level] - for block_idx, block in enumerate(stage.block): - hidden_features = block(hidden_features, temb=None) - if stage.attn: - hidden_features = stage.attn[block_idx](hidden_features) - - if level != 0 and hasattr(stage, "upsample"): - hidden_features = stage.upsample(hidden_features) - - if self.give_pre_end: - return hidden_features - - hidden = self.norm_out(hidden_features) - hidden = self.non_linearity(hidden) - decoded_output = self.conv_out(hidden) - decoded_output = torch.tanh(decoded_output) if self.tanh_out else decoded_output - - _, _, current_time, current_freq = decoded_output.shape - target_time = target_frames - target_freq = target_mel_bins - - decoded_output = decoded_output[ - :, :target_channels, : min(current_time, target_time), : min(current_freq, target_freq) - ] - - time_padding_needed = target_time - decoded_output.shape[2] - freq_padding_needed = target_freq - decoded_output.shape[3] - - if time_padding_needed > 0 or freq_padding_needed > 0: - padding = ( - 0, - max(freq_padding_needed, 0), - 0, - max(time_padding_needed, 0), - ) - decoded_output = F.pad(decoded_output, padding) - - decoded_output = decoded_output[:, :target_channels, :target_time, :target_freq] - - return decoded_output - - -class AutoencoderKLLTX2Audio(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - LTX2 audio VAE for encoding and decoding audio latent representations. - """ - - _supports_gradient_checkpointing = False - - @register_to_config - def __init__( - self, - base_channels: int = 128, - output_channels: int = 2, - ch_mult: tuple[int, ...] = (1, 2, 4), - num_res_blocks: int = 2, - attn_resolutions: tuple[int, ...] | None = None, - in_channels: int = 2, - resolution: int = 256, - latent_channels: int = 8, - norm_type: str = "pixel", - causality_axis: str | None = "height", - dropout: float = 0.0, - mid_block_add_attention: bool = False, - sample_rate: int = 16000, - mel_hop_length: int = 160, - is_causal: bool = True, - mel_bins: int | None = 64, - double_z: bool = True, - ) -> None: - super().__init__() - - supported_causality_axes = {"none", "width", "height", "width-compatibility"} - if causality_axis not in supported_causality_axes: - raise ValueError(f"{causality_axis=} is not valid. Supported values: {supported_causality_axes}") - - attn_resolution_set = set(attn_resolutions) if attn_resolutions else attn_resolutions - - self.encoder = LTX2AudioEncoder( - base_channels=base_channels, - output_channels=output_channels, - ch_mult=ch_mult, - num_res_blocks=num_res_blocks, - attn_resolutions=attn_resolution_set, - in_channels=in_channels, - resolution=resolution, - latent_channels=latent_channels, - norm_type=norm_type, - causality_axis=causality_axis, - dropout=dropout, - mid_block_add_attention=mid_block_add_attention, - sample_rate=sample_rate, - mel_hop_length=mel_hop_length, - is_causal=is_causal, - mel_bins=mel_bins, - double_z=double_z, - ) - - self.decoder = LTX2AudioDecoder( - base_channels=base_channels, - output_channels=output_channels, - ch_mult=ch_mult, - num_res_blocks=num_res_blocks, - attn_resolutions=attn_resolution_set, - in_channels=in_channels, - resolution=resolution, - latent_channels=latent_channels, - norm_type=norm_type, - causality_axis=causality_axis, - dropout=dropout, - mid_block_add_attention=mid_block_add_attention, - sample_rate=sample_rate, - mel_hop_length=mel_hop_length, - is_causal=is_causal, - mel_bins=mel_bins, - ) - - # Per-channel statistics for normalizing and denormalizing the latent representation. This statics is computed over - # the entire dataset and stored in model's checkpoint under AudioVAE state_dict - latents_std = torch.ones((base_channels,)) - latents_mean = torch.zeros((base_channels,)) - self.register_buffer("latents_mean", latents_mean, persistent=True) - self.register_buffer("latents_std", latents_std, persistent=True) - - # TODO: calculate programmatically instead of hardcoding - self.temporal_compression_ratio = LATENT_DOWNSAMPLE_FACTOR # 4 - # TODO: confirm whether the mel compression ratio below is correct - self.mel_compression_ratio = LATENT_DOWNSAMPLE_FACTOR - self.use_slicing = False - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - return self.encoder(x) - - @apply_forward_hook - def encode(self, x: torch.Tensor, return_dict: bool = True): - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor) -> torch.Tensor: - return self.decoder(z) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice) for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z) - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`DecoderOutput`] is returned, otherwise a plain `tuple` is returned. - """ - posterior = self.encode(sample).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z) - if not return_dict: - return (dec.sample,) - return dec diff --git a/diffusers/models/autoencoders/autoencoder_kl_magvit.py b/diffusers/models/autoencoders/autoencoder_kl_magvit.py deleted file mode 100644 index 9f9718e135840654def9fc4042b2387eada5aaad..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_magvit.py +++ /dev/null @@ -1,1080 +0,0 @@ -# Copyright 2025 The EasyAnimate team and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class EasyAnimateCausalConv3d(nn.Conv3d): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, ...] = 3, - stride: int | tuple[int, ...] = 1, - padding: int | tuple[int, ...] = 1, - dilation: int | tuple[int, ...] = 1, - groups: int = 1, - bias: bool = True, - padding_mode: str = "zeros", - ): - # Ensure kernel_size, stride, and dilation are tuples of length 3 - kernel_size = kernel_size if isinstance(kernel_size, tuple) else (kernel_size,) * 3 - assert len(kernel_size) == 3, f"Kernel size must be a 3-tuple, got {kernel_size} instead." - - stride = stride if isinstance(stride, tuple) else (stride,) * 3 - assert len(stride) == 3, f"Stride must be a 3-tuple, got {stride} instead." - - dilation = dilation if isinstance(dilation, tuple) else (dilation,) * 3 - assert len(dilation) == 3, f"Dilation must be a 3-tuple, got {dilation} instead." - - # Unpack kernel size, stride, and dilation for temporal, height, and width dimensions - t_ks, h_ks, w_ks = kernel_size - self.t_stride, h_stride, w_stride = stride - t_dilation, h_dilation, w_dilation = dilation - - # Calculate padding for temporal dimension to maintain causality - t_pad = (t_ks - 1) * t_dilation - - # Calculate padding for height and width dimensions based on the padding parameter - if padding is None: - h_pad = math.ceil(((h_ks - 1) * h_dilation + (1 - h_stride)) / 2) - w_pad = math.ceil(((w_ks - 1) * w_dilation + (1 - w_stride)) / 2) - elif isinstance(padding, int): - h_pad = w_pad = padding - else: - assert NotImplementedError - - # Store temporal padding and initialize flags and previous features cache - self.temporal_padding = t_pad - self.temporal_padding_origin = math.ceil(((t_ks - 1) * w_dilation + (1 - w_stride)) / 2) - - self.prev_features = None - - # Initialize the parent class with modified padding - super().__init__( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=kernel_size, - stride=stride, - dilation=dilation, - padding=(0, h_pad, w_pad), - groups=groups, - bias=bias, - padding_mode=padding_mode, - ) - - def _clear_conv_cache(self): - del self.prev_features - self.prev_features = None - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - # Ensure input tensor is of the correct type - dtype = hidden_states.dtype - if self.prev_features is None: - # Pad the input tensor in the temporal dimension to maintain causality - hidden_states = F.pad( - hidden_states, - pad=(0, 0, 0, 0, self.temporal_padding, 0), - mode="replicate", # TODO: check if this is necessary - ) - hidden_states = hidden_states.to(dtype=dtype) - - # Clear cache before processing and store previous features for causality - self._clear_conv_cache() - self.prev_features = hidden_states[:, :, -self.temporal_padding :].clone() - - # Process the input tensor in chunks along the temporal dimension - num_frames = hidden_states.size(2) - outputs = [] - i = 0 - while i + self.temporal_padding + 1 <= num_frames: - out = super().forward(hidden_states[:, :, i : i + self.temporal_padding + 1]) - i += self.t_stride - outputs.append(out) - return torch.concat(outputs, 2) - else: - # Concatenate previous features with the input tensor for continuous temporal processing - if self.t_stride == 2: - hidden_states = torch.concat( - [self.prev_features[:, :, -(self.temporal_padding - 1) :], hidden_states], dim=2 - ) - else: - hidden_states = torch.concat([self.prev_features, hidden_states], dim=2) - hidden_states = hidden_states.to(dtype=dtype) - - # Clear cache and update previous features - self._clear_conv_cache() - self.prev_features = hidden_states[:, :, -self.temporal_padding :].clone() - - # Process the concatenated tensor in chunks along the temporal dimension - num_frames = hidden_states.size(2) - outputs = [] - i = 0 - while i + self.temporal_padding + 1 <= num_frames: - out = super().forward(hidden_states[:, :, i : i + self.temporal_padding + 1]) - i += self.t_stride - outputs.append(out) - return torch.concat(outputs, 2) - - -class EasyAnimateResidualBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - non_linearity: str = "silu", - norm_num_groups: int = 32, - norm_eps: float = 1e-6, - spatial_group_norm: bool = True, - dropout: float = 0.0, - output_scale_factor: float = 1.0, - ): - super().__init__() - - self.output_scale_factor = output_scale_factor - - # Group normalization for input tensor - self.norm1 = nn.GroupNorm( - num_groups=norm_num_groups, - num_channels=in_channels, - eps=norm_eps, - affine=True, - ) - self.nonlinearity = get_activation(non_linearity) - self.conv1 = EasyAnimateCausalConv3d(in_channels, out_channels, kernel_size=3) - - self.norm2 = nn.GroupNorm(num_groups=norm_num_groups, num_channels=out_channels, eps=norm_eps, affine=True) - self.dropout = nn.Dropout(dropout) - self.conv2 = EasyAnimateCausalConv3d(out_channels, out_channels, kernel_size=3) - - if in_channels != out_channels: - self.shortcut = nn.Conv3d(in_channels, out_channels, kernel_size=1) - else: - self.shortcut = nn.Identity() - - self.spatial_group_norm = spatial_group_norm - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - shortcut = self.shortcut(hidden_states) - - if self.spatial_group_norm: - batch_size = hidden_states.size(0) - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) # [B, C, T, H, W] -> [B * T, C, H, W] - hidden_states = self.norm1(hidden_states) - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute( - 0, 2, 1, 3, 4 - ) # [B * T, C, H, W] -> [B, C, T, H, W] - else: - hidden_states = self.norm1(hidden_states) - - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv1(hidden_states) - - if self.spatial_group_norm: - batch_size = hidden_states.size(0) - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) # [B, C, T, H, W] -> [B * T, C, H, W] - hidden_states = self.norm2(hidden_states) - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute( - 0, 2, 1, 3, 4 - ) # [B * T, C, H, W] -> [B, C, T, H, W] - else: - hidden_states = self.norm2(hidden_states) - - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.dropout(hidden_states) - hidden_states = self.conv2(hidden_states) - - return (hidden_states + shortcut) / self.output_scale_factor - - -class EasyAnimateDownsampler3D(nn.Module): - def __init__(self, in_channels: int, out_channels: int, kernel_size: int = 3, stride: tuple = (2, 2, 2)): - super().__init__() - - self.conv = EasyAnimateCausalConv3d( - in_channels=in_channels, out_channels=out_channels, kernel_size=kernel_size, stride=stride, padding=0 - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = F.pad(hidden_states, (0, 1, 0, 1)) - hidden_states = self.conv(hidden_states) - return hidden_states - - -class EasyAnimateUpsampler3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int = 3, - temporal_upsample: bool = False, - spatial_group_norm: bool = True, - ): - super().__init__() - out_channels = out_channels or in_channels - - self.temporal_upsample = temporal_upsample - self.spatial_group_norm = spatial_group_norm - - self.conv = EasyAnimateCausalConv3d( - in_channels=in_channels, out_channels=out_channels, kernel_size=kernel_size - ) - self.prev_features = None - - def _clear_conv_cache(self): - del self.prev_features - self.prev_features = None - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = F.interpolate(hidden_states, scale_factor=(1, 2, 2), mode="nearest") - hidden_states = self.conv(hidden_states) - - if self.temporal_upsample: - if self.prev_features is None: - self.prev_features = hidden_states - else: - hidden_states = F.interpolate( - hidden_states, - scale_factor=(2, 1, 1), - mode="trilinear" if not self.spatial_group_norm else "nearest", - ) - return hidden_states - - -class EasyAnimateDownBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int = 1, - act_fn: str = "silu", - norm_num_groups: int = 32, - norm_eps: float = 1e-6, - spatial_group_norm: bool = True, - dropout: float = 0.0, - output_scale_factor: float = 1.0, - add_downsample: bool = True, - add_temporal_downsample: bool = True, - ): - super().__init__() - - self.convs = nn.ModuleList([]) - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - self.convs.append( - EasyAnimateResidualBlock3D( - in_channels=in_channels, - out_channels=out_channels, - non_linearity=act_fn, - norm_num_groups=norm_num_groups, - norm_eps=norm_eps, - spatial_group_norm=spatial_group_norm, - dropout=dropout, - output_scale_factor=output_scale_factor, - ) - ) - - if add_downsample and add_temporal_downsample: - self.downsampler = EasyAnimateDownsampler3D(out_channels, out_channels, kernel_size=3, stride=(2, 2, 2)) - self.spatial_downsample_factor = 2 - self.temporal_downsample_factor = 2 - elif add_downsample and not add_temporal_downsample: - self.downsampler = EasyAnimateDownsampler3D(out_channels, out_channels, kernel_size=3, stride=(1, 2, 2)) - self.spatial_downsample_factor = 2 - self.temporal_downsample_factor = 1 - else: - self.downsampler = None - self.spatial_downsample_factor = 1 - self.temporal_downsample_factor = 1 - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - for conv in self.convs: - hidden_states = conv(hidden_states) - if self.downsampler is not None: - hidden_states = self.downsampler(hidden_states) - return hidden_states - - -class EasyAnimateUpBlock3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int = 1, - act_fn: str = "silu", - norm_num_groups: int = 32, - norm_eps: float = 1e-6, - spatial_group_norm: bool = False, - dropout: float = 0.0, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - add_temporal_upsample: bool = True, - ): - super().__init__() - - self.convs = nn.ModuleList([]) - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - self.convs.append( - EasyAnimateResidualBlock3D( - in_channels=in_channels, - out_channels=out_channels, - non_linearity=act_fn, - norm_num_groups=norm_num_groups, - norm_eps=norm_eps, - spatial_group_norm=spatial_group_norm, - dropout=dropout, - output_scale_factor=output_scale_factor, - ) - ) - - if add_upsample: - self.upsampler = EasyAnimateUpsampler3D( - in_channels, - in_channels, - temporal_upsample=add_temporal_upsample, - spatial_group_norm=spatial_group_norm, - ) - else: - self.upsampler = None - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - for conv in self.convs: - hidden_states = conv(hidden_states) - if self.upsampler is not None: - hidden_states = self.upsampler(hidden_states) - return hidden_states - - -class EasyAnimateMidBlock3d(nn.Module): - def __init__( - self, - in_channels: int, - num_layers: int = 1, - act_fn: str = "silu", - norm_num_groups: int = 32, - norm_eps: float = 1e-6, - spatial_group_norm: bool = True, - dropout: float = 0.0, - output_scale_factor: float = 1.0, - ): - super().__init__() - - norm_num_groups = norm_num_groups if norm_num_groups is not None else min(in_channels // 4, 32) - - self.convs = nn.ModuleList( - [ - EasyAnimateResidualBlock3D( - in_channels=in_channels, - out_channels=in_channels, - non_linearity=act_fn, - norm_num_groups=norm_num_groups, - norm_eps=norm_eps, - spatial_group_norm=spatial_group_norm, - dropout=dropout, - output_scale_factor=output_scale_factor, - ) - ] - ) - - for _ in range(num_layers - 1): - self.convs.append( - EasyAnimateResidualBlock3D( - in_channels=in_channels, - out_channels=in_channels, - non_linearity=act_fn, - norm_num_groups=norm_num_groups, - norm_eps=norm_eps, - spatial_group_norm=spatial_group_norm, - dropout=dropout, - output_scale_factor=output_scale_factor, - ) - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.convs[0](hidden_states) - for resnet in self.convs[1:]: - hidden_states = resnet(hidden_states) - return hidden_states - - -class EasyAnimateEncoder(nn.Module): - r""" - Causal encoder for 3D video-like data used in [EasyAnimate](https://huggingface.co/papers/2405.18991). - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 8, - down_block_types: tuple[str, ...] = ( - "SpatialDownBlock3D", - "SpatialTemporalDownBlock3D", - "SpatialTemporalDownBlock3D", - "SpatialTemporalDownBlock3D", - ), - block_out_channels: tuple[int, ...] = [128, 256, 512, 512], - layers_per_block: int = 2, - norm_num_groups: int = 32, - act_fn: str = "silu", - double_z: bool = True, - spatial_group_norm: bool = False, - ): - super().__init__() - - # 1. Input convolution - self.conv_in = EasyAnimateCausalConv3d(in_channels, block_out_channels[0], kernel_size=3) - - # 2. Down blocks - self.down_blocks = nn.ModuleList([]) - output_channels = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - input_channels = output_channels - output_channels = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - if down_block_type == "SpatialDownBlock3D": - down_block = EasyAnimateDownBlock3D( - in_channels=input_channels, - out_channels=output_channels, - num_layers=layers_per_block, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - norm_eps=1e-6, - spatial_group_norm=spatial_group_norm, - add_downsample=not is_final_block, - add_temporal_downsample=False, - ) - elif down_block_type == "SpatialTemporalDownBlock3D": - down_block = EasyAnimateDownBlock3D( - in_channels=input_channels, - out_channels=output_channels, - num_layers=layers_per_block, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - norm_eps=1e-6, - spatial_group_norm=spatial_group_norm, - add_downsample=not is_final_block, - add_temporal_downsample=True, - ) - else: - raise ValueError(f"Unknown up block type: {down_block_type}") - self.down_blocks.append(down_block) - - # 3. Middle block - self.mid_block = EasyAnimateMidBlock3d( - in_channels=block_out_channels[-1], - num_layers=layers_per_block, - act_fn=act_fn, - spatial_group_norm=spatial_group_norm, - norm_num_groups=norm_num_groups, - norm_eps=1e-6, - dropout=0, - output_scale_factor=1, - ) - - # 4. Output normalization & convolution - self.spatial_group_norm = spatial_group_norm - self.conv_norm_out = nn.GroupNorm( - num_channels=block_out_channels[-1], - num_groups=norm_num_groups, - eps=1e-6, - ) - self.conv_act = get_activation(act_fn) - - # Initialize the output convolution layer - conv_out_channels = 2 * out_channels if double_z else out_channels - self.conv_out = EasyAnimateCausalConv3d(block_out_channels[-1], conv_out_channels, kernel_size=3) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - # hidden_states: (B, C, T, H, W) - hidden_states = self.conv_in(hidden_states) - - for down_block in self.down_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(down_block, hidden_states) - else: - hidden_states = down_block(hidden_states) - - hidden_states = self.mid_block(hidden_states) - - if self.spatial_group_norm: - batch_size = hidden_states.size(0) - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) - hidden_states = self.conv_norm_out(hidden_states) - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - else: - hidden_states = self.conv_norm_out(hidden_states) - - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states) - return hidden_states - - -class EasyAnimateDecoder(nn.Module): - r""" - Causal decoder for 3D video-like data used in [EasyAnimate](https://huggingface.co/papers/2405.18991). - """ - - _supports_gradient_checkpointing = True - - def __init__( - self, - in_channels: int = 8, - out_channels: int = 3, - up_block_types: tuple[str, ...] = ( - "SpatialUpBlock3D", - "SpatialTemporalUpBlock3D", - "SpatialTemporalUpBlock3D", - "SpatialTemporalUpBlock3D", - ), - block_out_channels: tuple[int, ...] = [128, 256, 512, 512], - layers_per_block: int = 2, - norm_num_groups: int = 32, - act_fn: str = "silu", - spatial_group_norm: bool = False, - ): - super().__init__() - - # 1. Input convolution - self.conv_in = EasyAnimateCausalConv3d(in_channels, block_out_channels[-1], kernel_size=3) - - # 2. Middle block - self.mid_block = EasyAnimateMidBlock3d( - in_channels=block_out_channels[-1], - num_layers=layers_per_block, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - norm_eps=1e-6, - dropout=0, - output_scale_factor=1, - ) - - # 3. Up blocks - self.up_blocks = nn.ModuleList([]) - reversed_block_out_channels = list(reversed(block_out_channels)) - output_channels = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - input_channels = output_channels - output_channels = reversed_block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - # Create and append up block to up_blocks - if up_block_type == "SpatialUpBlock3D": - up_block = EasyAnimateUpBlock3d( - in_channels=input_channels, - out_channels=output_channels, - num_layers=layers_per_block + 1, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - norm_eps=1e-6, - spatial_group_norm=spatial_group_norm, - add_upsample=not is_final_block, - add_temporal_upsample=False, - ) - elif up_block_type == "SpatialTemporalUpBlock3D": - up_block = EasyAnimateUpBlock3d( - in_channels=input_channels, - out_channels=output_channels, - num_layers=layers_per_block + 1, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - norm_eps=1e-6, - spatial_group_norm=spatial_group_norm, - add_upsample=not is_final_block, - add_temporal_upsample=True, - ) - else: - raise ValueError(f"Unknown up block type: {up_block_type}") - self.up_blocks.append(up_block) - - # Output normalization and activation - self.spatial_group_norm = spatial_group_norm - self.conv_norm_out = nn.GroupNorm( - num_channels=block_out_channels[0], - num_groups=norm_num_groups, - eps=1e-6, - ) - self.conv_act = get_activation(act_fn) - - # Output convolution layer - self.conv_out = EasyAnimateCausalConv3d(block_out_channels[0], out_channels, kernel_size=3) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - # hidden_states: (B, C, T, H, W) - hidden_states = self.conv_in(hidden_states) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states) - else: - hidden_states = self.mid_block(hidden_states) - - for up_block in self.up_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(up_block, hidden_states) - else: - hidden_states = up_block(hidden_states) - - if self.spatial_group_norm: - batch_size = hidden_states.size(0) - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) # [B, C, T, H, W] -> [B * T, C, H, W] - hidden_states = self.conv_norm_out(hidden_states) - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute( - 0, 2, 1, 3, 4 - ) # [B * T, C, H, W] -> [B, C, T, H, W] - else: - hidden_states = self.conv_norm_out(hidden_states) - - hidden_states = self.conv_act(hidden_states) - hidden_states = self.conv_out(hidden_states) - return hidden_states - - -class AutoencoderKLMagvit(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - A VAE model with KL loss for encoding images into latents and decoding latent representations into images. This - model is used in [EasyAnimate](https://huggingface.co/papers/2405.18991). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - latent_channels: int = 16, - out_channels: int = 3, - block_out_channels: tuple[int, ...] = [128, 256, 512, 512], - down_block_types: tuple[str, ...] = [ - "SpatialDownBlock3D", - "SpatialTemporalDownBlock3D", - "SpatialTemporalDownBlock3D", - "SpatialTemporalDownBlock3D", - ], - up_block_types: tuple[str, ...] = [ - "SpatialUpBlock3D", - "SpatialTemporalUpBlock3D", - "SpatialTemporalUpBlock3D", - "SpatialTemporalUpBlock3D", - ], - layers_per_block: int = 2, - act_fn: str = "silu", - norm_num_groups: int = 32, - scaling_factor: float = 0.7125, - spatial_group_norm: bool = True, - ): - super().__init__() - - # Initialize the encoder - self.encoder = EasyAnimateEncoder( - in_channels=in_channels, - out_channels=latent_channels, - down_block_types=down_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - norm_num_groups=norm_num_groups, - act_fn=act_fn, - double_z=True, - spatial_group_norm=spatial_group_norm, - ) - - # Initialize the decoder - self.decoder = EasyAnimateDecoder( - in_channels=latent_channels, - out_channels=out_channels, - up_block_types=up_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - norm_num_groups=norm_num_groups, - act_fn=act_fn, - spatial_group_norm=spatial_group_norm, - ) - - # Initialize convolution layers for quantization and post-quantization - self.quant_conv = nn.Conv3d(2 * latent_channels, 2 * latent_channels, kernel_size=1) - self.post_quant_conv = nn.Conv3d(latent_channels, latent_channels, kernel_size=1) - - self.spatial_compression_ratio = 2 ** (len(block_out_channels) - 1) - self.temporal_compression_ratio = 2 ** (len(block_out_channels) - 2) - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent - # frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the - # intermediate tiles together, the memory requirement can be lowered. - self.use_tiling = False - - # When decoding temporally long video latents, the memory requirement is very high. By decoding latent frames - # at a fixed frame batch size (based on `self.num_latent_frames_batch_size`), the memory requirement can be lowered. - self.use_framewise_encoding = False - self.use_framewise_decoding = False - - # Assign mini-batch sizes for encoder and decoder - self.num_sample_frames_batch_size = 4 - self.num_latent_frames_batch_size = 1 - - # The minimal tile height and width for spatial tiling to be used - self.tile_sample_min_height = 512 - self.tile_sample_min_width = 512 - self.tile_sample_min_num_frames = 4 - - # The minimal distance between two spatial tiles - self.tile_sample_stride_height = 448 - self.tile_sample_stride_width = 448 - self.tile_sample_stride_num_frames = 8 - - def _clear_conv_cache(self): - # Clear cache for convolutional layers if needed - for name, module in self.named_modules(): - if isinstance(module, EasyAnimateCausalConv3d): - module._clear_conv_cache() - if isinstance(module, EasyAnimateUpsampler3D): - module._clear_conv_cache() - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_sample_min_num_frames: int | None = None, - tile_sample_stride_height: float | None = None, - tile_sample_stride_width: float | None = None, - tile_sample_stride_num_frames: float | None = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_sample_stride_height (`int`, *optional*): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. - tile_sample_stride_width (`int`, *optional*): - The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling - artifacts produced across the width dimension. - """ - self.use_tiling = True - self.use_framewise_decoding = True - self.use_framewise_encoding = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_sample_min_num_frames = tile_sample_min_num_frames or self.tile_sample_min_num_frames - self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height - self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width - self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames - - @apply_forward_hook - def _encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded images. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_tiling and (x.shape[-1] > self.tile_sample_min_height or x.shape[-2] > self.tile_sample_min_width): - return self.tiled_encode(x, return_dict=return_dict) - - first_frames = self.encoder(x[:, :, :1, :, :]) - h = [first_frames] - for i in range(1, x.shape[2], self.num_sample_frames_batch_size): - next_frames = self.encoder(x[:, :, i : i + self.num_sample_frames_batch_size, :, :]) - h.append(next_frames) - h = torch.cat(h, dim=2) - moments = self.quant_conv(h) - - self._clear_conv_cache() - return moments - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - batch_size, num_channels, num_frames, height, width = z.shape - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - - if self.use_tiling and (z.shape[-1] > tile_latent_min_height or z.shape[-2] > tile_latent_min_width): - return self.tiled_decode(z, return_dict=return_dict) - - z = self.post_quant_conv(z) - - # Process the first frame and save the result - first_frames = self.decoder(z[:, :, :1, :, :]) - # Initialize the list to store the processed frames, starting with the first frame - dec = [first_frames] - # Process the remaining frames, with the number of frames processed at a time determined by mini_batch_decoder - for i in range(1, z.shape[2], self.num_latent_frames_batch_size): - next_frames = self.decoder(z[:, :, i : i + self.num_latent_frames_batch_size, :, :]) - dec.append(next_frames) - # Concatenate all processed frames along the channel dimension - dec = torch.cat(dec, dim=2) - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - """ - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - self._clear_conv_cache() - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[3], b.shape[3], blend_extent) - for y in range(blend_extent): - b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[4], b.shape[4], blend_extent) - for x in range(blend_extent): - b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( - x / blend_extent - ) - return b - - def tiled_encode(self, x: torch.Tensor, return_dict: bool = True) -> AutoencoderKLOutput: - batch_size, num_channels, num_frames, height, width = x.shape - latent_height = height // self.spatial_compression_ratio - latent_width = width // self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = tile_latent_min_height - tile_latent_stride_height - blend_width = tile_latent_min_width - tile_latent_stride_width - - # Split the image into 512x512 tiles and encode them separately. - rows = [] - for i in range(0, height, self.tile_sample_stride_height): - row = [] - for j in range(0, width, self.tile_sample_stride_width): - tile = x[ - :, - :, - :, - i : i + self.tile_sample_min_height, - j : j + self.tile_sample_min_width, - ] - - first_frames = self.encoder(tile[:, :, 0:1, :, :]) - tile_h = [first_frames] - for k in range(1, num_frames, self.num_sample_frames_batch_size): - next_frames = self.encoder(tile[:, :, k : k + self.num_sample_frames_batch_size, :, :]) - tile_h.append(next_frames) - tile = torch.cat(tile_h, dim=2) - tile = self.quant_conv(tile) - self._clear_conv_cache() - row.append(tile) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :latent_height, :latent_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - moments = torch.cat(result_rows, dim=3)[:, :, :, :latent_height, :latent_width] - return moments - - def tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - batch_size, num_channels, num_frames, height, width = z.shape - sample_height = height * self.spatial_compression_ratio - sample_width = width * self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = self.tile_sample_min_height - self.tile_sample_stride_height - blend_width = self.tile_sample_min_width - self.tile_sample_stride_width - - # Split z into overlapping 64x64 tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, tile_latent_stride_height): - row = [] - for j in range(0, width, tile_latent_stride_width): - tile = z[ - :, - :, - :, - i : i + tile_latent_min_height, - j : j + tile_latent_min_width, - ] - tile = self.post_quant_conv(tile) - - # Process the first frame and save the result - first_frames = self.decoder(tile[:, :, :1, :, :]) - # Initialize the list to store the processed frames, starting with the first frame - tile_dec = [first_frames] - # Process the remaining frames, with the number of frames processed at a time determined by mini_batch_decoder - for k in range(1, num_frames, self.num_latent_frames_batch_size): - next_frames = self.decoder(tile[:, :, k : k + self.num_latent_frames_batch_size, :, :]) - tile_dec.append(next_frames) - # Concatenate all processed frames along the channel dimension - decoded = torch.cat(tile_dec, dim=2) - self._clear_conv_cache() - row.append(decoded) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, : self.tile_sample_stride_height, : self.tile_sample_stride_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - dec = torch.cat(result_rows, dim=3)[:, :, :, :sample_height, :sample_width] - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z).sample - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) diff --git a/diffusers/models/autoencoders/autoencoder_kl_minimax_h3.py b/diffusers/models/autoencoders/autoencoder_kl_minimax_h3.py deleted file mode 100644 index 586138fc884e94d09318da23e48eb641d72f9c80..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_minimax_h3.py +++ /dev/null @@ -1,922 +0,0 @@ -# Copyright 2026 The MiniMax and HuggingFace Teams. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class MiniMaxH3VideoCausalConv3d(nn.Conv3d): - r""" - 3D convolution used throughout the MiniMax-H3 video encoder. - - Spatial padding is symmetric and uses `spatial_padding_mode` (`"reflect"` in the released checkpoint); temporal - padding is causal, i.e. `kernel_size_t - 1` zero frames are prepended and nothing is appended. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int], - stride: int | tuple[int, int, int] = 1, - spatial_padding: int = 0, - temporal_padding: int = 0, - spatial_padding_mode: str = "reflect", - ) -> None: - super().__init__(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=0) - self.spatial_padding = spatial_padding - self.temporal_padding = temporal_padding - self.spatial_padding_mode = spatial_padding_mode - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.spatial_padding > 0: - padding = self.spatial_padding - hidden_states = F.pad( - hidden_states, (padding, padding, padding, padding, 0, 0), mode=self.spatial_padding_mode - ) - if self.temporal_padding > 0: - hidden_states = F.pad(hidden_states, (0, 0, 0, 0, self.temporal_padding, 0), mode="constant") - return F.conv3d(hidden_states, self.weight, self.bias, stride=self.stride, padding=0, dilation=self.dilation) - - -class MiniMaxH3VideoGroupNorm(nn.GroupNorm): - r""" - Group normalization applied to each latent frame in isolation (`use_t_isolated_gn` in the original config): the - temporal axis is folded into the batch axis so statistics never mix across frames. - """ - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).contiguous() - hidden_states = hidden_states.view(batch_size * num_frames, num_channels, 1, height, width) - hidden_states = super().forward(hidden_states) - hidden_states = hidden_states.view(batch_size, num_frames, num_channels, height, width) - return hidden_states.permute(0, 2, 1, 3, 4).contiguous() - - -class MiniMaxH3VideoResnetBlock3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - norm_num_groups: int = 32, - norm_eps: float = 1e-6, - spatial_padding_mode: str = "reflect", - ) -> None: - super().__init__() - self.in_channels = in_channels - self.out_channels = out_channels - - self.norm1 = MiniMaxH3VideoGroupNorm(norm_num_groups, in_channels, eps=norm_eps, affine=True) - self.conv1 = MiniMaxH3VideoCausalConv3d( - in_channels, - out_channels, - kernel_size=3, - spatial_padding=1, - temporal_padding=2, - spatial_padding_mode=spatial_padding_mode, - ) - self.norm2 = MiniMaxH3VideoGroupNorm(norm_num_groups, out_channels, eps=norm_eps, affine=True) - self.conv2 = MiniMaxH3VideoCausalConv3d( - out_channels, - out_channels, - kernel_size=3, - spatial_padding=1, - temporal_padding=2, - spatial_padding_mode=spatial_padding_mode, - ) - self.conv_shortcut = None - if in_channels != out_channels: - self.conv_shortcut = MiniMaxH3VideoCausalConv3d(in_channels, out_channels, kernel_size=1) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - residual = hidden_states - hidden_states = F.silu(self.norm1(hidden_states)) - hidden_states = self.conv1(hidden_states) - hidden_states = F.silu(self.norm2(hidden_states)) - hidden_states = self.conv2(hidden_states) - if self.conv_shortcut is not None: - residual = self.conv_shortcut(residual) - return residual + hidden_states - - -class MiniMaxH3VideoDownsample3d(nn.Module): - r""" - Strided 3x3x3 downsampling convolution. A spatial stride of 2 is preceded by an asymmetric bottom/right pad of 1 - (the convolution itself carries no spatial padding), so the output is exactly `ceil(size / 2)`. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - temporal_stride: int = 1, - spatial_stride: int = 2, - spatial_padding_mode: str = "reflect", - ) -> None: - super().__init__() - self.spatial_stride = spatial_stride - self.spatial_padding_mode = spatial_padding_mode - self.conv = MiniMaxH3VideoCausalConv3d( - in_channels, - out_channels, - kernel_size=3, - stride=(temporal_stride, spatial_stride, spatial_stride), - spatial_padding=0, - temporal_padding=2, - spatial_padding_mode=spatial_padding_mode, - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.spatial_stride == 2: - hidden_states = F.pad(hidden_states, (0, 1, 0, 1, 0, 0), mode=self.spatial_padding_mode) - return self.conv(hidden_states) - - -class MiniMaxH3VideoDownBlock3d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int, - temporal_downsample_factor: int, - spatial_downsample_factor: int, - norm_num_groups: int = 32, - norm_eps: float = 1e-6, - spatial_padding_mode: str = "reflect", - ) -> None: - super().__init__() - self.resnets = nn.ModuleList( - [ - MiniMaxH3VideoResnetBlock3d( - in_channels=in_channels if i == 0 else out_channels, - out_channels=out_channels, - norm_num_groups=norm_num_groups, - norm_eps=norm_eps, - spatial_padding_mode=spatial_padding_mode, - ) - for i in range(num_layers) - ] - ) - self.downsamplers = None - if temporal_downsample_factor * spatial_downsample_factor > 1: - self.downsamplers = nn.ModuleList( - [ - MiniMaxH3VideoDownsample3d( - out_channels, - out_channels, - temporal_stride=temporal_downsample_factor, - spatial_stride=spatial_downsample_factor, - spatial_padding_mode=spatial_padding_mode, - ) - ] - ) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - for resnet in self.resnets: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states) - else: - hidden_states = resnet(hidden_states) - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - return hidden_states - - -class MiniMaxH3VideoEncoder3d(nn.Module): - r""" - Causal 3D CNN encoder. `block_out_channels` gives the channel count of every level; the per-level - `spatial_downsample_factors` / `temporal_downsample_factors` multiply out to the total compression ratios. - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 48, - block_out_channels: tuple[int, ...] = (128, 256, 256, 512, 512, 1024), - layers_per_block: int = 2, - spatial_downsample_factors: tuple[int, ...] = (2, 2, 2, 2, 1, 1), - temporal_downsample_factors: tuple[int, ...] = (1, 2, 2, 1, 1, 1), - norm_num_groups: int = 32, - norm_eps: float = 1e-6, - spatial_padding_mode: str = "reflect", - ) -> None: - super().__init__() - - self.conv_in = MiniMaxH3VideoCausalConv3d( - in_channels, - block_out_channels[0], - kernel_size=3, - spatial_padding=1, - temporal_padding=2, - spatial_padding_mode=spatial_padding_mode, - ) - - block_in_channels = (block_out_channels[0],) + tuple(block_out_channels[:-1]) - self.down_blocks = nn.ModuleList( - [ - MiniMaxH3VideoDownBlock3d( - in_channels=block_in_channels[i], - out_channels=block_out_channels[i], - num_layers=layers_per_block, - temporal_downsample_factor=temporal_downsample_factors[i], - spatial_downsample_factor=spatial_downsample_factors[i], - norm_num_groups=norm_num_groups, - norm_eps=norm_eps, - spatial_padding_mode=spatial_padding_mode, - ) - for i in range(len(block_out_channels)) - ] - ) - - self.norm_out = MiniMaxH3VideoGroupNorm(norm_num_groups, block_out_channels[-1], eps=norm_eps, affine=True) - self.conv_out = MiniMaxH3VideoCausalConv3d( - block_out_channels[-1], - out_channels, - kernel_size=3, - spatial_padding=1, - temporal_padding=2, - spatial_padding_mode=spatial_padding_mode, - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv_in(hidden_states) - for down_block in self.down_blocks: - hidden_states = down_block(hidden_states) - hidden_states = F.silu(self.norm_out(hidden_states)) - return self.conv_out(hidden_states) - - -class MiniMaxH3VideoRotaryPosEmbed(nn.Module): - r""" - 3-axis rotary embedding for the ViT decoder. Coordinates are length-normalized to `[-1, 1)` per axis and scaled by - `2 * pi`, and the resulting `(t, h, w)` angles are concatenated and then duplicated, so the first - `rope_dim_ratio * attention_head_dim` channels of every head are rotated. - """ - - def __init__(self, dim: int, theta: float = 100.0, num_axes: int = 3) -> None: - super().__init__() - if dim % (2 * num_axes) != 0: - raise ValueError(f"`dim` {dim} must be divisible by `2 * num_axes` {2 * num_axes}.") - inv_freq = 1.0 / theta ** torch.arange(0, 1, 2 * num_axes / dim, dtype=torch.float32) - self.register_buffer("inv_freq", inv_freq, persistent=False) - - def forward(self, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - angles = 2.0 * math.pi * position_ids[:, :, :, None] * self.inv_freq[None, None, None, :] - angles = angles.flatten(2, 3).tile(2).unsqueeze(2) - return angles.cos(), angles.sin() - - -class MiniMaxH3VideoAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __call__( - self, - attn: "MiniMaxH3VideoAttention", - hidden_states: torch.Tensor, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - query = attn.to_q(hidden_states).unflatten(2, (attn.heads, -1)) - key = attn.to_k(hidden_states).unflatten(2, (attn.heads, -1)) - value = attn.to_v(hidden_states).unflatten(2, (attn.heads, -1)) - - # The reference normalizes Q/K in float32 regardless of the compute dtype. - query = attn.norm_q(query.float()).to(query.dtype) - key = attn.norm_k(key.float()).to(key.dtype) - - if rotary_emb is not None: - cos, sin = rotary_emb - cos = cos.to(query.dtype) - sin = sin.to(query.dtype) - rotary_dim = cos.shape[-1] - query_rotary, query_pass = query[..., :rotary_dim], query[..., rotary_dim:] - key_rotary, key_pass = key[..., :rotary_dim], key[..., rotary_dim:] - query_first, query_second = query_rotary.chunk(2, dim=-1) - key_first, key_second = key_rotary.chunk(2, dim=-1) - query_rotated = torch.cat([-query_second, query_first], dim=-1) - key_rotated = torch.cat([-key_second, key_first], dim=-1) - query = torch.cat([query_rotary * cos + query_rotated * sin, query_pass], dim=-1) - key = torch.cat([key_rotary * cos + key_rotated * sin, key_pass], dim=-1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=None, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - return attn.to_out[0](hidden_states) - - -class MiniMaxH3VideoAttention(nn.Module, AttentionModuleMixin): - _default_processor_cls = MiniMaxH3VideoAttnProcessor - _available_processors = [MiniMaxH3VideoAttnProcessor] - - def __init__(self, dim: int, heads: int, dim_head: int, eps: float = 1e-5, bias: bool = True) -> None: - super().__init__() - self.heads = heads - self.dim_head = dim_head - self.use_bias = bias - inner_dim = heads * dim_head - - self.norm_q = nn.RMSNorm(dim_head, eps=eps, elementwise_affine=False) - self.norm_k = nn.RMSNorm(dim_head, eps=eps, elementwise_affine=False) - self.to_q = nn.Linear(dim, inner_dim, bias=bias) - self.to_k = nn.Linear(dim, inner_dim, bias=bias) - self.to_v = nn.Linear(dim, inner_dim, bias=bias) - self.to_out = nn.ModuleList([nn.Linear(inner_dim, dim, bias=bias), nn.Dropout(0.0)]) - - self.set_processor(MiniMaxH3VideoAttnProcessor()) - - def forward( - self, hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None - ) -> torch.Tensor: - return self.processor(self, hidden_states, rotary_emb) - - -class MiniMaxH3VideoTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - heads: int, - dim_head: int, - ffn_mult: int = 4, - eps: float = 1e-5, - bias: bool = True, - ) -> None: - super().__init__() - self.norm1 = nn.RMSNorm(dim, eps=eps, elementwise_affine=True) - self.attn = MiniMaxH3VideoAttention(dim=dim, heads=heads, dim_head=dim_head, eps=eps, bias=bias) - self.scale1 = nn.Parameter(torch.zeros(dim)) - self.norm2 = nn.RMSNorm(dim, eps=eps, elementwise_affine=True) - self.ff = FeedForward(dim, mult=ffn_mult, activation_fn="swiglu", bias=bias) - self.scale2 = nn.Parameter(torch.zeros(dim)) - - def forward( - self, hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None - ) -> torch.Tensor: - # The reference normalizes in float32 regardless of the compute dtype. - norm_hidden_states = self.norm1(hidden_states.float()).to(hidden_states.dtype) - hidden_states = hidden_states + self.attn(norm_hidden_states, rotary_emb) * self.scale1 - norm_hidden_states = self.norm2(hidden_states.float()).to(hidden_states.dtype) - hidden_states = hidden_states + self.ff(norm_hidden_states) * self.scale2 - return hidden_states - - -class MiniMaxH3VideoViTDecoder3d(nn.Module): - r""" - Non-causal ViT decoder. Every latent voxel becomes one token; `num_register_tokens` learned register tokens plus a - single all-zero token are appended (all at position `0`), attended over with full self-attention, and dropped - again before the patch projection expands each token into a `patch_size_t x patch_size x patch_size` pixel block. - """ - - def __init__( - self, - in_channels: int = 24, - out_channels: int = 3, - patch_size: int = 16, - patch_size_t: int = 4, - num_layers: int = 36, - num_attention_heads: int = 32, - attention_head_dim: int = 64, - num_register_tokens: int = 4, - ffn_mult: int = 4, - rope_theta: float = 100.0, - rope_dim_ratio: float = 0.75, - norm_eps: float = 1e-5, - ) -> None: - super().__init__() - dim = num_attention_heads * attention_head_dim - self.patch_size = patch_size - self.patch_size_t = patch_size_t - self.out_channels = out_channels - self.num_register_tokens = num_register_tokens - - self.rope = MiniMaxH3VideoRotaryPosEmbed(int(attention_head_dim * rope_dim_ratio), theta=rope_theta) - self.proj_in = nn.Linear(in_channels, dim) - self.register_tokens = nn.Parameter(torch.zeros(1, num_register_tokens, dim)) - self.transformer_blocks = nn.ModuleList( - [ - MiniMaxH3VideoTransformerBlock( - dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - ffn_mult=ffn_mult, - eps=norm_eps, - ) - for _ in range(num_layers) - ] - ) - self.norm_out = nn.LayerNorm(dim, elementwise_affine=True, eps=norm_eps) - self.proj_out = nn.Linear(dim, out_channels * patch_size_t * patch_size * patch_size) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - - hidden_states = hidden_states.permute(0, 2, 3, 4, 1).reshape( - batch_size, num_frames * height * width, num_channels - ) - hidden_states = self.proj_in(hidden_states) - num_patches = hidden_states.shape[1] - - register_tokens = self.register_tokens.expand(batch_size, -1, -1) - cls_token = torch.zeros_like(hidden_states[:, :1, :]) - hidden_states = torch.cat([hidden_states, register_tokens, cls_token], dim=1) - - grids = [ - 2.0 * (torch.arange(0.5, size, dtype=torch.float32, device=hidden_states.device) / size) - 1.0 - for size in (num_frames, height, width) - ] - position_ids = torch.stack(torch.meshgrid(*grids, indexing="ij"), dim=-1).flatten(0, 2) - position_ids = position_ids.unsqueeze(0).expand(batch_size, -1, -1) - suffix_ids = position_ids.new_zeros((batch_size, self.num_register_tokens + 1, 3)) - position_ids = torch.cat([position_ids, suffix_ids], dim=1) - rotary_emb = self.rope(position_ids) - - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(block, hidden_states, rotary_emb) - else: - hidden_states = block(hidden_states, rotary_emb) - - hidden_states = self.norm_out(hidden_states) - hidden_states = self.proj_out(hidden_states) - hidden_states = hidden_states[:, :num_patches, :] - - patch_size, patch_size_t = self.patch_size, self.patch_size_t - hidden_states = hidden_states.view( - batch_size, - num_frames, - height, - width, - self.out_channels, - patch_size_t, - patch_size, - patch_size, - ) - hidden_states = hidden_states.permute(0, 4, 1, 5, 2, 6, 3, 7).contiguous() - return hidden_states.reshape( - batch_size, - self.out_channels, - num_frames * patch_size_t, - height * patch_size, - width * patch_size, - ) - - -class AutoencoderKLMiniMaxH3(ModelMixin, ConfigMixin, AttentionMixin, AutoencoderMixin): - r""" - A VAE model with a causal 3D CNN encoder and a non-causal ViT decoder, used in - [MiniMax-H3](https://huggingface.co/MiniMaxAI). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Latents are normalized with per-channel `latents_mean` / `latents_std` rather than a `scaling_factor`; a pipeline - encodes with `(latent - latents_mean) / latents_std` and decodes with `latent * latents_std + latents_mean`. - - The pixel convention is ImageNet-normalized RGB over a `[0, 1]` base range, not the usual `[-1, 1]`: `encode` - expects `(pixel - imagenet_mean) / imagenet_std` and `decode` returns values in that same space, so a pipeline has - to apply `sample * imagenet_std + imagenet_mean` (mean `(0.485, 0.456, 0.406)`, std `(0.229, 0.224, 0.225)`) and - clamp to `[0, 1]` before postprocessing. - - The temporal geometry is fixed by `clip_length` (17 pixel frames per encoder chunk) and `token_drop` (3 trailing - latent frames dropped per encode): `17 * n + 5` pixel frames map to `5 * n + 2` latent frames. - - Unlike most autoencoders in the library, spatial tiling is **on by default**: MiniMax-H3 was released with tiling - enabled for both encoding and decoding, and the released frames are the blended-tile ones, so disabling tiling - changes the output. Use `enable_tiling` to change the tile geometry, `disable_tiling` to turn it off. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["MiniMaxH3VideoResnetBlock3d", "MiniMaxH3VideoTransformerBlock"] - _repeated_blocks = ["MiniMaxH3VideoTransformerBlock"] - _skip_layerwise_casting_patterns = ["norm"] - # The released checkpoint is float32 and the verified decode recipe is float16 *autocast over float32 weights* - # (see `decode`). A pipeline-level `torch_dtype=torch.bfloat16` must therefore not downcast the weights, so every - # top-level module is pinned, mirroring the transformer's mixed-precision contract. - _keep_in_fp32_modules = ["encoder", "decoder", "quant_conv", "post_quant_conv"] - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - latent_channels: int = 24, - block_out_channels: tuple[int, ...] = (128, 256, 256, 512, 512, 1024), - layers_per_block: int = 2, - spatial_downsample_factors: tuple[int, ...] = (2, 2, 2, 2, 1, 1), - temporal_downsample_factors: tuple[int, ...] = (1, 2, 2, 1, 1, 1), - norm_num_groups: int = 32, - norm_eps: float = 1e-6, - spatial_padding_mode: str = "reflect", - decoder_num_layers: int = 36, - decoder_num_attention_heads: int = 32, - decoder_attention_head_dim: int = 64, - decoder_num_register_tokens: int = 4, - decoder_ffn_mult: int = 4, - decoder_rope_theta: float = 100.0, - decoder_rope_dim_ratio: float = 0.75, - decoder_norm_eps: float = 1e-5, - clip_length: int = 17, - token_drop: int = 3, - latents_mean: tuple[float, ...] = (0.0,) * 24, - latents_std: tuple[float, ...] = (1.0,) * 24, - ) -> None: - super().__init__() - - self.spatial_compression_ratio = math.prod(spatial_downsample_factors) - self.temporal_compression_ratio = math.prod(temporal_downsample_factors) - - self.encoder = MiniMaxH3VideoEncoder3d( - in_channels=in_channels, - out_channels=2 * latent_channels, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - spatial_downsample_factors=spatial_downsample_factors, - temporal_downsample_factors=temporal_downsample_factors, - norm_num_groups=norm_num_groups, - norm_eps=norm_eps, - spatial_padding_mode=spatial_padding_mode, - ) - self.quant_conv = nn.Conv3d(2 * latent_channels, 2 * latent_channels, kernel_size=1) - self.post_quant_conv = nn.Conv3d(latent_channels, latent_channels, kernel_size=1) - self.decoder = MiniMaxH3VideoViTDecoder3d( - in_channels=latent_channels, - out_channels=out_channels, - patch_size=self.spatial_compression_ratio, - patch_size_t=self.temporal_compression_ratio, - num_layers=decoder_num_layers, - num_attention_heads=decoder_num_attention_heads, - attention_head_dim=decoder_attention_head_dim, - num_register_tokens=decoder_num_register_tokens, - ffn_mult=decoder_ffn_mult, - rope_theta=decoder_rope_theta, - rope_dim_ratio=decoder_rope_dim_ratio, - norm_eps=decoder_norm_eps, - ) - - # Derived temporal-chunking geometry. `clip_length` pixel frames are encoded at a time; because - # `clip_length` is not a multiple of `temporal_compression_ratio`, the decoder has to re-derive the - # implicit leading pad (`frame_pre_padding`) and the overlap that `token_drop` leaves behind. - self.frame_pre_padding = (-clip_length) % self.temporal_compression_ratio - self.tokens_chunk_size = math.ceil(clip_length / self.temporal_compression_ratio) - self.token_overlap = (-token_drop) % self.tokens_chunk_size - self.frame_overlap = max(self.token_overlap * self.temporal_compression_ratio - self.frame_pre_padding, 0) - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When encoding/decoding spatially large videos, the memory requirement is very high. By splitting the frames - # into smaller tiles, running the encoder/decoder per tile and blending the overlaps, the memory requirement - # can be lowered. MiniMax-H3 ships with tiling enabled. - self.use_tiling = True - - # The tile size in pixel space, and the minimum overlap between two neighbouring tiles. The actual overlaps are - # widened (in multiples of `spatial_compression_ratio`) so that the tiles cover the frame exactly. - self.tile_sample_min_height = 256 - self.tile_sample_min_width = 256 - self.tile_sample_min_overlap_height = 64 - self.tile_sample_min_overlap_width = 64 - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_sample_min_overlap_height: int | None = None, - tile_sample_min_overlap_width: int | None = None, - ) -> None: - r""" - Enable tiled VAE encoding/decoding. When this option is enabled, the VAE splits the frames into tiles, encodes - or decodes each tile separately and linearly blends the overlaps back together. This lowers the memory - requirement and allows processing larger frames. - - Args: - tile_sample_min_height (`int`, *optional*): - The tile height in pixel space. Frames taller than this are split along the height dimension. - tile_sample_min_width (`int`, *optional*): - The tile width in pixel space. Frames wider than this are split along the width dimension. - tile_sample_min_overlap_height (`int`, *optional*): - The minimum overlap, in pixels, between two consecutive vertical tiles. - tile_sample_min_overlap_width (`int`, *optional*): - The minimum overlap, in pixels, between two consecutive horizontal tiles. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_sample_min_overlap_height = tile_sample_min_overlap_height or self.tile_sample_min_overlap_height - self.tile_sample_min_overlap_width = tile_sample_min_overlap_width or self.tile_sample_min_overlap_width - - def _split_tiles(self, length: int, tile_size: int, min_overlap: int) -> tuple[list[int], list[int], list[int]]: - r""" - Lay `tile_size`-wide tiles over `length` pixels. The number of tiles is the smallest one whose union can cover - `length` while keeping every overlap at least `min_overlap`; the slack is then distributed round-robin over the - overlaps in whole `spatial_compression_ratio` steps so that every tile boundary stays latent-aligned. - """ - if tile_size >= length: - return [0], [length], [] - - num_tiles = math.ceil(length / tile_size) - while tile_size * num_tiles - min_overlap * (num_tiles - 1) - length < 0: - num_tiles += 1 - - overlaps = [min_overlap] * (num_tiles - 1) - remaining = tile_size * num_tiles - sum(overlaps) - length - for i in range(remaining // self.spatial_compression_ratio): - overlaps[i % (num_tiles - 1)] += self.spatial_compression_ratio - - tile_start_indices = [0] - for i in range(num_tiles - 1): - tile_start_indices.append(tile_start_indices[-1] + tile_size - overlaps[i]) - return tile_start_indices, [tile_size] * num_tiles, overlaps - - def _blend(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int, dim: int) -> torch.Tensor: - blend_extent = min(a.shape[dim], b.shape[dim], blend_extent) - positions = torch.arange(blend_extent, device=b.device, dtype=b.dtype) - shape = [1] * a.ndim - shape[dim] = blend_extent - weight_a = (1 - positions / blend_extent).view(shape) - weight_b = (positions / blend_extent).view(shape) - - slice_a = [slice(None)] * a.ndim - slice_a[dim] = slice(-blend_extent, None) - slice_b = [slice(None)] * b.ndim - slice_b[dim] = slice(0, blend_extent) - blended = a[tuple(slice_a)] * weight_a + b[tuple(slice_b)] * weight_b - - if blend_extent == b.shape[dim]: - return blended - slice_rest = [slice(None)] * b.ndim - slice_rest[dim] = slice(blend_extent, None) - return torch.cat([blended, b[tuple(slice_rest)]], dim=dim) - - def _stitch_tiles( - self, - tiles: list[list[torch.Tensor]], - height_overlaps: list[int], - width_overlaps: list[int], - ) -> torch.Tensor: - result_rows = [] - for i, row in enumerate(tiles): - result_row = [] - for j, tile in enumerate(row): - if i > 0: - tile = self._blend(tiles[i - 1][j], tile, height_overlaps[i - 1], dim=-2) - if j > 0: - tile = self._blend(row[j - 1], tile, width_overlaps[j - 1], dim=-1) - if i < len(tiles) - 1: - tile = tile[..., : -height_overlaps[i], :] - if j < len(row) - 1: - tile = tile[..., :, : -width_overlaps[j]] - result_row.append(tile) - result_rows.append(torch.cat(result_row, dim=-1)) - return torch.cat(result_rows, dim=-2) - - @apply_forward_hook - def _encode_clip(self, x: torch.Tensor) -> torch.Tensor: - r""" - Encode one temporal clip, spatially tiled when tiling is enabled. - - MiniMax-H3 encodes a keyframe or an image reference through this method rather than through [`~encode`], - because a single frame must not go through the temporal chunking, so it carries the offload hook too. - """ - if not self.use_tiling: - return self.quant_conv(self.encoder(x)) - - height, width = x.shape[-2], x.shape[-1] - y_indices, y_lengths, y_overlaps = self._split_tiles( - height, self.tile_sample_min_height, self.tile_sample_min_overlap_height - ) - x_indices, x_lengths, x_overlaps = self._split_tiles( - width, self.tile_sample_min_width, self.tile_sample_min_overlap_width - ) - - rows = [] - for i_pos, i_len in zip(y_indices, y_lengths): - row = [] - for j_pos, j_len in zip(x_indices, x_lengths): - tile = x[..., i_pos : i_pos + i_len, j_pos : j_pos + j_len] - row.append(self.quant_conv(self.encoder(tile))) - rows.append(row) - - latent_y_overlaps = [overlap // self.spatial_compression_ratio for overlap in y_overlaps] - latent_x_overlaps = [overlap // self.spatial_compression_ratio for overlap in x_overlaps] - return self._stitch_tiles(rows, latent_y_overlaps, latent_x_overlaps) - - def _decode_clip(self, z: torch.Tensor) -> torch.Tensor: - r"""Decode one temporal clip, spatially tiled when tiling is enabled.""" - if not self.use_tiling: - return self.decoder(self.post_quant_conv(z)) - - # Tiles are laid out in pixel space and then mapped back onto the latent grid. - height = z.shape[-2] * self.spatial_compression_ratio - width = z.shape[-1] * self.spatial_compression_ratio - y_indices, y_lengths, y_overlaps = self._split_tiles( - height, self.tile_sample_min_height, self.tile_sample_min_overlap_height - ) - x_indices, x_lengths, x_overlaps = self._split_tiles( - width, self.tile_sample_min_width, self.tile_sample_min_overlap_width - ) - - ratio = self.spatial_compression_ratio - rows = [] - for i_pos, i_len in zip(y_indices, y_lengths): - row = [] - for j_pos, j_len in zip(x_indices, x_lengths): - tile = z[ - ..., - i_pos // ratio : i_pos // ratio + i_len // ratio, - j_pos // ratio : j_pos // ratio + j_len // ratio, - ] - row.append(self.decoder(self.post_quant_conv(tile))) - rows.append(row) - - return self._stitch_tiles(rows, y_overlaps, x_overlaps) - - @apply_forward_hook - def _encode(self, x: torch.Tensor) -> torch.Tensor: - r""" - Encode a video in `clip_length`-frame chunks and drop the `token_drop` trailing latent frames. - - MiniMax-H3 encodes a video reference through this method rather than through [`~encode`], because the - posterior is sampled under a fixed generator rather than through the distribution object, so it carries the - offload hook too. - """ - clip_length = self.config.clip_length - num_frames = x.shape[2] - if num_frames % clip_length != 0: - pad_frames = x[:, :, -1:].repeat(1, 1, (-num_frames) % clip_length, 1, 1) - x = torch.cat([x, pad_frames], dim=2) - - moments = torch.cat( - [ - self._encode_clip(x[:, :, i * clip_length : (i + 1) * clip_length]) - for i in range(x.shape[2] // clip_length) - ], - dim=2, - ) - if self.config.token_drop > 0: - moments = moments[:, :, : -self.config.token_drop] - return moments - - def _decode(self, z: torch.Tensor) -> torch.Tensor: - r""" - Decode a latent video, mirroring the chunking that `_encode` applied. - - `token_drop` removed the tail of every encoded chunk, so consecutive decoded chunks overlap by - `frame_overlap` pixel frames and are linearly cross-faded. Latent frames are repeated at the end when the - length is not a whole number of chunks; the extra pixel frames are cut off again at the end. - """ - tokens_chunk_size = self.tokens_chunk_size - token_drop = self.config.token_drop - temporal_ratio = self.temporal_compression_ratio - chunk_num_frames = tokens_chunk_size * temporal_ratio - - num_tokens = z.shape[2] + token_drop - pad_tokens = (-num_tokens) % tokens_chunk_size - num_chunks = (num_tokens + pad_tokens) // tokens_chunk_size - int(token_drop > 0) - if pad_tokens > 0: - z = torch.cat([z, z[:, :, -1:].repeat(1, 1, pad_tokens, 1, 1)], dim=2) - - decoded_chunks = [] - overlap = None - for i in range(num_chunks): - start = i * tokens_chunk_size - clip = self._decode_clip(z[:, :, start : start + tokens_chunk_size + self.token_overlap]) - for j in range(int(token_drop > 0) + 1): - frame_start = j * chunk_num_frames - chunk = clip[:, :, frame_start : frame_start + chunk_num_frames] - chunk = chunk[:, :, self.frame_pre_padding :] - if j == 0: - if overlap is not None: - chunk = self._blend(overlap, chunk, self.frame_overlap, dim=-3) - decoded_chunks.append(chunk) - else: - overlap = chunk - if overlap is not None: - decoded_chunks.append(overlap) - - dec = torch.cat(decoded_chunks, dim=2) - - # `pad_tokens` repeated latent frames produced trailing pixel frames that were never requested. A chunk's - # last latent frame only covers `clip_length % temporal_ratio` pixel frames, the others cover `temporal_ratio`. - if pad_tokens > 0: - intra_tail = self.config.clip_length % temporal_ratio - num_tokens_before_pad = z.shape[2] - pad_tokens - pad_frames = sum( - intra_tail if intra_tail and (num_tokens_before_pad + k) % tokens_chunk_size == 0 else temporal_ratio - for k in range(pad_tokens) - ) - dec = dec[:, :, :-pad_frames] - return dec - - @apply_forward_hook - def encode(self, x: torch.Tensor, return_dict: bool = True) -> AutoencoderKLOutput | tuple[torch.Tensor]: - r""" - Encode a batch of videos into latents. - - Args: - x (`torch.Tensor`): - Input batch of videos, shape `(batch_size, in_channels, num_frames, height, width)`. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoders.autoencoder_kl.AutoencoderKLOutput`] instead of a plain - tuple. - - Returns: - The latent distribution of the encoded videos. Note that MiniMax-H3 normalizes the sampled latents with - `latents_mean` / `latents_std` afterwards. - """ - if self.use_slicing and x.shape[0] > 1: - moments = torch.cat([self._encode(x_slice) for x_slice in x.split(1)]) - else: - moments = self._encode(x) - posterior = DiagonalGaussianDistribution(moments) - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | tuple[torch.Tensor]: - r""" - Decode a batch of latent videos. - - Args: - z (`torch.Tensor`): - Input batch of latent videos, shape `(batch_size, latent_channels, num_latent_frames, height, width)`. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoders.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.autoencoders.vae.DecoderOutput`] or `tuple`: - The decoded videos, shape `(batch_size, out_channels, num_frames, height, width)`. - """ - if self.use_slicing and z.shape[0] > 1: - decoded = torch.cat([self._decode(z_slice) for z_slice in z.split(1)]) - else: - decoded = self._decode(z) - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - generator: torch.Generator | None = None, - return_dict: bool = True, - ) -> DecoderOutput | tuple[torch.Tensor]: - r""" - Encode then decode a batch of videos. - - Args: - sample (`torch.Tensor`): - Input batch of videos, shape `(batch_size, in_channels, num_frames, height, width)`. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample the posterior instead of taking its mode. - generator (`torch.Generator`, *optional*): - Generator used when `sample_posterior=True`. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoders.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.autoencoders.vae.DecoderOutput`] or `tuple`: - The round-tripped videos, shape `(batch_size, out_channels, num_frames, height, width)`. - """ - posterior = self.encode(sample).latent_dist - z = posterior.sample(generator=generator) if sample_posterior else posterior.mode() - return self.decode(z, return_dict=return_dict) diff --git a/diffusers/models/autoencoders/autoencoder_kl_minimax_h3_audio.py b/diffusers/models/autoencoders/autoencoder_kl_minimax_h3_audio.py deleted file mode 100644 index c35c62e28ec31975afec5b5f57c8d1b549b26f64..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_minimax_h3_audio.py +++ /dev/null @@ -1,679 +0,0 @@ -# Copyright 2025 The MiniMax authors and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""MiniMax-H3 audio autoencoder. - -Waveform in / waveform out — there is no mel front-end and no separate vocoder: - -* the **encoder** is a DAC-lineage strided convolutional stack (Snake activations, weight-normed - `Conv1d`) that downsamples by `prod(encoder_rates) = 800`, i.e. 40 latents/s at 32 kHz; -* a **causal-attention projection** (`pre_block`) rewires the 2048-wide encoder trunk to the - 32-channel latent width, followed by the `mean_proj` / `logs_proj` posterior heads; -* the **decoder** is BigVGAN (anti-aliased SnakeBeta activations, transposed-conv upsamplers, AMP - residual blocks) preceded by `dec_in_proj`, upsampling by `prod(decoder_rates) = 800`. - -The autoencoder is **mono**. MiniMax-H3 carries stereo as two *batch* items — the pipeline decodes -`[2, 32, T]` into `[2, 1, samples]` and interleaves at the output boundary — so no stereo handling -belongs here. - -Latents are normalized with per-channel `latents_mean` / `latents_std` (32 floats each) rather than a -scalar `scaling_factor`; both live in the config and are applied by the pipeline. - -Module and parameter names are identical to the original checkpoint, so conversion is a passthrough. -That includes `torch.nn.utils.weight_norm` (the `weight_g` / `weight_v` spelling, as used by the -other diffusers audio autoencoders) and the registered Kaiser-window resampling `filter` buffers of -the anti-aliased activations. -""" - -import math -from dataclasses import dataclass - -import torch -import torch.nn as nn -import torch.nn.functional as F -from torch.nn.utils import weight_norm - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput -from ...utils.accelerate_utils import apply_forward_hook -from ...utils.torch_utils import randn_tensor -from ..attention import AttentionMixin, AttentionModuleMixin -from ..attention_dispatch import dispatch_attention_fn -from ..modeling_utils import ModelMixin, get_parameter_dtype -from .vae import DecoderOutput - - -class MiniMaxH3AudioDiagonalGaussianDistribution: - r"""Posterior of the MiniMax-H3 audio autoencoder, parameterized as `(mean, log_std)`. - - The checkpoint keeps two separate `Conv1d` heads (`mean_proj`, `logs_proj`) instead of one fused - moments projection, and the second head predicts the **log standard deviation**, not the log - variance. The two tensors are therefore stored as produced, and `mode()` is bit-for-bit - `mean_proj`'s output. - - Args: - mean (`torch.Tensor`): Posterior mean, `[batch_size, latent_channels, num_frames]`. - logs (`torch.Tensor`): Posterior log standard deviation, same shape as `mean`. - """ - - def __init__(self, mean: torch.Tensor, logs: torch.Tensor): - self.mean = mean - self.logs = logs - self.std = torch.exp(logs) - - def mode(self) -> torch.Tensor: - return self.mean - - def sample(self, generator: torch.Generator | None = None) -> torch.Tensor: - noise = randn_tensor(self.mean.shape, generator=generator, device=self.mean.device, dtype=self.mean.dtype) - return self.mean + self.std * noise - - -@dataclass -class MiniMaxH3AudioEncoderOutput(BaseOutput): - r""" - Output of [`AutoencoderKLMiniMaxH3Audio.encode`]. - - Args: - latent_dist (`MiniMaxH3AudioDiagonalGaussianDistribution`): - Posterior over the audio latents. MiniMax-H3 always consumes `latent_dist.mode()`. - """ - - latent_dist: MiniMaxH3AudioDiagonalGaussianDistribution - - -def _wn_conv1d(*args, **kwargs) -> nn.Module: - return weight_norm(nn.Conv1d(*args, **kwargs)) - - -def kaiser_sinc_filter1d(cutoff: float, half_width: float, kernel_size: int) -> torch.Tensor: - r"""Kaiser-windowed sinc low-pass filter of shape `[1, 1, kernel_size]`. - - Kept arithmetically identical to the `alias-free-torch` implementation the checkpoint was trained - with, because the resulting tensor is stored as a persistent buffer. - """ - half_size = kernel_size // 2 - - attenuation = 2.285 * (half_size - 1) * math.pi * (4 * half_width) + 7.95 - if attenuation > 50.0: - beta = 0.1102 * (attenuation - 8.7) - elif attenuation >= 21.0: - beta = 0.5842 * (attenuation - 21) ** 0.4 + 0.07886 * (attenuation - 21.0) - else: - beta = 0.0 - window = torch.kaiser_window(kernel_size, beta=beta, periodic=False) - - if kernel_size % 2 == 0: - time = torch.arange(-half_size, half_size) + 0.5 - else: - time = torch.arange(kernel_size) - half_size - - filter_ = 2 * cutoff * window * torch.sinc(2 * cutoff * time) - # Normalize to sum 1 so a constant input does not leak through the resampler. - filter_ /= filter_.sum() - return filter_.view(1, 1, kernel_size) - - -class MiniMaxH3AudioSnake1d(nn.Module): - r"""`x + (alpha + 1e-9)^-1 * sin(alpha * x)^2` over `[batch_size, channels, length]`, with a - per-channel learnable `alpha` of shape `[1, channels, 1]`. Used throughout the DAC encoder.""" - - def __init__(self, channels: int): - super().__init__() - self.alpha = nn.Parameter(torch.ones(1, channels, 1)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return hidden_states + (self.alpha + 1e-9).reciprocal() * torch.sin(self.alpha * hidden_states).pow(2) - - -class MiniMaxH3AudioSnakeBeta(nn.Module): - r"""`x + (exp(beta) + 1e-9)^-1 * sin(exp(alpha) * x)^2` over `[batch_size, channels, length]`. - - The BigVGAN decoder's activation: separate frequency (`alpha`) and magnitude (`beta`) parameters, - both stored in log space as `[channels]` vectors. - """ - - def __init__(self, channels: int): - super().__init__() - self.alpha = nn.Parameter(torch.zeros(channels)) - self.beta = nn.Parameter(torch.zeros(channels)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - alpha = torch.exp(self.alpha.unsqueeze(0).unsqueeze(-1)) - beta = torch.exp(self.beta.unsqueeze(0).unsqueeze(-1)) - return hidden_states + (beta + 1e-9).reciprocal() * torch.sin(alpha * hidden_states).pow(2) - - -class MiniMaxH3AudioLowPassFilter1d(nn.Module): - r"""Depthwise Kaiser-sinc low-pass filter with a stride, i.e. the anti-aliased downsampler.""" - - def __init__(self, cutoff: float, half_width: float, stride: int, kernel_size: int): - super().__init__() - even = kernel_size % 2 == 0 - self.pad_left = kernel_size // 2 - int(even) - self.pad_right = kernel_size // 2 - self.stride = stride - self.register_buffer("filter", kaiser_sinc_filter1d(cutoff, half_width, kernel_size)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - num_channels = hidden_states.shape[1] - hidden_states = F.pad(hidden_states, (self.pad_left, self.pad_right), mode="replicate") - return F.conv1d( - hidden_states, self.filter.expand(num_channels, -1, -1), stride=self.stride, groups=num_channels - ) - - -class MiniMaxH3AudioUpSample1d(nn.Module): - r"""Anti-aliased `ratio`x upsampler (transposed depthwise Kaiser-sinc convolution).""" - - def __init__(self, ratio: int, kernel_size: int): - super().__init__() - self.ratio = ratio - self.stride = ratio - self.pad = kernel_size // ratio - 1 - self.pad_left = self.pad * self.stride + (kernel_size - self.stride) // 2 - self.pad_right = self.pad * self.stride + (kernel_size - self.stride + 1) // 2 - self.register_buffer( - "filter", - kaiser_sinc_filter1d(cutoff=0.5 / ratio, half_width=0.6 / ratio, kernel_size=kernel_size), - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - num_channels = hidden_states.shape[1] - hidden_states = F.pad(hidden_states, (self.pad, self.pad), mode="replicate") - hidden_states = self.ratio * F.conv_transpose1d( - hidden_states, self.filter.expand(num_channels, -1, -1), stride=self.stride, groups=num_channels - ) - return hidden_states[..., self.pad_left : -self.pad_right] - - -class MiniMaxH3AudioDownSample1d(nn.Module): - r"""Anti-aliased `ratio`x downsampler.""" - - def __init__(self, ratio: int, kernel_size: int): - super().__init__() - self.lowpass = MiniMaxH3AudioLowPassFilter1d( - cutoff=0.5 / ratio, half_width=0.6 / ratio, stride=ratio, kernel_size=kernel_size - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return self.lowpass(hidden_states) - - -class MiniMaxH3AudioActivation1d(nn.Module): - r"""Upsample -> activation -> downsample: the alias-free activation wrapper used by BigVGAN.""" - - def __init__(self, activation: nn.Module, ratio: int = 2, kernel_size: int = 12): - super().__init__() - self.act = activation - self.upsample = MiniMaxH3AudioUpSample1d(ratio, kernel_size) - self.downsample = MiniMaxH3AudioDownSample1d(ratio, kernel_size) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.upsample(hidden_states) - hidden_states = self.act(hidden_states) - return self.downsample(hidden_states) - - -class MiniMaxH3AudioResidualUnit(nn.Module): - r"""DAC residual unit: `Snake -> dilated Conv1d(k=7) -> Snake -> Conv1d(k=1)`, plus a shortcut - that is center-cropped when the dilated convolution shrinks the time axis.""" - - def __init__(self, dim: int, dilation: int): - super().__init__() - self.block = nn.Sequential( - MiniMaxH3AudioSnake1d(dim), - _wn_conv1d(dim, dim, kernel_size=7, dilation=dilation, padding=((7 - 1) * dilation) // 2), - MiniMaxH3AudioSnake1d(dim), - _wn_conv1d(dim, dim, kernel_size=1), - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - residual = self.block(hidden_states) - pad = (hidden_states.shape[-1] - residual.shape[-1]) // 2 - if pad > 0: - hidden_states = hidden_states[..., pad:-pad] - return hidden_states + residual - - -class MiniMaxH3AudioEncoderBlock(nn.Module): - r"""Three residual units at dilations 1/3/9, then a strided channel-doubling convolution.""" - - def __init__(self, dim: int, stride: int): - super().__init__() - self.block = nn.Sequential( - MiniMaxH3AudioResidualUnit(dim // 2, dilation=1), - MiniMaxH3AudioResidualUnit(dim // 2, dilation=3), - MiniMaxH3AudioResidualUnit(dim // 2, dilation=9), - MiniMaxH3AudioSnake1d(dim // 2), - _wn_conv1d( - dim // 2, - dim, - kernel_size=2 * stride, - stride=stride, - padding=math.ceil(stride / 2), - ), - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return self.block(hidden_states) - - -class MiniMaxH3AudioEncoder(nn.Module): - r"""DAC waveform encoder: `[batch_size, 1, samples] -> [batch_size, latent_dim, samples / 800]`.""" - - def __init__(self, d_model: int, strides: tuple[int, ...], d_latent: int): - super().__init__() - block: list[nn.Module] = [_wn_conv1d(1, d_model, kernel_size=7, padding=3)] - for stride in strides: - d_model *= 2 - block.append(MiniMaxH3AudioEncoderBlock(d_model, stride=stride)) - block += [ - MiniMaxH3AudioSnake1d(d_model), - _wn_conv1d(d_model, d_latent, kernel_size=3, padding=1), - ] - self.block = nn.Sequential(*block) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return self.block(hidden_states) - - -class MiniMaxH3AudioGeGluMlp(nn.Module): - r"""Pre-norm GeGLU MLP used inside the attention projection block.""" - - def __init__(self, in_features: int, hidden_features: int): - super().__init__() - self.norm = nn.LayerNorm(in_features) - self.act = nn.GELU(approximate="tanh") - self.w0 = nn.Linear(in_features, hidden_features) - self.w1 = nn.Linear(in_features, hidden_features) - self.w2 = nn.Linear(hidden_features, in_features) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.norm(hidden_states) - hidden_states = self.act(self.w0(hidden_states)) * self.w1(hidden_states) - return self.w2(hidden_states) - - -class MiniMaxH3AudioAttnProcessor: - r"""Processor of [`MiniMaxH3AudioCausalAttention`]. - - The causal mask is expressed as `is_causal=True` rather than as a materialized mask. Every - attention backend honours that flag, with two exceptions: `_native_npu`, whose kernel takes no - causal argument and would compute *bidirectional* attention, and context parallelism, which - raises for causal attention. - """ - - _attention_backend = None - _parallel_config = None - - def __call__(self, attn: "MiniMaxH3AudioCausalAttention", hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, seq_len, _ = hidden_states.shape - qkv = F.linear( - input=hidden_states, - weight=attn.qkv.weight, - bias=torch.cat((attn.q_bias, attn.zero_k_bias, attn.v_bias)), - ) - query, key, value = ( - qkv.reshape(batch_size, seq_len, 3, attn.num_heads, attn.head_dim).permute(2, 0, 1, 3, 4).unbind(0) - ) - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=None, - is_causal=True, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - # The heads are mean-pooled away instead of being concatenated, and the head dimension that - # remains is adaptively average-pooled down to `out_dim`. - hidden_states = torch.mean(hidden_states, dim=2) - hidden_states = F.adaptive_avg_pool1d(hidden_states, attn.out_dim) - return attn.proj(hidden_states) - - -class MiniMaxH3AudioCausalAttention(nn.Module, AttentionModuleMixin): - r"""Causal self-attention that narrows the feature width from `in_dim` to `out_dim`. - - QKV is a single bias-less `nn.Linear`; query and value biases are separate parameters and the key - bias is a frozen zero buffer (`zero_k_bias`), exactly as stored in the checkpoint. Heads are - `in_dim // num_heads` wide; instead of being concatenated they are **mean-pooled away**, and the - remaining head dimension is adaptively average-pooled down to `out_dim`. - """ - - _default_processor_cls = MiniMaxH3AudioAttnProcessor - _available_processors = [MiniMaxH3AudioAttnProcessor] - # The checkpoint stores one fused `qkv` projection, so there is nothing to fuse. - _supports_qkv_fusion = False - - def __init__(self, in_dim: int, out_dim: int, num_heads: int): - super().__init__() - self.out_dim = out_dim - self.num_heads = num_heads - self.head_dim = in_dim // num_heads - self.qkv = nn.Linear(in_dim, in_dim * 3, bias=False) - self.q_bias = nn.Parameter(torch.zeros(in_dim)) - self.v_bias = nn.Parameter(torch.zeros(in_dim)) - self.register_buffer("zero_k_bias", torch.zeros(in_dim)) - self.proj = nn.Linear(out_dim, out_dim) - - self.set_processor(MiniMaxH3AudioAttnProcessor()) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return self.processor(self, hidden_states) - - -class MiniMaxH3AudioAttnProjection(nn.Module): - r"""`pre_block`: residual causal-attention + GeGLU block that rewires `latent_dim` -> `latent_channels`.""" - - def __init__(self, in_dim: int, out_dim: int, num_heads: int, mlp_ratio: int = 2): - super().__init__() - self.norm1 = nn.LayerNorm(in_dim) - self.attn = MiniMaxH3AudioCausalAttention(in_dim, out_dim, num_heads) - self.proj = nn.Linear(in_dim, out_dim) - self.norm3 = nn.LayerNorm(in_dim) - self.norm2 = nn.LayerNorm(out_dim) - self.mlp = MiniMaxH3AudioGeGluMlp(in_features=out_dim, hidden_features=out_dim * mlp_ratio) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.proj(self.norm3(hidden_states)) + self.attn(self.norm1(hidden_states)) - return hidden_states + self.mlp(self.norm2(hidden_states)) - - -class MiniMaxH3AudioAMPBlock(nn.Module): - r"""BigVGAN anti-aliased multi-periodicity block (`AMPBlock1`). - - Each dilation contributes a `(dilated conv, dilation-1 conv)` pair, and every convolution is - preceded by its own alias-free SnakeBeta activation. - """ - - def __init__(self, channels: int, kernel_size: int, dilation: tuple[int, ...]): - super().__init__() - self.convs1 = nn.ModuleList( - [ - _wn_conv1d(channels, channels, kernel_size, dilation=d, padding=(kernel_size * d - d) // 2) - for d in dilation - ] - ) - self.convs2 = nn.ModuleList( - [_wn_conv1d(channels, channels, kernel_size, dilation=1, padding=(kernel_size - 1) // 2) for _ in dilation] - ) - self.activations = nn.ModuleList( - [ - MiniMaxH3AudioActivation1d(activation=MiniMaxH3AudioSnakeBeta(channels)) - for _ in range(2 * len(dilation)) - ] - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - acts1, acts2 = self.activations[::2], self.activations[1::2] - for conv1, conv2, act1, act2 in zip(self.convs1, self.convs2, acts1, acts2): - residual = conv1(act1(hidden_states)) - residual = conv2(act2(residual)) - hidden_states = residual + hidden_states - return hidden_states - - -class MiniMaxH3AudioBigVGANDecoder(nn.Module): - r"""BigVGAN decoder: `[batch_size, latent_dim, num_frames] -> [batch_size, 1, num_frames * 800]`.""" - - def __init__( - self, - in_channels: int, - upsample_initial_channel: int, - upsample_rates: tuple[int, ...], - upsample_kernel_sizes: tuple[int, ...], - resblock_kernel_sizes: tuple[int, ...], - resblock_dilation_sizes: tuple[tuple[int, ...], ...], - ): - super().__init__() - self.num_kernels = len(resblock_kernel_sizes) - self.num_upsamples = len(upsample_rates) - - self.conv_pre = _wn_conv1d(in_channels, upsample_initial_channel, 7, 1, padding=3) - - # Each upsampler is wrapped in a one-element `ModuleList` in the original checkpoint - # (`ups..0`); the extra nesting is kept so the state dict stays a passthrough. - self.ups = nn.ModuleList() - for i, (rate, kernel) in enumerate(zip(upsample_rates, upsample_kernel_sizes)): - self.ups.append( - nn.ModuleList( - [ - weight_norm( - nn.ConvTranspose1d( - upsample_initial_channel // (2**i), - upsample_initial_channel // (2 ** (i + 1)), - kernel, - rate, - padding=(kernel - rate) // 2, - ) - ) - ] - ) - ) - - self.resblocks = nn.ModuleList() - for i in range(self.num_upsamples): - channels = upsample_initial_channel // (2 ** (i + 1)) - for kernel, dilation in zip(resblock_kernel_sizes, resblock_dilation_sizes): - self.resblocks.append(MiniMaxH3AudioAMPBlock(channels, kernel, tuple(dilation))) - - self.activation_post = MiniMaxH3AudioActivation1d(activation=MiniMaxH3AudioSnakeBeta(channels)) - self.conv_post = _wn_conv1d(channels, 1, 7, 1, padding=3, bias=False) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.conv_pre(hidden_states) - - for i in range(self.num_upsamples): - hidden_states = self.ups[i][0](hidden_states) - residual = None - for j in range(self.num_kernels): - block = self.resblocks[i * self.num_kernels + j](hidden_states) - residual = block if residual is None else residual + block - hidden_states = residual / self.num_kernels - - hidden_states = self.activation_post(hidden_states) - hidden_states = self.conv_post(hidden_states) - return torch.clamp(hidden_states, min=-1.0, max=1.0) - - -class AutoencoderKLMiniMaxH3Audio(ModelMixin, ConfigMixin, AttentionMixin): - r""" - The audio autoencoder used by [MiniMax-H3](https://huggingface.co/MiniMaxAI): a DAC-lineage - convolutional encoder and a BigVGAN decoder, operating directly on mono 32 kHz waveforms. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for the generic methods the library - implements for all models (such as downloading or saving). - - Args: - encoder_dim (`int`, defaults to `64`): - Channel width of the encoder's first convolution; doubles at every downsampling stage. - encoder_rates (`tuple[int]`, defaults to `(2, 4, 4, 5, 5)`): - Encoder strides. Their product (`800`) is the hop length, i.e. 40 latents/s at 32 kHz. - latent_dim (`int`, defaults to `2048`): - Width of the encoder trunk and of the decoder input, before/after the latent projections. - latent_channels (`int`, defaults to `32`): - Width of the diffusion latent, i.e. the `mean_proj` / `logs_proj` output channels. - num_attention_heads (`int`, defaults to `8`): - Number of heads in the causal-attention projection `pre_block`. - decoder_dim (`int`, defaults to `1024`): - BigVGAN initial channel count; halved at every upsampling stage. - decoder_rates (`tuple[int]`, defaults to `(5, 5, 2, 2, 2, 2, 2)`): - BigVGAN upsampling rates. Their product must equal `prod(encoder_rates)`. - decoder_kernel_sizes (`tuple[int]`, defaults to `(9, 9, 4, 4, 4, 4, 4)`): - Transposed-convolution kernel size per upsampling stage. - resblock_kernel_sizes (`tuple[int]`, defaults to `(3, 7, 11)`): - Kernel sizes of the parallel AMP residual blocks at each upsampling stage. - resblock_dilation_sizes (`tuple[tuple[int]]`, defaults to `((1, 3, 5), (1, 3, 5), (1, 3, 5))`): - Per-AMP-block dilations. - sampling_rate (`int`, defaults to `32000`): - Waveform sampling rate. - latents_mean (`list[float]`, *optional*): - Per-channel latent mean the pipeline uses to normalize / denormalize latents. - latents_std (`list[float]`, *optional*): - Per-channel latent standard deviation the pipeline uses to normalize / denormalize latents. - """ - - _supports_gradient_checkpointing = False - # The released checkpoint is float32 and the DAC/BigVGAN stack (weight-normalized convolutions, Snake - # activations) degrades audibly under bfloat16 (roughly 20 dB quieter decodes), so a pipeline-level - # `torch_dtype=torch.bfloat16` must not downcast the weights. - _keep_in_fp32_modules = ["encoder", "decoder", "pre_block", "dec_in_proj", "mean_proj", "logs_proj"] - - @register_to_config - def __init__( - self, - encoder_dim: int = 64, - encoder_rates: tuple[int, ...] = (2, 4, 4, 5, 5), - latent_dim: int = 2048, - latent_channels: int = 32, - num_attention_heads: int = 8, - decoder_dim: int = 1024, - decoder_rates: tuple[int, ...] = (5, 5, 2, 2, 2, 2, 2), - decoder_kernel_sizes: tuple[int, ...] = (9, 9, 4, 4, 4, 4, 4), - resblock_kernel_sizes: tuple[int, ...] = (3, 7, 11), - resblock_dilation_sizes: tuple[tuple[int, ...], ...] = ((1, 3, 5), (1, 3, 5), (1, 3, 5)), - sampling_rate: int = 32000, - latents_mean: list[float] | None = None, - latents_std: list[float] | None = None, - ): - super().__init__() - - encoder_rates = tuple(int(rate) for rate in encoder_rates) - decoder_rates = tuple(int(rate) for rate in decoder_rates) - self.hop_length = math.prod(encoder_rates) - if math.prod(decoder_rates) != self.hop_length: - raise ValueError( - f"`decoder_rates` must upsample by the encoder hop length {self.hop_length}, got " - f"{math.prod(decoder_rates)}." - ) - if latent_dim % latent_channels != 0: - raise ValueError( - f"`latent_dim` ({latent_dim}) must be a multiple of `latent_channels` ({latent_channels})." - ) - - self.encoder = MiniMaxH3AudioEncoder(d_model=encoder_dim, strides=encoder_rates, d_latent=latent_dim) - self.pre_block = MiniMaxH3AudioAttnProjection(latent_dim, latent_channels, num_heads=num_attention_heads) - self.mean_proj = nn.Conv1d(latent_channels, latent_channels, 1) - self.logs_proj = nn.Conv1d(latent_channels, latent_channels, 1) - - self.dec_in_proj = nn.Conv1d(latent_channels, latent_dim, 1) - self.decoder = MiniMaxH3AudioBigVGANDecoder( - in_channels=latent_dim, - upsample_initial_channel=decoder_dim, - upsample_rates=decoder_rates, - upsample_kernel_sizes=tuple(int(kernel) for kernel in decoder_kernel_sizes), - resblock_kernel_sizes=tuple(int(kernel) for kernel in resblock_kernel_sizes), - resblock_dilation_sizes=tuple(tuple(int(d) for d in dilation) for dilation in resblock_dilation_sizes), - ) - - @apply_forward_hook - def encode( - self, sample: torch.Tensor, return_dict: bool = True - ) -> MiniMaxH3AudioEncoderOutput | tuple[MiniMaxH3AudioDiagonalGaussianDistribution]: - r""" - Encode a waveform into the audio latent posterior. - - The waveform is right-padded to a multiple of `hop_length` (800 samples) first. MiniMax-H3 - always consumes the posterior **mean** (`latent_dist.mode()`) — the `logs_proj` head is never - evaluated by the reference pipeline. - - Args: - sample (`torch.Tensor`): - Mono waveform of shape `[batch_size, 1, samples]`. MiniMax-H3 passes the two stereo - channels of a reference clip as `batch_size = 2`. - return_dict (`bool`, defaults to `True`): - Whether to return a [`MiniMaxH3AudioEncoderOutput`] instead of a plain tuple. - - Returns: - [`MiniMaxH3AudioEncoderOutput`] or `tuple`: - The latent posterior over `[batch_size, latent_channels, samples / 800]`. - """ - if sample.ndim != 3 or sample.shape[1] != 1: - raise ValueError(f"`sample` must have shape [batch_size, 1, samples], got {tuple(sample.shape)}.") - - right_pad = math.ceil(sample.shape[-1] / self.hop_length) * self.hop_length - sample.shape[-1] - if right_pad > 0: - sample = F.pad(sample, (0, right_pad)) - - encoder_dtype = get_parameter_dtype(self.encoder) - hidden_states = self.encoder(sample.to(encoder_dtype)) - hidden_states = self.pre_block(hidden_states.transpose(1, 2)).transpose(1, 2) - mean, logs = self.mean_proj(hidden_states), self.logs_proj(hidden_states) - if encoder_dtype != torch.float32: - mean, logs = mean.float(), logs.float() - - posterior = MiniMaxH3AudioDiagonalGaussianDistribution(mean, logs) - if not return_dict: - return (posterior,) - return MiniMaxH3AudioEncoderOutput(latent_dist=posterior) - - @apply_forward_hook - def decode(self, latents: torch.Tensor, return_dict: bool = True) -> DecoderOutput | tuple[torch.Tensor]: - r""" - Decode audio latents into a waveform. - - Args: - latents (`torch.Tensor`): - Denormalized latents of shape `[batch_size, latent_channels, num_frames]`. MiniMax-H3 - passes the two stereo channels as `batch_size = 2`. - return_dict (`bool`, defaults to `True`): - Whether to return a [`~models.autoencoders.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.autoencoders.vae.DecoderOutput`] or `tuple`: - Waveform of shape `[batch_size, 1, num_frames * 800]`, clamped to `[-1, 1]`. - """ - if latents.ndim != 3: - raise ValueError( - f"`latents` must have shape [batch_size, latent_channels, num_frames], got {tuple(latents.shape)}." - ) - - decoder_dtype = get_parameter_dtype(self.decoder) - decoded = self.decoder(self.dec_in_proj(latents.to(decoder_dtype))) - if decoder_dtype != torch.float32: - decoded = decoded.float() - - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | tuple[torch.Tensor]: - r""" - Encode then decode a waveform. - - Args: - sample (`torch.Tensor`): - Mono waveform of shape `[batch_size, 1, samples]`. - sample_posterior (`bool`, defaults to `False`): - Whether to sample the posterior instead of taking its mode. MiniMax-H3 uses the mode. - return_dict (`bool`, defaults to `True`): - Whether to return a [`~models.autoencoders.vae.DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - Generator used when `sample_posterior=True`. - - Returns: - [`~models.autoencoders.vae.DecoderOutput`] or `tuple`: - The round-tripped waveform of shape `[batch_size, 1, num_frames * 800]`, clamped to `[-1, 1]`. - """ - posterior = self.encode(sample).latent_dist - latents = posterior.sample(generator=generator) if sample_posterior else posterior.mode() - return self.decode(latents, return_dict=return_dict) diff --git a/diffusers/models/autoencoders/autoencoder_kl_mochi.py b/diffusers/models/autoencoders/autoencoder_kl_mochi.py deleted file mode 100644 index bb447015c54ddeac59b09dcccbea6a394f84c334..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_mochi.py +++ /dev/null @@ -1,1119 +0,0 @@ -# Copyright 2025 The Mochi team and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import functools - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..attention_processor import Attention, MochiVaeAttnProcessor2_0 -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .autoencoder_kl_cogvideox import CogVideoXCausalConv3d -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class MochiChunkedGroupNorm3D(nn.Module): - r""" - Applies per-frame group normalization for 5D video inputs. It also supports memory-efficient chunked group - normalization. - - Args: - num_channels (int): Number of channels expected in input - num_groups (int, optional): Number of groups to separate the channels into. Default: 32 - affine (bool, optional): If True, this module has learnable affine parameters. Default: True - chunk_size (int, optional): Size of each chunk for processing. Default: 8 - - """ - - def __init__( - self, - num_channels: int, - num_groups: int = 32, - affine: bool = True, - chunk_size: int = 8, - ): - super().__init__() - self.norm_layer = nn.GroupNorm(num_channels=num_channels, num_groups=num_groups, affine=affine) - self.chunk_size = chunk_size - - def forward(self, x: torch.Tensor = None) -> torch.Tensor: - batch_size = x.size(0) - - x = x.permute(0, 2, 1, 3, 4).flatten(0, 1) - output = torch.cat([self.norm_layer(chunk) for chunk in x.split(self.chunk_size, dim=0)], dim=0) - output = output.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) - - return output - - -class MochiResnetBlock3D(nn.Module): - r""" - A 3D ResNet block used in the Mochi model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - non_linearity (`str`, defaults to `"swish"`): - Activation function to use. - """ - - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - act_fn: str = "swish", - ): - super().__init__() - - out_channels = out_channels or in_channels - - self.in_channels = in_channels - self.out_channels = out_channels - self.nonlinearity = get_activation(act_fn) - - self.norm1 = MochiChunkedGroupNorm3D(num_channels=in_channels) - self.conv1 = CogVideoXCausalConv3d( - in_channels=in_channels, out_channels=out_channels, kernel_size=3, stride=1, pad_mode="replicate" - ) - self.norm2 = MochiChunkedGroupNorm3D(num_channels=out_channels) - self.conv2 = CogVideoXCausalConv3d( - in_channels=out_channels, out_channels=out_channels, kernel_size=3, stride=1, pad_mode="replicate" - ) - - def forward( - self, - inputs: torch.Tensor, - conv_cache: dict[str, torch.Tensor] | None = None, - ) -> torch.Tensor: - new_conv_cache = {} - conv_cache = conv_cache or {} - - hidden_states = inputs - - hidden_states = self.norm1(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states, new_conv_cache["conv1"] = self.conv1(hidden_states, conv_cache=conv_cache.get("conv1")) - - hidden_states = self.norm2(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states, new_conv_cache["conv2"] = self.conv2(hidden_states, conv_cache=conv_cache.get("conv2")) - - hidden_states = hidden_states + inputs - return hidden_states, new_conv_cache - - -class MochiDownBlock3D(nn.Module): - r""" - An downsampling block used in the Mochi model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - num_layers (`int`, defaults to `1`): - Number of resnet blocks in the block. - temporal_expansion (`int`, defaults to `2`): - Temporal expansion factor. - spatial_expansion (`int`, defaults to `2`): - Spatial expansion factor. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int = 1, - temporal_expansion: int = 2, - spatial_expansion: int = 2, - add_attention: bool = True, - ): - super().__init__() - self.temporal_expansion = temporal_expansion - self.spatial_expansion = spatial_expansion - - self.conv_in = CogVideoXCausalConv3d( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=(temporal_expansion, spatial_expansion, spatial_expansion), - stride=(temporal_expansion, spatial_expansion, spatial_expansion), - pad_mode="replicate", - ) - - resnets = [] - norms = [] - attentions = [] - for _ in range(num_layers): - resnets.append(MochiResnetBlock3D(in_channels=out_channels)) - if add_attention: - norms.append(MochiChunkedGroupNorm3D(num_channels=out_channels)) - attentions.append( - Attention( - query_dim=out_channels, - heads=out_channels // 32, - dim_head=32, - qk_norm="l2", - is_causal=True, - processor=MochiVaeAttnProcessor2_0(), - ) - ) - else: - norms.append(None) - attentions.append(None) - - self.resnets = nn.ModuleList(resnets) - self.norms = nn.ModuleList(norms) - self.attentions = nn.ModuleList(attentions) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - conv_cache: dict[str, torch.Tensor] | None = None, - chunk_size: int = 2**15, - ) -> torch.Tensor: - r"""Forward method of the `MochiUpBlock3D` class.""" - - new_conv_cache = {} - conv_cache = conv_cache or {} - - hidden_states, new_conv_cache["conv_in"] = self.conv_in(hidden_states) - - for i, (resnet, norm, attn) in enumerate(zip(self.resnets, self.norms, self.attentions)): - conv_cache_key = f"resnet_{i}" - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, new_conv_cache[conv_cache_key] = self._gradient_checkpointing_func( - resnet, - hidden_states, - conv_cache.get(conv_cache_key), - ) - else: - hidden_states, new_conv_cache[conv_cache_key] = resnet( - hidden_states, conv_cache=conv_cache.get(conv_cache_key) - ) - - if attn is not None: - residual = hidden_states - hidden_states = norm(hidden_states) - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - hidden_states = hidden_states.permute(0, 3, 4, 2, 1).flatten(0, 2).contiguous() - - # Perform attention in chunks to avoid following error: - # RuntimeError: CUDA error: invalid configuration argument - if hidden_states.size(0) <= chunk_size: - hidden_states = attn(hidden_states) - else: - hidden_states_chunks = [] - for i in range(0, hidden_states.size(0), chunk_size): - hidden_states_chunk = hidden_states[i : i + chunk_size] - hidden_states_chunk = attn(hidden_states_chunk) - hidden_states_chunks.append(hidden_states_chunk) - hidden_states = torch.cat(hidden_states_chunks) - - hidden_states = hidden_states.unflatten(0, (batch_size, height, width)).permute(0, 4, 3, 1, 2) - - hidden_states = residual + hidden_states - - return hidden_states, new_conv_cache - - -class MochiMidBlock3D(nn.Module): - r""" - A middle block used in the Mochi model. - - Args: - in_channels (`int`): - Number of input channels. - num_layers (`int`, defaults to `3`): - Number of resnet blocks in the block. - """ - - def __init__( - self, - in_channels: int, # 768 - num_layers: int = 3, - add_attention: bool = True, - ): - super().__init__() - - resnets = [] - norms = [] - attentions = [] - - for _ in range(num_layers): - resnets.append(MochiResnetBlock3D(in_channels=in_channels)) - - if add_attention: - norms.append(MochiChunkedGroupNorm3D(num_channels=in_channels)) - attentions.append( - Attention( - query_dim=in_channels, - heads=in_channels // 32, - dim_head=32, - qk_norm="l2", - is_causal=True, - processor=MochiVaeAttnProcessor2_0(), - ) - ) - else: - norms.append(None) - attentions.append(None) - - self.resnets = nn.ModuleList(resnets) - self.norms = nn.ModuleList(norms) - self.attentions = nn.ModuleList(attentions) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - conv_cache: dict[str, torch.Tensor] | None = None, - ) -> torch.Tensor: - r"""Forward method of the `MochiMidBlock3D` class.""" - - new_conv_cache = {} - conv_cache = conv_cache or {} - - for i, (resnet, norm, attn) in enumerate(zip(self.resnets, self.norms, self.attentions)): - conv_cache_key = f"resnet_{i}" - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, new_conv_cache[conv_cache_key] = self._gradient_checkpointing_func( - resnet, hidden_states, conv_cache.get(conv_cache_key) - ) - else: - hidden_states, new_conv_cache[conv_cache_key] = resnet( - hidden_states, conv_cache=conv_cache.get(conv_cache_key) - ) - - if attn is not None: - residual = hidden_states - hidden_states = norm(hidden_states) - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - hidden_states = hidden_states.permute(0, 3, 4, 2, 1).flatten(0, 2).contiguous() - hidden_states = attn(hidden_states) - hidden_states = hidden_states.unflatten(0, (batch_size, height, width)).permute(0, 4, 3, 1, 2) - - hidden_states = residual + hidden_states - - return hidden_states, new_conv_cache - - -class MochiUpBlock3D(nn.Module): - r""" - An upsampling block used in the Mochi model. - - Args: - in_channels (`int`): - Number of input channels. - out_channels (`int`, *optional*): - Number of output channels. If None, defaults to `in_channels`. - num_layers (`int`, defaults to `1`): - Number of resnet blocks in the block. - temporal_expansion (`int`, defaults to `2`): - Temporal expansion factor. - spatial_expansion (`int`, defaults to `2`): - Spatial expansion factor. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int = 1, - temporal_expansion: int = 2, - spatial_expansion: int = 2, - ): - super().__init__() - self.temporal_expansion = temporal_expansion - self.spatial_expansion = spatial_expansion - - resnets = [] - for _ in range(num_layers): - resnets.append(MochiResnetBlock3D(in_channels=in_channels)) - self.resnets = nn.ModuleList(resnets) - - self.proj = nn.Linear(in_channels, out_channels * temporal_expansion * spatial_expansion**2) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - conv_cache: dict[str, torch.Tensor] | None = None, - ) -> torch.Tensor: - r"""Forward method of the `MochiUpBlock3D` class.""" - - new_conv_cache = {} - conv_cache = conv_cache or {} - - for i, resnet in enumerate(self.resnets): - conv_cache_key = f"resnet_{i}" - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, new_conv_cache[conv_cache_key] = self._gradient_checkpointing_func( - resnet, - hidden_states, - conv_cache.get(conv_cache_key), - ) - else: - hidden_states, new_conv_cache[conv_cache_key] = resnet( - hidden_states, conv_cache=conv_cache.get(conv_cache_key) - ) - - hidden_states = hidden_states.permute(0, 2, 3, 4, 1) - hidden_states = self.proj(hidden_states) - hidden_states = hidden_states.permute(0, 4, 1, 2, 3) - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - st = self.temporal_expansion - sh = self.spatial_expansion - sw = self.spatial_expansion - - # Reshape and unpatchify - hidden_states = hidden_states.view(batch_size, -1, st, sh, sw, num_frames, height, width) - hidden_states = hidden_states.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous() - hidden_states = hidden_states.view(batch_size, -1, num_frames * st, height * sh, width * sw) - - return hidden_states, new_conv_cache - - -class FourierFeatures(nn.Module): - def __init__(self, start: int = 6, stop: int = 8, step: int = 1): - super().__init__() - - self.start = start - self.stop = stop - self.step = step - - def forward(self, inputs: torch.Tensor) -> torch.Tensor: - r"""Forward method of the `FourierFeatures` class.""" - original_dtype = inputs.dtype - inputs = inputs.to(torch.float32) - num_channels = inputs.shape[1] - num_freqs = (self.stop - self.start) // self.step - - freqs = torch.arange(self.start, self.stop, self.step, dtype=inputs.dtype, device=inputs.device) - w = torch.pow(2.0, freqs) * (2 * torch.pi) # [num_freqs] - w = w.repeat(num_channels)[None, :, None, None, None] # [1, num_channels * num_freqs, 1, 1, 1] - - # Interleaved repeat of input channels to match w - h = inputs.repeat_interleave( - num_freqs, dim=1, output_size=inputs.shape[1] * num_freqs - ) # [B, C * num_freqs, T, H, W] - # Scale channels by frequency. - h = w * h - - return torch.cat([inputs, torch.sin(h), torch.cos(h)], dim=1).to(original_dtype) - - -class MochiEncoder3D(nn.Module): - r""" - The `MochiEncoder3D` layer of a variational autoencoder that encodes input video samples to its latent - representation. - - Args: - in_channels (`int`, *optional*): - The number of input channels. - out_channels (`int`, *optional*): - The number of output channels. - block_out_channels (`tuple[int, ...]`, *optional*, defaults to `(128, 256, 512, 768)`): - The number of output channels for each block. - layers_per_block (`tuple[int, ...]`, *optional*, defaults to `(3, 3, 4, 6, 3)`): - The number of resnet blocks for each block. - temporal_expansions (`tuple[int, ...]`, *optional*, defaults to `(1, 2, 3)`): - The temporal expansion factor for each of the up blocks. - spatial_expansions (`tuple[int, ...]`, *optional*, defaults to `(2, 2, 2)`): - The spatial expansion factor for each of the up blocks. - non_linearity (`str`, *optional*, defaults to `"swish"`): - The non-linearity to use in the decoder. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - block_out_channels: tuple[int, ...] = (128, 256, 512, 768), - layers_per_block: tuple[int, ...] = (3, 3, 4, 6, 3), - temporal_expansions: tuple[int, ...] = (1, 2, 3), - spatial_expansions: tuple[int, ...] = (2, 2, 2), - add_attention_block: tuple[bool, ...] = (False, True, True, True, True), - act_fn: str = "swish", - ): - super().__init__() - - self.nonlinearity = get_activation(act_fn) - - self.fourier_features = FourierFeatures() - self.proj_in = nn.Linear(in_channels, block_out_channels[0]) - self.block_in = MochiMidBlock3D( - in_channels=block_out_channels[0], num_layers=layers_per_block[0], add_attention=add_attention_block[0] - ) - - down_blocks = [] - for i in range(len(block_out_channels) - 1): - down_block = MochiDownBlock3D( - in_channels=block_out_channels[i], - out_channels=block_out_channels[i + 1], - num_layers=layers_per_block[i + 1], - temporal_expansion=temporal_expansions[i], - spatial_expansion=spatial_expansions[i], - add_attention=add_attention_block[i + 1], - ) - down_blocks.append(down_block) - self.down_blocks = nn.ModuleList(down_blocks) - - self.block_out = MochiMidBlock3D( - in_channels=block_out_channels[-1], num_layers=layers_per_block[-1], add_attention=add_attention_block[-1] - ) - self.norm_out = MochiChunkedGroupNorm3D(block_out_channels[-1]) - self.proj_out = nn.Linear(block_out_channels[-1], 2 * out_channels, bias=False) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor, conv_cache: dict[str, torch.Tensor] | None = None) -> torch.Tensor: - r"""Forward method of the `MochiEncoder3D` class.""" - - new_conv_cache = {} - conv_cache = conv_cache or {} - - hidden_states = self.fourier_features(hidden_states) - - hidden_states = hidden_states.permute(0, 2, 3, 4, 1) - hidden_states = self.proj_in(hidden_states) - hidden_states = hidden_states.permute(0, 4, 1, 2, 3) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, new_conv_cache["block_in"] = self._gradient_checkpointing_func( - self.block_in, hidden_states, conv_cache.get("block_in") - ) - - for i, down_block in enumerate(self.down_blocks): - conv_cache_key = f"down_block_{i}" - hidden_states, new_conv_cache[conv_cache_key] = self._gradient_checkpointing_func( - down_block, hidden_states, conv_cache.get(conv_cache_key) - ) - else: - hidden_states, new_conv_cache["block_in"] = self.block_in( - hidden_states, conv_cache=conv_cache.get("block_in") - ) - - for i, down_block in enumerate(self.down_blocks): - conv_cache_key = f"down_block_{i}" - hidden_states, new_conv_cache[conv_cache_key] = down_block( - hidden_states, conv_cache=conv_cache.get(conv_cache_key) - ) - - hidden_states, new_conv_cache["block_out"] = self.block_out( - hidden_states, conv_cache=conv_cache.get("block_out") - ) - - hidden_states = self.norm_out(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - - hidden_states = hidden_states.permute(0, 2, 3, 4, 1) - hidden_states = self.proj_out(hidden_states) - hidden_states = hidden_states.permute(0, 4, 1, 2, 3) - - return hidden_states, new_conv_cache - - -class MochiDecoder3D(nn.Module): - r""" - The `MochiDecoder3D` layer of a variational autoencoder that decodes its latent representation into an output - sample. - - Args: - in_channels (`int`, *optional*): - The number of input channels. - out_channels (`int`, *optional*): - The number of output channels. - block_out_channels (`tuple[int, ...]`, *optional*, defaults to `(128, 256, 512, 768)`): - The number of output channels for each block. - layers_per_block (`tuple[int, ...]`, *optional*, defaults to `(3, 3, 4, 6, 3)`): - The number of resnet blocks for each block. - temporal_expansions (`tuple[int, ...]`, *optional*, defaults to `(1, 2, 3)`): - The temporal expansion factor for each of the up blocks. - spatial_expansions (`tuple[int, ...]`, *optional*, defaults to `(2, 2, 2)`): - The spatial expansion factor for each of the up blocks. - non_linearity (`str`, *optional*, defaults to `"swish"`): - The non-linearity to use in the decoder. - """ - - def __init__( - self, - in_channels: int, # 12 - out_channels: int, # 3 - block_out_channels: tuple[int, ...] = (128, 256, 512, 768), - layers_per_block: tuple[int, ...] = (3, 3, 4, 6, 3), - temporal_expansions: tuple[int, ...] = (1, 2, 3), - spatial_expansions: tuple[int, ...] = (2, 2, 2), - act_fn: str = "swish", - ): - super().__init__() - - self.nonlinearity = get_activation(act_fn) - - self.conv_in = nn.Conv3d(in_channels, block_out_channels[-1], kernel_size=(1, 1, 1)) - self.block_in = MochiMidBlock3D( - in_channels=block_out_channels[-1], - num_layers=layers_per_block[-1], - add_attention=False, - ) - - up_blocks = [] - for i in range(len(block_out_channels) - 1): - up_block = MochiUpBlock3D( - in_channels=block_out_channels[-i - 1], - out_channels=block_out_channels[-i - 2], - num_layers=layers_per_block[-i - 2], - temporal_expansion=temporal_expansions[-i - 1], - spatial_expansion=spatial_expansions[-i - 1], - ) - up_blocks.append(up_block) - self.up_blocks = nn.ModuleList(up_blocks) - - self.block_out = MochiMidBlock3D( - in_channels=block_out_channels[0], - num_layers=layers_per_block[0], - add_attention=False, - ) - self.proj_out = nn.Linear(block_out_channels[0], out_channels) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor, conv_cache: dict[str, torch.Tensor] | None = None) -> torch.Tensor: - r"""Forward method of the `MochiDecoder3D` class.""" - - new_conv_cache = {} - conv_cache = conv_cache or {} - - hidden_states = self.conv_in(hidden_states) - - # 1. Mid - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, new_conv_cache["block_in"] = self._gradient_checkpointing_func( - self.block_in, hidden_states, conv_cache.get("block_in") - ) - - for i, up_block in enumerate(self.up_blocks): - conv_cache_key = f"up_block_{i}" - hidden_states, new_conv_cache[conv_cache_key] = self._gradient_checkpointing_func( - up_block, hidden_states, conv_cache.get(conv_cache_key) - ) - else: - hidden_states, new_conv_cache["block_in"] = self.block_in( - hidden_states, conv_cache=conv_cache.get("block_in") - ) - - for i, up_block in enumerate(self.up_blocks): - conv_cache_key = f"up_block_{i}" - hidden_states, new_conv_cache[conv_cache_key] = up_block( - hidden_states, conv_cache=conv_cache.get(conv_cache_key) - ) - - hidden_states, new_conv_cache["block_out"] = self.block_out( - hidden_states, conv_cache=conv_cache.get("block_out") - ) - - hidden_states = self.nonlinearity(hidden_states) - - hidden_states = hidden_states.permute(0, 2, 3, 4, 1) - hidden_states = self.proj_out(hidden_states) - hidden_states = hidden_states.permute(0, 4, 1, 2, 3) - - return hidden_states, new_conv_cache - - -class AutoencoderKLMochi(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - A VAE model with KL loss for encoding images into latents and decoding latent representations into images. Used in - [Mochi 1 preview](https://github.com/genmoai/models). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - in_channels (int, *optional*, defaults to 3): Number of channels in the input image. - out_channels (int, *optional*, defaults to 3): Number of channels in the output. - block_out_channels (`tuple[int]`, *optional*, defaults to `(64,)`): - tuple of block output channels. - act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use. - scaling_factor (`float`, *optional*, defaults to `1.15258426`): - The component-wise standard deviation of the trained latent space computed using the first batch of the - training set. This is used to scale the latent space to have unit variance when training the diffusion - model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the - diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 - / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image - Synthesis with Latent Diffusion Models](https://huggingface.co/papers/2112.10752) paper. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["MochiResnetBlock3D"] - - @register_to_config - def __init__( - self, - in_channels: int = 15, - out_channels: int = 3, - encoder_block_out_channels: tuple[int] = (64, 128, 256, 384), - decoder_block_out_channels: tuple[int] = (128, 256, 512, 768), - latent_channels: int = 12, - layers_per_block: tuple[int, ...] = (3, 3, 4, 6, 3), - act_fn: str = "silu", - temporal_expansions: tuple[int, ...] = (1, 2, 3), - spatial_expansions: tuple[int, ...] = (2, 2, 2), - add_attention_block: tuple[bool, ...] = (False, True, True, True, True), - latents_mean: tuple[float, ...] = ( - -0.06730895953510081, - -0.038011381506090416, - -0.07477820912866141, - -0.05565264470995561, - 0.012767231469026969, - -0.04703542746246419, - 0.043896967884726704, - -0.09346305707025976, - -0.09918314763016893, - -0.008729793427399178, - -0.011931556316503654, - -0.0321993391887285, - ), - latents_std: tuple[float, ...] = ( - 0.9263795028493863, - 0.9248894543193766, - 0.9393059390890617, - 0.959253732819592, - 0.8244560132752793, - 0.917259975397747, - 0.9294154431013696, - 1.3720942357788521, - 0.881393668867029, - 0.9168315692124348, - 0.9185249279345552, - 0.9274757570805041, - ), - scaling_factor: float = 1.0, - ): - super().__init__() - - self.encoder = MochiEncoder3D( - in_channels=in_channels, - out_channels=latent_channels, - block_out_channels=encoder_block_out_channels, - layers_per_block=layers_per_block, - temporal_expansions=temporal_expansions, - spatial_expansions=spatial_expansions, - add_attention_block=add_attention_block, - act_fn=act_fn, - ) - self.decoder = MochiDecoder3D( - in_channels=latent_channels, - out_channels=out_channels, - block_out_channels=decoder_block_out_channels, - layers_per_block=layers_per_block, - temporal_expansions=temporal_expansions, - spatial_expansions=spatial_expansions, - act_fn=act_fn, - ) - - self.spatial_compression_ratio = functools.reduce(lambda x, y: x * y, spatial_expansions, 1) - self.temporal_compression_ratio = functools.reduce(lambda x, y: x * y, temporal_expansions, 1) - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent - # frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the - # intermediate tiles together, the memory requirement can be lowered. - self.use_tiling = False - - # When decoding temporally long video latents, the memory requirement is very high. By decoding latent frames - # at a fixed frame batch size (based on `self.num_latent_frames_batch_sizes`), the memory requirement can be lowered. - self.use_framewise_encoding = False - self.use_framewise_decoding = False - - # This can be used to determine how the number of output frames in the final decoded video. To maintain consistency with - # the original implementation, this defaults to `True`. - # - Original implementation (drop_last_temporal_frames=True): - # Output frames = (latent_frames - 1) * temporal_compression_ratio + 1 - # - Without dropping additional temporal upscaled frames (drop_last_temporal_frames=False): - # Output frames = latent_frames * temporal_compression_ratio - # The latter case is useful for frame packing and some training/finetuning scenarios where the additional. - self.drop_last_temporal_frames = True - - # This can be configured based on the amount of GPU memory available. - # `12` for sample frames and `2` for latent frames are sensible defaults for consumer GPUs. - # Setting it to higher values results in higher memory usage. - self.num_sample_frames_batch_size = 12 - self.num_latent_frames_batch_size = 2 - - # The minimal tile height and width for spatial tiling to be used - self.tile_sample_min_height = 256 - self.tile_sample_min_width = 256 - - # The minimal distance between two spatial tiles - self.tile_sample_stride_height = 192 - self.tile_sample_stride_width = 192 - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_sample_stride_height: float | None = None, - tile_sample_stride_width: float | None = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_sample_stride_height (`int`, *optional*): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. - tile_sample_stride_width (`int`, *optional*): - The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling - artifacts produced across the width dimension. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height - self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width - - def _enable_framewise_encoding(self): - r""" - Enables the framewise VAE encoding implementation with past latent padding. By default, Diffusers uses the - oneshot encoding implementation without current latent replicate padding. - - Warning: Framewise encoding may not work as expected due to the causal attention layers. If you enable - framewise encoding, encode a video, and try to decode it, there will be noticeable jittering effect. - """ - self.use_framewise_encoding = True - for name, module in self.named_modules(): - if isinstance(module, CogVideoXCausalConv3d): - module.pad_mode = "constant" - - def _enable_framewise_decoding(self): - r""" - Enables the framewise VAE decoding implementation with past latent padding. By default, Diffusers uses the - oneshot decoding implementation without current latent replicate padding. - """ - self.use_framewise_decoding = True - for name, module in self.named_modules(): - if isinstance(module, CogVideoXCausalConv3d): - module.pad_mode = "constant" - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = x.shape - - if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): - return self.tiled_encode(x) - - if self.use_framewise_encoding: - raise NotImplementedError( - "Frame-wise encoding does not work with the Mochi VAE Encoder due to the presence of attention layers. " - "As intermediate frames are not independent from each other, they cannot be encoded frame-wise." - ) - else: - enc, _ = self.encoder(x) - - return enc - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - batch_size, num_channels, num_frames, height, width = z.shape - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - - if self.use_tiling and (width > tile_latent_min_width or height > tile_latent_min_height): - return self.tiled_decode(z, return_dict=return_dict) - - if self.use_framewise_decoding: - conv_cache = None - dec = [] - - for i in range(0, num_frames, self.num_latent_frames_batch_size): - z_intermediate = z[:, :, i : i + self.num_latent_frames_batch_size] - z_intermediate, conv_cache = self.decoder(z_intermediate, conv_cache=conv_cache) - dec.append(z_intermediate) - - dec = torch.cat(dec, dim=2) - else: - dec, _ = self.decoder(z) - - if self.drop_last_temporal_frames and dec.size(2) >= self.temporal_compression_ratio: - dec = dec[:, :, self.temporal_compression_ratio - 1 :] - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - """ - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[3], b.shape[3], blend_extent) - for y in range(blend_extent): - b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[4], b.shape[4], blend_extent) - for x in range(blend_extent): - b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( - x / blend_extent - ) - return b - - def tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - r"""Encode a batch of images using a tiled encoder. - - Args: - x (`torch.Tensor`): Input batch of videos. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - batch_size, num_channels, num_frames, height, width = x.shape - latent_height = height // self.spatial_compression_ratio - latent_width = width // self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = tile_latent_min_height - tile_latent_stride_height - blend_width = tile_latent_min_width - tile_latent_stride_width - - # Split x into overlapping tiles and encode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, self.tile_sample_stride_height): - row = [] - for j in range(0, width, self.tile_sample_stride_width): - if self.use_framewise_encoding: - raise NotImplementedError( - "Frame-wise encoding does not work with the Mochi VAE Encoder due to the presence of attention layers. " - "As intermediate frames are not independent from each other, they cannot be encoded frame-wise." - ) - else: - time, _ = self.encoder( - x[:, :, :, i : i + self.tile_sample_min_height, j : j + self.tile_sample_min_width] - ) - - row.append(time) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :tile_latent_stride_height, :tile_latent_stride_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - enc = torch.cat(result_rows, dim=3)[:, :, :, :latent_height, :latent_width] - return enc - - def tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - - batch_size, num_channels, num_frames, height, width = z.shape - sample_height = height * self.spatial_compression_ratio - sample_width = width * self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = self.tile_sample_min_height - self.tile_sample_stride_height - blend_width = self.tile_sample_min_width - self.tile_sample_stride_width - - # Split z into overlapping tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, tile_latent_stride_height): - row = [] - for j in range(0, width, tile_latent_stride_width): - if self.use_framewise_decoding: - time = [] - conv_cache = None - - for k in range(0, num_frames, self.num_latent_frames_batch_size): - tile = z[ - :, - :, - k : k + self.num_latent_frames_batch_size, - i : i + tile_latent_min_height, - j : j + tile_latent_min_width, - ] - tile, conv_cache = self.decoder(tile, conv_cache=conv_cache) - time.append(tile) - - time = torch.cat(time, dim=2) - else: - time, _ = self.decoder(z[:, :, :, i : i + tile_latent_min_height, j : j + tile_latent_min_width]) - - if self.drop_last_temporal_frames and time.size(2) >= self.temporal_compression_ratio: - time = time[:, :, self.temporal_compression_ratio - 1 :] - - row.append(time) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, : self.tile_sample_stride_height, : self.tile_sample_stride_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - dec = torch.cat(result_rows, dim=3)[:, :, :, :sample_height, :sample_width] - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z, return_dict=return_dict) - return dec diff --git a/diffusers/models/autoencoders/autoencoder_kl_qwenimage.py b/diffusers/models/autoencoders/autoencoder_kl_qwenimage.py deleted file mode 100644 index 220520a12e68a8d10160c2fc0156e5a5b0309336..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_qwenimage.py +++ /dev/null @@ -1,1066 +0,0 @@ -# Copyright 2025 The Qwen-Image Team, Wan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -# -# We gratefully acknowledge the Wan Team for their outstanding contributions. -# QwenImageVAE is further fine-tuned from the Wan Video VAE to achieve improved performance. -# For more information about the Wan VAE, please refer to: -# - GitHub: https://github.com/Wan-Video/Wan2.1 -# - Paper: https://huggingface.co/papers/2503.20314 - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - -CACHE_T = 2 - - -class QwenImageCausalConv3d(nn.Conv3d): - r""" - A custom 3D causal convolution layer with feature caching support. - - This layer extends the standard Conv3D layer by ensuring causality in the time dimension and handling feature - caching for efficient inference. - - Args: - in_channels (int): Number of channels in the input image - out_channels (int): Number of channels produced by the convolution - kernel_size (int or tuple): Size of the convolving kernel - stride (int or tuple, optional): Stride of the convolution. Default: 1 - padding (int or tuple, optional): Zero-padding added to all three sides of the input. Default: 0 - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int], - stride: int | tuple[int, int, int] = 1, - padding: int | tuple[int, int, int] = 0, - ) -> None: - super().__init__( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=kernel_size, - stride=stride, - padding=padding, - ) - - # Set up causal padding - self._padding = (self.padding[2], self.padding[2], self.padding[1], self.padding[1], 2 * self.padding[0], 0) - self.padding = (0, 0, 0) - - def forward(self, x, cache_x=None): - padding = list(self._padding) - if cache_x is not None and self._padding[4] > 0: - cache_x = cache_x.to(x.device) - x = torch.cat([cache_x, x], dim=2) - padding[4] -= cache_x.shape[2] - x = F.pad(x, padding) - return super().forward(x) - - -class QwenImageRMS_norm(nn.Module): - r""" - A custom RMS normalization layer. - - Args: - dim (int): The number of dimensions to normalize over. - channel_first (bool, optional): Whether the input tensor has channels as the first dimension. - Default is True. - images (bool, optional): Whether the input represents image data. Default is True. - bias (bool, optional): Whether to include a learnable bias term. Default is False. - """ - - def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bias: bool = False) -> None: - super().__init__() - broadcastable_dims = (1, 1, 1) if not images else (1, 1) - shape = (dim, *broadcastable_dims) if channel_first else (dim,) - - self.channel_first = channel_first - self.scale = dim**0.5 - self.gamma = nn.Parameter(torch.ones(shape)) - self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - - def forward(self, x): - needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( - t in str(x.dtype) for t in ("float4_", "float8_") - ) - normalized = F.normalize(x.float() if needs_fp32_normalize else x, dim=(1 if self.channel_first else -1)).to( - x.dtype - ) - - return normalized * self.scale * self.gamma + self.bias - - -class QwenImageUpsample(nn.Upsample): - r""" - Perform upsampling while ensuring the output tensor has the same data type as the input. - - Args: - x (torch.Tensor): Input tensor to be upsampled. - - Returns: - torch.Tensor: Upsampled tensor with the same data type as the input. - """ - - def forward(self, x): - return super().forward(x.float()).type_as(x) - - -class QwenImageResample(nn.Module): - r""" - A custom resampling module for 2D and 3D data. - - Args: - dim (int): The number of input/output channels. - mode (str): The resampling mode. Must be one of: - - 'none': No resampling (identity operation). - - 'upsample2d': 2D upsampling with nearest-exact interpolation and convolution. - - 'upsample3d': 3D upsampling with nearest-exact interpolation, convolution, and causal 3D convolution. - - 'downsample2d': 2D downsampling with zero-padding and convolution. - - 'downsample3d': 3D downsampling with zero-padding, convolution, and causal 3D convolution. - """ - - def __init__(self, dim: int, mode: str) -> None: - super().__init__() - self.dim = dim - self.mode = mode - - # layers - if mode == "upsample2d": - self.resample = nn.Sequential( - QwenImageUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), - nn.Conv2d(dim, dim // 2, 3, padding=1), - ) - elif mode == "upsample3d": - self.resample = nn.Sequential( - QwenImageUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), - nn.Conv2d(dim, dim // 2, 3, padding=1), - ) - self.time_conv = QwenImageCausalConv3d(dim, dim * 2, (3, 1, 1), padding=(1, 0, 0)) - - elif mode == "downsample2d": - self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2))) - elif mode == "downsample3d": - self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2))) - self.time_conv = QwenImageCausalConv3d(dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0)) - - else: - self.resample = nn.Identity() - - def forward(self, x, feat_cache=None, feat_idx=[0]): - b, c, t, h, w = x.size() - if self.mode == "upsample3d": - if feat_cache is not None: - idx = feat_idx[0] - if feat_cache[idx] is None: - feat_cache[idx] = "Rep" - feat_idx[0] += 1 - else: - cache_x = x[:, :, -min(CACHE_T, x.shape[2]) :, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] != "Rep": - # cache last frame of last two chunk - cache_x = torch.cat( - [feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2 - ) - if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] == "Rep": - cache_x = torch.cat([torch.zeros_like(cache_x).to(cache_x.device), cache_x], dim=2) - if feat_cache[idx] == "Rep": - x = self.time_conv(x) - else: - x = self.time_conv(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - - x = x.reshape(b, 2, c, t, h, w) - x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3) - x = x.reshape(b, c, t * 2, h, w) - t = x.shape[2] - x = x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w) - x = self.resample(x) - x = x.view(b, t, x.size(1), x.size(2), x.size(3)).permute(0, 2, 1, 3, 4) - - if self.mode == "downsample3d": - if feat_cache is not None: - idx = feat_idx[0] - if feat_cache[idx] is None: - feat_cache[idx] = x.clone() - feat_idx[0] += 1 - else: - cache_x = x[:, :, -1:, :, :].clone() - x = self.time_conv(torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2)) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - return x - - -class QwenImageResidualBlock(nn.Module): - r""" - A custom residual block module. - - Args: - in_dim (int): Number of input channels. - out_dim (int): Number of output channels. - dropout (float, optional): Dropout rate for the dropout layer. Default is 0.0. - non_linearity (str, optional): Type of non-linearity to use. Default is "silu". - """ - - def __init__( - self, - in_dim: int, - out_dim: int, - dropout: float = 0.0, - non_linearity: str = "silu", - ) -> None: - super().__init__() - self.in_dim = in_dim - self.out_dim = out_dim - self.nonlinearity = get_activation(non_linearity) - - # layers - self.norm1 = QwenImageRMS_norm(in_dim, images=False) - self.conv1 = QwenImageCausalConv3d(in_dim, out_dim, 3, padding=1) - self.norm2 = QwenImageRMS_norm(out_dim, images=False) - self.dropout = nn.Dropout(dropout) - self.conv2 = QwenImageCausalConv3d(out_dim, out_dim, 3, padding=1) - self.conv_shortcut = QwenImageCausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity() - - def forward(self, x, feat_cache=None, feat_idx=[0]): - # Apply shortcut connection - h = self.conv_shortcut(x) - - # First normalization and activation - x = self.norm1(x) - x = self.nonlinearity(x) - - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -min(CACHE_T, x.shape[2]) :, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - - x = self.conv1(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv1(x) - - # Second normalization and activation - x = self.norm2(x) - x = self.nonlinearity(x) - - # Dropout - x = self.dropout(x) - - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -min(CACHE_T, x.shape[2]) :, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - - x = self.conv2(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv2(x) - - # Add residual connection - return x + h - - -class QwenImageAttentionBlock(nn.Module): - r""" - Causal self-attention with a single head. - - Args: - dim (int): The number of channels in the input tensor. - """ - - def __init__(self, dim): - super().__init__() - self.dim = dim - - # layers - self.norm = QwenImageRMS_norm(dim) - self.to_qkv = nn.Conv2d(dim, dim * 3, 1) - self.proj = nn.Conv2d(dim, dim, 1) - - def forward(self, x): - identity = x - batch_size, channels, time, height, width = x.size() - - x = x.permute(0, 2, 1, 3, 4).reshape(batch_size * time, channels, height, width) - x = self.norm(x) - - # compute query, key, value - qkv = self.to_qkv(x) - qkv = qkv.reshape(batch_size * time, 1, channels * 3, -1) - qkv = qkv.permute(0, 1, 3, 2).contiguous() - q, k, v = qkv.chunk(3, dim=-1) - - # apply attention - x = F.scaled_dot_product_attention(q, k, v) - - x = x.squeeze(1).permute(0, 2, 1).reshape(batch_size * time, channels, height, width) - - # output projection - x = self.proj(x) - - # Reshape back: [(b*t), c, h, w] -> [b, c, t, h, w] - x = x.view(batch_size, time, channels, height, width) - x = x.permute(0, 2, 1, 3, 4) - - return x + identity - - -class QwenImageMidBlock(nn.Module): - """ - Middle block for QwenImageVAE encoder and decoder. - - Args: - dim (int): Number of input/output channels. - dropout (float): Dropout rate. - non_linearity (str): Type of non-linearity to use. - """ - - def __init__(self, dim: int, dropout: float = 0.0, non_linearity: str = "silu", num_layers: int = 1): - super().__init__() - self.dim = dim - - # Create the components - resnets = [QwenImageResidualBlock(dim, dim, dropout, non_linearity)] - attentions = [] - for _ in range(num_layers): - attentions.append(QwenImageAttentionBlock(dim)) - resnets.append(QwenImageResidualBlock(dim, dim, dropout, non_linearity)) - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward(self, x, feat_cache=None, feat_idx=[0]): - # First residual block - x = self.resnets[0](x, feat_cache, feat_idx) - - # Process through attention and residual blocks - for attn, resnet in zip(self.attentions, self.resnets[1:]): - if attn is not None: - x = attn(x) - - x = resnet(x, feat_cache, feat_idx) - - return x - - -class QwenImageEncoder3d(nn.Module): - r""" - A 3D encoder module. - - Args: - dim (int): The base number of channels in the first layer. - z_dim (int): The dimensionality of the latent space. - dim_mult (list of int): Multipliers for the number of channels in each block. - num_res_blocks (int): Number of residual blocks in each block. - attn_scales (list of float): Scales at which to apply attention mechanisms. - temperal_downsample (list of bool): Whether to downsample temporally in each block. - dropout (float): Dropout rate for the dropout layers. - non_linearity (str): Type of non-linearity to use. - """ - - def __init__( - self, - dim=128, - z_dim=4, - dim_mult=[1, 2, 4, 4], - num_res_blocks=2, - attn_scales=[], - temperal_downsample=[True, True, False], - dropout=0.0, - input_channels=3, - non_linearity: str = "silu", - ): - super().__init__() - self.dim = dim - self.z_dim = z_dim - self.dim_mult = dim_mult - self.num_res_blocks = num_res_blocks - self.attn_scales = attn_scales - self.temperal_downsample = temperal_downsample - self.nonlinearity = get_activation(non_linearity) - - # dimensions - dims = [dim * u for u in [1] + dim_mult] - scale = 1.0 - - # init block - self.conv_in = QwenImageCausalConv3d(input_channels, dims[0], 3, padding=1) - - # downsample blocks - self.down_blocks = nn.ModuleList([]) - for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): - # residual (+attention) blocks - for _ in range(num_res_blocks): - self.down_blocks.append(QwenImageResidualBlock(in_dim, out_dim, dropout)) - if scale in attn_scales: - self.down_blocks.append(QwenImageAttentionBlock(out_dim)) - in_dim = out_dim - - # downsample block - if i != len(dim_mult) - 1: - mode = "downsample3d" if temperal_downsample[i] else "downsample2d" - self.down_blocks.append(QwenImageResample(out_dim, mode=mode)) - scale /= 2.0 - - # middle blocks - self.mid_block = QwenImageMidBlock(out_dim, dropout, non_linearity, num_layers=1) - - # output blocks - self.norm_out = QwenImageRMS_norm(out_dim, images=False) - self.conv_out = QwenImageCausalConv3d(out_dim, z_dim, 3, padding=1) - - self.gradient_checkpointing = False - - def forward(self, x, feat_cache=None, feat_idx=[0]): - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -min(CACHE_T, x.shape[2]) :, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - # cache last frame of last two chunk - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - x = self.conv_in(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv_in(x) - - ## downsamples - for layer in self.down_blocks: - if feat_cache is not None: - x = layer(x, feat_cache, feat_idx) - else: - x = layer(x) - - ## middle - x = self.mid_block(x, feat_cache, feat_idx) - - ## head - x = self.norm_out(x) - x = self.nonlinearity(x) - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -min(CACHE_T, x.shape[2]) :, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - # cache last frame of last two chunk - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - x = self.conv_out(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv_out(x) - return x - - -class QwenImageUpBlock(nn.Module): - """ - A block that handles upsampling for the QwenImageVAE decoder. - - Args: - in_dim (int): Input dimension - out_dim (int): Output dimension - num_res_blocks (int): Number of residual blocks - dropout (float): Dropout rate - upsample_mode (str, optional): Mode for upsampling ('upsample2d' or 'upsample3d') - non_linearity (str): Type of non-linearity to use - """ - - def __init__( - self, - in_dim: int, - out_dim: int, - num_res_blocks: int, - dropout: float = 0.0, - upsample_mode: str | None = None, - non_linearity: str = "silu", - ): - super().__init__() - self.in_dim = in_dim - self.out_dim = out_dim - - # Create layers list - resnets = [] - # Add residual blocks and attention if needed - current_dim = in_dim - for _ in range(num_res_blocks + 1): - resnets.append(QwenImageResidualBlock(current_dim, out_dim, dropout, non_linearity)) - current_dim = out_dim - - self.resnets = nn.ModuleList(resnets) - - # Add upsampling layer if needed - self.upsamplers = None - if upsample_mode is not None: - self.upsamplers = nn.ModuleList([QwenImageResample(out_dim, mode=upsample_mode)]) - - self.gradient_checkpointing = False - - def forward(self, x, feat_cache=None, feat_idx=[0]): - """ - Forward pass through the upsampling block. - - Args: - x (torch.Tensor): Input tensor - feat_cache (list, optional): Feature cache for causal convolutions - feat_idx (list, optional): Feature index for cache management - - Returns: - torch.Tensor: Output tensor - """ - for resnet in self.resnets: - if feat_cache is not None: - x = resnet(x, feat_cache, feat_idx) - else: - x = resnet(x) - - if self.upsamplers is not None: - if feat_cache is not None: - x = self.upsamplers[0](x, feat_cache, feat_idx) - else: - x = self.upsamplers[0](x) - return x - - -class QwenImageDecoder3d(nn.Module): - r""" - A 3D decoder module. - - Args: - dim (int): The base number of channels in the first layer. - z_dim (int): The dimensionality of the latent space. - dim_mult (list of int): Multipliers for the number of channels in each block. - num_res_blocks (int): Number of residual blocks in each block. - attn_scales (list of float): Scales at which to apply attention mechanisms. - temperal_upsample (list of bool): Whether to upsample temporally in each block. - dropout (float): Dropout rate for the dropout layers. - non_linearity (str): Type of non-linearity to use. - """ - - def __init__( - self, - dim=128, - z_dim=4, - dim_mult=[1, 2, 4, 4], - num_res_blocks=2, - attn_scales=[], - temperal_upsample=[False, True, True], - dropout=0.0, - input_channels=3, - non_linearity: str = "silu", - ): - super().__init__() - self.dim = dim - self.z_dim = z_dim - self.dim_mult = dim_mult - self.num_res_blocks = num_res_blocks - self.attn_scales = attn_scales - self.temperal_upsample = temperal_upsample - - self.nonlinearity = get_activation(non_linearity) - - # dimensions - dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]] - scale = 1.0 / 2 ** (len(dim_mult) - 2) - - # init block - self.conv_in = QwenImageCausalConv3d(z_dim, dims[0], 3, padding=1) - - # middle blocks - self.mid_block = QwenImageMidBlock(dims[0], dropout, non_linearity, num_layers=1) - - # upsample blocks - self.up_blocks = nn.ModuleList([]) - for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): - # residual (+attention) blocks - if i > 0: - in_dim = in_dim // 2 - - # Determine if we need upsampling - upsample_mode = None - if i != len(dim_mult) - 1: - upsample_mode = "upsample3d" if temperal_upsample[i] else "upsample2d" - - # Create and add the upsampling block - up_block = QwenImageUpBlock( - in_dim=in_dim, - out_dim=out_dim, - num_res_blocks=num_res_blocks, - dropout=dropout, - upsample_mode=upsample_mode, - non_linearity=non_linearity, - ) - self.up_blocks.append(up_block) - - # Update scale for next iteration - if upsample_mode is not None: - scale *= 2.0 - - # output blocks - self.norm_out = QwenImageRMS_norm(out_dim, images=False) - self.conv_out = QwenImageCausalConv3d(out_dim, input_channels, 3, padding=1) - - self.gradient_checkpointing = False - - def forward(self, x, feat_cache=None, feat_idx=[0]): - ## conv1 - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -min(CACHE_T, x.shape[2]) :, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - # cache last frame of last two chunk - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - x = self.conv_in(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv_in(x) - - ## middle - x = self.mid_block(x, feat_cache, feat_idx) - - ## upsamples - for up_block in self.up_blocks: - x = up_block(x, feat_cache, feat_idx) - - ## head - x = self.norm_out(x) - x = self.nonlinearity(x) - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -min(CACHE_T, x.shape[2]) :, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - # cache last frame of last two chunk - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - x = self.conv_out(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv_out(x) - return x - - -class AutoencoderKLQwenImage(ModelMixin, AutoencoderMixin, ConfigMixin, FromOriginalModelMixin): - r""" - A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - """ - - _supports_gradient_checkpointing = False - - # fmt: off - @register_to_config - def __init__( - self, - base_dim: int = 96, - z_dim: int = 16, - dim_mult: list[int] = [1, 2, 4, 4], - num_res_blocks: int = 2, - attn_scales: list[float] = [], - temperal_downsample: list[bool] = [False, True, True], - dropout: float = 0.0, - input_channels: int = 3, - latents_mean: list[float] = [-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508, 0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921], - latents_std: list[float] = [2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743, 3.2687, 2.1526, 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.9160], - ) -> None: - # fmt: on - super().__init__() - - self.z_dim = z_dim - self.temperal_downsample = temperal_downsample - self.temperal_upsample = temperal_downsample[::-1] - - self.encoder = QwenImageEncoder3d( - base_dim, z_dim * 2, dim_mult, num_res_blocks, attn_scales, self.temperal_downsample, dropout, input_channels - ) - self.quant_conv = QwenImageCausalConv3d(z_dim * 2, z_dim * 2, 1) - self.post_quant_conv = QwenImageCausalConv3d(z_dim, z_dim, 1) - - self.decoder = QwenImageDecoder3d( - base_dim, z_dim, dim_mult, num_res_blocks, attn_scales, self.temperal_upsample, dropout, input_channels - ) - - self.spatial_compression_ratio = 2 ** len(self.temperal_downsample) - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent - # frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the - # intermediate tiles together, the memory requirement can be lowered. - self.use_tiling = False - - # The minimal tile height and width for spatial tiling to be used - self.tile_sample_min_height = 256 - self.tile_sample_min_width = 256 - - # The minimal distance between two spatial tiles - self.tile_sample_stride_height = 192 - self.tile_sample_stride_width = 192 - - # Precompute and cache conv counts for encoder and decoder for clear_cache speedup - self._cached_conv_counts = { - "decoder": sum(isinstance(m, QwenImageCausalConv3d) for m in self.decoder.modules()) - if self.decoder is not None - else 0, - "encoder": sum(isinstance(m, QwenImageCausalConv3d) for m in self.encoder.modules()) - if self.encoder is not None - else 0, - } - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_sample_stride_height: float | None = None, - tile_sample_stride_width: float | None = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_sample_stride_height (`int`, *optional*): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. - tile_sample_stride_width (`int`, *optional*): - The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling - artifacts produced across the width dimension. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height - self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width - - def clear_cache(self): - def _count_conv3d(model): - count = 0 - for m in model.modules(): - if isinstance(m, QwenImageCausalConv3d): - count += 1 - return count - - self._conv_num = _count_conv3d(self.decoder) - self._conv_idx = [0] - self._feat_map = [None] * self._conv_num - # cache encode - self._enc_conv_num = _count_conv3d(self.encoder) - self._enc_conv_idx = [0] - self._enc_feat_map = [None] * self._enc_conv_num - - def _encode(self, x: torch.Tensor): - _, _, num_frame, height, width = x.shape - - if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): - return self.tiled_encode(x) - - self.clear_cache() - iter_ = 1 + (num_frame - 1) // 4 - for i in range(iter_): - self._enc_conv_idx = [0] - if i == 0: - out = self.encoder(x[:, :, :1, :, :], feat_cache=self._enc_feat_map, feat_idx=self._enc_conv_idx) - else: - out_ = self.encoder( - x[:, :, 1 + 4 * (i - 1) : 1 + 4 * i, :, :], - feat_cache=self._enc_feat_map, - feat_idx=self._enc_conv_idx, - ) - out = torch.cat([out, out_], 2) - - enc = self.quant_conv(out) - self.clear_cache() - return enc - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - r""" - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, return_dict: bool = True): - _, _, num_frame, height, width = z.shape - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - - if self.use_tiling and (width > tile_latent_min_width or height > tile_latent_min_height): - return self.tiled_decode(z, return_dict=return_dict) - - self.clear_cache() - x = self.post_quant_conv(z) - for i in range(num_frame): - self._conv_idx = [0] - if i == 0: - out = self.decoder(x[:, :, i : i + 1, :, :], feat_cache=self._feat_map, feat_idx=self._conv_idx) - else: - out_ = self.decoder(x[:, :, i : i + 1, :, :], feat_cache=self._feat_map, feat_idx=self._conv_idx) - out = torch.cat([out, out_], 2) - - out = torch.clamp(out, min=-1.0, max=1.0) - self.clear_cache() - if not return_dict: - return (out,) - - return DecoderOutput(sample=out) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-2], b.shape[-2], blend_extent) - for y in range(blend_extent): - b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-1], b.shape[-1], blend_extent) - for x in range(blend_extent): - b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( - x / blend_extent - ) - return b - - def tiled_encode(self, x: torch.Tensor) -> AutoencoderKLOutput: - r"""Encode a batch of images using a tiled encoder. - - Args: - x (`torch.Tensor`): Input batch of videos. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - _, _, num_frames, height, width = x.shape - latent_height = height // self.spatial_compression_ratio - latent_width = width // self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = tile_latent_min_height - tile_latent_stride_height - blend_width = tile_latent_min_width - tile_latent_stride_width - - # Split x into overlapping tiles and encode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, self.tile_sample_stride_height): - row = [] - for j in range(0, width, self.tile_sample_stride_width): - self.clear_cache() - time = [] - frame_range = 1 + (num_frames - 1) // 4 - for k in range(frame_range): - self._enc_conv_idx = [0] - if k == 0: - tile = x[:, :, :1, i : i + self.tile_sample_min_height, j : j + self.tile_sample_min_width] - else: - tile = x[ - :, - :, - 1 + 4 * (k - 1) : 1 + 4 * k, - i : i + self.tile_sample_min_height, - j : j + self.tile_sample_min_width, - ] - tile = self.encoder(tile, feat_cache=self._enc_feat_map, feat_idx=self._enc_conv_idx) - tile = self.quant_conv(tile) - time.append(tile) - row.append(torch.cat(time, dim=2)) - rows.append(row) - self.clear_cache() - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :tile_latent_stride_height, :tile_latent_stride_width]) - result_rows.append(torch.cat(result_row, dim=-1)) - - enc = torch.cat(result_rows, dim=3)[:, :, :, :latent_height, :latent_width] - return enc - - def tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - _, _, num_frames, height, width = z.shape - sample_height = height * self.spatial_compression_ratio - sample_width = width * self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - - blend_height = self.tile_sample_min_height - self.tile_sample_stride_height - blend_width = self.tile_sample_min_width - self.tile_sample_stride_width - - # Split z into overlapping tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, tile_latent_stride_height): - row = [] - for j in range(0, width, tile_latent_stride_width): - self.clear_cache() - time = [] - for k in range(num_frames): - self._conv_idx = [0] - tile = z[:, :, k : k + 1, i : i + tile_latent_min_height, j : j + tile_latent_min_width] - tile = self.post_quant_conv(tile) - decoded = self.decoder(tile, feat_cache=self._feat_map, feat_idx=self._conv_idx) - time.append(decoded) - row.append(torch.cat(time, dim=2)) - rows.append(row) - self.clear_cache() - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, : self.tile_sample_stride_height, : self.tile_sample_stride_width]) - result_rows.append(torch.cat(result_row, dim=-1)) - - dec = torch.cat(result_rows, dim=3)[:, :, :, :sample_height, :sample_width] - - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - """ - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z, return_dict=return_dict) - return dec diff --git a/diffusers/models/autoencoders/autoencoder_kl_temporal_decoder.py b/diffusers/models/autoencoders/autoencoder_kl_temporal_decoder.py deleted file mode 100644 index 8b0e5806d8efb4bc7e38e02b6f6f166ba7669c77..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_temporal_decoder.py +++ /dev/null @@ -1,313 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import itertools - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils.accelerate_utils import apply_forward_hook -from ..attention import AttentionMixin -from ..attention_processor import CROSS_ATTENTION_PROCESSORS, AttnProcessor -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from ..unets.unet_3d_blocks import MidBlockTemporalDecoder, UpBlockTemporalDecoder -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution, Encoder - - -class TemporalDecoder(nn.Module): - def __init__( - self, - in_channels: int = 4, - out_channels: int = 3, - block_out_channels: tuple[int] = (128, 256, 512, 512), - layers_per_block: int = 2, - ): - super().__init__() - self.layers_per_block = layers_per_block - - self.conv_in = nn.Conv2d(in_channels, block_out_channels[-1], kernel_size=3, stride=1, padding=1) - self.mid_block = MidBlockTemporalDecoder( - num_layers=self.layers_per_block, - in_channels=block_out_channels[-1], - out_channels=block_out_channels[-1], - attention_head_dim=block_out_channels[-1], - ) - - # up - self.up_blocks = nn.ModuleList([]) - reversed_block_out_channels = list(reversed(block_out_channels)) - output_channel = reversed_block_out_channels[0] - for i in range(len(block_out_channels)): - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - - is_final_block = i == len(block_out_channels) - 1 - up_block = UpBlockTemporalDecoder( - num_layers=self.layers_per_block + 1, - in_channels=prev_output_channel, - out_channels=output_channel, - add_upsample=not is_final_block, - ) - self.up_blocks.append(up_block) - prev_output_channel = output_channel - - self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=32, eps=1e-6) - - self.conv_act = nn.SiLU() - self.conv_out = torch.nn.Conv2d( - in_channels=block_out_channels[0], - out_channels=out_channels, - kernel_size=3, - padding=1, - ) - - conv_out_kernel_size = (3, 1, 1) - padding = [int(k // 2) for k in conv_out_kernel_size] - self.time_conv_out = torch.nn.Conv3d( - in_channels=out_channels, - out_channels=out_channels, - kernel_size=conv_out_kernel_size, - padding=padding, - ) - - self.gradient_checkpointing = False - - def forward( - self, - sample: torch.Tensor, - image_only_indicator: torch.Tensor, - num_frames: int = 1, - ) -> torch.Tensor: - r"""The forward method of the `Decoder` class.""" - - sample = self.conv_in(sample) - - upscale_dtype = next(itertools.chain(self.up_blocks.parameters(), self.up_blocks.buffers())).dtype - if torch.is_grad_enabled() and self.gradient_checkpointing: - # middle - sample = self._gradient_checkpointing_func( - self.mid_block, - sample, - image_only_indicator, - ) - sample = sample.to(upscale_dtype) - - # up - for up_block in self.up_blocks: - sample = self._gradient_checkpointing_func( - up_block, - sample, - image_only_indicator, - ) - else: - # middle - sample = self.mid_block(sample, image_only_indicator=image_only_indicator) - sample = sample.to(upscale_dtype) - - # up - for up_block in self.up_blocks: - sample = up_block(sample, image_only_indicator=image_only_indicator) - - # post-process - sample = self.conv_norm_out(sample) - sample = self.conv_act(sample) - sample = self.conv_out(sample) - - batch_frames, channels, height, width = sample.shape - batch_size = batch_frames // num_frames - sample = sample[None, :].reshape(batch_size, num_frames, channels, height, width).permute(0, 2, 1, 3, 4) - sample = self.time_conv_out(sample) - - sample = sample.permute(0, 2, 1, 3, 4).reshape(batch_frames, channels, height, width) - - return sample - - -class AutoencoderKLTemporalDecoder(ModelMixin, AttentionMixin, AutoencoderMixin, ConfigMixin): - r""" - A VAE model with KL loss for encoding images into latents and decoding latent representations into images. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - in_channels (int, *optional*, defaults to 3): Number of channels in the input image. - out_channels (int, *optional*, defaults to 3): Number of channels in the output. - down_block_types (`tuple[str]`, *optional*, defaults to `("DownEncoderBlock2D",)`): - tuple of downsample block types. - block_out_channels (`tuple[int]`, *optional*, defaults to `(64,)`): - tuple of block output channels. - layers_per_block: (`int`, *optional*, defaults to 1): Number of layers per block. - latent_channels (`int`, *optional*, defaults to 4): Number of channels in the latent space. - sample_size (`int`, *optional*, defaults to `32`): Sample input size. - scaling_factor (`float`, *optional*, defaults to 0.18215): - The component-wise standard deviation of the trained latent space computed using the first batch of the - training set. This is used to scale the latent space to have unit variance when training the diffusion - model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the - diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 - / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image - Synthesis with Latent Diffusion Models](https://huggingface.co/papers/2112.10752) paper. - force_upcast (`bool`, *optional*, default to `True`): - If enabled it will force the VAE to run in float32 for high image resolution pipelines, such as SD-XL. VAE - can be fine-tuned / trained to a lower range without losing too much precision in which case `force_upcast` - can be set to `False` - see: https://huggingface.co/madebyollin/sdxl-vae-fp16-fix - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - down_block_types: tuple[str] = ("DownEncoderBlock2D",), - block_out_channels: tuple[int] = (64,), - layers_per_block: int = 1, - latent_channels: int = 4, - sample_size: int = 32, - scaling_factor: float = 0.18215, - force_upcast: float = True, - ): - super().__init__() - - # pass init params to Encoder - self.encoder = Encoder( - in_channels=in_channels, - out_channels=latent_channels, - down_block_types=down_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - double_z=True, - ) - - # pass init params to Decoder - self.decoder = TemporalDecoder( - in_channels=latent_channels, - out_channels=out_channels, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - ) - - self.quant_conv = nn.Conv2d(2 * latent_channels, 2 * latent_channels, 1) - - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoders.autoencoder_kl.AutoencoderKLOutput`] instead of a plain - tuple. - - Returns: - The latent representations of the encoded images. If `return_dict` is True, a - [`~models.autoencoders.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - h = self.encoder(x) - moments = self.quant_conv(h) - posterior = DiagonalGaussianDistribution(moments) - - if not return_dict: - return (posterior,) - - return AutoencoderKLOutput(latent_dist=posterior) - - @apply_forward_hook - def decode( - self, - z: torch.Tensor, - num_frames: int, - return_dict: bool = True, - ) -> DecoderOutput | torch.Tensor: - """ - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - - """ - batch_size = z.shape[0] // num_frames - image_only_indicator = torch.zeros(batch_size, num_frames, dtype=z.dtype, device=z.device) - decoded = self.decoder(z, num_frames=num_frames, image_only_indicator=image_only_indicator) - - if not return_dict: - return (decoded,) - - return DecoderOutput(sample=decoded) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - num_frames: int = 1, - ) -> DecoderOutput | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - num_frames (`int`, *optional*, defaults to 1): - The number of frames to decode per batch. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - - dec = self.decode(z, num_frames=num_frames).sample - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) diff --git a/diffusers/models/autoencoders/autoencoder_kl_wan.py b/diffusers/models/autoencoders/autoencoder_kl_wan.py deleted file mode 100644 index de8a56edc20edccd9c8d64d8e7a8961cf7f4ea14..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_kl_wan.py +++ /dev/null @@ -1,1440 +0,0 @@ -# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..activations import get_activation -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - -CACHE_T = 2 - - -class AvgDown3D(nn.Module): - def __init__( - self, - in_channels, - out_channels, - factor_t, - factor_s=1, - ): - super().__init__() - self.in_channels = in_channels - self.out_channels = out_channels - self.factor_t = factor_t - self.factor_s = factor_s - self.factor = self.factor_t * self.factor_s * self.factor_s - - assert in_channels * self.factor % out_channels == 0 - self.group_size = in_channels * self.factor // out_channels - - def forward(self, x: torch.Tensor) -> torch.Tensor: - pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t - pad = (0, 0, 0, 0, pad_t, 0) - x = F.pad(x, pad) - B, C, T, H, W = x.shape - x = x.view( - B, - C, - T // self.factor_t, - self.factor_t, - H // self.factor_s, - self.factor_s, - W // self.factor_s, - self.factor_s, - ) - x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous() - x = x.view( - B, - C * self.factor, - T // self.factor_t, - H // self.factor_s, - W // self.factor_s, - ) - x = x.view( - B, - self.out_channels, - self.group_size, - T // self.factor_t, - H // self.factor_s, - W // self.factor_s, - ) - x = x.mean(dim=2) - return x - - -class DupUp3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - factor_t, - factor_s=1, - ): - super().__init__() - self.in_channels = in_channels - self.out_channels = out_channels - - self.factor_t = factor_t - self.factor_s = factor_s - self.factor = self.factor_t * self.factor_s * self.factor_s - - assert out_channels * self.factor % in_channels == 0 - self.repeats = out_channels * self.factor // in_channels - - def forward(self, x: torch.Tensor, first_chunk=False) -> torch.Tensor: - x = x.repeat_interleave(self.repeats, dim=1) - x = x.view( - x.size(0), - self.out_channels, - self.factor_t, - self.factor_s, - self.factor_s, - x.size(2), - x.size(3), - x.size(4), - ) - x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous() - x = x.view( - x.size(0), - self.out_channels, - x.size(2) * self.factor_t, - x.size(4) * self.factor_s, - x.size(6) * self.factor_s, - ) - if first_chunk: - x = x[:, :, self.factor_t - 1 :, :, :] - return x - - -class WanCausalConv3d(nn.Conv3d): - r""" - A custom 3D causal convolution layer with feature caching support. - - This layer extends the standard Conv3D layer by ensuring causality in the time dimension and handling feature - caching for efficient inference. - - Args: - in_channels (int): Number of channels in the input image - out_channels (int): Number of channels produced by the convolution - kernel_size (int or tuple): Size of the convolving kernel - stride (int or tuple, optional): Stride of the convolution. Default: 1 - padding (int or tuple, optional): Zero-padding added to all three sides of the input. Default: 0 - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int], - stride: int | tuple[int, int, int] = 1, - padding: int | tuple[int, int, int] = 0, - ) -> None: - super().__init__( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=kernel_size, - stride=stride, - padding=padding, - ) - - # Set up causal padding - self._padding = (self.padding[2], self.padding[2], self.padding[1], self.padding[1], 2 * self.padding[0], 0) - self.padding = (0, 0, 0) - - def forward(self, x, cache_x=None): - padding = list(self._padding) - if cache_x is not None and self._padding[4] > 0: - cache_x = cache_x.to(x.device) - x = torch.cat([cache_x, x], dim=2) - padding[4] -= cache_x.shape[2] - x = F.pad(x, padding) - return super().forward(x) - - -class WanRMS_norm(nn.Module): - r""" - A custom RMS normalization layer. - - Args: - dim (int): The number of dimensions to normalize over. - channel_first (bool, optional): Whether the input tensor has channels as the first dimension. - Default is True. - images (bool, optional): Whether the input represents image data. Default is True. - bias (bool, optional): Whether to include a learnable bias term. Default is False. - """ - - def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bias: bool = False) -> None: - super().__init__() - broadcastable_dims = (1, 1, 1) if not images else (1, 1) - shape = (dim, *broadcastable_dims) if channel_first else (dim,) - - self.channel_first = channel_first - self.scale = dim**0.5 - self.gamma = nn.Parameter(torch.ones(shape)) - self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - - def forward(self, x): - needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( - t in str(x.dtype) for t in ("float4_", "float8_") - ) - normalized = F.normalize(x.float() if needs_fp32_normalize else x, dim=(1 if self.channel_first else -1)).to( - x.dtype - ) - - return normalized * self.scale * self.gamma + self.bias - - -class WanUpsample(nn.Upsample): - r""" - Perform upsampling while ensuring the output tensor has the same data type as the input. - - Args: - x (torch.Tensor): Input tensor to be upsampled. - - Returns: - torch.Tensor: Upsampled tensor with the same data type as the input. - """ - - def forward(self, x): - return super().forward(x.float()).type_as(x) - - -class WanResample(nn.Module): - r""" - A custom resampling module for 2D and 3D data. - - Args: - dim (int): The number of input/output channels. - mode (str): The resampling mode. Must be one of: - - 'none': No resampling (identity operation). - - 'upsample2d': 2D upsampling with nearest-exact interpolation and convolution. - - 'upsample3d': 3D upsampling with nearest-exact interpolation, convolution, and causal 3D convolution. - - 'downsample2d': 2D downsampling with zero-padding and convolution. - - 'downsample3d': 3D downsampling with zero-padding, convolution, and causal 3D convolution. - """ - - def __init__(self, dim: int, mode: str, upsample_out_dim: int = None) -> None: - super().__init__() - self.dim = dim - self.mode = mode - - # default to dim //2 - if upsample_out_dim is None: - upsample_out_dim = dim // 2 - - # layers - if mode == "upsample2d": - self.resample = nn.Sequential( - WanUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), - nn.Conv2d(dim, upsample_out_dim, 3, padding=1), - ) - elif mode == "upsample3d": - self.resample = nn.Sequential( - WanUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), - nn.Conv2d(dim, upsample_out_dim, 3, padding=1), - ) - self.time_conv = WanCausalConv3d(dim, dim * 2, (3, 1, 1), padding=(1, 0, 0)) - - elif mode == "downsample2d": - self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2))) - elif mode == "downsample3d": - self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2))) - self.time_conv = WanCausalConv3d(dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0)) - - else: - self.resample = nn.Identity() - - def forward(self, x, feat_cache=None, feat_idx=[0]): - b, c, t, h, w = x.size() - if self.mode == "upsample3d": - if feat_cache is not None: - idx = feat_idx[0] - if feat_cache[idx] is None: - feat_cache[idx] = "Rep" - feat_idx[0] += 1 - else: - cache_x = x[:, :, -CACHE_T:, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] != "Rep": - # cache last frame of last two chunk - cache_x = torch.cat( - [feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2 - ) - if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] == "Rep": - cache_x = torch.cat([torch.zeros_like(cache_x).to(cache_x.device), cache_x], dim=2) - if feat_cache[idx] == "Rep": - x = self.time_conv(x) - else: - x = self.time_conv(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - - x = x.reshape(b, 2, c, t, h, w) - x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3) - x = x.reshape(b, c, t * 2, h, w) - t = x.shape[2] - x = x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w) - x = self.resample(x) - x = x.view(b, t, x.size(1), x.size(2), x.size(3)).permute(0, 2, 1, 3, 4) - - if self.mode == "downsample3d": - if feat_cache is not None: - idx = feat_idx[0] - if feat_cache[idx] is None: - feat_cache[idx] = x.clone() - feat_idx[0] += 1 - else: - cache_x = x[:, :, -1:, :, :].clone() - x = self.time_conv(torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2)) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - return x - - -class WanResidualBlock(nn.Module): - r""" - A custom residual block module. - - Args: - in_dim (int): Number of input channels. - out_dim (int): Number of output channels. - dropout (float, optional): Dropout rate for the dropout layer. Default is 0.0. - non_linearity (str, optional): Type of non-linearity to use. Default is "silu". - """ - - def __init__( - self, - in_dim: int, - out_dim: int, - dropout: float = 0.0, - non_linearity: str = "silu", - ) -> None: - super().__init__() - self.in_dim = in_dim - self.out_dim = out_dim - self.nonlinearity = get_activation(non_linearity) - - # layers - self.norm1 = WanRMS_norm(in_dim, images=False) - self.conv1 = WanCausalConv3d(in_dim, out_dim, 3, padding=1) - self.norm2 = WanRMS_norm(out_dim, images=False) - self.dropout = nn.Dropout(dropout) - self.conv2 = WanCausalConv3d(out_dim, out_dim, 3, padding=1) - self.conv_shortcut = WanCausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity() - - def forward(self, x, feat_cache=None, feat_idx=[0]): - # Apply shortcut connection - h = self.conv_shortcut(x) - - # First normalization and activation - x = self.norm1(x) - x = self.nonlinearity(x) - - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -CACHE_T:, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - - x = self.conv1(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv1(x) - - # Second normalization and activation - x = self.norm2(x) - x = self.nonlinearity(x) - - # Dropout - x = self.dropout(x) - - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -CACHE_T:, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - - x = self.conv2(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv2(x) - - # Add residual connection - return x + h - - -class WanAttentionBlock(nn.Module): - r""" - Causal self-attention with a single head. - - Args: - dim (int): The number of channels in the input tensor. - """ - - def __init__(self, dim): - super().__init__() - self.dim = dim - - # layers - self.norm = WanRMS_norm(dim) - self.to_qkv = nn.Conv2d(dim, dim * 3, 1) - self.proj = nn.Conv2d(dim, dim, 1) - - def forward(self, x): - identity = x - batch_size, channels, time, height, width = x.size() - - x = x.permute(0, 2, 1, 3, 4).reshape(batch_size * time, channels, height, width) - x = self.norm(x) - - # compute query, key, value - qkv = self.to_qkv(x) - qkv = qkv.reshape(batch_size * time, 1, channels * 3, -1) - qkv = qkv.permute(0, 1, 3, 2).contiguous() - q, k, v = qkv.chunk(3, dim=-1) - - # apply attention - x = F.scaled_dot_product_attention(q, k, v) - - x = x.squeeze(1).permute(0, 2, 1).reshape(batch_size * time, channels, height, width) - - # output projection - x = self.proj(x) - - # Reshape back: [(b*t), c, h, w] -> [b, c, t, h, w] - x = x.view(batch_size, time, channels, height, width) - x = x.permute(0, 2, 1, 3, 4) - - return x + identity - - -class WanMidBlock(nn.Module): - """ - Middle block for WanVAE encoder and decoder. - - Args: - dim (int): Number of input/output channels. - dropout (float): Dropout rate. - non_linearity (str): Type of non-linearity to use. - """ - - def __init__(self, dim: int, dropout: float = 0.0, non_linearity: str = "silu", num_layers: int = 1): - super().__init__() - self.dim = dim - - # Create the components - resnets = [WanResidualBlock(dim, dim, dropout, non_linearity)] - attentions = [] - for _ in range(num_layers): - attentions.append(WanAttentionBlock(dim)) - resnets.append(WanResidualBlock(dim, dim, dropout, non_linearity)) - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward(self, x, feat_cache=None, feat_idx=[0]): - # First residual block - x = self.resnets[0](x, feat_cache=feat_cache, feat_idx=feat_idx) - - # Process through attention and residual blocks - for attn, resnet in zip(self.attentions, self.resnets[1:]): - if attn is not None: - x = attn(x) - - x = resnet(x, feat_cache=feat_cache, feat_idx=feat_idx) - - return x - - -class WanResidualDownBlock(nn.Module): - def __init__(self, in_dim, out_dim, dropout, num_res_blocks, temperal_downsample=False, down_flag=False): - super().__init__() - - # Shortcut path with downsample - self.avg_shortcut = AvgDown3D( - in_dim, - out_dim, - factor_t=2 if temperal_downsample else 1, - factor_s=2 if down_flag else 1, - ) - - # Main path with residual blocks and downsample - resnets = [] - for _ in range(num_res_blocks): - resnets.append(WanResidualBlock(in_dim, out_dim, dropout)) - in_dim = out_dim - self.resnets = nn.ModuleList(resnets) - - # Add the final downsample block - if down_flag: - mode = "downsample3d" if temperal_downsample else "downsample2d" - self.downsampler = WanResample(out_dim, mode=mode) - else: - self.downsampler = None - - def forward(self, x, feat_cache=None, feat_idx=[0]): - x_copy = x.clone() - for resnet in self.resnets: - x = resnet(x, feat_cache=feat_cache, feat_idx=feat_idx) - if self.downsampler is not None: - x = self.downsampler(x, feat_cache=feat_cache, feat_idx=feat_idx) - - return x + self.avg_shortcut(x_copy) - - -class WanEncoder3d(nn.Module): - r""" - A 3D encoder module. - - Args: - dim (int): The base number of channels in the first layer. - z_dim (int): The dimensionality of the latent space. - dim_mult (list of int): Multipliers for the number of channels in each block. - num_res_blocks (int): Number of residual blocks in each block. - attn_scales (list of float): Scales at which to apply attention mechanisms. - temperal_downsample (list of bool): Whether to downsample temporally in each block. - dropout (float): Dropout rate for the dropout layers. - non_linearity (str): Type of non-linearity to use. - """ - - def __init__( - self, - in_channels: int = 3, - dim=128, - z_dim=4, - dim_mult=[1, 2, 4, 4], - num_res_blocks=2, - attn_scales=[], - temperal_downsample=[True, True, False], - dropout=0.0, - non_linearity: str = "silu", - is_residual: bool = False, # wan 2.2 vae use a residual downblock - ): - super().__init__() - self.dim = dim - self.z_dim = z_dim - self.dim_mult = dim_mult - self.num_res_blocks = num_res_blocks - self.attn_scales = attn_scales - self.temperal_downsample = temperal_downsample - self.nonlinearity = get_activation(non_linearity) - - # dimensions - dims = [dim * u for u in [1] + dim_mult] - scale = 1.0 - - # init block - self.conv_in = WanCausalConv3d(in_channels, dims[0], 3, padding=1) - - # downsample blocks - self.down_blocks = nn.ModuleList([]) - for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): - # residual (+attention) blocks - if is_residual: - self.down_blocks.append( - WanResidualDownBlock( - in_dim, - out_dim, - dropout, - num_res_blocks, - temperal_downsample=temperal_downsample[i] if i != len(dim_mult) - 1 else False, - down_flag=i != len(dim_mult) - 1, - ) - ) - else: - for _ in range(num_res_blocks): - self.down_blocks.append(WanResidualBlock(in_dim, out_dim, dropout)) - if scale in attn_scales: - self.down_blocks.append(WanAttentionBlock(out_dim)) - in_dim = out_dim - - # downsample block - if i != len(dim_mult) - 1: - mode = "downsample3d" if temperal_downsample[i] else "downsample2d" - self.down_blocks.append(WanResample(out_dim, mode=mode)) - scale /= 2.0 - - # middle blocks - self.mid_block = WanMidBlock(out_dim, dropout, non_linearity, num_layers=1) - - # output blocks - self.norm_out = WanRMS_norm(out_dim, images=False) - self.conv_out = WanCausalConv3d(out_dim, z_dim, 3, padding=1) - - self.gradient_checkpointing = False - - def forward(self, x, feat_cache=None, feat_idx=[0]): - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -CACHE_T:, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - # cache last frame of last two chunk - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - x = self.conv_in(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv_in(x) - - ## downsamples - for layer in self.down_blocks: - if feat_cache is not None: - x = layer(x, feat_cache=feat_cache, feat_idx=feat_idx) - else: - x = layer(x) - - ## middle - x = self.mid_block(x, feat_cache=feat_cache, feat_idx=feat_idx) - - ## head - x = self.norm_out(x) - x = self.nonlinearity(x) - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -CACHE_T:, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - # cache last frame of last two chunk - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - x = self.conv_out(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv_out(x) - - return x - - -class WanResidualUpBlock(nn.Module): - """ - A block that handles upsampling for the WanVAE decoder. - - Args: - in_dim (int): Input dimension - out_dim (int): Output dimension - num_res_blocks (int): Number of residual blocks - dropout (float): Dropout rate - temperal_upsample (bool): Whether to upsample on temporal dimension - up_flag (bool): Whether to upsample or not - non_linearity (str): Type of non-linearity to use - """ - - def __init__( - self, - in_dim: int, - out_dim: int, - num_res_blocks: int, - dropout: float = 0.0, - temperal_upsample: bool = False, - up_flag: bool = False, - non_linearity: str = "silu", - ): - super().__init__() - self.in_dim = in_dim - self.out_dim = out_dim - - if up_flag: - self.avg_shortcut = DupUp3D( - in_dim, - out_dim, - factor_t=2 if temperal_upsample else 1, - factor_s=2, - ) - else: - self.avg_shortcut = None - - # create residual blocks - resnets = [] - current_dim = in_dim - for _ in range(num_res_blocks + 1): - resnets.append(WanResidualBlock(current_dim, out_dim, dropout, non_linearity)) - current_dim = out_dim - - self.resnets = nn.ModuleList(resnets) - - # Add upsampling layer if needed - if up_flag: - upsample_mode = "upsample3d" if temperal_upsample else "upsample2d" - self.upsampler = WanResample(out_dim, mode=upsample_mode, upsample_out_dim=out_dim) - else: - self.upsampler = None - - self.gradient_checkpointing = False - - def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False): - """ - Forward pass through the upsampling block. - - Args: - x (torch.Tensor): Input tensor - feat_cache (list, optional): Feature cache for causal convolutions - feat_idx (list, optional): Feature index for cache management - - Returns: - torch.Tensor: Output tensor - """ - x_copy = x.clone() - - for resnet in self.resnets: - if feat_cache is not None: - x = resnet(x, feat_cache=feat_cache, feat_idx=feat_idx) - else: - x = resnet(x) - - if self.upsampler is not None: - if feat_cache is not None: - x = self.upsampler(x, feat_cache=feat_cache, feat_idx=feat_idx) - else: - x = self.upsampler(x) - - if self.avg_shortcut is not None: - x = x + self.avg_shortcut(x_copy, first_chunk=first_chunk) - - return x - - -class WanUpBlock(nn.Module): - """ - A block that handles upsampling for the WanVAE decoder. - - Args: - in_dim (int): Input dimension - out_dim (int): Output dimension - num_res_blocks (int): Number of residual blocks - dropout (float): Dropout rate - upsample_mode (str, optional): Mode for upsampling ('upsample2d' or 'upsample3d') - non_linearity (str): Type of non-linearity to use - """ - - def __init__( - self, - in_dim: int, - out_dim: int, - num_res_blocks: int, - dropout: float = 0.0, - upsample_mode: str | None = None, - non_linearity: str = "silu", - ): - super().__init__() - self.in_dim = in_dim - self.out_dim = out_dim - - # Create layers list - resnets = [] - # Add residual blocks and attention if needed - current_dim = in_dim - for _ in range(num_res_blocks + 1): - resnets.append(WanResidualBlock(current_dim, out_dim, dropout, non_linearity)) - current_dim = out_dim - - self.resnets = nn.ModuleList(resnets) - - # Add upsampling layer if needed - self.upsamplers = None - if upsample_mode is not None: - self.upsamplers = nn.ModuleList([WanResample(out_dim, mode=upsample_mode)]) - - self.gradient_checkpointing = False - - def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=None): - """ - Forward pass through the upsampling block. - - Args: - x (torch.Tensor): Input tensor - feat_cache (list, optional): Feature cache for causal convolutions - feat_idx (list, optional): Feature index for cache management - - Returns: - torch.Tensor: Output tensor - """ - for resnet in self.resnets: - if feat_cache is not None: - x = resnet(x, feat_cache=feat_cache, feat_idx=feat_idx) - else: - x = resnet(x) - - if self.upsamplers is not None: - if feat_cache is not None: - x = self.upsamplers[0](x, feat_cache=feat_cache, feat_idx=feat_idx) - else: - x = self.upsamplers[0](x) - return x - - -class WanDecoder3d(nn.Module): - r""" - A 3D decoder module. - - Args: - dim (int): The base number of channels in the first layer. - z_dim (int): The dimensionality of the latent space. - dim_mult (list of int): Multipliers for the number of channels in each block. - num_res_blocks (int): Number of residual blocks in each block. - attn_scales (list of float): Scales at which to apply attention mechanisms. - temperal_upsample (list of bool): Whether to upsample temporally in each block. - dropout (float): Dropout rate for the dropout layers. - non_linearity (str): Type of non-linearity to use. - """ - - def __init__( - self, - dim=128, - z_dim=4, - dim_mult=[1, 2, 4, 4], - num_res_blocks=2, - attn_scales=[], - temperal_upsample=[False, True, True], - dropout=0.0, - non_linearity: str = "silu", - out_channels: int = 3, - is_residual: bool = False, - ): - super().__init__() - self.dim = dim - self.z_dim = z_dim - self.dim_mult = dim_mult - self.num_res_blocks = num_res_blocks - self.attn_scales = attn_scales - self.temperal_upsample = temperal_upsample - - self.nonlinearity = get_activation(non_linearity) - - # dimensions - dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]] - - # init block - self.conv_in = WanCausalConv3d(z_dim, dims[0], 3, padding=1) - - # middle blocks - self.mid_block = WanMidBlock(dims[0], dropout, non_linearity, num_layers=1) - - # upsample blocks - self.up_blocks = nn.ModuleList([]) - for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): - # residual (+attention) blocks - if i > 0 and not is_residual: - # wan vae 2.1 - in_dim = in_dim // 2 - - # determine if we need upsampling - up_flag = i != len(dim_mult) - 1 - # determine upsampling mode, if not upsampling, set to None - upsample_mode = None - if up_flag and temperal_upsample[i]: - upsample_mode = "upsample3d" - elif up_flag: - upsample_mode = "upsample2d" - # Create and add the upsampling block - if is_residual: - up_block = WanResidualUpBlock( - in_dim=in_dim, - out_dim=out_dim, - num_res_blocks=num_res_blocks, - dropout=dropout, - temperal_upsample=temperal_upsample[i] if up_flag else False, - up_flag=up_flag, - non_linearity=non_linearity, - ) - else: - up_block = WanUpBlock( - in_dim=in_dim, - out_dim=out_dim, - num_res_blocks=num_res_blocks, - dropout=dropout, - upsample_mode=upsample_mode, - non_linearity=non_linearity, - ) - self.up_blocks.append(up_block) - - # output blocks - self.norm_out = WanRMS_norm(out_dim, images=False) - self.conv_out = WanCausalConv3d(out_dim, out_channels, 3, padding=1) - - self.gradient_checkpointing = False - - def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False): - ## conv1 - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -CACHE_T:, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - # cache last frame of last two chunk - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - x = self.conv_in(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv_in(x) - - ## middle - x = self.mid_block(x, feat_cache=feat_cache, feat_idx=feat_idx) - - ## upsamples - for up_block in self.up_blocks: - x = up_block(x, feat_cache=feat_cache, feat_idx=feat_idx, first_chunk=first_chunk) - - ## head - x = self.norm_out(x) - x = self.nonlinearity(x) - if feat_cache is not None: - idx = feat_idx[0] - cache_x = x[:, :, -CACHE_T:, :, :].clone() - if cache_x.shape[2] < 2 and feat_cache[idx] is not None: - # cache last frame of last two chunk - cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) - x = self.conv_out(x, feat_cache[idx]) - feat_cache[idx] = cache_x - feat_idx[0] += 1 - else: - x = self.conv_out(x) - return x - - -def patchify(x, patch_size): - if patch_size == 1: - return x - - if x.dim() != 5: - raise ValueError(f"Invalid input shape: {x.shape}") - # x shape: [batch_size, channels, frames, height, width] - batch_size, channels, frames, height, width = x.shape - - # Ensure height and width are divisible by patch_size - if height % patch_size != 0 or width % patch_size != 0: - raise ValueError(f"Height ({height}) and width ({width}) must be divisible by patch_size ({patch_size})") - - # Reshape to [batch_size, channels, frames, height//patch_size, patch_size, width//patch_size, patch_size] - x = x.view(batch_size, channels, frames, height // patch_size, patch_size, width // patch_size, patch_size) - - # Rearrange to [batch_size, channels * patch_size * patch_size, frames, height//patch_size, width//patch_size] - x = x.permute(0, 1, 6, 4, 2, 3, 5).contiguous() - x = x.view(batch_size, channels * patch_size * patch_size, frames, height // patch_size, width // patch_size) - - return x - - -def unpatchify(x, patch_size): - if patch_size == 1: - return x - - if x.dim() != 5: - raise ValueError(f"Invalid input shape: {x.shape}") - # x shape: [batch_size, (channels * patch_size * patch_size), frame, height, width] - batch_size, c_patches, frames, height, width = x.shape - channels = c_patches // (patch_size * patch_size) - - # Reshape to [b, c, patch_size, patch_size, f, h, w] - x = x.view(batch_size, channels, patch_size, patch_size, frames, height, width) - - # Rearrange to [b, c, f, h * patch_size, w * patch_size] - x = x.permute(0, 1, 4, 5, 3, 6, 2).contiguous() - x = x.view(batch_size, channels, frames, height * patch_size, width * patch_size) - - return x - - -class AutoencoderKLWan(ModelMixin, AutoencoderMixin, ConfigMixin, FromOriginalModelMixin): - r""" - A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. - Introduced in [Wan 2.1]. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - """ - - _supports_gradient_checkpointing = False - _group_offload_block_modules = ["quant_conv", "post_quant_conv", "encoder", "decoder"] - # keys toignore when AlignDeviceHook moves inputs/outputs between devices - # these are shared mutable state modified in-place - _skip_keys = ["feat_cache", "feat_idx"] - - @register_to_config - def __init__( - self, - base_dim: int = 96, - decoder_base_dim: int | None = None, - z_dim: int = 16, - dim_mult: list[int] = [1, 2, 4, 4], - num_res_blocks: int = 2, - attn_scales: list[float] = [], - temperal_downsample: list[bool] = [False, True, True], - dropout: float = 0.0, - latents_mean: list[float] = [ - -0.7571, - -0.7089, - -0.9113, - 0.1075, - -0.1745, - 0.9653, - -0.1517, - 1.5508, - 0.4134, - -0.0715, - 0.5517, - -0.3632, - -0.1922, - -0.9497, - 0.2503, - -0.2921, - ], - latents_std: list[float] = [ - 2.8184, - 1.4541, - 2.3275, - 2.6558, - 1.2196, - 1.7708, - 2.6052, - 2.0743, - 3.2687, - 2.1526, - 2.8652, - 1.5579, - 1.6382, - 1.1253, - 2.8251, - 1.9160, - ], - is_residual: bool = False, - in_channels: int = 3, - out_channels: int = 3, - patch_size: int | None = None, - scale_factor_temporal: int | None = 4, - scale_factor_spatial: int | None = 8, - ) -> None: - super().__init__() - - self.z_dim = z_dim - self.temperal_downsample = temperal_downsample - self.temperal_upsample = temperal_downsample[::-1] - - if decoder_base_dim is None: - decoder_base_dim = base_dim - - self.encoder = WanEncoder3d( - in_channels=in_channels, - dim=base_dim, - z_dim=z_dim * 2, - dim_mult=dim_mult, - num_res_blocks=num_res_blocks, - attn_scales=attn_scales, - temperal_downsample=temperal_downsample, - dropout=dropout, - is_residual=is_residual, - ) - self.quant_conv = WanCausalConv3d(z_dim * 2, z_dim * 2, 1) - self.post_quant_conv = WanCausalConv3d(z_dim, z_dim, 1) - - self.decoder = WanDecoder3d( - dim=decoder_base_dim, - z_dim=z_dim, - dim_mult=dim_mult, - num_res_blocks=num_res_blocks, - attn_scales=attn_scales, - temperal_upsample=self.temperal_upsample, - dropout=dropout, - out_channels=out_channels, - is_residual=is_residual, - ) - - self.spatial_compression_ratio = scale_factor_spatial - - # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension - # to perform decoding of a single video latent at a time. - self.use_slicing = False - - # When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent - # frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the - # intermediate tiles together, the memory requirement can be lowered. - self.use_tiling = False - - # The minimal tile height and width for spatial tiling to be used - self.tile_sample_min_height = 256 - self.tile_sample_min_width = 256 - - # The minimal distance between two spatial tiles - self.tile_sample_stride_height = 192 - self.tile_sample_stride_width = 192 - - # Precompute and cache conv counts for encoder and decoder for clear_cache speedup - self._cached_conv_counts = { - "decoder": sum(isinstance(m, WanCausalConv3d) for m in self.decoder.modules()) - if self.decoder is not None - else 0, - "encoder": sum(isinstance(m, WanCausalConv3d) for m in self.encoder.modules()) - if self.encoder is not None - else 0, - } - - def enable_tiling( - self, - tile_sample_min_height: int | None = None, - tile_sample_min_width: int | None = None, - tile_sample_stride_height: float | None = None, - tile_sample_stride_width: float | None = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_sample_stride_height (`int`, *optional*): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. - tile_sample_stride_width (`int`, *optional*): - The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling - artifacts produced across the width dimension. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height - self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width - - def clear_cache(self): - # Use cached conv counts for decoder and encoder to avoid re-iterating modules each call - self._conv_num = self._cached_conv_counts["decoder"] - self._conv_idx = [0] - self._feat_map = [None] * self._conv_num - # cache encode - self._enc_conv_num = self._cached_conv_counts["encoder"] - self._enc_conv_idx = [0] - self._enc_feat_map = [None] * self._enc_conv_num - - def _encode(self, x: torch.Tensor): - _, _, num_frame, height, width = x.shape - - self.clear_cache() - if self.config.patch_size is not None: - x = patchify(x, patch_size=self.config.patch_size) - - if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): - return self.tiled_encode(x) - - iter_ = 1 + (num_frame - 1) // 4 - for i in range(iter_): - self._enc_conv_idx = [0] - if i == 0: - out = self.encoder(x[:, :, :1, :, :], feat_cache=self._enc_feat_map, feat_idx=self._enc_conv_idx) - else: - out_ = self.encoder( - x[:, :, 1 + 4 * (i - 1) : 1 + 4 * i, :, :], - feat_cache=self._enc_feat_map, - feat_idx=self._enc_conv_idx, - ) - out = torch.cat([out, out_], 2) - - enc = self.quant_conv(out) - self.clear_cache() - return enc - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]: - r""" - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded videos. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - posterior = DiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) - - def _decode(self, z: torch.Tensor, return_dict: bool = True): - _, _, num_frame, height, width = z.shape - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - - if self.use_tiling and (width > tile_latent_min_width or height > tile_latent_min_height): - return self.tiled_decode(z, return_dict=return_dict) - - self.clear_cache() - x = self.post_quant_conv(z) - for i in range(num_frame): - self._conv_idx = [0] - if i == 0: - out = self.decoder( - x[:, :, i : i + 1, :, :], feat_cache=self._feat_map, feat_idx=self._conv_idx, first_chunk=True - ) - else: - out_ = self.decoder(x[:, :, i : i + 1, :, :], feat_cache=self._feat_map, feat_idx=self._conv_idx) - out = torch.cat([out, out_], 2) - - if self.config.patch_size is not None: - out = unpatchify(out, patch_size=self.config.patch_size) - - out = torch.clamp(out, min=-1.0, max=1.0) - - self.clear_cache() - if not return_dict: - return (out,) - - return DecoderOutput(sample=out) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-2], b.shape[-2], blend_extent) - for y in range(blend_extent): - b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[-1], b.shape[-1], blend_extent) - for x in range(blend_extent): - b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( - x / blend_extent - ) - return b - - def tiled_encode(self, x: torch.Tensor) -> AutoencoderKLOutput: - r"""Encode a batch of images using a tiled encoder. - - Args: - x (`torch.Tensor`): Input batch of videos. - - Returns: - `torch.Tensor`: - The latent representation of the encoded videos. - """ - - _, _, num_frames, height, width = x.shape - encode_spatial_compression_ratio = self.spatial_compression_ratio - if self.config.patch_size is not None: - assert encode_spatial_compression_ratio % self.config.patch_size == 0 - encode_spatial_compression_ratio = self.spatial_compression_ratio // self.config.patch_size - - latent_height = height // encode_spatial_compression_ratio - latent_width = width // encode_spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // encode_spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // encode_spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // encode_spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // encode_spatial_compression_ratio - - blend_height = tile_latent_min_height - tile_latent_stride_height - blend_width = tile_latent_min_width - tile_latent_stride_width - - # Split x into overlapping tiles and encode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, self.tile_sample_stride_height): - row = [] - for j in range(0, width, self.tile_sample_stride_width): - self.clear_cache() - time = [] - frame_range = 1 + (num_frames - 1) // 4 - for k in range(frame_range): - self._enc_conv_idx = [0] - if k == 0: - tile = x[:, :, :1, i : i + self.tile_sample_min_height, j : j + self.tile_sample_min_width] - else: - tile = x[ - :, - :, - 1 + 4 * (k - 1) : 1 + 4 * k, - i : i + self.tile_sample_min_height, - j : j + self.tile_sample_min_width, - ] - tile = self.encoder(tile, feat_cache=self._enc_feat_map, feat_idx=self._enc_conv_idx) - tile = self.quant_conv(tile) - time.append(tile) - row.append(torch.cat(time, dim=2)) - rows.append(row) - self.clear_cache() - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :tile_latent_stride_height, :tile_latent_stride_width]) - result_rows.append(torch.cat(result_row, dim=-1)) - - enc = torch.cat(result_rows, dim=3)[:, :, :, :latent_height, :latent_width] - return enc - - def tiled_decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - _, _, num_frames, height, width = z.shape - sample_height = height * self.spatial_compression_ratio - sample_width = width * self.spatial_compression_ratio - - tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio - tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio - tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio - tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio - tile_sample_stride_height = self.tile_sample_stride_height - tile_sample_stride_width = self.tile_sample_stride_width - if self.config.patch_size is not None: - sample_height = sample_height // self.config.patch_size - sample_width = sample_width // self.config.patch_size - tile_sample_stride_height = tile_sample_stride_height // self.config.patch_size - tile_sample_stride_width = tile_sample_stride_width // self.config.patch_size - blend_height = self.tile_sample_min_height // self.config.patch_size - tile_sample_stride_height - blend_width = self.tile_sample_min_width // self.config.patch_size - tile_sample_stride_width - else: - blend_height = self.tile_sample_min_height - tile_sample_stride_height - blend_width = self.tile_sample_min_width - tile_sample_stride_width - - # Split z into overlapping tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, tile_latent_stride_height): - row = [] - for j in range(0, width, tile_latent_stride_width): - self.clear_cache() - time = [] - for k in range(num_frames): - self._conv_idx = [0] - tile = z[:, :, k : k + 1, i : i + tile_latent_min_height, j : j + tile_latent_min_width] - tile = self.post_quant_conv(tile) - decoded = self.decoder( - tile, feat_cache=self._feat_map, feat_idx=self._conv_idx, first_chunk=(k == 0) - ) - time.append(decoded) - row.append(torch.cat(time, dim=2)) - rows.append(row) - self.clear_cache() - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_width) - result_row.append(tile[:, :, :, :tile_sample_stride_height, :tile_sample_stride_width]) - result_rows.append(torch.cat(result_row, dim=-1)) - dec = torch.cat(result_rows, dim=3)[:, :, :, :sample_height, :sample_width] - - if self.config.patch_size is not None: - dec = unpatchify(dec, patch_size=self.config.patch_size) - - dec = torch.clamp(dec, min=-1.0, max=1.0) - - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | torch.Tensor: - """ - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - x = sample - posterior = self.encode(x).latent_dist - - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z, return_dict=return_dict) - return dec diff --git a/diffusers/models/autoencoders/autoencoder_longcat_audio_dit.py b/diffusers/models/autoencoders/autoencoder_longcat_audio_dit.py deleted file mode 100644 index 3b5e81d814c0a9faacef2abcfa0e0e8c561b18fb..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_longcat_audio_dit.py +++ /dev/null @@ -1,416 +0,0 @@ -# Copyright 2026 MeiTuan LongCat-AudioDiT Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -# Adapted from the LongCat-AudioDiT reference implementation: -# https://github.com/meituan-longcat/LongCat-AudioDiT - -import math -from dataclasses import dataclass - -import torch -import torch.nn as nn -import torch.nn.functional as F -from torch.nn.utils import weight_norm - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput -from ...utils.accelerate_utils import apply_forward_hook -from ...utils.torch_utils import randn_tensor -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin - - -def _wn_conv1d(in_channels, out_channels, kernel_size, stride=1, dilation=1, padding=0, bias=True): - return weight_norm(nn.Conv1d(in_channels, out_channels, kernel_size, stride, padding, dilation, bias=bias)) - - -def _wn_conv_transpose1d(*args, **kwargs): - return weight_norm(nn.ConvTranspose1d(*args, **kwargs)) - - -class Snake1d(nn.Module): - def __init__(self, channels: int, alpha_logscale: bool = True): - super().__init__() - self.alpha_logscale = alpha_logscale - self.alpha = nn.Parameter(torch.zeros(channels)) - self.beta = nn.Parameter(torch.zeros(channels)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - alpha = self.alpha[None, :, None] - beta = self.beta[None, :, None] - if self.alpha_logscale: - alpha = torch.exp(alpha) - beta = torch.exp(beta) - return hidden_states + (1.0 / (beta + 1e-9)) * torch.sin(hidden_states * alpha).pow(2) - - -def _get_vae_activation(name: str, channels: int = 0) -> nn.Module: - if name == "elu": - act = nn.ELU() - elif name == "snake": - act = Snake1d(channels) - else: - raise ValueError(f"Unknown activation: {name}") - return act - - -def _pixel_shuffle_1d(hidden_states: torch.Tensor, factor: int) -> torch.Tensor: - batch, channels, width = hidden_states.size() - return ( - hidden_states.view(batch, channels // factor, factor, width) - .permute(0, 1, 3, 2) - .contiguous() - .view(batch, channels // factor, width * factor) - ) - - -class DownsampleShortcut(nn.Module): - def __init__(self, in_channels: int, out_channels: int, factor: int): - super().__init__() - self.factor = factor - self.group_size = in_channels * factor // out_channels - self.out_channels = out_channels - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch, channels, width = hidden_states.shape - hidden_states = ( - hidden_states.view(batch, channels, width // self.factor, self.factor) - .permute(0, 1, 3, 2) - .contiguous() - .view(batch, channels * self.factor, width // self.factor) - ) - return hidden_states.view(batch, self.out_channels, self.group_size, width // self.factor).mean(dim=2) - - -class UpsampleShortcut(nn.Module): - def __init__(self, in_channels: int, out_channels: int, factor: int): - super().__init__() - self.factor = factor - self.repeats = out_channels * factor // in_channels - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = hidden_states.repeat_interleave(self.repeats, dim=1) - return _pixel_shuffle_1d(hidden_states, self.factor) - - -class VaeResidualUnit(nn.Module): - def __init__( - self, in_channels: int, out_channels: int, dilation: int, kernel_size: int = 7, act_fn: str = "snake" - ): - super().__init__() - padding = (dilation * (kernel_size - 1)) // 2 - self.layers = nn.Sequential( - _get_vae_activation(act_fn, channels=out_channels), - _wn_conv1d(in_channels, out_channels, kernel_size, dilation=dilation, padding=padding), - _get_vae_activation(act_fn, channels=out_channels), - _wn_conv1d(out_channels, out_channels, kernel_size=1), - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return hidden_states + self.layers(hidden_states) - - -class VaeEncoderBlock(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - stride: int, - act_fn: str = "snake", - downsample_shortcut: str = "none", - ): - super().__init__() - layers = [ - VaeResidualUnit(in_channels, in_channels, dilation=1, act_fn=act_fn), - VaeResidualUnit(in_channels, in_channels, dilation=3, act_fn=act_fn), - VaeResidualUnit(in_channels, in_channels, dilation=9, act_fn=act_fn), - ] - layers.append(_get_vae_activation(act_fn, channels=in_channels)) - layers.append( - _wn_conv1d(in_channels, out_channels, kernel_size=2 * stride, stride=stride, padding=math.ceil(stride / 2)) - ) - self.layers = nn.Sequential(*layers) - self.residual = ( - DownsampleShortcut(in_channels, out_channels, stride) if downsample_shortcut == "averaging" else None - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - output_hidden_states = self.layers(hidden_states) - if self.residual is not None: - residual = self.residual(hidden_states) - output_hidden_states = output_hidden_states + residual - return output_hidden_states - - -class VaeDecoderBlock(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - stride: int, - act_fn: str = "snake", - upsample_shortcut: str = "none", - ): - super().__init__() - layers = [ - _get_vae_activation(act_fn, channels=in_channels), - _wn_conv_transpose1d( - in_channels, out_channels, kernel_size=2 * stride, stride=stride, padding=math.ceil(stride / 2) - ), - VaeResidualUnit(out_channels, out_channels, dilation=1, act_fn=act_fn), - VaeResidualUnit(out_channels, out_channels, dilation=3, act_fn=act_fn), - VaeResidualUnit(out_channels, out_channels, dilation=9, act_fn=act_fn), - ] - self.layers = nn.Sequential(*layers) - self.residual = ( - UpsampleShortcut(in_channels, out_channels, stride) if upsample_shortcut == "duplicating" else None - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - output_hidden_states = self.layers(hidden_states) - if self.residual is not None: - residual = self.residual(hidden_states) - output_hidden_states = output_hidden_states + residual - return output_hidden_states - - -class AudioDiTVaeEncoder(nn.Module): - def __init__( - self, - in_channels: int = 1, - channels: int = 128, - c_mults: list[int] | None = None, - strides: list[int] | None = None, - latent_dim: int = 64, - encoder_latent_dim: int = 128, - act_fn: str = "snake", - downsample_shortcut: str = "averaging", - out_shortcut: str = "averaging", - ): - super().__init__() - c_mults = [1] + (c_mults or [1, 2, 4, 8, 16]) - strides = list(strides or [2] * (len(c_mults) - 1)) - if len(strides) < len(c_mults) - 1: - strides.extend([strides[-1] if strides else 2] * (len(c_mults) - 1 - len(strides))) - else: - strides = strides[: len(c_mults) - 1] - channels_base = channels - layers = [_wn_conv1d(in_channels, c_mults[0] * channels_base, kernel_size=7, padding=3)] - for idx in range(len(c_mults) - 1): - layers.append( - VaeEncoderBlock( - c_mults[idx] * channels_base, - c_mults[idx + 1] * channels_base, - strides[idx], - act_fn=act_fn, - downsample_shortcut=downsample_shortcut, - ) - ) - layers.append(_wn_conv1d(c_mults[-1] * channels_base, encoder_latent_dim, kernel_size=3, padding=1)) - self.layers = nn.Sequential(*layers) - self.shortcut = ( - DownsampleShortcut(c_mults[-1] * channels_base, encoder_latent_dim, 1) - if out_shortcut == "averaging" - else None - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.layers[:-1](hidden_states) - output_hidden_states = self.layers[-1](hidden_states) - if self.shortcut is not None: - shortcut = self.shortcut(hidden_states) - output_hidden_states = output_hidden_states + shortcut - return output_hidden_states - - -class AudioDiTVaeDecoder(nn.Module): - def __init__( - self, - in_channels: int = 1, - channels: int = 128, - c_mults: list[int] | None = None, - strides: list[int] | None = None, - latent_dim: int = 64, - act_fn: str = "snake", - in_shortcut: str = "duplicating", - final_tanh: bool = False, - upsample_shortcut: str = "duplicating", - ): - super().__init__() - c_mults = [1] + (c_mults or [1, 2, 4, 8, 16]) - strides = list(strides or [2] * (len(c_mults) - 1)) - if len(strides) < len(c_mults) - 1: - strides.extend([strides[-1] if strides else 2] * (len(c_mults) - 1 - len(strides))) - else: - strides = strides[: len(c_mults) - 1] - channels_base = channels - - self.shortcut = ( - UpsampleShortcut(latent_dim, c_mults[-1] * channels_base, 1) if in_shortcut == "duplicating" else None - ) - - layers = [_wn_conv1d(latent_dim, c_mults[-1] * channels_base, kernel_size=7, padding=3)] - for idx in range(len(c_mults) - 1, 0, -1): - layers.append( - VaeDecoderBlock( - c_mults[idx] * channels_base, - c_mults[idx - 1] * channels_base, - strides[idx - 1], - act_fn=act_fn, - upsample_shortcut=upsample_shortcut, - ) - ) - layers.append(_get_vae_activation(act_fn, channels=c_mults[0] * channels_base)) - layers.append(_wn_conv1d(c_mults[0] * channels_base, in_channels, kernel_size=7, padding=3, bias=False)) - layers.append(nn.Tanh() if final_tanh else nn.Identity()) - self.layers = nn.Sequential(*layers) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.shortcut is None: - return self.layers(hidden_states) - hidden_states = self.shortcut(hidden_states) + self.layers[0](hidden_states) - return self.layers[1:](hidden_states) - - -@dataclass -class LongCatAudioDiTVaeEncoderOutput(BaseOutput): - latents: torch.Tensor - - -@dataclass -class LongCatAudioDiTVaeDecoderOutput(BaseOutput): - sample: torch.Tensor - - -class LongCatAudioDiTVae(ModelMixin, AutoencoderMixin, ConfigMixin): - _supports_group_offloading = False - - @register_to_config - def __init__( - self, - in_channels: int = 1, - channels: int = 128, - c_mults: list[int] | None = None, - strides: list[int] | None = None, - latent_dim: int = 64, - encoder_latent_dim: int = 128, - act_fn: str | None = None, - use_snake: bool | None = None, - downsample_shortcut: str = "averaging", - upsample_shortcut: str = "duplicating", - out_shortcut: str = "averaging", - in_shortcut: str = "duplicating", - final_tanh: bool = False, - downsampling_ratio: int = 2048, - sample_rate: int = 24000, - scale: float = 0.71, - ): - super().__init__() - if act_fn is None: - if use_snake is None: - act_fn = "snake" - else: - act_fn = "snake" if use_snake else "elu" - self.encoder = AudioDiTVaeEncoder( - in_channels=in_channels, - channels=channels, - c_mults=c_mults, - strides=strides, - latent_dim=latent_dim, - encoder_latent_dim=encoder_latent_dim, - act_fn=act_fn, - downsample_shortcut=downsample_shortcut, - out_shortcut=out_shortcut, - ) - self.decoder = AudioDiTVaeDecoder( - in_channels=in_channels, - channels=channels, - c_mults=c_mults, - strides=strides, - latent_dim=latent_dim, - act_fn=act_fn, - in_shortcut=in_shortcut, - final_tanh=final_tanh, - upsample_shortcut=upsample_shortcut, - ) - - @apply_forward_hook - def encode( - self, - sample: torch.Tensor, - sample_posterior: bool = True, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> LongCatAudioDiTVaeEncoderOutput | tuple[torch.Tensor]: - encoder_dtype = next(self.encoder.parameters()).dtype - if sample.dtype != encoder_dtype: - sample = sample.to(encoder_dtype) - encoded = self.encoder(sample) - mean, scale_param = encoded.chunk(2, dim=1) - std = F.softplus(scale_param) + 1e-4 - if sample_posterior: - noise = randn_tensor(mean.shape, generator=generator, device=mean.device, dtype=mean.dtype) - latents = mean + std * noise - else: - latents = mean - latents = latents / self.config.scale - if encoder_dtype != torch.float32: - latents = latents.float() - if not return_dict: - return (latents,) - return LongCatAudioDiTVaeEncoderOutput(latents=latents) - - @apply_forward_hook - def decode( - self, latents: torch.Tensor, return_dict: bool = True - ) -> LongCatAudioDiTVaeDecoderOutput | tuple[torch.Tensor]: - decoder_dtype = next(self.decoder.parameters()).dtype - latents = latents * self.config.scale - if latents.dtype != decoder_dtype: - latents = latents.to(decoder_dtype) - decoded = self.decoder(latents) - if decoder_dtype != torch.float32: - decoded = decoded.float() - if not return_dict: - return (decoded,) - return LongCatAudioDiTVaeDecoderOutput(sample=decoded) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> LongCatAudioDiTVaeDecoderOutput | tuple[torch.Tensor]: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`LongCatAudioDiTVaeDecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`LongCatAudioDiTVaeDecoderOutput`] or `tuple`: - If `return_dict` is True, a [`LongCatAudioDiTVaeDecoderOutput`] is returned, otherwise a plain `tuple` - is returned. - """ - latents = self.encode(sample, sample_posterior=sample_posterior, return_dict=True, generator=generator).latents - decoded = self.decode(latents, return_dict=True).sample - if not return_dict: - return (decoded,) - return LongCatAudioDiTVaeDecoderOutput(sample=decoded) diff --git a/diffusers/models/autoencoders/autoencoder_oobleck.py b/diffusers/models/autoencoders/autoencoder_oobleck.py deleted file mode 100644 index d4251fd9f1a98eb1b3501811349c5e8e2d536224..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_oobleck.py +++ /dev/null @@ -1,551 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import math -from dataclasses import dataclass - -import numpy as np -import torch -import torch.nn as nn -from torch.nn.utils import weight_norm - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput -from ...utils.accelerate_utils import apply_forward_hook -from ...utils.torch_utils import randn_tensor -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin - - -class Snake1d(nn.Module): - """ - A 1-dimensional Snake activation function module. - """ - - def __init__(self, hidden_dim, logscale=True): - super().__init__() - self.alpha = nn.Parameter(torch.zeros(1, hidden_dim, 1)) - self.beta = nn.Parameter(torch.zeros(1, hidden_dim, 1)) - - self.alpha.requires_grad = True - self.beta.requires_grad = True - self.logscale = logscale - - def forward(self, hidden_states): - shape = hidden_states.shape - - alpha = self.alpha if not self.logscale else torch.exp(self.alpha) - beta = self.beta if not self.logscale else torch.exp(self.beta) - - hidden_states = hidden_states.reshape(shape[0], shape[1], -1) - hidden_states = hidden_states + (beta + 1e-9).reciprocal() * torch.sin(alpha * hidden_states).pow(2) - hidden_states = hidden_states.reshape(shape) - return hidden_states - - -class OobleckResidualUnit(nn.Module): - """ - A residual unit composed of Snake1d and weight-normalized Conv1d layers with dilations. - """ - - def __init__(self, dimension: int = 16, dilation: int = 1): - super().__init__() - pad = ((7 - 1) * dilation) // 2 - - self.snake1 = Snake1d(dimension) - self.conv1 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=7, dilation=dilation, padding=pad)) - self.snake2 = Snake1d(dimension) - self.conv2 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=1)) - - def forward(self, hidden_state): - """ - Forward pass through the residual unit. - - Args: - hidden_state (`torch.Tensor` of shape `(batch_size, channels, time_steps)`): - Input tensor . - - Returns: - output_tensor (`torch.Tensor` of shape `(batch_size, channels, time_steps)`) - Input tensor after passing through the residual unit. - """ - output_tensor = hidden_state - output_tensor = self.conv1(self.snake1(output_tensor)) - output_tensor = self.conv2(self.snake2(output_tensor)) - - padding = (hidden_state.shape[-1] - output_tensor.shape[-1]) // 2 - if padding > 0: - hidden_state = hidden_state[..., padding:-padding] - output_tensor = hidden_state + output_tensor - return output_tensor - - -class OobleckEncoderBlock(nn.Module): - """Encoder block used in Oobleck encoder.""" - - def __init__(self, input_dim, output_dim, stride: int = 1): - super().__init__() - - self.res_unit1 = OobleckResidualUnit(input_dim, dilation=1) - self.res_unit2 = OobleckResidualUnit(input_dim, dilation=3) - self.res_unit3 = OobleckResidualUnit(input_dim, dilation=9) - self.snake1 = Snake1d(input_dim) - self.conv1 = weight_norm( - nn.Conv1d(input_dim, output_dim, kernel_size=2 * stride, stride=stride, padding=math.ceil(stride / 2)) - ) - - def forward(self, hidden_state): - hidden_state = self.res_unit1(hidden_state) - hidden_state = self.res_unit2(hidden_state) - hidden_state = self.snake1(self.res_unit3(hidden_state)) - hidden_state = self.conv1(hidden_state) - - return hidden_state - - -class OobleckDecoderBlock(nn.Module): - """Decoder block used in Oobleck decoder.""" - - def __init__(self, input_dim, output_dim, stride: int = 1): - super().__init__() - - self.snake1 = Snake1d(input_dim) - self.conv_t1 = weight_norm( - nn.ConvTranspose1d( - input_dim, - output_dim, - kernel_size=2 * stride, - stride=stride, - padding=math.ceil(stride / 2), - ) - ) - self.res_unit1 = OobleckResidualUnit(output_dim, dilation=1) - self.res_unit2 = OobleckResidualUnit(output_dim, dilation=3) - self.res_unit3 = OobleckResidualUnit(output_dim, dilation=9) - - def forward(self, hidden_state): - hidden_state = self.snake1(hidden_state) - hidden_state = self.conv_t1(hidden_state) - hidden_state = self.res_unit1(hidden_state) - hidden_state = self.res_unit2(hidden_state) - hidden_state = self.res_unit3(hidden_state) - - return hidden_state - - -class OobleckDiagonalGaussianDistribution(object): - def __init__(self, parameters: torch.Tensor, deterministic: bool = False): - self.parameters = parameters - self.mean, self.scale = parameters.chunk(2, dim=1) - self.std = nn.functional.softplus(self.scale) + 1e-4 - self.var = self.std * self.std - self.logvar = torch.log(self.var) - self.deterministic = deterministic - - def sample(self, generator: torch.Generator | None = None) -> torch.Tensor: - # make sure sample is on the same device as the parameters and has same dtype - sample = randn_tensor( - self.mean.shape, - generator=generator, - device=self.parameters.device, - dtype=self.parameters.dtype, - ) - x = self.mean + self.std * sample - return x - - def kl(self, other: "OobleckDiagonalGaussianDistribution" = None) -> torch.Tensor: - if self.deterministic: - return torch.Tensor([0.0]) - else: - if other is None: - return (self.mean * self.mean + self.var - self.logvar - 1.0).sum(1).mean() - else: - normalized_diff = torch.pow(self.mean - other.mean, 2) / other.var - var_ratio = self.var / other.var - logvar_diff = self.logvar - other.logvar - - kl = normalized_diff + var_ratio + logvar_diff - 1 - - kl = kl.sum(1).mean() - return kl - - def mode(self) -> torch.Tensor: - return self.mean - - -@dataclass -class AutoencoderOobleckOutput(BaseOutput): - """ - Output of AutoencoderOobleck encoding method. - - Args: - latent_dist (`OobleckDiagonalGaussianDistribution`): - Encoded outputs of `Encoder` represented as the mean and standard deviation of - `OobleckDiagonalGaussianDistribution`. `OobleckDiagonalGaussianDistribution` allows for sampling latents - from the distribution. - """ - - latent_dist: "OobleckDiagonalGaussianDistribution" # noqa: F821 - - -@dataclass -class OobleckDecoderOutput(BaseOutput): - r""" - Output of decoding method. - - Args: - sample (`torch.Tensor` of shape `(batch_size, audio_channels, sequence_length)`): - The decoded output sample from the last layer of the model. - """ - - sample: torch.Tensor - - -class OobleckEncoder(nn.Module): - """Oobleck Encoder""" - - def __init__(self, encoder_hidden_size, audio_channels, downsampling_ratios, channel_multiples): - super().__init__() - - strides = downsampling_ratios - channel_multiples = [1] + channel_multiples - - # Create first convolution - self.conv1 = weight_norm(nn.Conv1d(audio_channels, encoder_hidden_size, kernel_size=7, padding=3)) - - self.block = [] - # Create EncoderBlocks that double channels as they downsample by `stride` - for stride_index, stride in enumerate(strides): - self.block += [ - OobleckEncoderBlock( - input_dim=encoder_hidden_size * channel_multiples[stride_index], - output_dim=encoder_hidden_size * channel_multiples[stride_index + 1], - stride=stride, - ) - ] - - self.block = nn.ModuleList(self.block) - d_model = encoder_hidden_size * channel_multiples[-1] - self.snake1 = Snake1d(d_model) - self.conv2 = weight_norm(nn.Conv1d(d_model, encoder_hidden_size, kernel_size=3, padding=1)) - - def forward(self, hidden_state): - hidden_state = self.conv1(hidden_state) - - for module in self.block: - hidden_state = module(hidden_state) - - hidden_state = self.snake1(hidden_state) - hidden_state = self.conv2(hidden_state) - - return hidden_state - - -class OobleckDecoder(nn.Module): - """Oobleck Decoder""" - - def __init__(self, channels, input_channels, audio_channels, upsampling_ratios, channel_multiples): - super().__init__() - - strides = upsampling_ratios - channel_multiples = [1] + channel_multiples - - # Add first conv layer - self.conv1 = weight_norm(nn.Conv1d(input_channels, channels * channel_multiples[-1], kernel_size=7, padding=3)) - - # Add upsampling + MRF blocks - block = [] - for stride_index, stride in enumerate(strides): - block += [ - OobleckDecoderBlock( - input_dim=channels * channel_multiples[len(strides) - stride_index], - output_dim=channels * channel_multiples[len(strides) - stride_index - 1], - stride=stride, - ) - ] - - self.block = nn.ModuleList(block) - output_dim = channels - self.snake1 = Snake1d(output_dim) - self.conv2 = weight_norm(nn.Conv1d(channels, audio_channels, kernel_size=7, padding=3, bias=False)) - - def forward(self, hidden_state): - hidden_state = self.conv1(hidden_state) - - for layer in self.block: - hidden_state = layer(hidden_state) - - hidden_state = self.snake1(hidden_state) - hidden_state = self.conv2(hidden_state) - - return hidden_state - - -class AutoencoderOobleck(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - An autoencoder for encoding waveforms into latents and decoding latent representations into waveforms. First - introduced in Stable Audio. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - encoder_hidden_size (`int`, *optional*, defaults to 128): - Intermediate representation dimension for the encoder. - downsampling_ratios (`list[int]`, *optional*, defaults to `[2, 4, 4, 8, 8]`): - Ratios for downsampling in the encoder. These are used in reverse order for upsampling in the decoder. - channel_multiples (`list[int]`, *optional*, defaults to `[1, 2, 4, 8, 16]`): - Multiples used to determine the hidden sizes of the hidden layers. - decoder_channels (`int`, *optional*, defaults to 128): - Intermediate representation dimension for the decoder. - decoder_input_channels (`int`, *optional*, defaults to 64): - Input dimension for the decoder. Corresponds to the latent dimension. - audio_channels (`int`, *optional*, defaults to 2): - Number of channels in the audio data. Either 1 for mono or 2 for stereo. - sampling_rate (`int`, *optional*, defaults to 44100): - The sampling rate at which the audio waveform should be digitalized expressed in hertz (Hz). - """ - - _supports_gradient_checkpointing = False - _supports_group_offloading = False - - @register_to_config - def __init__( - self, - encoder_hidden_size=128, - downsampling_ratios=[2, 4, 4, 8, 8], - channel_multiples=[1, 2, 4, 8, 16], - decoder_channels=128, - decoder_input_channels=64, - audio_channels=2, - sampling_rate=44100, - ): - super().__init__() - - self.encoder_hidden_size = encoder_hidden_size - self.downsampling_ratios = downsampling_ratios - self.decoder_channels = decoder_channels - self.upsampling_ratios = downsampling_ratios[::-1] - self.hop_length = int(np.prod(downsampling_ratios)) - self.sampling_rate = sampling_rate - - self.encoder = OobleckEncoder( - encoder_hidden_size=encoder_hidden_size, - audio_channels=audio_channels, - downsampling_ratios=downsampling_ratios, - channel_multiples=channel_multiples, - ) - - self.decoder = OobleckDecoder( - channels=decoder_channels, - input_channels=decoder_input_channels, - audio_channels=audio_channels, - upsampling_ratios=self.upsampling_ratios, - channel_multiples=channel_multiples, - ) - - self.use_slicing = False - self.use_tiling = False - - # 1D time-axis tiling defaults. `tile_sample_min_length` is the raw-audio - # threshold (in samples) above which `encode` splits the input; chunks are - # `tile_sample_min_length` wide with `tile_sample_overlap` samples of overlap - # on each side, trimmed back out after decoding. `tile_latent_min_length` - # is the equivalent threshold on the decode side, expressed in latent frames. - self.tile_sample_min_length = sampling_rate * 30 # 30 seconds - self.tile_sample_overlap = sampling_rate * 2 # 2 seconds per side - # Decode chunk is smaller than encode chunk because the decoder upsamples - # back to raw audio and is more VRAM-heavy per frame. - self.tile_latent_min_length = 512 - self.tile_latent_overlap = 64 - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - if self.use_tiling and x.shape[-1] > self.tile_sample_min_length: - return self._tiled_encode(x) - return self.encoder(x) - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> AutoencoderOobleckOutput | tuple[OobleckDiagonalGaussianDistribution]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple. - - Returns: - The latent representations of the encoded images. If `return_dict` is True, a - [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self._encode(x) - - posterior = OobleckDiagonalGaussianDistribution(h) - - if not return_dict: - return (posterior,) - - return AutoencoderOobleckOutput(latent_dist=posterior) - - def _tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - r"""Encode a long audio waveform by splitting it into overlapping tiles along - the time axis and concatenating the resulting encoder features. Used to keep memory bounded regardless of clip - length. Not bit-identical to a single unsplit encode — each tile has its own receptive-field boundary — but the - overlap/trim scheme keeps the joined feature map smooth. - """ - _B, _C, S = x.shape - chunk = self.tile_sample_min_length - overlap = self.tile_sample_overlap - stride = chunk - 2 * overlap - if stride <= 0: - raise ValueError( - f"tile_sample_min_length ({chunk}) must be greater than 2 * tile_sample_overlap ({overlap})" - ) - - num_steps = math.ceil(S / stride) - tiles = [] - hop = None - - for i in range(num_steps): - core_start = i * stride - core_end = min(core_start + stride, S) - win_start = max(0, core_start - overlap) - win_end = min(S, core_end + overlap) - - tile = self.encoder(x[:, :, win_start:win_end]) - - if hop is None: - hop = (win_end - win_start) / tile.shape[-1] - - trim_l = int(round((core_start - win_start) / hop)) - trim_r = int(round((win_end - core_end) / hop)) - end_idx = tile.shape[-1] - trim_r if trim_r > 0 else tile.shape[-1] - tiles.append(tile[:, :, trim_l:end_idx]) - - return torch.cat(tiles, dim=-1) - - def _decode(self, z: torch.Tensor, return_dict: bool = True) -> OobleckDecoderOutput | torch.Tensor: - if self.use_tiling and z.shape[-1] > self.tile_latent_min_length: - dec = self._tiled_decode(z) - else: - dec = self.decoder(z) - - if not return_dict: - return (dec,) - - return OobleckDecoderOutput(sample=dec) - - def _tiled_decode(self, z: torch.Tensor) -> torch.Tensor: - r"""Decode a long latent by splitting it into overlapping tiles along the - time axis, decoding each, and concatenating the audio tiles back together.""" - _B, _C, T = z.shape - chunk = self.tile_latent_min_length - overlap = self.tile_latent_overlap - stride = chunk - 2 * overlap - if stride <= 0: - raise ValueError( - f"tile_latent_min_length ({chunk}) must be greater than 2 * tile_latent_overlap ({overlap})" - ) - - num_steps = math.ceil(T / stride) - tiles = [] - upsample = None - - for i in range(num_steps): - core_start = i * stride - core_end = min(core_start + stride, T) - win_start = max(0, core_start - overlap) - win_end = min(T, core_end + overlap) - - tile = self.decoder(z[:, :, win_start:win_end]) - - if upsample is None: - upsample = tile.shape[-1] / (win_end - win_start) - - trim_l = int(round((core_start - win_start) * upsample)) - trim_r = int(round((win_end - core_end) * upsample)) - end_idx = tile.shape[-1] - trim_r if trim_r > 0 else tile.shape[-1] - tiles.append(tile[:, :, trim_l:end_idx]) - - return torch.cat(tiles, dim=-1) - - @apply_forward_hook - def decode( - self, z: torch.FloatTensor, return_dict: bool = True, generator=None - ) -> OobleckDecoderOutput | torch.FloatTensor: - """ - Decode a batch of images. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.vae.OobleckDecoderOutput`] instead of a plain tuple. - - Returns: - [`~models.vae.OobleckDecoderOutput`] or `tuple`: - If return_dict is True, a [`~models.vae.OobleckDecoderOutput`] is returned, otherwise a plain `tuple` - is returned. - - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z).sample - - if not return_dict: - return (decoded,) - - return OobleckDecoderOutput(sample=decoded) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> OobleckDecoderOutput | torch.Tensor: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`OobleckDecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.OobleckDecoderOutput`] or `tuple`: - If `return_dict` is True, a [`~models.vae.OobleckDecoderOutput`] is returned, otherwise a plain `tuple` - is returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z).sample - - if not return_dict: - return (dec,) - - return OobleckDecoderOutput(sample=dec) diff --git a/diffusers/models/autoencoders/autoencoder_rae.py b/diffusers/models/autoencoders/autoencoder_rae.py deleted file mode 100644 index 35a96e6f67bccfa135f784446560ae29cde6cb91..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_rae.py +++ /dev/null @@ -1,702 +0,0 @@ -# Copyright 2026 The NYU Vision-X and HuggingFace Teams. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass -from math import sqrt -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput, logging -from ...utils.accelerate_utils import apply_forward_hook -from ...utils.import_utils import is_transformers_available -from ...utils.torch_utils import randn_tensor - - -if is_transformers_available(): - from transformers import ( - Dinov2WithRegistersConfig, - Dinov2WithRegistersModel, - SiglipVisionConfig, - SiglipVisionModel, - ViTMAEConfig, - ViTMAEModel, - ) - -from ..activations import get_activation -from ..attention import AttentionMixin -from ..attention_processor import Attention -from ..embeddings import get_2d_sincos_pos_embed -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, EncoderOutput - - -logger = logging.get_logger(__name__) - - -# --------------------------------------------------------------------------- -# Per-encoder forward functions -# --------------------------------------------------------------------------- -# Each function takes the raw transformers model + images and returns patch -# tokens of shape (B, N, C), stripping CLS / register tokens as needed. - - -def _dinov2_encoder_forward(model: nn.Module, images: torch.Tensor) -> torch.Tensor: - outputs = model(images, output_hidden_states=True) - unused_token_num = 5 # 1 CLS + 4 register tokens - return outputs.last_hidden_state[:, unused_token_num:] - - -def _siglip2_encoder_forward(model: nn.Module, images: torch.Tensor) -> torch.Tensor: - outputs = model(images, output_hidden_states=True, interpolate_pos_encoding=True) - return outputs.last_hidden_state - - -def _mae_encoder_forward(model: nn.Module, images: torch.Tensor, patch_size: int) -> torch.Tensor: - h, w = images.shape[2], images.shape[3] - patch_num = int(h * w // patch_size**2) - if patch_num * patch_size**2 != h * w: - raise ValueError("Image size should be divisible by patch size.") - noise = torch.arange(patch_num).unsqueeze(0).expand(images.shape[0], -1).to(images.device).to(images.dtype) - outputs = model(images, noise, interpolate_pos_encoding=True) - return outputs.last_hidden_state[:, 1:] # remove cls token - - -# --------------------------------------------------------------------------- -# Encoder construction helpers -# --------------------------------------------------------------------------- - - -def _build_encoder( - encoder_type: str, hidden_size: int, patch_size: int, num_hidden_layers: int, head_dim: int = 64 -) -> nn.Module: - """Build a frozen encoder from config (no pretrained download).""" - num_attention_heads = hidden_size // head_dim # all supported encoders use head_dim=64 - - if encoder_type == "dinov2": - config = Dinov2WithRegistersConfig( - hidden_size=hidden_size, - patch_size=patch_size, - image_size=518, - num_attention_heads=num_attention_heads, - num_hidden_layers=num_hidden_layers, - ) - model = Dinov2WithRegistersModel(config) - # RAE strips the final layernorm affine params (identity LN). Remove them from - # the architecture so `from_pretrained` doesn't leave them on the meta device. - model.layernorm.weight = None - model.layernorm.bias = None - elif encoder_type == "siglip2": - config = SiglipVisionConfig( - hidden_size=hidden_size, - patch_size=patch_size, - image_size=256, - num_attention_heads=num_attention_heads, - num_hidden_layers=num_hidden_layers, - ) - model = SiglipVisionModel(config) - # See dinov2 comment above. - model.vision_model.post_layernorm.weight = None - model.vision_model.post_layernorm.bias = None - elif encoder_type == "mae": - config = ViTMAEConfig( - hidden_size=hidden_size, - patch_size=patch_size, - image_size=224, - num_attention_heads=num_attention_heads, - num_hidden_layers=num_hidden_layers, - mask_ratio=0.0, - ) - model = ViTMAEModel(config) - # See dinov2 comment above. - model.layernorm.weight = None - model.layernorm.bias = None - else: - raise ValueError(f"Unknown encoder_type='{encoder_type}'. Available: dinov2, siglip2, mae") - - model.requires_grad_(False) - return model - - -_ENCODER_FORWARD_FNS = { - "dinov2": _dinov2_encoder_forward, - "siglip2": _siglip2_encoder_forward, - "mae": _mae_encoder_forward, -} - - -@dataclass -class RAEDecoderOutput(BaseOutput): - """ - Output of `RAEDecoder`. - - Args: - logits (`torch.Tensor`): - Patch reconstruction logits of shape `(batch_size, num_patches, patch_size**2 * num_channels)`. - """ - - logits: torch.Tensor - - -class ViTMAEIntermediate(nn.Module): - def __init__(self, hidden_size: int, intermediate_size: int, hidden_act: str = "gelu"): - super().__init__() - self.dense = nn.Linear(hidden_size, intermediate_size) - self.intermediate_act_fn = get_activation(hidden_act) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.dense(hidden_states) - hidden_states = self.intermediate_act_fn(hidden_states) - return hidden_states - - -class ViTMAEOutput(nn.Module): - def __init__(self, hidden_size: int, intermediate_size: int, hidden_dropout_prob: float = 0.0): - super().__init__() - self.dense = nn.Linear(intermediate_size, hidden_size) - self.dropout = nn.Dropout(hidden_dropout_prob) - - def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor: - hidden_states = self.dense(hidden_states) - hidden_states = self.dropout(hidden_states) - hidden_states = hidden_states + input_tensor - return hidden_states - - -class ViTMAELayer(nn.Module): - """ - This matches the naming/parameter structure used in RAE-main (ViTMAE decoder block). - """ - - def __init__( - self, - *, - hidden_size: int, - num_attention_heads: int, - intermediate_size: int, - qkv_bias: bool = True, - layer_norm_eps: float = 1e-12, - hidden_dropout_prob: float = 0.0, - attention_probs_dropout_prob: float = 0.0, - hidden_act: str = "gelu", - ): - super().__init__() - if hidden_size % num_attention_heads != 0: - raise ValueError( - f"hidden_size={hidden_size} must be divisible by num_attention_heads={num_attention_heads}" - ) - self.attention = Attention( - query_dim=hidden_size, - heads=num_attention_heads, - dim_head=hidden_size // num_attention_heads, - dropout=attention_probs_dropout_prob, - bias=qkv_bias, - ) - self.intermediate = ViTMAEIntermediate( - hidden_size=hidden_size, intermediate_size=intermediate_size, hidden_act=hidden_act - ) - self.output = ViTMAEOutput( - hidden_size=hidden_size, intermediate_size=intermediate_size, hidden_dropout_prob=hidden_dropout_prob - ) - self.layernorm_before = nn.LayerNorm(hidden_size, eps=layer_norm_eps) - self.layernorm_after = nn.LayerNorm(hidden_size, eps=layer_norm_eps) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - attention_output = self.attention(self.layernorm_before(hidden_states)) - hidden_states = attention_output + hidden_states - - layer_output = self.layernorm_after(hidden_states) - layer_output = self.intermediate(layer_output) - layer_output = self.output(layer_output, hidden_states) - return layer_output - - -class RAEDecoder(nn.Module): - """ - Decoder implementation ported from RAE-main to keep checkpoint compatibility. - - Key attributes (must match checkpoint keys): - - decoder_embed - - decoder_pos_embed - - decoder_layers - - decoder_norm - - decoder_pred - - trainable_cls_token - """ - - def __init__( - self, - hidden_size: int = 768, - decoder_hidden_size: int = 512, - decoder_num_hidden_layers: int = 8, - decoder_num_attention_heads: int = 16, - decoder_intermediate_size: int = 2048, - num_patches: int = 256, - patch_size: int = 16, - num_channels: int = 3, - image_size: int = 256, - qkv_bias: bool = True, - layer_norm_eps: float = 1e-12, - hidden_dropout_prob: float = 0.0, - attention_probs_dropout_prob: float = 0.0, - hidden_act: str = "gelu", - ): - super().__init__() - self.decoder_hidden_size = decoder_hidden_size - self.patch_size = patch_size - self.num_channels = num_channels - self.image_size = image_size - self.num_patches = num_patches - - self.decoder_embed = nn.Linear(hidden_size, decoder_hidden_size, bias=True) - grid_size = int(num_patches**0.5) - pos_embed = get_2d_sincos_pos_embed( - decoder_hidden_size, grid_size, cls_token=True, extra_tokens=1, output_type="pt" - ) - self.register_buffer("decoder_pos_embed", pos_embed.unsqueeze(0).float(), persistent=False) - - self.decoder_layers = nn.ModuleList( - [ - ViTMAELayer( - hidden_size=decoder_hidden_size, - num_attention_heads=decoder_num_attention_heads, - intermediate_size=decoder_intermediate_size, - qkv_bias=qkv_bias, - layer_norm_eps=layer_norm_eps, - hidden_dropout_prob=hidden_dropout_prob, - attention_probs_dropout_prob=attention_probs_dropout_prob, - hidden_act=hidden_act, - ) - for _ in range(decoder_num_hidden_layers) - ] - ) - - self.decoder_norm = nn.LayerNorm(decoder_hidden_size, eps=layer_norm_eps) - self.decoder_pred = nn.Linear(decoder_hidden_size, patch_size**2 * num_channels, bias=True) - self.gradient_checkpointing = False - - self.trainable_cls_token = nn.Parameter(torch.zeros(1, 1, decoder_hidden_size)) - - def interpolate_pos_encoding(self, embeddings: torch.Tensor) -> torch.Tensor: - embeddings_positions = embeddings.shape[1] - 1 - num_positions = self.decoder_pos_embed.shape[1] - 1 - - class_pos_embed = self.decoder_pos_embed[:, 0, :] - patch_pos_embed = self.decoder_pos_embed[:, 1:, :] - dim = self.decoder_pos_embed.shape[-1] - - patch_pos_embed = patch_pos_embed.reshape(1, 1, -1, dim).permute(0, 3, 1, 2) - patch_pos_embed = F.interpolate( - patch_pos_embed, - scale_factor=(1, embeddings_positions / num_positions), - mode="bicubic", - align_corners=False, - ) - patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).view(1, -1, dim) - return torch.cat((class_pos_embed.unsqueeze(0), patch_pos_embed), dim=1) - - def interpolate_latent(self, x: torch.Tensor) -> torch.Tensor: - b, l, c = x.shape - if l == self.num_patches: - return x - h = w = int(l**0.5) - x = x.reshape(b, h, w, c).permute(0, 3, 1, 2) - target_size = (int(self.num_patches**0.5), int(self.num_patches**0.5)) - x = F.interpolate(x, size=target_size, mode="bilinear", align_corners=False) - x = x.permute(0, 2, 3, 1).contiguous().view(b, self.num_patches, c) - return x - - def unpatchify(self, patchified_pixel_values: torch.Tensor, original_image_size: tuple[int, int] | None = None): - patch_size, num_channels = self.patch_size, self.num_channels - original_image_size = ( - original_image_size if original_image_size is not None else (self.image_size, self.image_size) - ) - original_height, original_width = original_image_size - num_patches_h = original_height // patch_size - num_patches_w = original_width // patch_size - if num_patches_h * num_patches_w != patchified_pixel_values.shape[1]: - raise ValueError( - f"The number of patches in the patchified pixel values {patchified_pixel_values.shape[1]}, does not match the number of patches on original image {num_patches_h}*{num_patches_w}" - ) - - batch_size = patchified_pixel_values.shape[0] - patchified_pixel_values = patchified_pixel_values.reshape( - batch_size, - num_patches_h, - num_patches_w, - patch_size, - patch_size, - num_channels, - ) - patchified_pixel_values = torch.einsum("nhwpqc->nchpwq", patchified_pixel_values) - pixel_values = patchified_pixel_values.reshape( - batch_size, - num_channels, - num_patches_h * patch_size, - num_patches_w * patch_size, - ) - return pixel_values - - def forward( - self, - hidden_states: torch.Tensor, - *, - interpolate_pos_encoding: bool = False, - drop_cls_token: bool = False, - return_dict: bool = True, - ) -> RAEDecoderOutput | tuple[torch.Tensor]: - x = self.decoder_embed(hidden_states) - if drop_cls_token: - x_ = x[:, 1:, :] - x_ = self.interpolate_latent(x_) - else: - x_ = self.interpolate_latent(x) - - cls_token = self.trainable_cls_token.expand(x_.shape[0], -1, -1) - x = torch.cat([cls_token, x_], dim=1) - - if interpolate_pos_encoding: - if not drop_cls_token: - raise ValueError("interpolate_pos_encoding only supports drop_cls_token=True") - decoder_pos_embed = self.interpolate_pos_encoding(x) - else: - decoder_pos_embed = self.decoder_pos_embed - - hidden_states = x + decoder_pos_embed.to(device=x.device, dtype=x.dtype) - - for layer_module in self.decoder_layers: - hidden_states = layer_module(hidden_states) - - hidden_states = self.decoder_norm(hidden_states) - logits = self.decoder_pred(hidden_states) - logits = logits[:, 1:, :] - - if not return_dict: - return (logits,) - return RAEDecoderOutput(logits=logits) - - -class AutoencoderRAE(ModelMixin, AttentionMixin, AutoencoderMixin, ConfigMixin): - r""" - Representation Autoencoder (RAE) model for encoding images to latents and decoding latents to images. - - This model uses a frozen pretrained encoder (DINOv2, SigLIP2, or MAE) with a trainable ViT decoder to reconstruct - images from learned representations. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for its generic methods implemented for - all models (such as downloading or saving). - - Args: - encoder_type (`str`, *optional*, defaults to `"dinov2"`): - Type of frozen encoder to use. One of `"dinov2"`, `"siglip2"`, or `"mae"`. - encoder_hidden_size (`int`, *optional*, defaults to `768`): - Hidden size of the encoder model. - encoder_patch_size (`int`, *optional*, defaults to `14`): - Patch size of the encoder model. - encoder_num_hidden_layers (`int`, *optional*, defaults to `12`): - Number of hidden layers in the encoder model. - patch_size (`int`, *optional*, defaults to `16`): - Decoder patch size (used for unpatchify and decoder head). - encoder_input_size (`int`, *optional*, defaults to `224`): - Input size expected by the encoder. - image_size (`int`, *optional*): - Decoder output image size. If `None`, it is derived from encoder token count and `patch_size` like - RAE-main: `image_size = patch_size * sqrt(num_patches)`, where `num_patches = (encoder_input_size // - encoder_patch_size) ** 2`. - num_channels (`int`, *optional*, defaults to `3`): - Number of input/output channels. - encoder_norm_mean (`list`, *optional*, defaults to `[0.485, 0.456, 0.406]`): - Channel-wise mean for encoder input normalization (ImageNet defaults). - encoder_norm_std (`list`, *optional*, defaults to `[0.229, 0.224, 0.225]`): - Channel-wise std for encoder input normalization (ImageNet defaults). - latents_mean (`list` or `tuple`, *optional*): - Optional mean for latent normalization. Tensor inputs are accepted and converted to config-serializable - lists. - latents_std (`list` or `tuple`, *optional*): - Optional standard deviation for latent normalization. Tensor inputs are accepted and converted to - config-serializable lists. - noise_tau (`float`, *optional*, defaults to `0.0`): - Noise level for training (adds noise to latents during training). - reshape_to_2d (`bool`, *optional*, defaults to `True`): - Whether to reshape latents to 2D (B, C, H, W) format. - use_encoder_loss (`bool`, *optional*, defaults to `False`): - Whether to use encoder hidden states in the loss (for advanced training). - """ - - # NOTE: gradient checkpointing is not wired up for this model yet. - _supports_gradient_checkpointing = False - _no_split_modules = ["ViTMAELayer"] - _keys_to_ignore_on_load_unexpected = ["decoder.decoder_pos_embed"] - - @register_to_config - def __init__( - self, - encoder_type: str = "dinov2", - encoder_hidden_size: int = 768, - encoder_patch_size: int = 14, - encoder_num_hidden_layers: int = 12, - decoder_hidden_size: int = 512, - decoder_num_hidden_layers: int = 8, - decoder_num_attention_heads: int = 16, - decoder_intermediate_size: int = 2048, - patch_size: int = 16, - encoder_input_size: int = 224, - image_size: int | None = None, - num_channels: int = 3, - encoder_norm_mean: list | None = None, - encoder_norm_std: list | None = None, - latents_mean: list | tuple | torch.Tensor | None = None, - latents_std: list | tuple | torch.Tensor | None = None, - noise_tau: float = 0.0, - reshape_to_2d: bool = True, - use_encoder_loss: bool = False, - scaling_factor: float = 1.0, - ): - super().__init__() - - if encoder_type not in _ENCODER_FORWARD_FNS: - raise ValueError( - f"Unknown encoder_type='{encoder_type}'. Available: {sorted(_ENCODER_FORWARD_FNS.keys())}" - ) - - def _to_config_compatible(value: Any) -> Any: - if isinstance(value, torch.Tensor): - return value.detach().cpu().tolist() - if isinstance(value, tuple): - return [_to_config_compatible(v) for v in value] - if isinstance(value, list): - return [_to_config_compatible(v) for v in value] - return value - - def _as_optional_tensor(value: torch.Tensor | list | tuple | None) -> torch.Tensor | None: - if value is None: - return None - if isinstance(value, torch.Tensor): - return value.detach().clone() - return torch.tensor(value, dtype=torch.float32) - - latents_std_tensor = _as_optional_tensor(latents_std) - - # Ensure config values are JSON-serializable (list/None), even if caller passes torch.Tensors. - self.register_to_config( - latents_mean=_to_config_compatible(latents_mean), - latents_std=_to_config_compatible(latents_std), - ) - - self.encoder_input_size = encoder_input_size - self.noise_tau = float(noise_tau) - self.reshape_to_2d = bool(reshape_to_2d) - self.use_encoder_loss = bool(use_encoder_loss) - - # Validate early, before building the (potentially large) encoder/decoder. - encoder_patch_size = int(encoder_patch_size) - if self.encoder_input_size % encoder_patch_size != 0: - raise ValueError( - f"encoder_input_size={self.encoder_input_size} must be divisible by encoder_patch_size={encoder_patch_size}." - ) - decoder_patch_size = int(patch_size) - if decoder_patch_size <= 0: - raise ValueError("patch_size must be a positive integer (this is decoder_patch_size).") - - # Frozen representation encoder (built from config, no downloads) - self.encoder: nn.Module = _build_encoder( - encoder_type=encoder_type, - hidden_size=encoder_hidden_size, - patch_size=encoder_patch_size, - num_hidden_layers=encoder_num_hidden_layers, - ) - self._encoder_forward_fn = _ENCODER_FORWARD_FNS[encoder_type] - num_patches = (self.encoder_input_size // encoder_patch_size) ** 2 - - grid = int(sqrt(num_patches)) - if grid * grid != num_patches: - raise ValueError(f"Computed num_patches={num_patches} must be a perfect square.") - - derived_image_size = decoder_patch_size * grid - if image_size is None: - image_size = derived_image_size - else: - image_size = int(image_size) - if image_size != derived_image_size: - raise ValueError( - f"image_size={image_size} must equal decoder_patch_size*sqrt(num_patches)={derived_image_size} " - f"for patch_size={decoder_patch_size} and computed num_patches={num_patches}." - ) - - # Encoder input normalization stats (ImageNet defaults) - if encoder_norm_mean is None: - encoder_norm_mean = [0.485, 0.456, 0.406] - if encoder_norm_std is None: - encoder_norm_std = [0.229, 0.224, 0.225] - encoder_mean_tensor = torch.tensor(encoder_norm_mean, dtype=torch.float32).view(1, 3, 1, 1) - encoder_std_tensor = torch.tensor(encoder_norm_std, dtype=torch.float32).view(1, 3, 1, 1) - - self.register_buffer("encoder_mean", encoder_mean_tensor, persistent=True) - self.register_buffer("encoder_std", encoder_std_tensor, persistent=True) - - # Latent normalization buffers (defaults are no-ops; actual values come from checkpoint) - latents_mean_tensor = _as_optional_tensor(latents_mean) - if latents_mean_tensor is None: - latents_mean_tensor = torch.zeros(1) - self.register_buffer("_latents_mean", latents_mean_tensor, persistent=True) - - if latents_std_tensor is None: - latents_std_tensor = torch.ones(1) - self.register_buffer("_latents_std", latents_std_tensor, persistent=True) - - # ViT-MAE style decoder - self.decoder = RAEDecoder( - hidden_size=int(encoder_hidden_size), - decoder_hidden_size=int(decoder_hidden_size), - decoder_num_hidden_layers=int(decoder_num_hidden_layers), - decoder_num_attention_heads=int(decoder_num_attention_heads), - decoder_intermediate_size=int(decoder_intermediate_size), - num_patches=int(num_patches), - patch_size=int(decoder_patch_size), - num_channels=int(num_channels), - image_size=int(image_size), - ) - self.num_patches = int(num_patches) - self.decoder_patch_size = int(decoder_patch_size) - self.decoder_image_size = int(image_size) - - # Slicing support (batch dimension) similar to other diffusers autoencoders - self.use_slicing = False - - def _noising(self, x: torch.Tensor, generator: torch.Generator | None = None) -> torch.Tensor: - # Per-sample random sigma in [0, noise_tau] - noise_sigma = self.noise_tau * torch.rand( - (x.size(0),) + (1,) * (x.ndim - 1), device=x.device, dtype=x.dtype, generator=generator - ) - return x + noise_sigma * randn_tensor(x.shape, generator=generator, device=x.device, dtype=x.dtype) - - def _resize_and_normalize(self, x: torch.Tensor) -> torch.Tensor: - _, _, h, w = x.shape - if h != self.encoder_input_size or w != self.encoder_input_size: - x = F.interpolate( - x, size=(self.encoder_input_size, self.encoder_input_size), mode="bicubic", align_corners=False - ) - mean = self.encoder_mean.to(device=x.device, dtype=x.dtype) - std = self.encoder_std.to(device=x.device, dtype=x.dtype) - return (x - mean) / std - - def _denormalize_image(self, x: torch.Tensor) -> torch.Tensor: - mean = self.encoder_mean.to(device=x.device, dtype=x.dtype) - std = self.encoder_std.to(device=x.device, dtype=x.dtype) - return x * std + mean - - def _normalize_latents(self, z: torch.Tensor) -> torch.Tensor: - latents_mean = self._latents_mean.to(device=z.device, dtype=z.dtype) - latents_std = self._latents_std.to(device=z.device, dtype=z.dtype) - return (z - latents_mean) / (latents_std + 1e-5) - - def _denormalize_latents(self, z: torch.Tensor) -> torch.Tensor: - latents_mean = self._latents_mean.to(device=z.device, dtype=z.dtype) - latents_std = self._latents_std.to(device=z.device, dtype=z.dtype) - return z * (latents_std + 1e-5) + latents_mean - - def _encode(self, x: torch.Tensor, generator: torch.Generator | None = None) -> torch.Tensor: - x = self._resize_and_normalize(x) - - if self.config.encoder_type == "mae": - tokens = self._encoder_forward_fn(self.encoder, x, self.config.encoder_patch_size) - else: - tokens = self._encoder_forward_fn(self.encoder, x) # (B, N, C) - - if self.training and self.noise_tau > 0: - tokens = self._noising(tokens, generator=generator) - - if self.reshape_to_2d: - b, n, c = tokens.shape - side = int(sqrt(n)) - if side * side != n: - raise ValueError(f"Token length n={n} is not a perfect square; cannot reshape to 2D.") - z = tokens.transpose(1, 2).contiguous().view(b, c, side, side) # (B, C, h, w) - else: - z = tokens - - z = self._normalize_latents(z) - - # Follow diffusers convention: optionally scale latents for diffusion - if self.config.scaling_factor != 1.0: - z = z * self.config.scaling_factor - - return z - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True, generator: torch.Generator | None = None - ) -> EncoderOutput | tuple[torch.Tensor]: - if self.use_slicing and x.shape[0] > 1: - latents = torch.cat([self._encode(x_slice, generator=generator) for x_slice in x.split(1)], dim=0) - else: - latents = self._encode(x, generator=generator) - - if not return_dict: - return (latents,) - return EncoderOutput(latent=latents) - - def _decode(self, z: torch.Tensor) -> torch.Tensor: - # Undo scaling factor if applied at encode time - if self.config.scaling_factor != 1.0: - z = z / self.config.scaling_factor - - z = self._denormalize_latents(z) - - if self.reshape_to_2d: - b, c, h, w = z.shape - tokens = z.view(b, c, h * w).transpose(1, 2).contiguous() # (B, N, C) - else: - tokens = z - - logits = self.decoder(tokens, return_dict=True).logits - x_rec = self.decoder.unpatchify(logits) - x_rec = self._denormalize_image(x_rec) - return x_rec.to(device=z.device) - - @apply_forward_hook - def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | tuple[torch.Tensor]: - if self.use_slicing and z.shape[0] > 1: - decoded = torch.cat([self._decode(z_slice) for z_slice in z.split(1)], dim=0) - else: - decoded = self._decode(z) - - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) - - def forward( - self, sample: torch.Tensor, return_dict: bool = True, generator: torch.Generator | None = None - ) -> DecoderOutput | tuple[torch.Tensor]: - r""" - Args: - sample (`torch.Tensor`): Input sample. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`DecoderOutput`] is returned, otherwise a plain `tuple` is returned. - """ - latents = self.encode(sample, return_dict=False, generator=generator)[0] - decoded = self.decode(latents, return_dict=False)[0] - if not return_dict: - return (decoded,) - return DecoderOutput(sample=decoded) diff --git a/diffusers/models/autoencoders/autoencoder_tiny.py b/diffusers/models/autoencoders/autoencoder_tiny.py deleted file mode 100644 index 5647203e02e1b62bcb196faeb6c77cf295e6558b..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_tiny.py +++ /dev/null @@ -1,320 +0,0 @@ -# Copyright 2025 Ollin Boer Bohan and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from dataclasses import dataclass - -import torch - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput -from ...utils.accelerate_utils import apply_forward_hook -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin, DecoderOutput, DecoderTiny, EncoderTiny - - -@dataclass -class AutoencoderTinyOutput(BaseOutput): - """ - Output of AutoencoderTiny encoding method. - - Args: - latents (`torch.Tensor`): Encoded outputs of the `Encoder`. - - """ - - latents: torch.Tensor - - -class AutoencoderTiny(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - A tiny distilled VAE model for encoding images into latents and decoding latent representations into images. - - [`AutoencoderTiny`] is a wrapper around the original implementation of `TAESD`. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for its generic methods implemented for - all models (such as downloading or saving). - - Parameters: - in_channels (`int`, *optional*, defaults to 3): Number of channels in the input image. - out_channels (`int`, *optional*, defaults to 3): Number of channels in the output. - encoder_block_out_channels (`tuple[int]`, *optional*, defaults to `(64, 64, 64, 64)`): - tuple of integers representing the number of output channels for each encoder block. The length of the - tuple should be equal to the number of encoder blocks. - decoder_block_out_channels (`tuple[int]`, *optional*, defaults to `(64, 64, 64, 64)`): - tuple of integers representing the number of output channels for each decoder block. The length of the - tuple should be equal to the number of decoder blocks. - act_fn (`str`, *optional*, defaults to `"relu"`): - Activation function to be used throughout the model. - latent_channels (`int`, *optional*, defaults to 4): - Number of channels in the latent representation. The latent space acts as a compressed representation of - the input image. - upsampling_scaling_factor (`int`, *optional*, defaults to 2): - Scaling factor for upsampling in the decoder. It determines the size of the output image during the - upsampling process. - num_encoder_blocks (`tuple[int]`, *optional*, defaults to `(1, 3, 3, 3)`): - tuple of integers representing the number of encoder blocks at each stage of the encoding process. The - length of the tuple should be equal to the number of stages in the encoder. Each stage has a different - number of encoder blocks. - num_decoder_blocks (`tuple[int]`, *optional*, defaults to `(3, 3, 3, 1)`): - tuple of integers representing the number of decoder blocks at each stage of the decoding process. The - length of the tuple should be equal to the number of stages in the decoder. Each stage has a different - number of decoder blocks. - latent_magnitude (`float`, *optional*, defaults to 3.0): - Magnitude of the latent representation. This parameter scales the latent representation values to control - the extent of information preservation. - latent_shift (float, *optional*, defaults to 0.5): - Shift applied to the latent representation. This parameter controls the center of the latent space. - scaling_factor (`float`, *optional*, defaults to 1.0): - The component-wise standard deviation of the trained latent space computed using the first batch of the - training set. This is used to scale the latent space to have unit variance when training the diffusion - model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the - diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 - / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image - Synthesis with Latent Diffusion Models](https://huggingface.co/papers/2112.10752) paper. For this - Autoencoder, however, no such scaling factor was used, hence the value of 1.0 as the default. - force_upcast (`bool`, *optional*, default to `False`): - If enabled it will force the VAE to run in float32 for high image resolution pipelines, such as SD-XL. VAE - can be fine-tuned / trained to a lower range without losing too much precision, in which case - `force_upcast` can be set to `False` (see this fp16-friendly - [AutoEncoder](https://huggingface.co/madebyollin/sdxl-vae-fp16-fix)). - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - encoder_block_out_channels: tuple[int, ...] = (64, 64, 64, 64), - decoder_block_out_channels: tuple[int, ...] = (64, 64, 64, 64), - act_fn: str = "relu", - upsample_fn: str = "nearest", - latent_channels: int = 4, - upsampling_scaling_factor: int = 2, - num_encoder_blocks: tuple[int, ...] = (1, 3, 3, 3), - num_decoder_blocks: tuple[int, ...] = (3, 3, 3, 1), - latent_magnitude: int = 3, - latent_shift: float = 0.5, - force_upcast: bool = False, - scaling_factor: float = 1.0, - shift_factor: float = 0.0, - ): - super().__init__() - - if len(encoder_block_out_channels) != len(num_encoder_blocks): - raise ValueError("`encoder_block_out_channels` should have the same length as `num_encoder_blocks`.") - if len(decoder_block_out_channels) != len(num_decoder_blocks): - raise ValueError("`decoder_block_out_channels` should have the same length as `num_decoder_blocks`.") - - self.encoder = EncoderTiny( - in_channels=in_channels, - out_channels=latent_channels, - num_blocks=num_encoder_blocks, - block_out_channels=encoder_block_out_channels, - act_fn=act_fn, - ) - - self.decoder = DecoderTiny( - in_channels=latent_channels, - out_channels=out_channels, - num_blocks=num_decoder_blocks, - block_out_channels=decoder_block_out_channels, - upsampling_scaling_factor=upsampling_scaling_factor, - act_fn=act_fn, - upsample_fn=upsample_fn, - ) - - self.latent_magnitude = latent_magnitude - self.latent_shift = latent_shift - self.scaling_factor = scaling_factor - - self.use_slicing = False - self.use_tiling = False - - # only relevant if vae tiling is enabled - self.spatial_scale_factor = 2**out_channels - self.tile_overlap_factor = 0.125 - self.tile_sample_min_size = 512 - self.tile_latent_min_size = self.tile_sample_min_size // self.spatial_scale_factor - - self.register_to_config(block_out_channels=decoder_block_out_channels) - self.register_to_config(force_upcast=False) - - def scale_latents(self, x: torch.Tensor) -> torch.Tensor: - """raw latents -> [0, 1]""" - return x.div(2 * self.latent_magnitude).add(self.latent_shift).clamp(0, 1) - - def unscale_latents(self, x: torch.Tensor) -> torch.Tensor: - """[0, 1] -> raw latents""" - return x.sub(self.latent_shift).mul(2 * self.latent_magnitude) - - def _tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - r"""Encode a batch of images using a tiled encoder. - - When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several - steps. This is useful to keep memory use constant regardless of image size. To avoid tiling artifacts, the - tiles overlap and are blended together to form a smooth output. - - Args: - x (`torch.Tensor`): Input batch of images. - - Returns: - `torch.Tensor`: Encoded batch of images. - """ - # scale of encoder output relative to input - sf = self.spatial_scale_factor - tile_size = self.tile_sample_min_size - - # number of pixels to blend and to traverse between tile - blend_size = int(tile_size * self.tile_overlap_factor) - traverse_size = tile_size - blend_size - - # tiles index (up/left) - ti = range(0, x.shape[-2], traverse_size) - tj = range(0, x.shape[-1], traverse_size) - - # mask for blending - blend_masks = torch.stack( - torch.meshgrid([torch.arange(tile_size / sf) / (blend_size / sf - 1)] * 2, indexing="ij") - ) - blend_masks = blend_masks.clamp(0, 1).to(x.device) - - # output array - out = torch.zeros(x.shape[0], 4, x.shape[-2] // sf, x.shape[-1] // sf, device=x.device) - for i in ti: - for j in tj: - tile_in = x[..., i : i + tile_size, j : j + tile_size] - # tile result - tile_out = out[..., i // sf : (i + tile_size) // sf, j // sf : (j + tile_size) // sf] - tile = self.encoder(tile_in) - h, w = tile.shape[-2], tile.shape[-1] - # blend tile result into output - blend_mask_i = torch.ones_like(blend_masks[0]) if i == 0 else blend_masks[0] - blend_mask_j = torch.ones_like(blend_masks[1]) if j == 0 else blend_masks[1] - blend_mask = blend_mask_i * blend_mask_j - tile, blend_mask = tile[..., :h, :w], blend_mask[..., :h, :w] - tile_out.copy_(blend_mask * tile + (1 - blend_mask) * tile_out) - return out - - def _tiled_decode(self, x: torch.Tensor) -> torch.Tensor: - r"""Encode a batch of images using a tiled encoder. - - When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several - steps. This is useful to keep memory use constant regardless of image size. To avoid tiling artifacts, the - tiles overlap and are blended together to form a smooth output. - - Args: - x (`torch.Tensor`): Input batch of images. - - Returns: - `torch.Tensor`: Encoded batch of images. - """ - # scale of decoder output relative to input - sf = self.spatial_scale_factor - tile_size = self.tile_latent_min_size - - # number of pixels to blend and to traverse between tiles - blend_size = int(tile_size * self.tile_overlap_factor) - traverse_size = tile_size - blend_size - - # tiles index (up/left) - ti = range(0, x.shape[-2], traverse_size) - tj = range(0, x.shape[-1], traverse_size) - - # mask for blending - blend_masks = torch.stack( - torch.meshgrid([torch.arange(tile_size * sf) / (blend_size * sf - 1)] * 2, indexing="ij") - ) - blend_masks = blend_masks.clamp(0, 1).to(x.device) - - # output array - out = torch.zeros(x.shape[0], 3, x.shape[-2] * sf, x.shape[-1] * sf, device=x.device) - for i in ti: - for j in tj: - tile_in = x[..., i : i + tile_size, j : j + tile_size] - # tile result - tile_out = out[..., i * sf : (i + tile_size) * sf, j * sf : (j + tile_size) * sf] - tile = self.decoder(tile_in) - h, w = tile.shape[-2], tile.shape[-1] - # blend tile result into output - blend_mask_i = torch.ones_like(blend_masks[0]) if i == 0 else blend_masks[0] - blend_mask_j = torch.ones_like(blend_masks[1]) if j == 0 else blend_masks[1] - blend_mask = (blend_mask_i * blend_mask_j)[..., :h, :w] - tile_out.copy_(blend_mask * tile + (1 - blend_mask) * tile_out) - return out - - @apply_forward_hook - def encode(self, x: torch.Tensor, return_dict: bool = True) -> AutoencoderTinyOutput | tuple[torch.Tensor]: - if self.use_slicing and x.shape[0] > 1: - output = [ - self._tiled_encode(x_slice) if self.use_tiling else self.encoder(x_slice) for x_slice in x.split(1) - ] - output = torch.cat(output) - else: - output = self._tiled_encode(x) if self.use_tiling else self.encoder(x) - - if not return_dict: - return (output,) - - return AutoencoderTinyOutput(latents=output) - - @apply_forward_hook - def decode( - self, x: torch.Tensor, generator: torch.Generator | None = None, return_dict: bool = True - ) -> DecoderOutput | tuple[torch.Tensor]: - if self.use_slicing and x.shape[0] > 1: - output = [ - self._tiled_decode(x_slice) if self.use_tiling else self.decoder(x_slice) for x_slice in x.split(1) - ] - output = torch.cat(output) - else: - output = self._tiled_decode(x) if self.use_tiling else self.decoder(x) - - if not return_dict: - return (output,) - - return DecoderOutput(sample=output) - - def forward( - self, - sample: torch.Tensor, - return_dict: bool = True, - ) -> DecoderOutput | tuple[torch.Tensor]: - r""" - Args: - sample (`torch.Tensor`): Input sample. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - - Returns: - [`DecoderOutput`] or `tuple`: - If `return_dict` is True, a [`DecoderOutput`] is returned, otherwise a plain `tuple` is returned. - """ - enc = self.encode(sample).latents - - # scale latents to be in [0, 1], then quantize latents to a byte tensor, - # as if we were storing the latents in an RGBA uint8 image. - scaled_enc = self.scale_latents(enc).mul_(255).round_().byte() - - # unquantize latents back into [0, 1], then unscale latents back to their original range, - # as if we were loading the latents from an RGBA uint8 image. - unscaled_enc = self.unscale_latents(scaled_enc / 255.0) - - dec = self.decode(unscaled_enc).sample - - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) diff --git a/diffusers/models/autoencoders/autoencoder_vidtok.py b/diffusers/models/autoencoders/autoencoder_vidtok.py deleted file mode 100644 index 296c7bd8d85a43c4c675c3e6e0e8e73ec75219f3..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/autoencoder_vidtok.py +++ /dev/null @@ -1,1506 +0,0 @@ -# Copyright 2025 The VidTok team, MSRA & Shanghai Jiao Tong University and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import List, Optional, Tuple, Union - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.accelerate_utils import apply_forward_hook -from ..modeling_outputs import AutoencoderKLOutput -from ..modeling_utils import ModelMixin -from .vae import DecoderOutput, DiagonalGaussianDistribution - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class FSQRegularizer(nn.Module): - r""" - Finite Scalar Quantization: VQ-VAE Made Simple - https://arxiv.org/abs/2309.15505 Code adapted from - https://github.com/lucidrains/vector-quantize-pytorch/blob/master/vector_quantize_pytorch/finite_scalar_quantization.py - - Args: - levels (`List[int]`): - A list of quantization levels. - dim (`int`, *optional*, defaults to `None`): - The dimension of latent codes. - num_codebooks (`int`, defaults to 1): - The number of codebooks. - keep_num_codebooks_dim (`bool`, *optional*, defaults to `None`): - Whether to keep the number of codebook dim. - """ - - def __init__( - self, - levels: List[int], - dim: Optional[int] = None, - num_codebooks: int = 1, - keep_num_codebooks_dim: Optional[bool] = None, - ): - super().__init__() - - _levels = torch.tensor(levels, dtype=torch.int32) - self.register_buffer("_levels", _levels, persistent=False) - - _basis = torch.cumprod(torch.tensor([1] + levels[:-1]), dim=0, dtype=torch.int32) - self.register_buffer("_basis", _basis, persistent=False) - - codebook_dim = len(levels) - self.codebook_dim = codebook_dim - - effective_codebook_dim = codebook_dim * num_codebooks - self.num_codebooks = num_codebooks - self.effective_codebook_dim = effective_codebook_dim - - if keep_num_codebooks_dim is None: - keep_num_codebooks_dim = num_codebooks > 1 - self.keep_num_codebooks_dim = keep_num_codebooks_dim - self.dim = len(_levels) * num_codebooks if dim is None else dim - - has_projections = self.dim != effective_codebook_dim - self.project_in = nn.Linear(self.dim, effective_codebook_dim) if has_projections else nn.Identity() - self.project_out = nn.Linear(effective_codebook_dim, self.dim) if has_projections else nn.Identity() - self.has_projections = has_projections - - self.codebook_size = self._levels.prod().item() - - implicit_codebook = self.indices_to_codes(torch.arange(self.codebook_size), project_out=False) - self.register_buffer("implicit_codebook", implicit_codebook, persistent=False) - self.register_buffer("zero", torch.tensor(0.0), persistent=False) - - self.global_codebook_usage = torch.zeros([2**self.codebook_dim, self.num_codebooks], dtype=torch.long) - - def quantize(self, z: torch.Tensor, eps: float = 1e-3) -> torch.Tensor: - r"""Quantizes z, returns quantized zhat, same shape as z.""" - half_l = (self._levels - 1) * (1 + eps) / 2 - offset = torch.where(self._levels % 2 == 0, 0.5, 0.0) - shift = (offset / half_l).atanh() - z = (z + shift).tanh() * half_l - offset - zhat = z.round() - quantized = z + (zhat - z).detach() - half_width = self._levels // 2 - return quantized / half_width - - def codes_to_indices(self, zhat: torch.Tensor) -> torch.Tensor: - r"""Converts a `code` to an index in the codebook.""" - half_width = self._levels // 2 - zhat = (zhat * half_width) + half_width - return (zhat * self._basis).sum(dim=-1).to(torch.int32) - - def indices_to_codes(self, indices: torch.Tensor, project_out: bool = True) -> torch.Tensor: - r"""Inverse of `codes_to_indices`.""" - is_img_or_video = indices.ndim >= (3 + int(self.keep_num_codebooks_dim)) - indices = indices.unsqueeze(-1) - codes_non_centered = (indices // self._basis) % self._levels - half_width = self._levels // 2 - codes = (codes_non_centered - half_width) / half_width - if self.keep_num_codebooks_dim: - codes = codes.reshape(*codes.shape[:-2], -1) - if project_out: - codes = self.project_out(codes) - if is_img_or_video: - codes = codes.permute(0, -1, *range(1, codes.dim() - 1)) - return codes - - def forward(self, z: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: - r""" - einstein notation b - batch n - sequence (or flattened spatial dimensions) d - feature dimension c - number of - codebook dim - """ - is_img_or_video = z.ndim >= 4 - - if is_img_or_video: - if z.ndim == 5: - b, d, t, h, w = z.shape - is_video = True - else: - b, d, h, w = z.shape - is_video = False - z = z.reshape(b, d, -1).permute(0, 2, 1) - - z = self.project_in(z) - b, n, _ = z.shape - z = z.reshape(b, n, self.num_codebooks, -1) - - orig_dtype = z.dtype - z = z.float() - codes = self.quantize(z) - indices = self.codes_to_indices(codes) - codes = codes.type(orig_dtype) - - codes = codes.reshape(b, n, -1) - out = self.project_out(codes) - - # reconstitute image or video dimensions - if is_img_or_video: - if is_video: - out = out.reshape(b, t, h, w, d).permute(0, 4, 1, 2, 3) - indices = indices.reshape(b, t, h, w, 1) - else: - out = out.reshape(b, h, w, d).permute(0, 3, 1, 2) - indices = indices.reshape(b, h, w, 1) - - if not self.keep_num_codebooks_dim: - indices = indices.squeeze(-1) - - return out, indices - - -class VidTokDownsample2D(nn.Module): - r"""A 2D downsampling layer used in VidTok Model.""" - - def __init__(self, in_channels: int): - super().__init__() - - self.in_channels = in_channels - self.conv = nn.Conv2d(in_channels, in_channels, kernel_size=3, stride=2, padding=0) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - pad = (0, 1, 0, 1) - x = F.pad(x, pad, mode="constant", value=0) - x = self.conv(x) - return x - - -class VidTokUpsample2D(nn.Module): - r"""A 2D upsampling layer used in VidTok Model.""" - - def __init__(self, in_channels: int): - super().__init__() - - self.in_channels = in_channels - self.conv = nn.Conv2d(in_channels, in_channels, kernel_size=3, stride=1, padding=1) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = F.interpolate(x.to(torch.float32), scale_factor=2.0, mode="nearest").to(x.dtype) - x = self.conv(x) - return x - - -class VidTokLayerNorm(nn.Module): - def __init__(self, dim: int, eps: float = 1e-6): - super().__init__() - - self.norm = nn.LayerNorm(dim, eps=eps, elementwise_affine=True) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - if x.dim() == 5: - x = x.permute(0, 2, 3, 4, 1) - x = self.norm(x) - x = x.permute(0, 4, 1, 2, 3) - elif x.dim() == 4: - x = x.permute(0, 2, 3, 1) - x = self.norm(x) - x = x.permute(0, 3, 1, 2) - else: - x = x.permute(0, 2, 1) - x = self.norm(x) - x = x.permute(0, 2, 1) - return x - - -class VidTokCausalConv1d(nn.Module): - r"""A 1D causal convolution layer that pads the input tensor to ensure causality in VidTok Model.""" - - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int, - stride: int = 1, - dilation: int = 1, - padding: int = 0, - ): - super().__init__() - - self.time_pad = dilation * (kernel_size - 1) + (1 - stride) - - self.conv = nn.Conv1d(in_channels, out_channels, kernel_size, stride=stride, dilation=dilation) - - self.is_first_chunk = True - self.causal_cache = None - self.cache_offset = 0 - - def forward(self, x: torch.Tensor) -> torch.Tensor: - if self.is_first_chunk: - first_frame_pad = x[:, :, :1].repeat((1, 1, self.time_pad)) - else: - first_frame_pad = self.causal_cache - if self.time_pad != 0: - first_frame_pad = first_frame_pad[:, :, -self.time_pad :] - else: - first_frame_pad = first_frame_pad[:, :, 0:0] - x = torch.concatenate((first_frame_pad, x), dim=2) - if self.cache_offset == 0: - self.causal_cache = x.clone() - else: - self.causal_cache = x[:, :, : -self.cache_offset].clone() - return self.conv(x) - - -class VidTokCausalConv3d(nn.Module): - r"""A 3D causal convolution layer that pads the input tensor to ensure causality in VidTok Model.""" - - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: Union[int, Tuple[int, int, int]], - stride: Union[int, Tuple[int, int, int]] = 1, - dilation: Union[int, Tuple[int, int, int]] = 1, - padding: Union[int, Tuple[int, int, int]] = 0, - pad_mode: str = "constant", - ): - super().__init__() - self.pad_mode = pad_mode - if isinstance(kernel_size, int): - kernel_size = (kernel_size,) * 3 - if isinstance(dilation, int): - dilation = (dilation,) * 3 - if isinstance(stride, int): - stride = (stride,) * 3 - time_kernel_size, height_kernel_size, width_kernel_size = kernel_size - time_pad = dilation[0] * (time_kernel_size - 1) + (1 - stride[0]) - height_pad = dilation[1] * (height_kernel_size - 1) + (1 - stride[1]) - width_pad = dilation[2] * (width_kernel_size - 1) + (1 - stride[2]) - - self.time_pad = time_pad - self.spatial_padding = ( - width_pad // 2, - width_pad - width_pad // 2, - height_pad // 2, - height_pad - height_pad // 2, - 0, - 0, - ) - self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride=stride, dilation=dilation) - - self.is_first_chunk = True - self.causal_cache = None - self.cache_offset = 0 - - def forward(self, x: torch.Tensor) -> torch.Tensor: - if self.is_first_chunk: - first_frame_pad = x[:, :, :1, :, :].repeat((1, 1, self.time_pad, 1, 1)) - else: - first_frame_pad = self.causal_cache - if self.time_pad != 0: - first_frame_pad = first_frame_pad[:, :, -self.time_pad :] - else: - first_frame_pad = first_frame_pad[:, :, 0:0] - x = torch.concatenate((first_frame_pad, x), dim=2) - if self.cache_offset == 0: - self.causal_cache = x.clone() - else: - self.causal_cache = x[:, :, : -self.cache_offset].clone() - x = F.pad(x, self.spatial_padding, mode=self.pad_mode) - return self.conv(x) - - -class VidTokDownsample3D(nn.Module): - r"""A 3D downsampling layer used in VidTok Model.""" - - def __init__(self, in_channels: int, out_channels: int, mix_factor: float = 2.0, is_causal: bool = True): - super().__init__() - self.is_causal = is_causal - self.kernel_size = (3, 3, 3) - self.avg_pool = nn.AvgPool3d((3, 1, 1), stride=(2, 1, 1)) - make_conv_cls = VidTokCausalConv3d if self.is_causal else nn.Conv3d - self.conv = make_conv_cls(in_channels, out_channels, 3, stride=(2, 1, 1), padding=(0, 1, 1)) - self.mix_factor = nn.Parameter(torch.Tensor([mix_factor])) - if self.is_causal: - self.is_first_chunk = True - self.causal_cache = None - - def forward(self, x: torch.Tensor) -> torch.Tensor: - alpha = torch.sigmoid(self.mix_factor) - if self.is_causal: - pad = (0, 0, 0, 0, 1, 0) - if self.is_first_chunk: - x_pad = torch.nn.functional.pad(x, pad, mode="replicate") - else: - x_pad = torch.concatenate((self.causal_cache, x), dim=2) - self.causal_cache = x_pad[:, :, -1:].clone() - if x_pad.device.type == "cpu" and x_pad.dtype == torch.bfloat16: - # PyTorch's avg_pool3d lacks CPU support for BFloat16. - # To avoid errors, we cast to float32, perform the pooling, - # and then cast back to BFloat16 to maintain the expected dtype. - x1 = self.avg_pool(x_pad.float()).to(torch.bfloat16) - else: - x1 = self.avg_pool(x_pad) - else: - pad = (0, 0, 0, 0, 0, 1) - x = F.pad(x, pad, mode="constant", value=0) - if x.device.type == "cpu" and x.dtype == torch.bfloat16: - # PyTorch's avg_pool3d lacks CPU support for BFloat16. - # To avoid errors, we cast to float32, perform the pooling, - # and then cast back to BFloat16 to maintain the expected dtype. - x1 = self.avg_pool(x.float()).to(torch.bfloat16) - else: - x1 = self.avg_pool(x) - x2 = self.conv(x) - return alpha * x1 + (1 - alpha) * x2 - - -class VidTokUpsample3D(nn.Module): - r"""A 3D upsampling layer used in VidTok Model.""" - - def __init__( - self, - in_channels: int, - out_channels: int, - mix_factor: float = 2.0, - num_temp_upsample: int = 1, - is_causal: bool = True, - ): - super().__init__() - make_conv_cls = VidTokCausalConv3d if is_causal else nn.Conv3d - self.conv = make_conv_cls(in_channels, out_channels, 3, padding=1) - self.mix_factor = nn.Parameter(torch.Tensor([mix_factor])) - - self.is_causal = is_causal - if self.is_causal: - self.enable_cached = True - self.interpolation_mode = "trilinear" - self.is_first_chunk = True - self.causal_cache = None - self.num_temp_upsample = num_temp_upsample - else: - self.enable_cached = False - self.interpolation_mode = "nearest" - - def forward(self, x: torch.Tensor) -> torch.Tensor: - alpha = torch.sigmoid(self.mix_factor) - if not self.is_causal: - xlst = [ - F.interpolate( - sx.unsqueeze(0).to(torch.float32), scale_factor=[2.0, 1.0, 1.0], mode=self.interpolation_mode - ).to(x.dtype) - for sx in x - ] - x = torch.cat(xlst, dim=0) - else: - if not self.enable_cached: - x = F.interpolate(x.to(torch.float32), scale_factor=[2.0, 1.0, 1.0], mode=self.interpolation_mode).to( - x.dtype - ) - elif not self.is_first_chunk: - x = torch.cat([self.causal_cache, x], dim=2) - self.causal_cache = x[:, :, -2 * self.num_temp_upsample : -self.num_temp_upsample].clone() - x = F.interpolate(x.to(torch.float32), scale_factor=[2.0, 1.0, 1.0], mode=self.interpolation_mode).to( - x.dtype - ) - x = x[:, :, 2 * self.num_temp_upsample :] - else: - self.causal_cache = x[:, :, -self.num_temp_upsample :].clone() - x, _x = x[:, :, : self.num_temp_upsample], x[:, :, self.num_temp_upsample :] - x = F.interpolate(x.to(torch.float32), scale_factor=[2.0, 1.0, 1.0], mode=self.interpolation_mode).to( - x.dtype - ) - if _x.shape[-3] > 0: - _x = F.interpolate( - _x.to(torch.float32), scale_factor=[2.0, 1.0, 1.0], mode=self.interpolation_mode - ).to(_x.dtype) - x = torch.concat([x, _x], dim=2) - x_ = self.conv(x) - return alpha * x + (1 - alpha) * x_ - - -class VidTokAttnBlock(nn.Module): - r"""A 3D self-attention block used in VidTok Model.""" - - def __init__(self, in_channels: int, is_causal: bool = True): - super().__init__() - make_conv_cls = VidTokCausalConv3d if is_causal else nn.Conv3d - self.norm = VidTokLayerNorm(dim=in_channels, eps=1e-6) - self.q = make_conv_cls(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.k = make_conv_cls(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.v = make_conv_cls(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.proj_out = make_conv_cls(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - - def attention(self, hidden_states: torch.Tensor) -> torch.Tensor: - r"""Implement self-attention.""" - hidden_states = self.norm(hidden_states) - q = self.q(hidden_states) - k = self.k(hidden_states) - v = self.v(hidden_states) - b, c, t, h, w = q.shape - q, k, v = [x.permute(0, 2, 3, 4, 1).reshape(b, t, -1, c).contiguous() for x in [q, k, v]] - hidden_states = F.scaled_dot_product_attention(q, k, v) # scale is dim ** -0.5 per default - return hidden_states.reshape(b, t, h, w, c).permute(0, 4, 1, 2, 3) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - hidden_states = x - hidden_states = self.attention(hidden_states) - hidden_states = self.proj_out(hidden_states) - return x + hidden_states - - -class VidTokResnetBlock(nn.Module): - r"""A versatile ResNet block used in VidTok Model.""" - - def __init__( - self, - in_channels: int, - out_channels: Optional[int] = None, - conv_shortcut: bool = False, - dropout: float = 0.0, - temb_channels: int = 512, - btype: str = "3d", - is_causal: bool = True, - ): - super().__init__() - assert btype in ["1d", "2d", "3d"], f"Invalid btype: {btype}" - if btype == "2d": - make_conv_cls = nn.Conv2d - elif btype == "1d": - make_conv_cls = VidTokCausalConv1d if is_causal else nn.Conv1d - else: - make_conv_cls = VidTokCausalConv3d if is_causal else nn.Conv3d - - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - self.use_conv_shortcut = conv_shortcut - self.nonlinearity = nn.SiLU() - - self.norm1 = VidTokLayerNorm(dim=in_channels, eps=1e-6) - self.conv1 = make_conv_cls(in_channels, out_channels, kernel_size=3, stride=1, padding=1) - if temb_channels > 0: - self.temb_proj = nn.Linear(temb_channels, out_channels) - self.norm2 = VidTokLayerNorm(dim=out_channels, eps=1e-6) - self.dropout = nn.Dropout(dropout) - self.conv2 = make_conv_cls(out_channels, out_channels, kernel_size=3, stride=1, padding=1) - if self.in_channels != self.out_channels: - if self.use_conv_shortcut: - self.conv_shortcut = make_conv_cls(in_channels, out_channels, kernel_size=3, stride=1, padding=1) - else: - self.nin_shortcut = make_conv_cls(in_channels, out_channels, kernel_size=1, stride=1, padding=0) - - def forward(self, x: torch.Tensor, temb: Optional[torch.Tensor]) -> torch.Tensor: - hidden_states = x - hidden_states = self.norm1(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv1(hidden_states) - - if temb is not None: - hidden_states = hidden_states + self.temb_proj(self.nonlinearity(temb))[:, :, None, None] - - hidden_states = self.norm2(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.dropout(hidden_states) - hidden_states = self.conv2(hidden_states) - - if self.in_channels != self.out_channels: - if self.use_conv_shortcut: - x = self.conv_shortcut(x) - else: - x = self.nin_shortcut(x) - return x + hidden_states - - -class VidTokEncoder3D(nn.Module): - r""" - The `VidTokEncoder3D` layer of a variational autoencoder that encodes its input into a latent representation. - - Args: - in_channels (`int`): - The number of input channels. - ch (`int`): - The number of the basic channel. - ch_mult (`List[int]`, defaults to `[1, 2, 4, 8]`): - The multiple of the basic channel for each block. - num_res_blocks (`int`, defaults to 2): - The number of resblocks. - dropout (`float`, defaults to 0.0): - Dropout rate. - z_channels (`int`, defaults to 4): - The number of latent channels. - double_z (`bool`, defaults to `True`): - Whether or not to double the z_channels. - spatial_ds (`List`, *optional*, defaults to `None`): - Spatial downsample layers. - tempo_ds (`List`, *optional*, defaults to `None`): - Temporal downsample layers. - is_causal (`bool`, defaults to `True`): - Whether it is a causal module. - """ - - def __init__( - self, - in_channels: int, - ch: int, - ch_mult: List[int] = [1, 2, 4, 8], - num_res_blocks: int = 2, - dropout: float = 0.0, - z_channels: int = 4, - double_z: bool = True, - spatial_ds: Optional[List] = None, - tempo_ds: Optional[List] = None, - is_causal: bool = True, - ): - super().__init__() - self.is_causal = is_causal - - self.ch = ch - self.temb_ch = 0 - self.num_resolutions = len(ch_mult) - self.num_res_blocks = num_res_blocks - self.in_channels = in_channels - self.nonlinearity = nn.SiLU() - - make_conv_cls = VidTokCausalConv3d if self.is_causal else nn.Conv3d - - self.conv_in = make_conv_cls(in_channels, self.ch, kernel_size=3, stride=1, padding=1) - - in_ch_mult = (1,) + tuple(ch_mult) - self.in_ch_mult = in_ch_mult - self.spatial_ds = list(range(0, self.num_resolutions - 1)) if spatial_ds is None else spatial_ds - self.tempo_ds = [self.num_resolutions - 2, self.num_resolutions - 3] if tempo_ds is None else tempo_ds - self.down = nn.ModuleList() - self.down_temporal = nn.ModuleList() - for i_level in range(self.num_resolutions): - block_in = ch * in_ch_mult[i_level] - block_out = ch * ch_mult[i_level] - - block = nn.ModuleList() - attn = nn.ModuleList() - block_temporal = nn.ModuleList() - attn_temporal = nn.ModuleList() - - for i_block in range(self.num_res_blocks): - block.append( - VidTokResnetBlock( - in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - btype="2d", - ) - ) - block_temporal.append( - VidTokResnetBlock( - in_channels=block_out, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - btype="1d", - is_causal=self.is_causal, - ) - ) - block_in = block_out - - down = nn.Module() - down.block = block - down.attn = attn - - down_temporal = nn.Module() - down_temporal.block = block_temporal - down_temporal.attn = attn_temporal - - if i_level in self.spatial_ds: - down.downsample = VidTokDownsample2D(block_in) - if i_level in self.tempo_ds: - down_temporal.downsample = VidTokDownsample3D(block_in, block_in, is_causal=self.is_causal) - - self.down.append(down) - self.down_temporal.append(down_temporal) - - # middle - self.mid = nn.Module() - self.mid.block_1 = VidTokResnetBlock( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - btype="3d", - is_causal=self.is_causal, - ) - self.mid.attn_1 = VidTokAttnBlock(block_in, is_causal=self.is_causal) - self.mid.block_2 = VidTokResnetBlock( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - btype="3d", - is_causal=self.is_causal, - ) - - # end - self.norm_out = VidTokLayerNorm(dim=block_in, eps=1e-6) - self.conv_out = make_conv_cls( - block_in, - 2 * z_channels if double_z else z_channels, - kernel_size=3, - stride=1, - padding=1, - ) - - self.gradient_checkpointing = False - - def forward(self, x: torch.Tensor) -> torch.Tensor: - temb = None - B, _, T, H, W = x.shape - hs = [self.conv_in(x)] - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for i_level in range(self.num_resolutions): - for i_block in range(self.num_res_blocks): - hidden_states = hs[-1].permute(0, 2, 1, 3, 4).reshape(B * T, -1, H, W) - hidden_states = self._gradient_checkpointing_func( - self.down[i_level].block[i_block], hidden_states, temb - ) - hidden_states = ( - hidden_states.reshape(B, T, -1, H, W).permute(0, 3, 4, 2, 1).reshape(B * H * W, -1, T) - ) - hidden_states = self._gradient_checkpointing_func( - self.down_temporal[i_level].block[i_block], hidden_states, temb - ) - hidden_states = hidden_states.reshape(B, H, W, -1, T).permute(0, 3, 4, 1, 2) - hs.append(hidden_states) - - if i_level in self.spatial_ds: - # spatial downsample - hidden_states = hs[-1].permute(0, 2, 1, 3, 4).reshape(B * T, -1, H, W) - hidden_states = self._gradient_checkpointing_func(self.down[i_level].downsample, hidden_states) - hidden_states = hidden_states.reshape(B, T, -1, *hidden_states.shape[-2:]).permute(0, 2, 1, 3, 4) - if i_level in self.tempo_ds: - # temporal downsample - hidden_states = self._gradient_checkpointing_func( - self.down_temporal[i_level].downsample, hidden_states - ) - hs.append(hidden_states) - B, _, T, H, W = hidden_states.shape - # middle - hidden_states = hs[-1] - hidden_states = self._gradient_checkpointing_func(self.mid.block_1, hidden_states, temb) - hidden_states = self._gradient_checkpointing_func(self.mid.attn_1, hidden_states) - hidden_states = self._gradient_checkpointing_func(self.mid.block_2, hidden_states, temb) - - else: - for i_level in range(self.num_resolutions): - for i_block in range(self.num_res_blocks): - hidden_states = hs[-1].permute(0, 2, 1, 3, 4).reshape(B * T, -1, H, W) - hidden_states = self.down[i_level].block[i_block](hidden_states, temb) - hidden_states = ( - hidden_states.reshape(B, T, -1, H, W).permute(0, 3, 4, 2, 1).reshape(B * H * W, -1, T) - ) - hidden_states = self.down_temporal[i_level].block[i_block](hidden_states, temb) - hidden_states = hidden_states.reshape(B, H, W, -1, T).permute(0, 3, 4, 1, 2) - hs.append(hidden_states) - - if i_level in self.spatial_ds: - # spatial downsample - hidden_states = hs[-1].permute(0, 2, 1, 3, 4).reshape(B * T, -1, H, W) - hidden_states = self.down[i_level].downsample(hidden_states) - hidden_states = hidden_states.reshape(B, T, -1, *hidden_states.shape[-2:]).permute(0, 2, 1, 3, 4) - if i_level in self.tempo_ds: - # temporal downsample - hidden_states = self.down_temporal[i_level].downsample(hidden_states) - hs.append(hidden_states) - B, _, T, H, W = hidden_states.shape - # middle - hidden_states = hs[-1] - hidden_states = self.mid.block_1(hidden_states, temb) - hidden_states = self.mid.attn_1(hidden_states) - hidden_states = self.mid.block_2(hidden_states, temb) - - # end - hidden_states = self.norm_out(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv_out(hidden_states) - return hidden_states - - -class VidTokDecoder3D(nn.Module): - r""" - The `VidTokDecoder3D` layer of a variational autoencoder that decodes its latent representation into an output - video. - - Args: - ch (`int`): - The number of the basic channel. - ch_mult (`List[int]`, defaults to `[1, 2, 4, 8]`): - The multiple of the basic channel for each block. - num_res_blocks (`int`, defaults to 2): - The number of resblocks. - dropout (`float`, defaults to 0.0): - Dropout rate. - z_channels (`int`, defaults to 4): - The number of latent channels. - out_channels (`int`, defaults to 3): - The number of output channels. - spatial_us (`List`, *optional*, defaults to `None`): - Spatial upsample layers. - tempo_us (`List`, *optional*, defaults to `None`): - Temporal upsample layers. - is_causal (`bool`, defaults to `True`): - Whether it is a causal module. - """ - - def __init__( - self, - ch: int, - ch_mult: List[int] = [1, 2, 4, 8], - num_res_blocks: int = 2, - dropout: float = 0.0, - z_channels: int = 4, - out_channels: int = 3, - spatial_us: Optional[List] = None, - tempo_us: Optional[List] = None, - is_causal: bool = True, - ): - super().__init__() - - self.is_causal = is_causal - self.ch = ch - self.temb_ch = 0 - self.num_resolutions = len(ch_mult) - self.num_res_blocks = num_res_blocks - self.nonlinearity = nn.SiLU() - - block_in = ch * ch_mult[self.num_resolutions - 1] - - make_conv_cls = VidTokCausalConv3d if self.is_causal else nn.Conv3d - - self.conv_in = make_conv_cls(z_channels, block_in, kernel_size=3, stride=1, padding=1) - - # middle - self.mid = nn.Module() - self.mid.block_1 = VidTokResnetBlock( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - btype="3d", - is_causal=self.is_causal, - ) - self.mid.attn_1 = VidTokAttnBlock(block_in, is_causal=self.is_causal) - self.mid.block_2 = VidTokResnetBlock( - in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - btype="3d", - is_causal=self.is_causal, - ) - - # upsampling - self.spatial_us = list(range(1, self.num_resolutions)) if spatial_us is None else spatial_us - self.tempo_us = [1, 2] if tempo_us is None else tempo_us - self.up = nn.ModuleList() - for i_level in reversed(range(self.num_resolutions)): - block = nn.ModuleList() - attn = nn.ModuleList() - block_out = ch * ch_mult[i_level] - for i_block in range(self.num_res_blocks + 1): - block.append( - VidTokResnetBlock( - in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - btype="2d", - ) - ) - block_in = block_out - - up = nn.Module() - up.block = block - up.attn = attn - if i_level in self.spatial_us: - up.upsample = VidTokUpsample2D(block_in) - self.up.insert(0, up) - - num_temp_upsample = 1 - self.up_temporal = nn.ModuleList() - for i_level in reversed(range(self.num_resolutions)): - block = nn.ModuleList() - attn = nn.ModuleList() - block_in = ch * ch_mult[i_level] - block_out = ch * ch_mult[i_level] - for i_block in range(self.num_res_blocks + 1): - block.append( - VidTokResnetBlock( - in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - btype="1d", - is_causal=self.is_causal, - ) - ) - block_in = block_out - up_temporal = nn.Module() - up_temporal.block = block - up_temporal.attn = attn - if i_level in self.tempo_us: - up_temporal.upsample = VidTokUpsample3D( - block_in, block_in, num_temp_upsample=num_temp_upsample, is_causal=self.is_causal - ) - num_temp_upsample *= 2 - - self.up_temporal.insert(0, up_temporal) - - # end - self.norm_out = VidTokLayerNorm(dim=block_in, eps=1e-6) - self.conv_out = make_conv_cls(block_in, out_channels, kernel_size=3, stride=1, padding=1) - - self.gradient_checkpointing = False - - def forward(self, z: torch.Tensor) -> torch.Tensor: - temb = None - B, _, T, H, W = z.shape - hidden_states = self.conv_in(z) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - # middle - hidden_states = self._gradient_checkpointing_func(self.mid.block_1, hidden_states, temb) - hidden_states = self._gradient_checkpointing_func(self.mid.attn_1, hidden_states) - hidden_states = self._gradient_checkpointing_func(self.mid.block_2, hidden_states, temb) - - for i_level in reversed(range(self.num_resolutions)): - for i_block in range(self.num_res_blocks + 1): - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).reshape(B * T, -1, H, W) - hidden_states = self._gradient_checkpointing_func( - self.up[i_level].block[i_block], hidden_states, temb - ) - hidden_states = ( - hidden_states.reshape(B, T, -1, H, W).permute(0, 3, 4, 2, 1).reshape(B * H * W, -1, T) - ) - hidden_states = self._gradient_checkpointing_func( - self.up_temporal[i_level].block[i_block], hidden_states, temb - ) - hidden_states = hidden_states.reshape(B, H, W, -1, T).permute(0, 3, 4, 1, 2) - - if i_level in self.spatial_us: - # spatial upsample - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).reshape(B * T, -1, H, W) - hidden_states = self._gradient_checkpointing_func(self.up[i_level].upsample, hidden_states) - hidden_states = hidden_states.reshape(B, T, -1, *hidden_states.shape[-2:]).permute(0, 2, 1, 3, 4) - if i_level in self.tempo_us: - # temporal upsample - hidden_states = self._gradient_checkpointing_func( - self.up_temporal[i_level].upsample, hidden_states - ) - B, _, T, H, W = hidden_states.shape - - else: - # middle - hidden_states = self.mid.block_1(hidden_states, temb) - hidden_states = self.mid.attn_1(hidden_states) - hidden_states = self.mid.block_2(hidden_states, temb) - - for i_level in reversed(range(self.num_resolutions)): - for i_block in range(self.num_res_blocks + 1): - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).reshape(B * T, -1, H, W) - hidden_states = self.up[i_level].block[i_block](hidden_states, temb) - hidden_states = ( - hidden_states.reshape(B, T, -1, H, W).permute(0, 3, 4, 2, 1).reshape(B * H * W, -1, T) - ) - hidden_states = self.up_temporal[i_level].block[i_block](hidden_states, temb) - hidden_states = hidden_states.reshape(B, H, W, -1, T).permute(0, 3, 4, 1, 2) - - if i_level in self.spatial_us: - # spatial upsample - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).reshape(B * T, -1, H, W) - hidden_states = self.up[i_level].upsample(hidden_states) - hidden_states = hidden_states.reshape(B, T, -1, *hidden_states.shape[-2:]).permute(0, 2, 1, 3, 4) - if i_level in self.tempo_us: - # temporal upsample - hidden_states = self.up_temporal[i_level].upsample(hidden_states) - B, _, T, H, W = hidden_states.shape - - # end - hidden_states = self.norm_out(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - out = self.conv_out(hidden_states) - return out - - -class AutoencoderVidTok(ModelMixin, ConfigMixin): - r""" - A VAE model for encoding videos into latents and decoding latent representations into videos, supporting both - continuous and discrete latent representations. Used in [VidTok](https://github.com/microsoft/VidTok). - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Args: - in_channels (`int`, defaults to 3): - The number of input channels. - out_channels (`int`, defaults to 3): - The number of output channels. - ch (`int`, defaults to 128): - The number of the basic channel. - ch_mult (`List[int]`, defaults to `[1, 2, 4, 4]`): - The multiple of the basic channel for each block. - z_channels (`int`, defaults to 4): - The number of latent channels. - double_z (`bool`, defaults to `True`): - Whether or not to double the z_channels. - num_res_blocks (`int`, defaults to 2): - The number of resblocks. - spatial_ds (`List`, *optional*, defaults to `None`): - Spatial downsample layers. - spatial_us (`List`, *optional*, defaults to `None`): - Spatial upsample layers. - tempo_ds (`List`, *optional*, defaults to `None`): - Temporal downsample layers. - tempo_us (`List`, *optional*, defaults to `None`): - Temporal upsample layers. - dropout (`float`, defaults to 0.0): - Dropout rate. - regularizer (`str`, defaults to `"kl"`): - The regularizer type - "kl" for continuous cases and "fsq" for discrete cases. - codebook_size (`int`, defaults to 262144): - The codebook size used only in discrete cases. - is_causal (`bool`, defaults to `True`): - Whether it is a causal module. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - ch: int = 128, - ch_mult: List[int] = [1, 2, 4, 4], - z_channels: int = 4, - double_z: bool = True, - num_res_blocks: int = 2, - spatial_ds: Optional[List] = None, - spatial_us: Optional[List] = None, - tempo_ds: Optional[List] = None, - tempo_us: Optional[List] = None, - dropout: float = 0.0, - regularizer: str = "kl", - codebook_size: int = 262144, - is_causal: bool = True, - ): - super().__init__() - self.is_causal = is_causal - - self.encoder = VidTokEncoder3D( - in_channels=in_channels, - ch=ch, - ch_mult=ch_mult, - num_res_blocks=num_res_blocks, - dropout=dropout, - z_channels=z_channels, - double_z=double_z, - spatial_ds=spatial_ds, - tempo_ds=tempo_ds, - is_causal=self.is_causal, - ) - self.decoder = VidTokDecoder3D( - ch=ch, - ch_mult=ch_mult, - num_res_blocks=num_res_blocks, - dropout=dropout, - z_channels=z_channels, - out_channels=out_channels, - spatial_us=spatial_us, - tempo_us=tempo_us, - is_causal=self.is_causal, - ) - self.temporal_compression_ratio = 2 ** len(self.encoder.tempo_ds) - - self.regularizer = regularizer - if self.regularizer not in ["kl", "fsq"]: - raise ValueError(f"Invalid regularizer: {self.regularizer}. Only `kl` and `fsq` are supported.") - - if self.regularizer == "fsq": - if z_channels != int(math.log(codebook_size, 8)): - raise ValueError( - f"When using the `fsq` regularizer, `z_channels` must be {int(math.log(codebook_size, 8))}, the" - f" log base 8 of the `codebook_size` {codebook_size}, but got {z_channels}." - ) - if double_z: - raise ValueError("When using the `fsq` regularizer, `double_z` must be `False`.") - - self.regularization = FSQRegularizer(levels=[8] * z_channels) - - self.use_slicing = False - self.use_tiling = False - - # Decode more latent frames at once - self.num_sample_frames_batch_size = 16 - self.num_latent_frames_batch_size = self.num_sample_frames_batch_size // self.temporal_compression_ratio - - # We make the minimum height and width of sample for tiling half that of the generally supported - self.tile_sample_min_height = 256 - self.tile_sample_min_width = 256 - self.tile_latent_min_height = int(self.tile_sample_min_height / (2 ** len(self.encoder.spatial_ds))) - self.tile_latent_min_width = int(self.tile_sample_min_width / (2 ** len(self.encoder.spatial_ds))) - self.tile_overlap_factor_height = 0.0 # 1 / 8 - self.tile_overlap_factor_width = 0.0 # 1 / 8 - - @staticmethod - def _pad_at_dim( - t: torch.Tensor, pad: Tuple[int], dim: int = -1, pad_mode: str = "constant", value: float = 0.0 - ) -> torch.Tensor: - r"""Pad function. Supported pad_mode: `constant`, `replicate`, `reflect`.""" - dims_from_right = (-dim - 1) if dim < 0 else (t.ndim - dim - 1) - zeros = (0, 0) * dims_from_right - if pad_mode == "constant": - return F.pad(t, (*zeros, *pad), value=value) - return F.pad(t, (*zeros, *pad), mode=pad_mode) - - def enable_tiling( - self, - tile_sample_min_height: Optional[int] = None, - tile_sample_min_width: Optional[int] = None, - tile_overlap_factor_height: Optional[float] = None, - tile_overlap_factor_width: Optional[float] = None, - ) -> None: - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - - Args: - tile_sample_min_height (`int`, *optional*, defaults to `None`): - The minimum height required for a sample to be separated into tiles across the height dimension. - tile_sample_min_width (`int`, *optional*, defaults to `None`): - The minimum width required for a sample to be separated into tiles across the width dimension. - tile_overlap_factor_height (`float`, *optional*, defaults to `None`): - The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are - no tiling artifacts produced across the height dimension. Must be between 0 and 1. Setting a higher - value might cause more tiles to be processed leading to slow down of the decoding process. - tile_overlap_factor_width (`float`, *optional*, defaults to `None`): - The minimum amount of overlap between two consecutive horizontal tiles. This is to ensure that there - are no tiling artifacts produced across the width dimension. Must be between 0 and 1. Setting a higher - value might cause more tiles to be processed leading to slow down of the decoding process. - """ - self.use_tiling = True - self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height - self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width - self.tile_latent_min_height = int(self.tile_sample_min_height / (2 ** len(self.encoder.spatial_ds))) - self.tile_latent_min_width = int(self.tile_sample_min_width / (2 ** len(self.encoder.spatial_ds))) - self.tile_overlap_factor_height = tile_overlap_factor_height or self.tile_overlap_factor_height - self.tile_overlap_factor_width = tile_overlap_factor_width or self.tile_overlap_factor_width - - def disable_tiling(self) -> None: - r""" - Disable tiled VAE decoding. If `enable_tiling` was previously enabled, this method will go back to computing - decoding in one step. - """ - self.use_tiling = False - - def enable_slicing(self) -> None: - r""" - Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to - compute decoding in several steps. This is useful to save some memory and allow larger batch sizes. - """ - self.use_slicing = True - - def disable_slicing(self) -> None: - r""" - Disable sliced VAE decoding. If `enable_slicing` was previously enabled, this method will go back to computing - decoding in one step. - """ - self.use_slicing = False - - def _encode(self, x: torch.Tensor) -> torch.Tensor: - self._empty_causal_cached(self.encoder) - self._set_first_chunk(True) - - if self.use_tiling: - return self.tiled_encode(x) - return self.encoder(x) - - @apply_forward_hook - def encode(self, x: torch.Tensor) -> Union[AutoencoderKLOutput, Tuple[torch.Tensor, torch.Tensor]]: - r""" - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - - Returns: - `AutoencoderKLOutput` or `Tuple[torch.Tensor]`: - The latent representations of the encoded videos. If the regularizer is `kl`, an `AutoencoderKLOutput` - is returned, otherwise a tuple of `torch.Tensor` is returned. - """ - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] - z = torch.cat(encoded_slices) - else: - z = self._encode(x) - - if self.regularizer == "kl": - posterior = DiagonalGaussianDistribution(z) - return AutoencoderKLOutput(latent_dist=posterior) - else: - quant_z, indices = self.regularization(z) - return quant_z, indices - - def _decode(self, z: torch.Tensor, decode_from_indices: bool = False) -> torch.Tensor: - self._empty_causal_cached(self.decoder) - self._set_first_chunk(True) - if not self.is_causal and z.shape[-3] % self.num_latent_frames_batch_size != 0: - assert z.shape[-3] >= self.num_latent_frames_batch_size, ( - f"Too short latent frames. At least {self.num_latent_frames_batch_size} frames." - ) - z = z[..., : (z.shape[-3] // self.num_latent_frames_batch_size * self.num_latent_frames_batch_size), :, :] - if decode_from_indices: - z = self.tile_indices_to_latent(z) if self.use_tiling else self.indices_to_latent(z) - dec = self.tiled_decode(z) if self.use_tiling else self.decoder(z) - return dec - - @apply_forward_hook - def decode(self, z: torch.Tensor, decode_from_indices: bool = False) -> torch.Tensor: - r""" - Decode a batch of images from latents. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - decode_from_indices (`bool`): If decode from indices or decode from latent code. - Returns: - `torch.Tensor`: The decoded images. - """ - if self.use_slicing and z.shape[0] > 1: - decoded_slices = [self._decode(z_slice, decode_from_indices=decode_from_indices) for z_slice in z.split(1)] - decoded = torch.cat(decoded_slices) - else: - decoded = self._decode(z, decode_from_indices=decode_from_indices) - if self.is_causal: - decoded = decoded[:, :, self.temporal_compression_ratio - 1 :, :, :] - return decoded - - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[3], b.shape[3], blend_extent) - for y in range(blend_extent): - b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( - y / blend_extent - ) - return b - - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[4], b.shape[4], blend_extent) - for x in range(blend_extent): - b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( - x / blend_extent - ) - return b - - def build_chunk_start_end(self, t, decoder_mode=False): - if self.is_causal: - start_end = [[0, self.temporal_compression_ratio]] if not decoder_mode else [[0, 1]] - start = start_end[0][-1] - else: - start_end, start = [], 0 - end = start - while True: - if start >= t: - break - end = min( - t, end + (self.num_latent_frames_batch_size if decoder_mode else self.num_sample_frames_batch_size) - ) - start_end.append([start, end]) - start = end - if len(start_end) > (2 if self.is_causal else 1): - if start_end[-1][1] - start_end[-1][0] < ( - self.num_latent_frames_batch_size if decoder_mode else self.num_sample_frames_batch_size - ): - start_end[-2] = [start_end[-2][0], start_end[-1][1]] - start_end = start_end[:-1] - return start_end - - def _set_first_chunk(self, is_first_chunk=True): - for module in self.modules(): - if hasattr(module, "is_first_chunk"): - module.is_first_chunk = is_first_chunk - - def _empty_causal_cached(self, parent): - for name, module in parent.named_modules(): - if hasattr(module, "causal_cache"): - module.causal_cache = None - - def _set_cache_offset(self, modules, cache_offset=0): - for module in modules: - for submodule in module.modules(): - if hasattr(submodule, "cache_offset"): - submodule.cache_offset = cache_offset - - def tiled_encode(self, x: torch.Tensor) -> torch.Tensor: - r""" - Encode a batch of images using a tiled encoder. - - When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several - steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is - different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the - tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the - output, but they should be much less noticeable. - - Args: - x (`torch.Tensor`): Input batch of videos. - - Returns: - `torch.Tensor`: The latent representation of the encoded videos. - """ - num_frames, height, width = x.shape[-3:] - - overlap_height = int(self.tile_sample_min_height * (1 - self.tile_overlap_factor_height)) - overlap_width = int(self.tile_sample_min_width * (1 - self.tile_overlap_factor_width)) - blend_extent_height = int(self.tile_latent_min_height * self.tile_overlap_factor_height) - blend_extent_width = int(self.tile_latent_min_width * self.tile_overlap_factor_width) - row_limit_height = self.tile_latent_min_height - blend_extent_height - row_limit_width = self.tile_latent_min_width - blend_extent_width - - # Split x into overlapping tiles and encode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, overlap_height): - row = [] - for j in range(0, width, overlap_width): - start_end = self.build_chunk_start_end(num_frames) - time = [] - for idx, (start_frame, end_frame) in enumerate(start_end): - self._set_first_chunk(idx == 0) - tile = x[ - :, - :, - start_frame:end_frame, - i : i + self.tile_sample_min_height, - j : j + self.tile_sample_min_width, - ] - tile = self.encoder(tile) - time.append(tile) - row.append(torch.cat(time, dim=2)) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent_width) - result_row.append(tile[:, :, :, :row_limit_height, :row_limit_width]) - result_rows.append(torch.cat(result_row, dim=4)) - enc = torch.cat(result_rows, dim=3) - return enc - - def indices_to_latent(self, token_indices: torch.Tensor) -> torch.Tensor: - r""" - Transform indices to latent code. - - Args: - token_indices (`torch.Tensor`): Token indices. - - Returns: - `torch.Tensor`: Latent code corresponding to the input token indices. - """ - b, t, h, w = token_indices.shape - token_indices = token_indices.unsqueeze(-1).reshape(b, -1, 1) - codes = self.regularization.indices_to_codes(token_indices) - codes = codes.permute(0, 2, 3, 1).reshape(b, codes.shape[2], -1) - z = self.regularization.project_out(codes) - return z.reshape(b, t, h, w, -1).permute(0, 4, 1, 2, 3) - - def tile_indices_to_latent(self, token_indices: torch.Tensor) -> torch.Tensor: - r""" - Transform indices to latent code with tiling inference. - - Args: - token_indices (`torch.Tensor`): Token indices. - - Returns: - `torch.Tensor`: Latent code corresponding to the input token indices. - """ - num_frames = token_indices.shape[1] - start_end = self.build_chunk_start_end(num_frames, decoder_mode=True) - result_z = [] - for start, end in start_end: - chunk_z = self.indices_to_latent(token_indices[:, start:end, :, :]) - result_z.append(chunk_z.clone()) - return torch.cat(result_z, dim=2) - - def tiled_decode(self, z: torch.Tensor) -> torch.Tensor: - r""" - Decode a batch of images using a tiled decoder. - - Args: - z (`torch.Tensor`): Input batch of latent vectors. - - Returns: - `torch.Tensor`: Reconstructed batch of videos. - """ - num_frames, height, width = z.shape[-3:] - - overlap_height = int(self.tile_latent_min_height * (1 - self.tile_overlap_factor_height)) - overlap_width = int(self.tile_latent_min_width * (1 - self.tile_overlap_factor_width)) - blend_extent_height = int(self.tile_sample_min_height * self.tile_overlap_factor_height) - blend_extent_width = int(self.tile_sample_min_width * self.tile_overlap_factor_width) - row_limit_height = self.tile_sample_min_height - blend_extent_height - row_limit_width = self.tile_sample_min_width - blend_extent_width - - # Split z into overlapping tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. - rows = [] - for i in range(0, height, overlap_height): - row = [] - for j in range(0, width, overlap_width): - if self.is_causal: - assert self.temporal_compression_ratio in [ - 2, - 4, - 8, - ], "Only support 2x, 4x or 8x temporal downsampling now." - if self.temporal_compression_ratio == 4: - self._set_cache_offset([self.decoder], 1) - self._set_cache_offset([self.decoder.up_temporal[2].upsample, self.decoder.up_temporal[1]], 2) - self._set_cache_offset( - [self.decoder.up_temporal[1].upsample, self.decoder.up_temporal[0], self.decoder.conv_out], - 4, - ) - elif self.temporal_compression_ratio == 2: - self._set_cache_offset([self.decoder], 1) - self._set_cache_offset( - [ - self.decoder.up_temporal[2].upsample, - self.decoder.up_temporal[1], - self.decoder.up_temporal[0], - self.decoder.conv_out, - ], - 2, - ) - else: - self._set_cache_offset([self.decoder], 1) - self._set_cache_offset([self.decoder.up_temporal[3].upsample, self.decoder.up_temporal[2]], 2) - self._set_cache_offset([self.decoder.up_temporal[2].upsample, self.decoder.up_temporal[1]], 4) - self._set_cache_offset( - [self.decoder.up_temporal[1].upsample, self.decoder.up_temporal[0], self.decoder.conv_out], - 8, - ) - - start_end = self.build_chunk_start_end(num_frames, decoder_mode=True) - time = [] - for idx, (start_frame, end_frame) in enumerate(start_end): - self._set_first_chunk(idx == 0) - tile = z[ - :, - :, - start_frame : (end_frame + 1 if self.is_causal and end_frame + 1 <= num_frames else end_frame), - i : i + self.tile_latent_min_height, - j : j + self.tile_latent_min_width, - ] - tile = self.decoder(tile) - if self.is_causal and end_frame + 1 <= num_frames: - tile = tile[:, :, : -self.temporal_compression_ratio] - time.append(tile) - row.append(torch.cat(time, dim=2)) - rows.append(row) - - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent_height) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent_width) - result_row.append(tile[:, :, :, :row_limit_height, :row_limit_width]) - result_rows.append(torch.cat(result_row, dim=4)) - - dec = torch.cat(result_rows, dim=3) - return dec - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = True, - encoder_mode: bool = False, - return_dict: bool = True, - generator: Optional[torch.Generator] = None, - ) -> Union[torch.Tensor, DecoderOutput]: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `True`): - Whether to sample from the posterior. - encoder_mode (`bool`, *optional*, defaults to `False`): - If `True`, only run the encoder and return the encoded latent without decoding. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make sampling - deterministic. - - Returns: - [`~models.vae.DecoderOutput`] or `torch.Tensor`: - If `return_dict` is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `torch.Tensor` - is returned. - """ - x = sample - res = 1 if self.is_causal else 0 - if self.is_causal: - if x.shape[2] % self.temporal_compression_ratio != res: - time_padding = self.temporal_compression_ratio - x.shape[2] % self.temporal_compression_ratio + res - x = self._pad_at_dim(x, (0, time_padding), dim=2, pad_mode="replicate") - else: - time_padding = 0 - else: - if x.shape[2] % self.num_sample_frames_batch_size != res: - if not encoder_mode: - time_padding = ( - self.num_sample_frames_batch_size - x.shape[2] % self.num_sample_frames_batch_size + res - ) - x = self._pad_at_dim(x, (0, time_padding), dim=2, pad_mode="replicate") - else: - assert x.shape[2] >= self.num_sample_frames_batch_size, ( - f"Too short video. At least {self.num_sample_frames_batch_size} frames." - ) - x = x[:, :, : x.shape[2] // self.num_sample_frames_batch_size * self.num_sample_frames_batch_size] - else: - time_padding = 0 - - if self.is_causal: - x = self._pad_at_dim(x, (self.temporal_compression_ratio - 1, 0), dim=2, pad_mode="replicate") - - if self.regularizer == "kl": - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - if encoder_mode: - return z - else: - z, indices = self.encode(x) - if encoder_mode: - return z, indices - - dec = self.decode(z) - if time_padding != 0: - dec = dec[:, :, :-time_padding, :, :] - - if not return_dict: - return (dec,) - return DecoderOutput(sample=dec) diff --git a/diffusers/models/autoencoders/consistency_decoder_vae.py b/diffusers/models/autoencoders/consistency_decoder_vae.py deleted file mode 100644 index dbe0f4c30541cda39e711975dd4d5af3aa525fe8..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/consistency_decoder_vae.py +++ /dev/null @@ -1,368 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from dataclasses import dataclass - -import torch -import torch.nn.functional as F -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...schedulers import ConsistencyDecoderScheduler -from ...utils import BaseOutput -from ...utils.accelerate_utils import apply_forward_hook -from ...utils.torch_utils import randn_tensor -from ..attention import AttentionMixin -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - AttnAddedKVProcessor, - AttnProcessor, -) -from ..modeling_utils import ModelMixin -from ..unets.unet_2d import UNet2DModel -from .vae import AutoencoderMixin, DecoderOutput, DiagonalGaussianDistribution, Encoder - - -@dataclass -class ConsistencyDecoderVAEOutput(BaseOutput): - """ - Output of encoding method. - - Args: - latent_dist (`DiagonalGaussianDistribution`): - Encoded outputs of `Encoder` represented as the mean and logvar of `DiagonalGaussianDistribution`. - `DiagonalGaussianDistribution` allows for sampling latents from the distribution. - """ - - latent_dist: "DiagonalGaussianDistribution" - - -class ConsistencyDecoderVAE(ModelMixin, AttentionMixin, AutoencoderMixin, ConfigMixin): - r""" - The consistency decoder used with DALL-E 3. - - Examples: - ```py - >>> import torch - >>> from diffusers import StableDiffusionPipeline, ConsistencyDecoderVAE - - >>> vae = ConsistencyDecoderVAE.from_pretrained("openai/consistency-decoder", torch_dtype=torch.float16) - >>> pipe = StableDiffusionPipeline.from_pretrained( - ... "stable-diffusion-v1-5/stable-diffusion-v1-5", vae=vae, torch_dtype=torch.float16 - ... ).to("cuda") - - >>> image = pipe("horse", generator=torch.manual_seed(0)).images[0] - >>> image - ``` - """ - - _supports_group_offloading = False - - @register_to_config - def __init__( - self, - scaling_factor: float = 0.18215, - latent_channels: int = 4, - sample_size: int = 32, - encoder_act_fn: str = "silu", - encoder_block_out_channels: tuple[int, ...] = (128, 256, 512, 512), - encoder_double_z: bool = True, - encoder_down_block_types: tuple[str, ...] = ( - "DownEncoderBlock2D", - "DownEncoderBlock2D", - "DownEncoderBlock2D", - "DownEncoderBlock2D", - ), - encoder_in_channels: int = 3, - encoder_layers_per_block: int = 2, - encoder_norm_num_groups: int = 32, - encoder_out_channels: int = 4, - decoder_add_attention: bool = False, - decoder_block_out_channels: tuple[int, ...] = (320, 640, 1024, 1024), - decoder_down_block_types: tuple[str, ...] = ( - "ResnetDownsampleBlock2D", - "ResnetDownsampleBlock2D", - "ResnetDownsampleBlock2D", - "ResnetDownsampleBlock2D", - ), - decoder_downsample_padding: int = 1, - decoder_in_channels: int = 7, - decoder_layers_per_block: int = 3, - decoder_norm_eps: float = 1e-05, - decoder_norm_num_groups: int = 32, - decoder_num_train_timesteps: int = 1024, - decoder_out_channels: int = 6, - decoder_resnet_time_scale_shift: str = "scale_shift", - decoder_time_embedding_type: str = "learned", - decoder_up_block_types: tuple[str, ...] = ( - "ResnetUpsampleBlock2D", - "ResnetUpsampleBlock2D", - "ResnetUpsampleBlock2D", - "ResnetUpsampleBlock2D", - ), - ): - super().__init__() - self.encoder = Encoder( - act_fn=encoder_act_fn, - block_out_channels=encoder_block_out_channels, - double_z=encoder_double_z, - down_block_types=encoder_down_block_types, - in_channels=encoder_in_channels, - layers_per_block=encoder_layers_per_block, - norm_num_groups=encoder_norm_num_groups, - out_channels=encoder_out_channels, - ) - - self.decoder_unet = UNet2DModel( - add_attention=decoder_add_attention, - block_out_channels=decoder_block_out_channels, - down_block_types=decoder_down_block_types, - downsample_padding=decoder_downsample_padding, - in_channels=decoder_in_channels, - layers_per_block=decoder_layers_per_block, - norm_eps=decoder_norm_eps, - norm_num_groups=decoder_norm_num_groups, - num_train_timesteps=decoder_num_train_timesteps, - out_channels=decoder_out_channels, - resnet_time_scale_shift=decoder_resnet_time_scale_shift, - time_embedding_type=decoder_time_embedding_type, - up_block_types=decoder_up_block_types, - ) - self.decoder_scheduler = ConsistencyDecoderScheduler() - self.register_to_config(block_out_channels=encoder_block_out_channels) - self.register_to_config(force_upcast=False) - self.register_buffer( - "means", - torch.tensor([0.38862467, 0.02253063, 0.07381133, -0.0171294])[None, :, None, None], - persistent=False, - ) - self.register_buffer( - "stds", torch.tensor([0.9654121, 1.0440036, 0.76147926, 0.77022034])[None, :, None, None], persistent=False - ) - - self.quant_conv = nn.Conv2d(2 * latent_channels, 2 * latent_channels, 1) - - self.use_slicing = False - self.use_tiling = False - - # only relevant if vae tiling is enabled - self.tile_sample_min_size = self.config.sample_size - sample_size = ( - self.config.sample_size[0] - if isinstance(self.config.sample_size, (list, tuple)) - else self.config.sample_size - ) - self.tile_latent_min_size = int(sample_size / (2 ** (len(self.config.block_out_channels) - 1))) - self.tile_overlap_factor = 0.25 - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - @apply_forward_hook - def encode( - self, x: torch.Tensor, return_dict: bool = True - ) -> ConsistencyDecoderVAEOutput | tuple[DiagonalGaussianDistribution]: - """ - Encode a batch of images into latents. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.autoencoders.consistency_decoder_vae.ConsistencyDecoderVAEOutput`] - instead of a plain tuple. - - Returns: - The latent representations of the encoded images. If `return_dict` is True, a - [`~models.autoencoders.consistency_decoder_vae.ConsistencyDecoderVAEOutput`] is returned, otherwise a - plain `tuple` is returned. - """ - if self.use_tiling and (x.shape[-1] > self.tile_sample_min_size or x.shape[-2] > self.tile_sample_min_size): - return self.tiled_encode(x, return_dict=return_dict) - - if self.use_slicing and x.shape[0] > 1: - encoded_slices = [self.encoder(x_slice) for x_slice in x.split(1)] - h = torch.cat(encoded_slices) - else: - h = self.encoder(x) - - moments = self.quant_conv(h) - posterior = DiagonalGaussianDistribution(moments) - - if not return_dict: - return (posterior,) - - return ConsistencyDecoderVAEOutput(latent_dist=posterior) - - @apply_forward_hook - def decode( - self, - z: torch.Tensor, - generator: torch.Generator | None = None, - return_dict: bool = True, - num_inference_steps: int = 2, - ) -> DecoderOutput | tuple[torch.Tensor]: - """ - Decodes the input latent vector `z` using the consistency decoder VAE model. - - Args: - z (torch.Tensor): The input latent vector. - generator (torch.Generator | None): The random number generator. Default is None. - return_dict (bool): Whether to return the output as a dictionary. Default is True. - num_inference_steps (int): The number of inference steps. Default is 2. - - Returns: - DecoderOutput | tuple[torch.Tensor]: The decoded output. - - """ - z = (z * self.config.scaling_factor - self.means) / self.stds - - scale_factor = 2 ** (len(self.config.block_out_channels) - 1) - z = F.interpolate(z, mode="nearest", scale_factor=scale_factor) - - batch_size, _, height, width = z.shape - - self.decoder_scheduler.set_timesteps(num_inference_steps, device=self.device) - - x_t = self.decoder_scheduler.init_noise_sigma * randn_tensor( - (batch_size, 3, height, width), generator=generator, dtype=z.dtype, device=z.device - ) - - for t in self.decoder_scheduler.timesteps: - model_input = torch.concat([self.decoder_scheduler.scale_model_input(x_t, t), z], dim=1) - model_output = self.decoder_unet(model_input, t).sample[:, :3, :, :] - prev_sample = self.decoder_scheduler.step(model_output, t, x_t, generator).prev_sample - x_t = prev_sample - - x_0 = x_t - - if not return_dict: - return (x_0,) - - return DecoderOutput(sample=x_0) - - # Copied from diffusers.models.autoencoders.autoencoder_kl.AutoencoderKL.blend_v - def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[2], b.shape[2], blend_extent) - for y in range(blend_extent): - b[:, :, y, :] = a[:, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, y, :] * (y / blend_extent) - return b - - # Copied from diffusers.models.autoencoders.autoencoder_kl.AutoencoderKL.blend_h - def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: - blend_extent = min(a.shape[3], b.shape[3], blend_extent) - for x in range(blend_extent): - b[:, :, :, x] = a[:, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, x] * (x / blend_extent) - return b - - def tiled_encode(self, x: torch.Tensor, return_dict: bool = True) -> ConsistencyDecoderVAEOutput | tuple: - r"""Encode a batch of images using a tiled encoder. - - When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several - steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is - different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the - tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the - output, but they should be much less noticeable. - - Args: - x (`torch.Tensor`): Input batch of images. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.autoencoders.consistency_decoder_vae.ConsistencyDecoderVAEOutput`] - instead of a plain tuple. - - Returns: - [`~models.autoencoders.consistency_decoder_vae.ConsistencyDecoderVAEOutput`] or `tuple`: - If return_dict is True, a [`~models.autoencoders.consistency_decoder_vae.ConsistencyDecoderVAEOutput`] - is returned, otherwise a plain `tuple` is returned. - """ - overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor)) - blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor) - row_limit = self.tile_latent_min_size - blend_extent - - # Split the image into 512x512 tiles and encode them separately. - rows = [] - for i in range(0, x.shape[2], overlap_size): - row = [] - for j in range(0, x.shape[3], overlap_size): - tile = x[:, :, i : i + self.tile_sample_min_size, j : j + self.tile_sample_min_size] - tile = self.encoder(tile) - tile = self.quant_conv(tile) - row.append(tile) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=3)) - - moments = torch.cat(result_rows, dim=2) - posterior = DiagonalGaussianDistribution(moments) - - if not return_dict: - return (posterior,) - - return ConsistencyDecoderVAEOutput(latent_dist=posterior) - - def forward( - self, - sample: torch.Tensor, - sample_posterior: bool = False, - return_dict: bool = True, - generator: torch.Generator | None = None, - ) -> DecoderOutput | tuple[torch.Tensor]: - r""" - Args: - sample (`torch.Tensor`): Input sample. - sample_posterior (`bool`, *optional*, defaults to `False`): - Whether to sample from the posterior. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`DecoderOutput`] instead of a plain tuple. - generator (`torch.Generator`, *optional*, defaults to `None`): - Generator to use for sampling. - - Returns: - [`DecoderOutput`] or `tuple`: - If return_dict is True, a [`DecoderOutput`] is returned, otherwise a plain `tuple` is returned. - """ - x = sample - posterior = self.encode(x).latent_dist - if sample_posterior: - z = posterior.sample(generator=generator) - else: - z = posterior.mode() - dec = self.decode(z, generator=generator).sample - - if not return_dict: - return (dec,) - - return DecoderOutput(sample=dec) diff --git a/diffusers/models/autoencoders/vae.py b/diffusers/models/autoencoders/vae.py deleted file mode 100644 index a65bca418175f5f704f1e37c1a9d4af346bcf9f5..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/vae.py +++ /dev/null @@ -1,927 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from dataclasses import dataclass - -import numpy as np -import torch -import torch.nn as nn - -from ...utils import BaseOutput -from ...utils.torch_utils import randn_tensor -from ..activations import get_activation -from ..attention_processor import SpatialNorm -from ..unets.unet_2d_blocks import ( - AutoencoderTinyBlock, - UNetMidBlock2D, - get_down_block, - get_up_block, -) - - -@dataclass -class EncoderOutput(BaseOutput): - r""" - Output of encoding method. - - Args: - latent (`torch.Tensor` of shape `(batch_size, num_channels, latent_height, latent_width)`): - The encoded latent. - """ - - latent: torch.Tensor - - -@dataclass -class DecoderOutput(BaseOutput): - r""" - Output of decoding method. - - Args: - sample (`torch.Tensor` of shape `(batch_size, num_channels, height, width)`): - The decoded output sample from the last layer of the model. - """ - - sample: torch.Tensor - commit_loss: torch.FloatTensor | None = None - - -class Encoder(nn.Module): - r""" - The `Encoder` layer of a variational autoencoder that encodes its input into a latent representation. - - Args: - in_channels (`int`, *optional*, defaults to 3): - The number of input channels. - out_channels (`int`, *optional*, defaults to 3): - The number of output channels. - down_block_types (`tuple[str, ...]`, *optional*, defaults to `("DownEncoderBlock2D",)`): - The types of down blocks to use. See `~diffusers.models.unet_2d_blocks.get_down_block` for available - options. - block_out_channels (`tuple[int, ...]`, *optional*, defaults to `(64,)`): - The number of output channels for each block. - layers_per_block (`int`, *optional*, defaults to 2): - The number of layers per block. - norm_num_groups (`int`, *optional*, defaults to 32): - The number of groups for normalization. - act_fn (`str`, *optional*, defaults to `"silu"`): - The activation function to use. See `~diffusers.models.activations.get_activation` for available options. - double_z (`bool`, *optional*, defaults to `True`): - Whether to double the number of output channels for the last block. - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - down_block_types: tuple[str, ...] = ("DownEncoderBlock2D",), - block_out_channels: tuple[int, ...] = (64,), - layers_per_block: int = 2, - norm_num_groups: int = 32, - act_fn: str = "silu", - double_z: bool = True, - mid_block_add_attention=True, - ): - super().__init__() - self.layers_per_block = layers_per_block - - self.conv_in = nn.Conv2d( - in_channels, - block_out_channels[0], - kernel_size=3, - stride=1, - padding=1, - ) - - self.down_blocks = nn.ModuleList([]) - - # down - output_channel = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - down_block = get_down_block( - down_block_type, - num_layers=self.layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - add_downsample=not is_final_block, - resnet_eps=1e-6, - downsample_padding=0, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - attention_head_dim=output_channel, - temb_channels=None, - ) - self.down_blocks.append(down_block) - - # mid - self.mid_block = UNetMidBlock2D( - in_channels=block_out_channels[-1], - resnet_eps=1e-6, - resnet_act_fn=act_fn, - output_scale_factor=1, - resnet_time_scale_shift="default", - attention_head_dim=block_out_channels[-1], - resnet_groups=norm_num_groups, - temb_channels=None, - add_attention=mid_block_add_attention, - ) - - # out - self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6) - self.conv_act = nn.SiLU() - - conv_out_channels = 2 * out_channels if double_z else out_channels - self.conv_out = nn.Conv2d(block_out_channels[-1], conv_out_channels, 3, padding=1) - - self.gradient_checkpointing = False - - def forward(self, sample: torch.Tensor) -> torch.Tensor: - r"""The forward method of the `Encoder` class.""" - - sample = self.conv_in(sample) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - # down - for down_block in self.down_blocks: - sample = self._gradient_checkpointing_func(down_block, sample) - # middle - sample = self._gradient_checkpointing_func(self.mid_block, sample) - - else: - # down - for down_block in self.down_blocks: - sample = down_block(sample) - - # middle - sample = self.mid_block(sample) - - # post-process - sample = self.conv_norm_out(sample) - sample = self.conv_act(sample) - sample = self.conv_out(sample) - - return sample - - -class Decoder(nn.Module): - r""" - The `Decoder` layer of a variational autoencoder that decodes its latent representation into an output sample. - - Args: - in_channels (`int`, *optional*, defaults to 3): - The number of input channels. - out_channels (`int`, *optional*, defaults to 3): - The number of output channels. - up_block_types (`tuple[str, ...]`, *optional*, defaults to `("UpDecoderBlock2D",)`): - The types of up blocks to use. See `~diffusers.models.unet_2d_blocks.get_up_block` for available options. - block_out_channels (`tuple[int, ...]`, *optional*, defaults to `(64,)`): - The number of output channels for each block. - layers_per_block (`int`, *optional*, defaults to 2): - The number of layers per block. - norm_num_groups (`int`, *optional*, defaults to 32): - The number of groups for normalization. - act_fn (`str`, *optional*, defaults to `"silu"`): - The activation function to use. See `~diffusers.models.activations.get_activation` for available options. - norm_type (`str`, *optional*, defaults to `"group"`): - The normalization type to use. Can be either `"group"` or `"spatial"`. - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - up_block_types: tuple[str, ...] = ("UpDecoderBlock2D",), - block_out_channels: tuple[int, ...] = (64,), - layers_per_block: int = 2, - norm_num_groups: int = 32, - act_fn: str = "silu", - norm_type: str = "group", # group, spatial - mid_block_add_attention=True, - ): - super().__init__() - self.layers_per_block = layers_per_block - - self.conv_in = nn.Conv2d( - in_channels, - block_out_channels[-1], - kernel_size=3, - stride=1, - padding=1, - ) - - self.up_blocks = nn.ModuleList([]) - - temb_channels = in_channels if norm_type == "spatial" else None - - # mid - self.mid_block = UNetMidBlock2D( - in_channels=block_out_channels[-1], - resnet_eps=1e-6, - resnet_act_fn=act_fn, - output_scale_factor=1, - resnet_time_scale_shift="default" if norm_type == "group" else norm_type, - attention_head_dim=block_out_channels[-1], - resnet_groups=norm_num_groups, - temb_channels=temb_channels, - add_attention=mid_block_add_attention, - ) - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - output_channel = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - - is_final_block = i == len(block_out_channels) - 1 - - up_block = get_up_block( - up_block_type, - num_layers=self.layers_per_block + 1, - in_channels=prev_output_channel, - out_channels=output_channel, - prev_output_channel=prev_output_channel, - add_upsample=not is_final_block, - resnet_eps=1e-6, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - attention_head_dim=output_channel, - temb_channels=temb_channels, - resnet_time_scale_shift=norm_type, - ) - self.up_blocks.append(up_block) - prev_output_channel = output_channel - - # out - if norm_type == "spatial": - self.conv_norm_out = SpatialNorm(block_out_channels[0], temb_channels) - else: - self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6) - self.conv_act = nn.SiLU() - self.conv_out = nn.Conv2d(block_out_channels[0], out_channels, 3, padding=1) - - self.gradient_checkpointing = False - - def forward( - self, - sample: torch.Tensor, - latent_embeds: torch.Tensor | None = None, - ) -> torch.Tensor: - r"""The forward method of the `Decoder` class.""" - - sample = self.conv_in(sample) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - # middle - sample = self._gradient_checkpointing_func(self.mid_block, sample, latent_embeds) - - # up - for up_block in self.up_blocks: - sample = self._gradient_checkpointing_func(up_block, sample, latent_embeds) - else: - # middle - sample = self.mid_block(sample, latent_embeds) - - # up - for up_block in self.up_blocks: - sample = up_block(sample, latent_embeds) - - # post-process - if latent_embeds is None: - sample = self.conv_norm_out(sample) - else: - sample = self.conv_norm_out(sample, latent_embeds) - sample = self.conv_act(sample) - sample = self.conv_out(sample) - - return sample - - -class UpSample(nn.Module): - r""" - The `UpSample` layer of a variational autoencoder that upsamples its input. - - Args: - in_channels (`int`, *optional*, defaults to 3): - The number of input channels. - out_channels (`int`, *optional*, defaults to 3): - The number of output channels. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - ) -> None: - super().__init__() - self.in_channels = in_channels - self.out_channels = out_channels - self.deconv = nn.ConvTranspose2d(in_channels, out_channels, kernel_size=4, stride=2, padding=1) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - r"""The forward method of the `UpSample` class.""" - x = torch.relu(x) - x = self.deconv(x) - return x - - -class MaskConditionEncoder(nn.Module): - """ - used in AsymmetricAutoencoderKL - """ - - def __init__( - self, - in_ch: int, - out_ch: int = 192, - res_ch: int = 768, - stride: int = 16, - ) -> None: - super().__init__() - - channels = [] - while stride > 1: - stride = stride // 2 - in_ch_ = out_ch * 2 - if out_ch > res_ch: - out_ch = res_ch - if stride == 1: - in_ch_ = res_ch - channels.append((in_ch_, out_ch)) - out_ch *= 2 - - out_channels = [] - for _in_ch, _out_ch in channels: - out_channels.append(_out_ch) - out_channels.append(channels[-1][0]) - - layers = [] - in_ch_ = in_ch - for l in range(len(out_channels)): - out_ch_ = out_channels[l] - if l == 0 or l == 1: - layers.append(nn.Conv2d(in_ch_, out_ch_, kernel_size=3, stride=1, padding=1)) - else: - layers.append(nn.Conv2d(in_ch_, out_ch_, kernel_size=4, stride=2, padding=1)) - in_ch_ = out_ch_ - - self.layers = nn.Sequential(*layers) - - def forward(self, x: torch.Tensor, mask=None) -> torch.Tensor: - r"""The forward method of the `MaskConditionEncoder` class.""" - out = {} - for l in range(len(self.layers)): - layer = self.layers[l] - x = layer(x) - out[str(tuple(x.shape))] = x - x = torch.relu(x) - return out - - -class MaskConditionDecoder(nn.Module): - r"""The `MaskConditionDecoder` should be used in combination with [`AsymmetricAutoencoderKL`] to enhance the model's - decoder with a conditioner on the mask and masked image. - - Args: - in_channels (`int`, *optional*, defaults to 3): - The number of input channels. - out_channels (`int`, *optional*, defaults to 3): - The number of output channels. - up_block_types (`tuple[str, ...]`, *optional*, defaults to `("UpDecoderBlock2D",)`): - The types of up blocks to use. See `~diffusers.models.unet_2d_blocks.get_up_block` for available options. - block_out_channels (`tuple[int, ...]`, *optional*, defaults to `(64,)`): - The number of output channels for each block. - layers_per_block (`int`, *optional*, defaults to 2): - The number of layers per block. - norm_num_groups (`int`, *optional*, defaults to 32): - The number of groups for normalization. - act_fn (`str`, *optional*, defaults to `"silu"`): - The activation function to use. See `~diffusers.models.activations.get_activation` for available options. - norm_type (`str`, *optional*, defaults to `"group"`): - The normalization type to use. Can be either `"group"` or `"spatial"`. - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - up_block_types: tuple[str, ...] = ("UpDecoderBlock2D",), - block_out_channels: tuple[int, ...] = (64,), - layers_per_block: int = 2, - norm_num_groups: int = 32, - act_fn: str = "silu", - norm_type: str = "group", # group, spatial - ): - super().__init__() - self.layers_per_block = layers_per_block - - self.conv_in = nn.Conv2d( - in_channels, - block_out_channels[-1], - kernel_size=3, - stride=1, - padding=1, - ) - - self.up_blocks = nn.ModuleList([]) - - temb_channels = in_channels if norm_type == "spatial" else None - - # mid - self.mid_block = UNetMidBlock2D( - in_channels=block_out_channels[-1], - resnet_eps=1e-6, - resnet_act_fn=act_fn, - output_scale_factor=1, - resnet_time_scale_shift="default" if norm_type == "group" else norm_type, - attention_head_dim=block_out_channels[-1], - resnet_groups=norm_num_groups, - temb_channels=temb_channels, - ) - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - output_channel = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - - is_final_block = i == len(block_out_channels) - 1 - - up_block = get_up_block( - up_block_type, - num_layers=self.layers_per_block + 1, - in_channels=prev_output_channel, - out_channels=output_channel, - prev_output_channel=None, - add_upsample=not is_final_block, - resnet_eps=1e-6, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - attention_head_dim=output_channel, - temb_channels=temb_channels, - resnet_time_scale_shift=norm_type, - ) - self.up_blocks.append(up_block) - prev_output_channel = output_channel - - # condition encoder - self.condition_encoder = MaskConditionEncoder( - in_ch=out_channels, - out_ch=block_out_channels[0], - res_ch=block_out_channels[-1], - ) - - # out - if norm_type == "spatial": - self.conv_norm_out = SpatialNorm(block_out_channels[0], temb_channels) - else: - self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6) - self.conv_act = nn.SiLU() - self.conv_out = nn.Conv2d(block_out_channels[0], out_channels, 3, padding=1) - - self.gradient_checkpointing = False - - def forward( - self, - z: torch.Tensor, - image: torch.Tensor | None = None, - mask: torch.Tensor | None = None, - latent_embeds: torch.Tensor | None = None, - ) -> torch.Tensor: - r"""The forward method of the `MaskConditionDecoder` class.""" - sample = z - sample = self.conv_in(sample) - - upscale_dtype = next(iter(self.up_blocks.parameters())).dtype - if torch.is_grad_enabled() and self.gradient_checkpointing: - # middle - sample = self._gradient_checkpointing_func(self.mid_block, sample, latent_embeds) - sample = sample.to(upscale_dtype) - - # condition encoder - if image is not None and mask is not None: - masked_image = (1 - mask) * image - im_x = self._gradient_checkpointing_func( - self.condition_encoder, - masked_image, - mask, - ) - - # up - for up_block in self.up_blocks: - if image is not None and mask is not None: - sample_ = im_x[str(tuple(sample.shape))] - mask_ = nn.functional.interpolate(mask, size=sample.shape[-2:], mode="nearest") - sample = sample * mask_ + sample_ * (1 - mask_) - sample = self._gradient_checkpointing_func(up_block, sample, latent_embeds) - if image is not None and mask is not None: - sample = sample * mask + im_x[str(tuple(sample.shape))] * (1 - mask) - else: - # middle - sample = self.mid_block(sample, latent_embeds) - sample = sample.to(upscale_dtype) - - # condition encoder - if image is not None and mask is not None: - masked_image = (1 - mask) * image - im_x = self.condition_encoder(masked_image, mask) - - # up - for up_block in self.up_blocks: - if image is not None and mask is not None: - sample_ = im_x[str(tuple(sample.shape))] - mask_ = nn.functional.interpolate(mask, size=sample.shape[-2:], mode="nearest") - sample = sample * mask_ + sample_ * (1 - mask_) - sample = up_block(sample, latent_embeds) - if image is not None and mask is not None: - sample = sample * mask + im_x[str(tuple(sample.shape))] * (1 - mask) - - # post-process - if latent_embeds is None: - sample = self.conv_norm_out(sample) - else: - sample = self.conv_norm_out(sample, latent_embeds) - sample = self.conv_act(sample) - sample = self.conv_out(sample) - - return sample - - -class VectorQuantizer(nn.Module): - """ - Improved version over VectorQuantizer, can be used as a drop-in replacement. Mostly avoids costly matrix - multiplications and allows for post-hoc remapping of indices. - """ - - # NOTE: due to a bug the beta term was applied to the wrong term. for - # backwards compatibility we use the buggy version by default, but you can - # specify legacy=False to fix it. - def __init__( - self, - n_e: int, - vq_embed_dim: int, - beta: float, - remap=None, - unknown_index: str = "random", - sane_index_shape: bool = False, - legacy: bool = True, - ): - super().__init__() - self.n_e = n_e - self.vq_embed_dim = vq_embed_dim - self.beta = beta - self.legacy = legacy - - self.embedding = nn.Embedding(self.n_e, self.vq_embed_dim) - self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e) - - self.remap = remap - if self.remap is not None: - self.register_buffer("used", torch.tensor(np.load(self.remap))) - self.used: torch.Tensor - self.re_embed = self.used.shape[0] - self.unknown_index = unknown_index # "random" or "extra" or integer - if self.unknown_index == "extra": - self.unknown_index = self.re_embed - self.re_embed = self.re_embed + 1 - print( - f"Remapping {self.n_e} indices to {self.re_embed} indices. " - f"Using {self.unknown_index} for unknown indices." - ) - else: - self.re_embed = n_e - - self.sane_index_shape = sane_index_shape - - def remap_to_used(self, inds: torch.LongTensor) -> torch.LongTensor: - ishape = inds.shape - assert len(ishape) > 1 - inds = inds.reshape(ishape[0], -1) - used = self.used.to(inds) - match = (inds[:, :, None] == used[None, None, ...]).long() - new = match.argmax(-1) - unknown = match.sum(2) < 1 - if self.unknown_index == "random": - new[unknown] = torch.randint(0, self.re_embed, size=new[unknown].shape).to(device=new.device) - else: - new[unknown] = self.unknown_index - return new.reshape(ishape) - - def unmap_to_all(self, inds: torch.LongTensor) -> torch.LongTensor: - ishape = inds.shape - assert len(ishape) > 1 - inds = inds.reshape(ishape[0], -1) - used = self.used.to(inds) - if self.re_embed > self.used.shape[0]: # extra token - inds[inds >= self.used.shape[0]] = 0 # simply set to zero - back = torch.gather(used[None, :][inds.shape[0] * [0], :], 1, inds) - return back.reshape(ishape) - - def forward(self, z: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, tuple]: - # reshape z -> (batch, height, width, channel) and flatten - z = z.permute(0, 2, 3, 1).contiguous() - z_flattened = z.view(-1, self.vq_embed_dim) - - # distances from z to embeddings e_j (z - e)^2 = z^2 + e^2 - 2 e * z - min_encoding_indices = torch.argmin(torch.cdist(z_flattened, self.embedding.weight), dim=1) - - z_q = self.embedding(min_encoding_indices).view(z.shape) - perplexity = None - min_encodings = None - - # compute loss for embedding - if not self.legacy: - loss = self.beta * torch.mean((z_q.detach() - z) ** 2) + torch.mean((z_q - z.detach()) ** 2) - else: - loss = torch.mean((z_q.detach() - z) ** 2) + self.beta * torch.mean((z_q - z.detach()) ** 2) - - # preserve gradients - z_q: torch.Tensor = z + (z_q - z).detach() - - # reshape back to match original input shape - z_q = z_q.permute(0, 3, 1, 2).contiguous() - - if self.remap is not None: - min_encoding_indices = min_encoding_indices.reshape(z.shape[0], -1) # add batch axis - min_encoding_indices = self.remap_to_used(min_encoding_indices) - min_encoding_indices = min_encoding_indices.reshape(-1, 1) # flatten - - if self.sane_index_shape: - min_encoding_indices = min_encoding_indices.reshape(z_q.shape[0], z_q.shape[2], z_q.shape[3]) - - return z_q, loss, (perplexity, min_encodings, min_encoding_indices) - - def get_codebook_entry(self, indices: torch.LongTensor, shape: tuple[int, ...]) -> torch.Tensor: - # shape specifying (batch, height, width, channel) - if self.remap is not None: - indices = indices.reshape(shape[0], -1) # add batch axis - indices = self.unmap_to_all(indices) - indices = indices.reshape(-1) # flatten again - - # get quantized latent vectors - z_q: torch.Tensor = self.embedding(indices) - - if shape is not None: - z_q = z_q.view(shape) - # reshape back to match original input shape - z_q = z_q.permute(0, 3, 1, 2).contiguous() - - return z_q - - -class DiagonalGaussianDistribution(object): - def __init__(self, parameters: torch.Tensor, deterministic: bool = False): - self.parameters = parameters - self.mean, self.logvar = torch.chunk(parameters, 2, dim=1) - self.logvar = torch.clamp(self.logvar, -30.0, 20.0) - self.deterministic = deterministic - self.std = torch.exp(0.5 * self.logvar) - self.var = torch.exp(self.logvar) - if self.deterministic: - self.var = self.std = torch.zeros_like( - self.mean, device=self.parameters.device, dtype=self.parameters.dtype - ) - - def sample(self, generator: torch.Generator | None = None) -> torch.Tensor: - # make sure sample is on the same device as the parameters and has same dtype - sample = randn_tensor( - self.mean.shape, - generator=generator, - device=self.parameters.device, - dtype=self.parameters.dtype, - ) - x = self.mean + self.std * sample - return x - - def kl(self, other: "DiagonalGaussianDistribution" = None) -> torch.Tensor: - if self.deterministic: - return torch.Tensor([0.0]) - else: - if other is None: - return 0.5 * torch.sum( - torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar, - dim=[1, 2, 3], - ) - else: - return 0.5 * torch.sum( - torch.pow(self.mean - other.mean, 2) / other.var - + self.var / other.var - - 1.0 - - self.logvar - + other.logvar, - dim=[1, 2, 3], - ) - - def nll(self, sample: torch.Tensor, dims: tuple[int, ...] = [1, 2, 3]) -> torch.Tensor: - if self.deterministic: - return torch.Tensor([0.0]) - logtwopi = np.log(2.0 * np.pi) - return 0.5 * torch.sum( - logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var, - dim=dims, - ) - - def mode(self) -> torch.Tensor: - return self.mean - - -class IdentityDistribution(object): - def __init__(self, parameters: torch.Tensor): - self.parameters = parameters - - def sample(self, generator: torch.Generator | None = None) -> torch.Tensor: - return self.parameters - - def mode(self) -> torch.Tensor: - return self.parameters - - -class EncoderTiny(nn.Module): - r""" - The `EncoderTiny` layer is a simpler version of the `Encoder` layer. - - Args: - in_channels (`int`): - The number of input channels. - out_channels (`int`): - The number of output channels. - num_blocks (`tuple[int, ...]`): - Each value of the tuple represents a Conv2d layer followed by `value` number of `AutoencoderTinyBlock`'s to - use. - block_out_channels (`tuple[int, ...]`): - The number of output channels for each block. - act_fn (`str`): - The activation function to use. See `~diffusers.models.activations.get_activation` for available options. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - num_blocks: tuple[int, ...], - block_out_channels: tuple[int, ...], - act_fn: str, - ): - super().__init__() - - layers = [] - for i, num_block in enumerate(num_blocks): - num_channels = block_out_channels[i] - - if i == 0: - layers.append(nn.Conv2d(in_channels, num_channels, kernel_size=3, padding=1)) - else: - layers.append( - nn.Conv2d( - num_channels, - num_channels, - kernel_size=3, - padding=1, - stride=2, - bias=False, - ) - ) - - for _ in range(num_block): - layers.append(AutoencoderTinyBlock(num_channels, num_channels, act_fn)) - - layers.append(nn.Conv2d(block_out_channels[-1], out_channels, kernel_size=3, padding=1)) - - self.layers = nn.Sequential(*layers) - self.gradient_checkpointing = False - - def forward(self, x: torch.Tensor) -> torch.Tensor: - r"""The forward method of the `EncoderTiny` class.""" - if torch.is_grad_enabled() and self.gradient_checkpointing: - x = self._gradient_checkpointing_func(self.layers, x) - - else: - # scale image from [-1, 1] to [0, 1] to match TAESD convention - x = self.layers(x.add(1).div(2)) - - return x - - -class DecoderTiny(nn.Module): - r""" - The `DecoderTiny` layer is a simpler version of the `Decoder` layer. - - Args: - in_channels (`int`): - The number of input channels. - out_channels (`int`): - The number of output channels. - num_blocks (`tuple[int, ...]`): - Each value of the tuple represents a Conv2d layer followed by `value` number of `AutoencoderTinyBlock`'s to - use. - block_out_channels (`tuple[int, ...]`): - The number of output channels for each block. - upsampling_scaling_factor (`int`): - The scaling factor to use for upsampling. - act_fn (`str`): - The activation function to use. See `~diffusers.models.activations.get_activation` for available options. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - num_blocks: tuple[int, ...], - block_out_channels: tuple[int, ...], - upsampling_scaling_factor: int, - act_fn: str, - upsample_fn: str, - ): - super().__init__() - - layers = [ - nn.Conv2d(in_channels, block_out_channels[0], kernel_size=3, padding=1), - get_activation(act_fn), - ] - - for i, num_block in enumerate(num_blocks): - is_final_block = i == (len(num_blocks) - 1) - num_channels = block_out_channels[i] - - for _ in range(num_block): - layers.append(AutoencoderTinyBlock(num_channels, num_channels, act_fn)) - - if not is_final_block: - layers.append(nn.Upsample(scale_factor=upsampling_scaling_factor, mode=upsample_fn)) - - conv_out_channel = num_channels if not is_final_block else out_channels - layers.append( - nn.Conv2d( - num_channels, - conv_out_channel, - kernel_size=3, - padding=1, - bias=is_final_block, - ) - ) - - self.layers = nn.Sequential(*layers) - self.gradient_checkpointing = False - - def forward(self, x: torch.Tensor) -> torch.Tensor: - r"""The forward method of the `DecoderTiny` class.""" - # Clamp. - x = torch.tanh(x / 3) * 3 - - if torch.is_grad_enabled() and self.gradient_checkpointing: - x = self._gradient_checkpointing_func(self.layers, x) - else: - x = self.layers(x) - - # scale image from [0, 1] to [-1, 1] to match diffusers convention - return x.mul(2).sub(1) - - -class AutoencoderMixin: - def enable_tiling(self): - r""" - Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to - compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow - processing larger images. - """ - if not hasattr(self, "use_tiling"): - raise NotImplementedError(f"Tiling doesn't seem to be implemented for {self.__class__.__name__}.") - self.use_tiling = True - - def disable_tiling(self): - r""" - Disable tiled VAE decoding. If `enable_tiling` was previously enabled, this method will go back to computing - decoding in one step. - """ - self.use_tiling = False - - def enable_slicing(self): - r""" - Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to - compute decoding in several steps. This is useful to save some memory and allow larger batch sizes. - """ - if not hasattr(self, "use_slicing"): - raise NotImplementedError(f"Slicing doesn't seem to be implemented for {self.__class__.__name__}.") - self.use_slicing = True - - def disable_slicing(self): - r""" - Disable sliced VAE decoding. If `enable_slicing` was previously enabled, this method will go back to computing - decoding in one step. - """ - self.use_slicing = False diff --git a/diffusers/models/autoencoders/vq_model.py b/diffusers/models/autoencoders/vq_model.py deleted file mode 100644 index 619327dde417a54b603382afe74e585f50037b25..0000000000000000000000000000000000000000 --- a/diffusers/models/autoencoders/vq_model.py +++ /dev/null @@ -1,183 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from dataclasses import dataclass - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput -from ...utils.accelerate_utils import apply_forward_hook -from ..autoencoders.vae import Decoder, DecoderOutput, Encoder, VectorQuantizer -from ..modeling_utils import ModelMixin -from .vae import AutoencoderMixin - - -@dataclass -class VQEncoderOutput(BaseOutput): - """ - Output of VQModel encoding method. - - Args: - latents (`torch.Tensor` of shape `(batch_size, num_channels, height, width)`): - The encoded output sample from the last layer of the model. - """ - - latents: torch.Tensor - - -class VQModel(ModelMixin, AutoencoderMixin, ConfigMixin): - r""" - A VQ-VAE model for decoding latent representations. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - in_channels (int, *optional*, defaults to 3): Number of channels in the input image. - out_channels (int, *optional*, defaults to 3): Number of channels in the output. - down_block_types (`tuple[str]`, *optional*, defaults to `("DownEncoderBlock2D",)`): - tuple of downsample block types. - up_block_types (`tuple[str]`, *optional*, defaults to `("UpDecoderBlock2D",)`): - tuple of upsample block types. - block_out_channels (`tuple[int]`, *optional*, defaults to `(64,)`): - tuple of block output channels. - layers_per_block (`int`, *optional*, defaults to `1`): Number of layers per block. - act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use. - latent_channels (`int`, *optional*, defaults to `3`): Number of channels in the latent space. - sample_size (`int`, *optional*, defaults to `32`): Sample input size. - num_vq_embeddings (`int`, *optional*, defaults to `256`): Number of codebook vectors in the VQ-VAE. - norm_num_groups (`int`, *optional*, defaults to `32`): Number of groups for normalization layers. - vq_embed_dim (`int`, *optional*): Hidden dim of codebook vectors in the VQ-VAE. - scaling_factor (`float`, *optional*, defaults to `0.18215`): - The component-wise standard deviation of the trained latent space computed using the first batch of the - training set. This is used to scale the latent space to have unit variance when training the diffusion - model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the - diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 - / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image - Synthesis with Latent Diffusion Models](https://huggingface.co/papers/2112.10752) paper. - norm_type (`str`, *optional*, defaults to `"group"`): - Type of normalization layer to use. Can be one of `"group"` or `"spatial"`. - """ - - _skip_layerwise_casting_patterns = ["quantize"] - _supports_group_offloading = False - - @register_to_config - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - down_block_types: tuple[str, ...] = ("DownEncoderBlock2D",), - up_block_types: tuple[str, ...] = ("UpDecoderBlock2D",), - block_out_channels: tuple[int, ...] = (64,), - layers_per_block: int = 1, - act_fn: str = "silu", - latent_channels: int = 3, - sample_size: int = 32, - num_vq_embeddings: int = 256, - norm_num_groups: int = 32, - vq_embed_dim: int | None = None, - scaling_factor: float = 0.18215, - norm_type: str = "group", # group, spatial - mid_block_add_attention=True, - lookup_from_codebook=False, - force_upcast=False, - ): - super().__init__() - - # pass init params to Encoder - self.encoder = Encoder( - in_channels=in_channels, - out_channels=latent_channels, - down_block_types=down_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - double_z=False, - mid_block_add_attention=mid_block_add_attention, - ) - - vq_embed_dim = vq_embed_dim if vq_embed_dim is not None else latent_channels - - self.quant_conv = nn.Conv2d(latent_channels, vq_embed_dim, 1) - self.quantize = VectorQuantizer(num_vq_embeddings, vq_embed_dim, beta=0.25, remap=None, sane_index_shape=False) - self.post_quant_conv = nn.Conv2d(vq_embed_dim, latent_channels, 1) - - # pass init params to Decoder - self.decoder = Decoder( - in_channels=latent_channels, - out_channels=out_channels, - up_block_types=up_block_types, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - norm_type=norm_type, - mid_block_add_attention=mid_block_add_attention, - ) - - @apply_forward_hook - def encode(self, x: torch.Tensor, return_dict: bool = True) -> VQEncoderOutput: - h = self.encoder(x) - h = self.quant_conv(h) - - if not return_dict: - return (h,) - - return VQEncoderOutput(latents=h) - - @apply_forward_hook - def decode( - self, h: torch.Tensor, force_not_quantize: bool = False, return_dict: bool = True, shape=None - ) -> DecoderOutput | torch.Tensor: - # also go through quantization layer - if not force_not_quantize: - quant, commit_loss, _ = self.quantize(h) - elif self.config.lookup_from_codebook: - quant = self.quantize.get_codebook_entry(h, shape) - commit_loss = torch.zeros((h.shape[0])).to(h.device, dtype=h.dtype) - else: - quant = h - commit_loss = torch.zeros((h.shape[0])).to(h.device, dtype=h.dtype) - quant2 = self.post_quant_conv(quant) - dec = self.decoder(quant2, quant if self.config.norm_type == "spatial" else None) - - if not return_dict: - return dec, commit_loss - - return DecoderOutput(sample=dec, commit_loss=commit_loss) - - def forward(self, sample: torch.Tensor, return_dict: bool = True) -> DecoderOutput | tuple[torch.Tensor, ...]: - r""" - The [`VQModel`] forward method. - - Args: - sample (`torch.Tensor`): Input sample. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`models.autoencoders.vq_model.VQEncoderOutput`] instead of a plain tuple. - - Returns: - [`~models.autoencoders.vq_model.VQEncoderOutput`] or `tuple`: - If return_dict is True, a [`~models.autoencoders.vq_model.VQEncoderOutput`] is returned, otherwise a - plain `tuple` is returned. - """ - - h = self.encode(sample).latents - dec = self.decode(h) - - if not return_dict: - return dec.sample, dec.commit_loss - return dec diff --git a/diffusers/models/cache_utils.py b/diffusers/models/cache_utils.py deleted file mode 100644 index 5aa189987ba21b6a2413b613a109cfe8430c0079..0000000000000000000000000000000000000000 --- a/diffusers/models/cache_utils.py +++ /dev/null @@ -1,164 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from contextlib import contextmanager - -from ..utils.logging import get_logger - - -logger = get_logger(__name__) # pylint: disable=invalid-name - - -class CacheMixin: - r""" - A class for enable/disabling caching techniques on diffusion models. - - Supported caching techniques: - - [Pyramid Attention Broadcast](https://huggingface.co/papers/2408.12588) - - [FasterCache](https://huggingface.co/papers/2410.19355) - - [FirstBlockCache](https://github.com/chengzeyi/ParaAttention/blob/7a266123671b55e7e5a2fe9af3121f07a36afc78/README.md#first-block-cache-our-dynamic-caching) - """ - - _cache_config = None - - @property - def is_cache_enabled(self) -> bool: - return self._cache_config is not None - - def enable_cache(self, config) -> None: - r""" - Enable caching techniques on the model. - - Args: - config (`PyramidAttentionBroadcastConfig | FasterCacheConfig | FirstBlockCacheConfig | TextKVCacheConfig`): - The configuration for applying the caching technique. Currently supported caching techniques are: - - [`~hooks.PyramidAttentionBroadcastConfig`] - - [`~hooks.FasterCacheConfig`] - - [`~hooks.FirstBlockCacheConfig`] - - [`~hooks.TextKVCacheConfig`] - - Example: - - ```python - >>> import torch - >>> from diffusers import CogVideoXPipeline, PyramidAttentionBroadcastConfig - - >>> pipe = CogVideoXPipeline.from_pretrained("THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16) - >>> pipe.to("cuda") - - >>> config = PyramidAttentionBroadcastConfig( - ... spatial_attention_block_skip_range=2, - ... spatial_attention_timestep_skip_range=(100, 800), - ... current_timestep_callback=lambda: pipe.current_timestep, - ... ) - >>> pipe.transformer.enable_cache(config) - ``` - """ - - from ..hooks import ( - FasterCacheConfig, - FirstBlockCacheConfig, - MagCacheConfig, - PyramidAttentionBroadcastConfig, - TaylorSeerCacheConfig, - TextKVCacheConfig, - apply_faster_cache, - apply_first_block_cache, - apply_mag_cache, - apply_pyramid_attention_broadcast, - apply_taylorseer_cache, - apply_text_kv_cache, - ) - - if self.is_cache_enabled: - raise ValueError( - f"Caching has already been enabled with {type(self._cache_config)}. To apply a new caching technique, please disable the existing one first." - ) - - if isinstance(config, FasterCacheConfig): - apply_faster_cache(self, config) - elif isinstance(config, FirstBlockCacheConfig): - apply_first_block_cache(self, config) - elif isinstance(config, MagCacheConfig): - apply_mag_cache(self, config) - elif isinstance(config, TextKVCacheConfig): - apply_text_kv_cache(self, config) - elif isinstance(config, PyramidAttentionBroadcastConfig): - apply_pyramid_attention_broadcast(self, config) - elif isinstance(config, TaylorSeerCacheConfig): - apply_taylorseer_cache(self, config) - else: - raise ValueError(f"Cache config {type(config)} is not supported.") - - self._cache_config = config - - def disable_cache(self) -> None: - from ..hooks import ( - FasterCacheConfig, - FirstBlockCacheConfig, - HookRegistry, - MagCacheConfig, - PyramidAttentionBroadcastConfig, - TaylorSeerCacheConfig, - TextKVCacheConfig, - ) - from ..hooks.faster_cache import _FASTER_CACHE_BLOCK_HOOK, _FASTER_CACHE_DENOISER_HOOK - from ..hooks.first_block_cache import _FBC_BLOCK_HOOK, _FBC_LEADER_BLOCK_HOOK - from ..hooks.mag_cache import _MAG_CACHE_BLOCK_HOOK, _MAG_CACHE_LEADER_BLOCK_HOOK - from ..hooks.pyramid_attention_broadcast import _PYRAMID_ATTENTION_BROADCAST_HOOK - from ..hooks.taylorseer_cache import _TAYLORSEER_CACHE_HOOK - from ..hooks.text_kv_cache import _TEXT_KV_CACHE_BLOCK_HOOK, _TEXT_KV_CACHE_TRANSFORMER_HOOK - - if self._cache_config is None: - logger.warning("Caching techniques have not been enabled, so there's nothing to disable.") - return - - registry = HookRegistry.check_if_exists_or_initialize(self) - if isinstance(self._cache_config, FasterCacheConfig): - registry.remove_hook(_FASTER_CACHE_DENOISER_HOOK, recurse=True) - registry.remove_hook(_FASTER_CACHE_BLOCK_HOOK, recurse=True) - elif isinstance(self._cache_config, FirstBlockCacheConfig): - registry.remove_hook(_FBC_LEADER_BLOCK_HOOK, recurse=True) - registry.remove_hook(_FBC_BLOCK_HOOK, recurse=True) - elif isinstance(self._cache_config, MagCacheConfig): - registry.remove_hook(_MAG_CACHE_LEADER_BLOCK_HOOK, recurse=True) - registry.remove_hook(_MAG_CACHE_BLOCK_HOOK, recurse=True) - elif isinstance(self._cache_config, PyramidAttentionBroadcastConfig): - registry.remove_hook(_PYRAMID_ATTENTION_BROADCAST_HOOK, recurse=True) - elif isinstance(self._cache_config, TextKVCacheConfig): - registry.remove_hook(_TEXT_KV_CACHE_TRANSFORMER_HOOK, recurse=True) - registry.remove_hook(_TEXT_KV_CACHE_BLOCK_HOOK, recurse=True) - elif isinstance(self._cache_config, TaylorSeerCacheConfig): - registry.remove_hook(_TAYLORSEER_CACHE_HOOK, recurse=True) - else: - raise ValueError(f"Cache config {type(self._cache_config)} is not supported.") - - self._cache_config = None - - def _reset_stateful_cache(self, recurse: bool = True) -> None: - from ..hooks import HookRegistry - - HookRegistry.check_if_exists_or_initialize(self).reset_stateful_hooks(recurse=recurse) - - @contextmanager - def cache_context(self, name: str): - r"""Context manager that provides additional methods for cache management.""" - from ..hooks import HookRegistry - - registry = HookRegistry.check_if_exists_or_initialize(self) - registry._set_context(name) - - yield - - registry._set_context(None) diff --git a/diffusers/models/condition_embedders/__init__.py b/diffusers/models/condition_embedders/__init__.py deleted file mode 100644 index 3a92469a13ce7d05016e7945e82dbc1c4f12be6e..0000000000000000000000000000000000000000 --- a/diffusers/models/condition_embedders/__init__.py +++ /dev/null @@ -1,5 +0,0 @@ -from ...utils import is_torch_available - - -if is_torch_available(): - from .condition_embedder_anima import AnimaTextConditioner diff --git a/diffusers/models/condition_embedders/condition_embedder_anima.py b/diffusers/models/condition_embedders/condition_embedder_anima.py deleted file mode 100644 index 40fda447ec685bff6a656b576282d7dfc82881ef..0000000000000000000000000000000000000000 --- a/diffusers/models/condition_embedders/condition_embedder_anima.py +++ /dev/null @@ -1,346 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ..attention import AttentionModuleMixin -from ..attention_dispatch import dispatch_attention_fn -from ..modeling_utils import ModelMixin - - -def _rotate_half(hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states_1 = hidden_states[..., : hidden_states.shape[-1] // 2] - hidden_states_2 = hidden_states[..., hidden_states.shape[-1] // 2 :] - return torch.cat((-hidden_states_2, hidden_states_1), dim=-1) - - -def _apply_rotary_pos_emb( - hidden_states: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, unsqueeze_dim: int = 1 -) -> torch.Tensor: - cos = cos.unsqueeze(unsqueeze_dim) - sin = sin.unsqueeze(unsqueeze_dim) - return (hidden_states * cos) + (_rotate_half(hidden_states) * sin) - - -class AnimaRotaryEmbedding(nn.Module): - def __init__(self, head_dim: int, rope_theta: float = 10000.0): - super().__init__() - inv_freq = 1.0 / ( - rope_theta ** (torch.arange(0, head_dim, 2, dtype=torch.int64).to(dtype=torch.float32) / head_dim) - ) - self.register_buffer("inv_freq", inv_freq, persistent=False) - - def forward(self, hidden_states: torch.Tensor, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) - inv_freq_expanded = inv_freq_expanded.to(hidden_states.device) - position_ids_expanded = position_ids[:, None, :].float() - - freqs = (inv_freq_expanded @ position_ids_expanded).transpose(1, 2) - emb = torch.cat((freqs, freqs), dim=-1) - cos = emb.cos() - sin = emb.sin() - - return cos.to(dtype=hidden_states.dtype), sin.to(dtype=hidden_states.dtype) - - -class AnimaTextConditionerAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __call__( - self, - attn: "AnimaTextConditionerAttention", - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, - encoder_position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - encoder_hidden_states = hidden_states if encoder_hidden_states is None else encoder_hidden_states - input_shape = hidden_states.shape[:-1] - encoder_input_shape = encoder_hidden_states.shape[:-1] - - query = attn.q_proj(hidden_states) - key = attn.k_proj(encoder_hidden_states) - value = attn.v_proj(encoder_hidden_states) - - query = query.view(*input_shape, attn.num_attention_heads, attn.attention_head_dim) - key = key.view(*encoder_input_shape, attn.num_attention_heads, attn.attention_head_dim) - value = value.view(*encoder_input_shape, attn.num_attention_heads, attn.attention_head_dim) - - query = attn.q_norm(query) - key = attn.k_norm(key) - - if position_embeddings is not None: - if encoder_position_embeddings is None: - raise ValueError("`encoder_position_embeddings` must be provided when using rotary embeddings.") - cos, sin = position_embeddings - query = _apply_rotary_pos_emb(query, cos, sin, unsqueeze_dim=2) - cos, sin = encoder_position_embeddings - key = _apply_rotary_pos_emb(key, cos, sin, unsqueeze_dim=2) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3).contiguous() - hidden_states = attn.o_proj(hidden_states) - return hidden_states - - -class AnimaTextConditionerAttention(nn.Module, AttentionModuleMixin): - _default_processor_cls = AnimaTextConditionerAttnProcessor - _available_processors = [AnimaTextConditionerAttnProcessor] - _supports_qkv_fusion = False - - def __init__( - self, - query_dim: int, - context_dim: int, - num_attention_heads: int, - attention_head_dim: int, - processor: AnimaTextConditionerAttnProcessor | None = None, - ): - super().__init__() - inner_dim = num_attention_heads * attention_head_dim - - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - self.q_proj = nn.Linear(query_dim, inner_dim, bias=False) - self.q_norm = nn.RMSNorm(attention_head_dim, eps=1e-6) - self.k_proj = nn.Linear(context_dim, inner_dim, bias=False) - self.k_norm = nn.RMSNorm(attention_head_dim, eps=1e-6) - self.v_proj = nn.Linear(context_dim, inner_dim, bias=False) - self.o_proj = nn.Linear(inner_dim, query_dim, bias=False) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, - encoder_position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - return self.processor( - self, - hidden_states, - attention_mask=attention_mask, - encoder_hidden_states=encoder_hidden_states, - position_embeddings=position_embeddings, - encoder_position_embeddings=encoder_position_embeddings, - ) - - -class AnimaTextConditionerBlock(nn.Module): - def __init__( - self, - source_dim: int, - model_dim: int, - num_attention_heads: int = 16, - mlp_ratio: float = 4.0, - use_self_attention: bool = True, - use_layer_norm: bool = False, - ): - super().__init__() - self.use_self_attention = use_self_attention - norm_cls = nn.LayerNorm if use_layer_norm else nn.RMSNorm - norm_kwargs = {} if use_layer_norm else {"eps": 1e-6} - - if use_self_attention: - self.norm_self_attn = norm_cls(model_dim, **norm_kwargs) - self.self_attn = AnimaTextConditionerAttention( - query_dim=model_dim, - context_dim=model_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=model_dim // num_attention_heads, - ) - - self.norm_cross_attn = norm_cls(model_dim, **norm_kwargs) - self.cross_attn = AnimaTextConditionerAttention( - query_dim=model_dim, - context_dim=source_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=model_dim // num_attention_heads, - ) - self.norm_mlp = norm_cls(model_dim, **norm_kwargs) - self.mlp = nn.Sequential( - nn.Linear(model_dim, int(model_dim * mlp_ratio)), - nn.GELU(), - nn.Linear(int(model_dim * mlp_ratio), model_dim), - ) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - target_attention_mask: torch.Tensor | None = None, - source_attention_mask: torch.Tensor | None = None, - position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, - source_position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - if self.use_self_attention: - norm_hidden_states = self.norm_self_attn(hidden_states) - attn_hidden_states = self.self_attn( - norm_hidden_states, - attention_mask=target_attention_mask, - position_embeddings=position_embeddings, - encoder_position_embeddings=position_embeddings, - ) - hidden_states = hidden_states + attn_hidden_states - - norm_hidden_states = self.norm_cross_attn(hidden_states) - attn_hidden_states = self.cross_attn( - norm_hidden_states, - attention_mask=source_attention_mask, - encoder_hidden_states=encoder_hidden_states, - position_embeddings=position_embeddings, - encoder_position_embeddings=source_position_embeddings, - ) - hidden_states = hidden_states + attn_hidden_states - hidden_states = hidden_states + self.mlp(self.norm_mlp(hidden_states)) - return hidden_states - - -class AnimaTextConditioner(ModelMixin, ConfigMixin, PeftAdapterMixin): - r""" - Text conditioner used by Anima to map Qwen3 hidden states and T5 token ids to Cosmos text embeddings. - - Anima reuses the Cosmos Predict2 DiT. The only model-specific conditioning module is this LLM adapter, which - cross-attends from learned T5 token embeddings to Qwen3 text encoder hidden states before the diffusion loop. - `target_dim` is the conditioner output dimension and must match the transformer's `text_embed_dim`. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["AnimaTextConditionerBlock"] - - @register_to_config - def __init__( - self, - source_dim: int = 1024, - target_dim: int = 1024, - model_dim: int = 1024, - num_layers: int = 6, - num_attention_heads: int = 16, - mlp_ratio: float = 4.0, - target_vocab_size: int = 32128, - use_self_attention: bool = True, - use_layer_norm: bool = False, - min_sequence_length: int = 512, - ): - super().__init__() - self.embed = nn.Embedding(target_vocab_size, target_dim) - self.in_proj = nn.Linear(target_dim, model_dim) if model_dim != target_dim else nn.Identity() - self.rotary_emb = AnimaRotaryEmbedding(model_dim // num_attention_heads) - self.blocks = nn.ModuleList( - [ - AnimaTextConditionerBlock( - source_dim=source_dim, - model_dim=model_dim, - num_attention_heads=num_attention_heads, - mlp_ratio=mlp_ratio, - use_self_attention=use_self_attention, - use_layer_norm=use_layer_norm, - ) - for _ in range(num_layers) - ] - ) - self.out_proj = nn.Linear(model_dim, target_dim) - self.norm = nn.RMSNorm(target_dim, eps=1e-6) - self.gradient_checkpointing = False - - @staticmethod - def _prepare_attention_mask(attention_mask: torch.Tensor | None) -> torch.Tensor | None: - if attention_mask is None: - return None - attention_mask = attention_mask.to(torch.bool) - if attention_mask.ndim == 2: - attention_mask = attention_mask.unsqueeze(1).unsqueeze(1) - return attention_mask - - def forward( - self, - source_hidden_states: torch.Tensor, - target_input_ids: torch.Tensor, - target_attention_mask: torch.Tensor | None = None, - source_attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - """ - Args: - source_hidden_states (`torch.Tensor` of shape `(batch_size, source_sequence_length, source_dim)`): - Qwen3 text encoder hidden states to condition on. - target_input_ids (`torch.Tensor` of shape `(batch_size, target_sequence_length)`): - T5 token ids used as learned query tokens. - target_attention_mask (`torch.Tensor`, *optional*): - Attention mask for the target T5 token ids. - source_attention_mask (`torch.Tensor`, *optional*): - Attention mask for the source Qwen3 hidden states. - - Returns: - `torch.Tensor`: Text conditioning embeddings for the Cosmos transformer. - """ - target_attention_mask = self._prepare_attention_mask(target_attention_mask) - source_attention_mask = self._prepare_attention_mask(source_attention_mask) - - hidden_states = self.embed(target_input_ids).to(dtype=source_hidden_states.dtype) - hidden_states = self.in_proj(hidden_states) - - position_ids = torch.arange(hidden_states.shape[1], device=hidden_states.device).unsqueeze(0) - source_position_ids = torch.arange(source_hidden_states.shape[1], device=hidden_states.device).unsqueeze(0) - position_embeddings = self.rotary_emb(hidden_states, position_ids) - source_position_embeddings = self.rotary_emb(hidden_states, source_position_ids) - - for block in self.blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - source_hidden_states, - target_attention_mask, - source_attention_mask, - position_embeddings, - source_position_embeddings, - ) - else: - hidden_states = block( - hidden_states, - source_hidden_states, - target_attention_mask=target_attention_mask, - source_attention_mask=source_attention_mask, - position_embeddings=position_embeddings, - source_position_embeddings=source_position_embeddings, - ) - - hidden_states = self.norm(self.out_proj(hidden_states)) - - if target_attention_mask is not None: - hidden_states = hidden_states * target_attention_mask.squeeze(1).squeeze(1).to(hidden_states).unsqueeze(-1) - - if hidden_states.shape[1] < self.config.min_sequence_length: - hidden_states = F.pad(hidden_states, (0, 0, 0, self.config.min_sequence_length - hidden_states.shape[1])) - - return hidden_states diff --git a/diffusers/models/controlnets/__init__.py b/diffusers/models/controlnets/__init__.py deleted file mode 100644 index 3f9c7337e0f2bda0be6f3c4aa34c16382ac05427..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/__init__.py +++ /dev/null @@ -1,25 +0,0 @@ -from ...utils import is_torch_available - - -if is_torch_available(): - from .controlnet import ControlNetModel, ControlNetOutput - from .controlnet_cosmos import CosmosControlNetModel - from .controlnet_flux import FluxControlNetModel, FluxControlNetOutput, FluxMultiControlNetModel - from .controlnet_hunyuan import ( - HunyuanControlNetOutput, - HunyuanDiT2DControlNetModel, - HunyuanDiT2DMultiControlNetModel, - ) - from .controlnet_qwenimage import QwenImageControlNetModel, QwenImageMultiControlNetModel - from .controlnet_sana import SanaControlNetModel - from .controlnet_sd3 import SD3ControlNetModel, SD3ControlNetOutput, SD3MultiControlNetModel - from .controlnet_sparsectrl import ( - SparseControlNetConditioningEmbedding, - SparseControlNetModel, - SparseControlNetOutput, - ) - from .controlnet_union import ControlNetUnionModel - from .controlnet_xs import ControlNetXSAdapter, ControlNetXSOutput, UNetControlNetXSModel - from .controlnet_z_image import ZImageControlNetModel - from .multicontrolnet import MultiControlNetModel - from .multicontrolnet_union import MultiControlNetUnionModel diff --git a/diffusers/models/controlnets/controlnet.py b/diffusers/models/controlnets/controlnet.py deleted file mode 100644 index acd88655c9fe2502c0981f071c2868d5ef8278ac..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/controlnet.py +++ /dev/null @@ -1,807 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from dataclasses import dataclass -from typing import Any - -import torch -from torch import nn -from torch.nn import functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...loaders.single_file_model import FromOriginalModelMixin -from ...utils import BaseOutput, apply_lora_scale, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..attention import AttentionMixin -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - AttnAddedKVProcessor, - AttnProcessor, -) -from ..embeddings import TextImageProjection, TextImageTimeEmbedding, TextTimeEmbedding, TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from ..unets.unet_2d_blocks import ( - UNetMidBlock2D, - UNetMidBlock2DCrossAttn, - get_down_block, -) -from ..unets.unet_2d_condition import UNet2DConditionModel - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class ControlNetOutput(BaseOutput): - """ - The output of [`ControlNetModel`]. - - Args: - down_block_res_samples (`tuple[torch.Tensor]`): - A tuple of downsample activations at different resolutions for each downsampling block. Each tensor should - be of shape `(batch_size, channel * resolution, height //resolution, width // resolution)`. Output can be - used to condition the original UNet's downsampling activations. - mid_down_block_re_sample (`torch.Tensor`): - The activation of the middle block (the lowest sample resolution). Each tensor should be of shape - `(batch_size, channel * lowest_resolution, height // lowest_resolution, width // lowest_resolution)`. - Output can be used to condition the original UNet's middle block activation. - """ - - down_block_res_samples: tuple[torch.Tensor] - mid_block_res_sample: torch.Tensor - - -class ControlNetConditioningEmbedding(nn.Module): - """ - Quoting from https://huggingface.co/papers/2302.05543: "Stable Diffusion uses a pre-processing method similar to - VQ-GAN [11] to convert the entire dataset of 512 × 512 images into smaller 64 × 64 “latent images” for stabilized - training. This requires ControlNets to convert image-based conditions to 64 × 64 feature space to match the - convolution size. We use a tiny network E(·) of four convolution layers with 4 × 4 kernels and 2 × 2 strides - (activated by ReLU, channels are 16, 32, 64, 128, initialized with Gaussian weights, trained jointly with the full - model) to encode image-space conditions ... into feature maps ..." - """ - - def __init__( - self, - conditioning_embedding_channels: int, - conditioning_channels: int = 3, - block_out_channels: tuple[int, ...] = (16, 32, 96, 256), - ): - super().__init__() - - self.conv_in = nn.Conv2d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1) - - self.blocks = nn.ModuleList([]) - - for i in range(len(block_out_channels) - 1): - channel_in = block_out_channels[i] - channel_out = block_out_channels[i + 1] - self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1)) - self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2)) - - self.conv_out = zero_module( - nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1) - ) - - def forward(self, conditioning): - embedding = self.conv_in(conditioning) - embedding = F.silu(embedding) - - for block in self.blocks: - embedding = block(embedding) - embedding = F.silu(embedding) - - embedding = self.conv_out(embedding) - - return embedding - - -class ControlNetModel(ModelMixin, AttentionMixin, ConfigMixin, FromOriginalModelMixin, PeftAdapterMixin): - """ - A ControlNet model. - - Args: - in_channels (`int`, defaults to 4): - The number of channels in the input sample. - flip_sin_to_cos (`bool`, defaults to `True`): - Whether to flip the sin to cos in the time embedding. - freq_shift (`int`, defaults to 0): - The frequency shift to apply to the time embedding. - down_block_types (`tuple[str]`, defaults to `("CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "DownBlock2D")`): - The tuple of downsample blocks to use. - only_cross_attention (`bool | tuple[bool]`, defaults to `False`): - block_out_channels (`tuple[int]`, defaults to `(320, 640, 1280, 1280)`): - The tuple of output channels for each block. - layers_per_block (`int`, defaults to 2): - The number of layers per block. - downsample_padding (`int`, defaults to 1): - The padding to use for the downsampling convolution. - mid_block_scale_factor (`float`, defaults to 1): - The scale factor to use for the mid block. - act_fn (`str`, defaults to "silu"): - The activation function to use. - norm_num_groups (`int`, *optional*, defaults to 32): - The number of groups to use for the normalization. If None, normalization and activation layers is skipped - in post-processing. - norm_eps (`float`, defaults to 1e-5): - The epsilon to use for the normalization. - cross_attention_dim (`int`, defaults to 1280): - The dimension of the cross attention features. - transformer_layers_per_block (`int` or `tuple[int]`, *optional*, defaults to 1): - The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`]. Only relevant for - [`~models.unet_2d_blocks.CrossAttnDownBlock2D`], [`~models.unet_2d_blocks.CrossAttnUpBlock2D`], - [`~models.unet_2d_blocks.UNetMidBlock2DCrossAttn`]. - encoder_hid_dim (`int`, *optional*, defaults to None): - If `encoder_hid_dim_type` is defined, `encoder_hidden_states` will be projected from `encoder_hid_dim` - dimension to `cross_attention_dim`. - encoder_hid_dim_type (`str`, *optional*, defaults to `None`): - If given, the `encoder_hidden_states` and potentially other embeddings are down-projected to text - embeddings of dimension `cross_attention` according to `encoder_hid_dim_type`. - attention_head_dim (`int | tuple[int]`, defaults to 8): - The dimension of the attention heads. - use_linear_projection (`bool`, defaults to `False`): - class_embed_type (`str`, *optional*, defaults to `None`): - The type of class embedding to use which is ultimately summed with the time embeddings. Choose from None, - `"timestep"`, `"identity"`, `"projection"`, or `"simple_projection"`. - addition_embed_type (`str`, *optional*, defaults to `None`): - Configures an optional embedding which will be summed with the time embeddings. Choose from `None` or - "text". "text" will use the `TextTimeEmbedding` layer. - num_class_embeds (`int`, *optional*, defaults to 0): - Input dimension of the learnable embedding matrix to be projected to `time_embed_dim`, when performing - class conditioning with `class_embed_type` equal to `None`. - upcast_attention (`bool`, defaults to `False`): - resnet_time_scale_shift (`str`, defaults to `"default"`): - Time scale shift config for ResNet blocks (see `ResnetBlock2D`). Choose from `default` or `scale_shift`. - projection_class_embeddings_input_dim (`int`, *optional*, defaults to `None`): - The dimension of the `class_labels` input when `class_embed_type="projection"`. Required when - `class_embed_type="projection"`. - controlnet_conditioning_channel_order (`str`, defaults to `"rgb"`): - The channel order of conditional image. Will convert to `rgb` if it's `bgr`. - conditioning_embedding_out_channels (`tuple[int]`, *optional*, defaults to `(16, 32, 96, 256)`): - The tuple of output channel for each block in the `conditioning_embedding` layer. - global_pool_conditions (`bool`, defaults to `False`): - TODO(Patrick) - unused parameter. - addition_embed_type_num_heads (`int`, defaults to 64): - The number of heads to use for the `TextTimeEmbedding` layer. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 4, - conditioning_channels: int = 3, - flip_sin_to_cos: bool = True, - freq_shift: int = 0, - down_block_types: tuple[str, ...] = ( - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "DownBlock2D", - ), - mid_block_type: str | None = "UNetMidBlock2DCrossAttn", - only_cross_attention: bool | tuple[bool] = False, - block_out_channels: tuple[int, ...] = (320, 640, 1280, 1280), - layers_per_block: int = 2, - downsample_padding: int = 1, - mid_block_scale_factor: float = 1, - act_fn: str = "silu", - norm_num_groups: int | None = 32, - norm_eps: float = 1e-5, - cross_attention_dim: int = 1280, - transformer_layers_per_block: int | tuple[int, ...] = 1, - encoder_hid_dim: int | None = None, - encoder_hid_dim_type: str | None = None, - attention_head_dim: int | tuple[int, ...] = 8, - num_attention_heads: int | tuple[int, ...] | None = None, - use_linear_projection: bool = False, - class_embed_type: str | None = None, - addition_embed_type: str | None = None, - addition_time_embed_dim: int | None = None, - num_class_embeds: int | None = None, - upcast_attention: bool = False, - resnet_time_scale_shift: str = "default", - projection_class_embeddings_input_dim: int | None = None, - controlnet_conditioning_channel_order: str = "rgb", - conditioning_embedding_out_channels: tuple[int, ...] | None = (16, 32, 96, 256), - global_pool_conditions: bool = False, - addition_embed_type_num_heads: int = 64, - ): - super().__init__() - - # If `num_attention_heads` is not defined (which is the case for most models) - # it will default to `attention_head_dim`. This looks weird upon first reading it and it is. - # The reason for this behavior is to correct for incorrectly named variables that were introduced - # when this library was created. The incorrect naming was only discovered much later in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131 - # Changing `attention_head_dim` to `num_attention_heads` for 40,000+ configurations is too backwards breaking - # which is why we correct for the naming here. - num_attention_heads = num_attention_heads or attention_head_dim - - # Check inputs - if len(block_out_channels) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(only_cross_attention, bool) and len(only_cross_attention) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `only_cross_attention` as `down_block_types`. `only_cross_attention`: {only_cross_attention}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}." - ) - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types) - - # input - conv_in_kernel = 3 - conv_in_padding = (conv_in_kernel - 1) // 2 - self.conv_in = nn.Conv2d( - in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding - ) - - # time - time_embed_dim = block_out_channels[0] * 4 - self.time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos, freq_shift) - timestep_input_dim = block_out_channels[0] - self.time_embedding = TimestepEmbedding( - timestep_input_dim, - time_embed_dim, - act_fn=act_fn, - ) - - if encoder_hid_dim_type is None and encoder_hid_dim is not None: - encoder_hid_dim_type = "text_proj" - self.register_to_config(encoder_hid_dim_type=encoder_hid_dim_type) - logger.info("encoder_hid_dim_type defaults to 'text_proj' as `encoder_hid_dim` is defined.") - - if encoder_hid_dim is None and encoder_hid_dim_type is not None: - raise ValueError( - f"`encoder_hid_dim` has to be defined when `encoder_hid_dim_type` is set to {encoder_hid_dim_type}." - ) - - if encoder_hid_dim_type == "text_proj": - self.encoder_hid_proj = nn.Linear(encoder_hid_dim, cross_attention_dim) - elif encoder_hid_dim_type == "text_image_proj": - # image_embed_dim DOESN'T have to be `cross_attention_dim`. To not clutter the __init__ too much - # they are set to `cross_attention_dim` here as this is exactly the required dimension for the currently only use - # case when `addition_embed_type == "text_image_proj"` (Kandinsky 2.1)` - self.encoder_hid_proj = TextImageProjection( - text_embed_dim=encoder_hid_dim, - image_embed_dim=cross_attention_dim, - cross_attention_dim=cross_attention_dim, - ) - - elif encoder_hid_dim_type is not None: - raise ValueError( - f"encoder_hid_dim_type: {encoder_hid_dim_type} must be None, 'text_proj' or 'text_image_proj'." - ) - else: - self.encoder_hid_proj = None - - # class embedding - if class_embed_type is None and num_class_embeds is not None: - self.class_embedding = nn.Embedding(num_class_embeds, time_embed_dim) - elif class_embed_type == "timestep": - self.class_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim) - elif class_embed_type == "identity": - self.class_embedding = nn.Identity(time_embed_dim, time_embed_dim) - elif class_embed_type == "projection": - if projection_class_embeddings_input_dim is None: - raise ValueError( - "`class_embed_type`: 'projection' requires `projection_class_embeddings_input_dim` be set" - ) - # The projection `class_embed_type` is the same as the timestep `class_embed_type` except - # 1. the `class_labels` inputs are not first converted to sinusoidal embeddings - # 2. it projects from an arbitrary input dimension. - # - # Note that `TimestepEmbedding` is quite general, being mainly linear layers and activations. - # When used for embedding actual timesteps, the timesteps are first converted to sinusoidal embeddings. - # As a result, `TimestepEmbedding` can be passed arbitrary vectors. - self.class_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) - else: - self.class_embedding = None - - if addition_embed_type == "text": - if encoder_hid_dim is not None: - text_time_embedding_from_dim = encoder_hid_dim - else: - text_time_embedding_from_dim = cross_attention_dim - - self.add_embedding = TextTimeEmbedding( - text_time_embedding_from_dim, time_embed_dim, num_heads=addition_embed_type_num_heads - ) - elif addition_embed_type == "text_image": - # text_embed_dim and image_embed_dim DON'T have to be `cross_attention_dim`. To not clutter the __init__ too much - # they are set to `cross_attention_dim` here as this is exactly the required dimension for the currently only use - # case when `addition_embed_type == "text_image"` (Kandinsky 2.1)` - self.add_embedding = TextImageTimeEmbedding( - text_embed_dim=cross_attention_dim, image_embed_dim=cross_attention_dim, time_embed_dim=time_embed_dim - ) - elif addition_embed_type == "text_time": - self.add_time_proj = Timesteps(addition_time_embed_dim, flip_sin_to_cos, freq_shift) - self.add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) - - elif addition_embed_type is not None: - raise ValueError(f"addition_embed_type: {addition_embed_type} must be None, 'text' or 'text_image'.") - - # control net conditioning embedding - self.controlnet_cond_embedding = ControlNetConditioningEmbedding( - conditioning_embedding_channels=block_out_channels[0], - block_out_channels=conditioning_embedding_out_channels, - conditioning_channels=conditioning_channels, - ) - - self.down_blocks = nn.ModuleList([]) - self.controlnet_down_blocks = nn.ModuleList([]) - - if isinstance(only_cross_attention, bool): - only_cross_attention = [only_cross_attention] * len(down_block_types) - - if isinstance(attention_head_dim, int): - attention_head_dim = (attention_head_dim,) * len(down_block_types) - - if isinstance(num_attention_heads, int): - num_attention_heads = (num_attention_heads,) * len(down_block_types) - - # down - output_channel = block_out_channels[0] - - controlnet_block = nn.Conv2d(output_channel, output_channel, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_down_blocks.append(controlnet_block) - - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - down_block = get_down_block( - down_block_type, - num_layers=layers_per_block, - transformer_layers_per_block=transformer_layers_per_block[i], - in_channels=input_channel, - out_channels=output_channel, - temb_channels=time_embed_dim, - add_downsample=not is_final_block, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads[i], - attention_head_dim=attention_head_dim[i] if attention_head_dim[i] is not None else output_channel, - downsample_padding=downsample_padding, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention[i], - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - ) - self.down_blocks.append(down_block) - - for _ in range(layers_per_block): - controlnet_block = nn.Conv2d(output_channel, output_channel, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_down_blocks.append(controlnet_block) - - if not is_final_block: - controlnet_block = nn.Conv2d(output_channel, output_channel, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_down_blocks.append(controlnet_block) - - # mid - mid_block_channel = block_out_channels[-1] - - controlnet_block = nn.Conv2d(mid_block_channel, mid_block_channel, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_mid_block = controlnet_block - - if mid_block_type == "UNetMidBlock2DCrossAttn": - self.mid_block = UNetMidBlock2DCrossAttn( - transformer_layers_per_block=transformer_layers_per_block[-1], - in_channels=mid_block_channel, - temb_channels=time_embed_dim, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - output_scale_factor=mid_block_scale_factor, - resnet_time_scale_shift=resnet_time_scale_shift, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads[-1], - resnet_groups=norm_num_groups, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - ) - elif mid_block_type == "UNetMidBlock2D": - self.mid_block = UNetMidBlock2D( - in_channels=block_out_channels[-1], - temb_channels=time_embed_dim, - num_layers=0, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - output_scale_factor=mid_block_scale_factor, - resnet_groups=norm_num_groups, - resnet_time_scale_shift=resnet_time_scale_shift, - add_attention=False, - ) - else: - raise ValueError(f"unknown mid_block_type : {mid_block_type}") - - @classmethod - def from_unet( - cls, - unet: UNet2DConditionModel, - controlnet_conditioning_channel_order: str = "rgb", - conditioning_embedding_out_channels: tuple[int, ...] | None = (16, 32, 96, 256), - load_weights_from_unet: bool = True, - conditioning_channels: int = 3, - ): - r""" - Instantiate a [`ControlNetModel`] from [`UNet2DConditionModel`]. - - Parameters: - unet (`UNet2DConditionModel`): - The UNet model weights to copy to the [`ControlNetModel`]. All configuration options are also copied - where applicable. - """ - transformer_layers_per_block = ( - unet.config.transformer_layers_per_block if "transformer_layers_per_block" in unet.config else 1 - ) - encoder_hid_dim = unet.config.encoder_hid_dim if "encoder_hid_dim" in unet.config else None - encoder_hid_dim_type = unet.config.encoder_hid_dim_type if "encoder_hid_dim_type" in unet.config else None - addition_embed_type = unet.config.addition_embed_type if "addition_embed_type" in unet.config else None - addition_time_embed_dim = ( - unet.config.addition_time_embed_dim if "addition_time_embed_dim" in unet.config else None - ) - - controlnet = cls( - encoder_hid_dim=encoder_hid_dim, - encoder_hid_dim_type=encoder_hid_dim_type, - addition_embed_type=addition_embed_type, - addition_time_embed_dim=addition_time_embed_dim, - transformer_layers_per_block=transformer_layers_per_block, - in_channels=unet.config.in_channels, - flip_sin_to_cos=unet.config.flip_sin_to_cos, - freq_shift=unet.config.freq_shift, - down_block_types=unet.config.down_block_types, - only_cross_attention=unet.config.only_cross_attention, - block_out_channels=unet.config.block_out_channels, - layers_per_block=unet.config.layers_per_block, - downsample_padding=unet.config.downsample_padding, - mid_block_scale_factor=unet.config.mid_block_scale_factor, - act_fn=unet.config.act_fn, - norm_num_groups=unet.config.norm_num_groups, - norm_eps=unet.config.norm_eps, - cross_attention_dim=unet.config.cross_attention_dim, - attention_head_dim=unet.config.attention_head_dim, - num_attention_heads=unet.config.num_attention_heads, - use_linear_projection=unet.config.use_linear_projection, - class_embed_type=unet.config.class_embed_type, - num_class_embeds=unet.config.num_class_embeds, - upcast_attention=unet.config.upcast_attention, - resnet_time_scale_shift=unet.config.resnet_time_scale_shift, - projection_class_embeddings_input_dim=unet.config.projection_class_embeddings_input_dim, - mid_block_type=unet.config.mid_block_type, - controlnet_conditioning_channel_order=controlnet_conditioning_channel_order, - conditioning_embedding_out_channels=conditioning_embedding_out_channels, - conditioning_channels=conditioning_channels, - ) - - if load_weights_from_unet: - controlnet.conv_in.load_state_dict(unet.conv_in.state_dict()) - controlnet.time_proj.load_state_dict(unet.time_proj.state_dict()) - controlnet.time_embedding.load_state_dict(unet.time_embedding.state_dict()) - - if controlnet.class_embedding: - controlnet.class_embedding.load_state_dict(unet.class_embedding.state_dict()) - - if hasattr(controlnet, "add_embedding"): - controlnet.add_embedding.load_state_dict(unet.add_embedding.state_dict()) - - controlnet.down_blocks.load_state_dict(unet.down_blocks.state_dict()) - controlnet.mid_block.load_state_dict(unet.mid_block.state_dict()) - - return controlnet - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attention_slice - def set_attention_slice(self, slice_size: str | int | list[int]) -> None: - r""" - Enable sliced attention computation. - - When this option is enabled, the attention module splits the input tensor in slices to compute attention in - several steps. This is useful for saving some memory in exchange for a small decrease in speed. - - Args: - slice_size (`str` or `int` or `list(int)`, *optional*, defaults to `"auto"`): - When `"auto"`, input to the attention heads is halved, so attention is computed in two steps. If - `"max"`, maximum amount of memory is saved by running only one slice at a time. If a number is - provided, uses as many slices as `attention_head_dim // slice_size`. In this case, `attention_head_dim` - must be a multiple of `slice_size`. - """ - sliceable_head_dims = [] - - def fn_recursive_retrieve_sliceable_dims(module: torch.nn.Module): - if hasattr(module, "set_attention_slice"): - sliceable_head_dims.append(module.sliceable_head_dim) - - for child in module.children(): - fn_recursive_retrieve_sliceable_dims(child) - - # retrieve number of attention layers - for module in self.children(): - fn_recursive_retrieve_sliceable_dims(module) - - num_sliceable_layers = len(sliceable_head_dims) - - if slice_size == "auto": - # half the attention head size is usually a good trade-off between - # speed and memory - slice_size = [dim // 2 for dim in sliceable_head_dims] - elif slice_size == "max": - # make smallest slice possible - slice_size = num_sliceable_layers * [1] - - slice_size = num_sliceable_layers * [slice_size] if not isinstance(slice_size, list) else slice_size - - if len(slice_size) != len(sliceable_head_dims): - raise ValueError( - f"You have provided {len(slice_size)}, but {self.config} has {len(sliceable_head_dims)} different" - f" attention layers. Make sure to match `len(slice_size)` to be {len(sliceable_head_dims)}." - ) - - for i in range(len(slice_size)): - size = slice_size[i] - dim = sliceable_head_dims[i] - if size is not None and size > dim: - raise ValueError(f"size {size} has to be smaller or equal to {dim}.") - - # Recursively walk through all the children. - # Any children which exposes the set_attention_slice method - # gets the message - def fn_recursive_set_attention_slice(module: torch.nn.Module, slice_size: list[int]): - if hasattr(module, "set_attention_slice"): - module.set_attention_slice(slice_size.pop()) - - for child in module.children(): - fn_recursive_set_attention_slice(child, slice_size) - - reversed_slice_size = list(reversed(slice_size)) - for module in self.children(): - fn_recursive_set_attention_slice(module, reversed_slice_size) - - @apply_lora_scale("cross_attention_kwargs") - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - encoder_hidden_states: torch.Tensor, - controlnet_cond: torch.Tensor, - conditioning_scale: float = 1.0, - class_labels: torch.Tensor | None = None, - timestep_cond: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - added_cond_kwargs: dict[str, torch.Tensor] | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - guess_mode: bool = False, - return_dict: bool = True, - ) -> ControlNetOutput | tuple[tuple[torch.Tensor, ...], torch.Tensor]: - """ - The [`ControlNetModel`] forward method. - - Args: - sample (`torch.Tensor`): - The noisy input tensor. - timestep (`torch.Tensor | float | int`): - The number of timesteps to denoise an input. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states. - controlnet_cond (`torch.Tensor`): - The conditional input tensor of shape `(batch_size, sequence_length, hidden_size)`. - conditioning_scale (`float`, defaults to `1.0`): - The scale factor for ControlNet outputs. - class_labels (`torch.Tensor`, *optional*, defaults to `None`): - Optional class labels for conditioning. Their embeddings will be summed with the timestep embeddings. - timestep_cond (`torch.Tensor`, *optional*, defaults to `None`): - Additional conditional embeddings for timestep. If provided, the embeddings will be summed with the - timestep_embedding passed through the `self.time_embedding` layer to obtain the final timestep - embeddings. - attention_mask (`torch.Tensor`, *optional*, defaults to `None`): - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask - is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large - negative values to the attention scores corresponding to "discard" tokens. - added_cond_kwargs (`dict`): - Additional conditions for the Stable Diffusion XL UNet. - cross_attention_kwargs (`dict[str]`, *optional*, defaults to `None`): - A kwargs dictionary that if specified is passed along to the `AttnProcessor`. - guess_mode (`bool`, defaults to `False`): - In this mode, the ControlNet encoder tries its best to recognize the input content of the input even if - you remove all prompts. A `guidance_scale` between 3.0 and 5.0 is recommended. - return_dict (`bool`, defaults to `True`): - Whether or not to return a [`~models.controlnets.controlnet.ControlNetOutput`] instead of a plain - tuple. - - Returns: - [`~models.controlnets.controlnet.ControlNetOutput`] **or** `tuple`: - If `return_dict` is `True`, a [`~models.controlnets.controlnet.ControlNetOutput`] is returned, - otherwise a tuple is returned where the first element is the sample tensor. - """ - # check channel order - channel_order = self.config.controlnet_conditioning_channel_order - - if channel_order == "rgb": - # in rgb order by default - ... - elif channel_order == "bgr": - controlnet_cond = torch.flip(controlnet_cond, dims=[1]) - else: - raise ValueError(f"unknown `controlnet_conditioning_channel_order`: {channel_order}") - - # prepare attention_mask - if attention_mask is not None: - attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # 1. time - timesteps = timestep - if not torch.is_tensor(timesteps): - # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can - # This would be a good case for the `match` statement (Python 3.10+) - dtype = maybe_adjust_dtype_for_device( - torch.float64 if isinstance(timestep, float) else torch.int64, sample.device - ) - timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) - elif len(timesteps.shape) == 0: - timesteps = timesteps[None].to(sample.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - timesteps = timesteps.expand(sample.shape[0]) - - t_emb = self.time_proj(timesteps) - - # timesteps does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=sample.dtype) - - emb = self.time_embedding(t_emb, timestep_cond) - aug_emb = None - - if self.class_embedding is not None: - if class_labels is None: - raise ValueError("class_labels should be provided when num_class_embeds > 0") - - if self.config.class_embed_type == "timestep": - class_labels = self.time_proj(class_labels) - - class_emb = self.class_embedding(class_labels).to(dtype=self.dtype) - emb = emb + class_emb - - if self.config.addition_embed_type is not None: - if self.config.addition_embed_type == "text": - aug_emb = self.add_embedding(encoder_hidden_states) - - elif self.config.addition_embed_type == "text_time": - if "text_embeds" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `text_embeds` to be passed in `added_cond_kwargs`" - ) - text_embeds = added_cond_kwargs.get("text_embeds") - if "time_ids" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `time_ids` to be passed in `added_cond_kwargs`" - ) - time_ids = added_cond_kwargs.get("time_ids") - time_embeds = self.add_time_proj(time_ids.flatten()) - time_embeds = time_embeds.reshape((text_embeds.shape[0], -1)) - - add_embeds = torch.concat([text_embeds, time_embeds], dim=-1) - add_embeds = add_embeds.to(emb.dtype) - aug_emb = self.add_embedding(add_embeds) - - emb = emb + aug_emb if aug_emb is not None else emb - - # 2. pre-process - sample = self.conv_in(sample) - - controlnet_cond = self.controlnet_cond_embedding(controlnet_cond) - sample = sample + controlnet_cond - - # 3. down - down_block_res_samples = (sample,) - for downsample_block in self.down_blocks: - if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention: - sample, res_samples = downsample_block( - hidden_states=sample, - temb=emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample, res_samples = downsample_block(hidden_states=sample, temb=emb) - - down_block_res_samples += res_samples - - # 4. mid - if self.mid_block is not None: - if hasattr(self.mid_block, "has_cross_attention") and self.mid_block.has_cross_attention: - sample = self.mid_block( - sample, - emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample = self.mid_block(sample, emb) - - # 5. Control net blocks - - controlnet_down_block_res_samples = () - - for down_block_res_sample, controlnet_block in zip(down_block_res_samples, self.controlnet_down_blocks): - down_block_res_sample = controlnet_block(down_block_res_sample) - controlnet_down_block_res_samples = controlnet_down_block_res_samples + (down_block_res_sample,) - - down_block_res_samples = controlnet_down_block_res_samples - - mid_block_res_sample = self.controlnet_mid_block(sample) - - # 6. scaling - if guess_mode and not self.config.global_pool_conditions: - scales = torch.logspace(-1, 0, len(down_block_res_samples) + 1, device=sample.device) # 0.1 to 1.0 - scales = scales * conditioning_scale - down_block_res_samples = [sample * scale for sample, scale in zip(down_block_res_samples, scales)] - mid_block_res_sample = mid_block_res_sample * scales[-1] # last one - else: - down_block_res_samples = [sample * conditioning_scale for sample in down_block_res_samples] - mid_block_res_sample = mid_block_res_sample * conditioning_scale - - if self.config.global_pool_conditions: - down_block_res_samples = [ - torch.mean(sample, dim=(2, 3), keepdim=True) for sample in down_block_res_samples - ] - mid_block_res_sample = torch.mean(mid_block_res_sample, dim=(2, 3), keepdim=True) - - if not return_dict: - return (down_block_res_samples, mid_block_res_sample) - - return ControlNetOutput( - down_block_res_samples=down_block_res_samples, mid_block_res_sample=mid_block_res_sample - ) - - -def zero_module(module): - for p in module.parameters(): - nn.init.zeros_(p) - return module diff --git a/diffusers/models/controlnets/controlnet_cosmos.py b/diffusers/models/controlnets/controlnet_cosmos.py deleted file mode 100644 index e39f8dfb568a02a74f076ab9212af25b8a59f816..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/controlnet_cosmos.py +++ /dev/null @@ -1,317 +0,0 @@ -from dataclasses import dataclass -from typing import List, Optional, Tuple, Union - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin -from ...utils import BaseOutput, is_torchvision_available, logging -from ..modeling_utils import ModelMixin -from ..transformers.transformer_cosmos import ( - CosmosEmbedding, - CosmosLearnablePositionalEmbed, - CosmosPatchEmbed, - CosmosRotaryPosEmbed, - CosmosTransformerBlock, -) - - -if is_torchvision_available(): - from torchvision import transforms - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class CosmosControlNetOutput(BaseOutput): - """ - Output of [`CosmosControlNetModel`]. - - Args: - control_block_samples (`list[torch.Tensor]`): - List of control block activations to be injected into transformer blocks. - """ - - control_block_samples: List[torch.Tensor] - - -class CosmosControlNetModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): - r""" - ControlNet for Cosmos Transfer2.5. - - This model duplicates the shared embedding modules from the transformer (patch_embed, time_embed, - learnable_pos_embed, img_context_proj) to enable proper CPU offloading. The forward() method computes everything - internally from raw inputs. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["patch_embed", "patch_embed_base", "time_embed"] - _no_split_modules = ["CosmosTransformerBlock"] - _keep_in_fp32_modules = ["learnable_pos_embed"] - - @register_to_config - def __init__( - self, - n_controlnet_blocks: int = 4, - in_channels: int = 130, - latent_channels: int = 18, # base latent channels (latents + condition_mask) + padding_mask - model_channels: int = 2048, - num_attention_heads: int = 32, - attention_head_dim: int = 128, - mlp_ratio: float = 4.0, - text_embed_dim: int = 1024, - adaln_lora_dim: int = 256, - patch_size: Tuple[int, int, int] = (1, 2, 2), - max_size: Tuple[int, int, int] = (128, 240, 240), - rope_scale: Tuple[float, float, float] = (2.0, 1.0, 1.0), - extra_pos_embed_type: str | None = None, - img_context_dim_in: int | None = None, - img_context_dim_out: int = 2048, - use_crossattn_projection: bool = False, - crossattn_proj_in_channels: int = 1024, - encoder_hidden_states_channels: int = 1024, - ): - super().__init__() - - self.patch_embed = CosmosPatchEmbed(in_channels, model_channels, patch_size, bias=False) - - self.patch_embed_base = CosmosPatchEmbed(latent_channels, model_channels, patch_size, bias=False) - self.time_embed = CosmosEmbedding(model_channels, model_channels) - - self.learnable_pos_embed = None - if extra_pos_embed_type == "learnable": - self.learnable_pos_embed = CosmosLearnablePositionalEmbed( - hidden_size=model_channels, - max_size=max_size, - patch_size=patch_size, - ) - - self.img_context_proj = None - if img_context_dim_in is not None and img_context_dim_in > 0: - self.img_context_proj = nn.Sequential( - nn.Linear(img_context_dim_in, img_context_dim_out, bias=True), - nn.GELU(), - ) - - # Cross-attention projection for text embeddings (same as transformer) - self.crossattn_proj = None - if use_crossattn_projection: - self.crossattn_proj = nn.Sequential( - nn.Linear(crossattn_proj_in_channels, encoder_hidden_states_channels, bias=True), - nn.GELU(), - ) - - # RoPE for both control and base latents - self.rope = CosmosRotaryPosEmbed( - hidden_size=attention_head_dim, max_size=max_size, patch_size=patch_size, rope_scale=rope_scale - ) - - self.control_blocks = nn.ModuleList( - [ - CosmosTransformerBlock( - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - cross_attention_dim=text_embed_dim, - mlp_ratio=mlp_ratio, - adaln_lora_dim=adaln_lora_dim, - qk_norm="rms_norm", - out_bias=False, - img_context=img_context_dim_in is not None and img_context_dim_in > 0, - before_proj=(block_idx == 0), - after_proj=True, - ) - for block_idx in range(n_controlnet_blocks) - ] - ) - - self.gradient_checkpointing = False - - def _expand_conditioning_scale(self, conditioning_scale: float | list[float]) -> List[float]: - if isinstance(conditioning_scale, list): - scales = conditioning_scale - else: - scales = [conditioning_scale] * len(self.control_blocks) - - if len(scales) < len(self.control_blocks): - logger.warning( - "Received %d control scales, but control network defines %d blocks. " - "Scales will be trimmed or repeated to match.", - len(scales), - len(self.control_blocks), - ) - scales = (scales * len(self.control_blocks))[: len(self.control_blocks)] - return scales - - def forward( - self, - controls_latents: torch.Tensor, - latents: torch.Tensor, - timestep: torch.Tensor, - encoder_hidden_states: Union[Optional[torch.Tensor], Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]], - condition_mask: torch.Tensor, - conditioning_scale: float | list[float] = 1.0, - padding_mask: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - fps: int | None = None, - return_dict: bool = True, - ) -> Union[CosmosControlNetOutput, Tuple[List[torch.Tensor]]]: - """ - Forward pass for the ControlNet. - - Args: - controls_latents: Control signal latents [B, C, T, H, W] - latents: Base latents from the noising process [B, C, T, H, W] - timestep: Diffusion timestep tensor - encoder_hidden_states: Tuple of (text_context, img_context) or text_context - condition_mask: Conditioning mask [B, 1, T, H, W] - conditioning_scale: Scale factor(s) for control outputs - padding_mask: Padding mask [B, 1, H, W] or None - attention_mask: Optional attention mask or None - fps: Frames per second for RoPE or None - return_dict: Whether to return a CosmosControlNetOutput or a tuple - - Returns: - CosmosControlNetOutput or tuple of control tensors - """ - B, C, T, H, W = controls_latents.shape - - # 1. Prepare control latents - control_hidden_states = controls_latents - vace_in_channels = self.config.in_channels - 1 - if control_hidden_states.shape[1] < vace_in_channels - 1: - pad_C = vace_in_channels - 1 - control_hidden_states.shape[1] - control_hidden_states = torch.cat( - [ - control_hidden_states, - torch.zeros( - (B, pad_C, T, H, W), dtype=control_hidden_states.dtype, device=control_hidden_states.device - ), - ], - dim=1, - ) - - if condition_mask is not None: - control_hidden_states = torch.cat([control_hidden_states, condition_mask], dim=1) - else: - control_hidden_states = torch.cat( - [control_hidden_states, torch.zeros_like(controls_latents[:, :1])], dim=1 - ) - - padding_mask_resized = transforms.functional.resize( - padding_mask, list(control_hidden_states.shape[-2:]), interpolation=transforms.InterpolationMode.NEAREST - ) - control_hidden_states = torch.cat( - [control_hidden_states, padding_mask_resized.unsqueeze(2).repeat(B, 1, T, 1, 1)], dim=1 - ) - - # 2. Prepare base latents (same processing as transformer.forward) - base_hidden_states = latents - if condition_mask is not None: - base_hidden_states = torch.cat([base_hidden_states, condition_mask], dim=1) - - base_padding_mask = transforms.functional.resize( - padding_mask, list(base_hidden_states.shape[-2:]), interpolation=transforms.InterpolationMode.NEAREST - ) - base_hidden_states = torch.cat( - [base_hidden_states, base_padding_mask.unsqueeze(2).repeat(B, 1, T, 1, 1)], dim=1 - ) - - # 3. Generate positional embeddings (shared for both) - image_rotary_emb = self.rope(control_hidden_states, fps=fps) - extra_pos_emb = self.learnable_pos_embed(control_hidden_states) if self.learnable_pos_embed else None - - # 4. Patchify control latents - control_hidden_states = self.patch_embed(control_hidden_states) - control_hidden_states = control_hidden_states.flatten(1, 3) - - # 5. Patchify base latents - p_t, p_h, p_w = self.config.patch_size - post_patch_num_frames = T // p_t - post_patch_height = H // p_h - post_patch_width = W // p_w - - base_hidden_states = self.patch_embed_base(base_hidden_states) - base_hidden_states = base_hidden_states.flatten(1, 3) - - # 6. Time embeddings - if timestep.ndim == 1: - temb, embedded_timestep = self.time_embed(base_hidden_states, timestep) - elif timestep.ndim == 5: - batch_size, _, num_frames, _, _ = latents.shape - assert timestep.shape == (batch_size, 1, num_frames, 1, 1), ( - f"Expected timestep to have shape [B, 1, T, 1, 1], but got {timestep.shape}" - ) - timestep_flat = timestep.flatten() - temb, embedded_timestep = self.time_embed(base_hidden_states, timestep_flat) - temb, embedded_timestep = ( - x.view(batch_size, post_patch_num_frames, 1, 1, -1) - .expand(-1, -1, post_patch_height, post_patch_width, -1) - .flatten(1, 3) - for x in (temb, embedded_timestep) - ) - else: - raise ValueError(f"Expected timestep to have shape [B, 1, T, 1, 1] or [T], but got {timestep.shape}") - - # 7. Process encoder hidden states - if isinstance(encoder_hidden_states, tuple): - text_context, img_context = encoder_hidden_states - else: - text_context = encoder_hidden_states - img_context = None - - # Apply cross-attention projection to text context - if self.crossattn_proj is not None: - text_context = self.crossattn_proj(text_context) - - # Apply cross-attention projection to image context (if provided) - if img_context is not None and self.img_context_proj is not None: - img_context = self.img_context_proj(img_context) - - # Combine text and image context into a single tuple - if self.config.img_context_dim_in is not None and self.config.img_context_dim_in > 0: - processed_encoder_hidden_states = (text_context, img_context) - else: - processed_encoder_hidden_states = text_context - - # 8. Prepare attention mask - if attention_mask is not None: - attention_mask = attention_mask.unsqueeze(1).unsqueeze(1) # [B, 1, 1, S] - - # 9. Run control blocks - scales = self._expand_conditioning_scale(conditioning_scale) - result = [] - for block_idx, (block, scale) in enumerate(zip(self.control_blocks, scales)): - if torch.is_grad_enabled() and self.gradient_checkpointing: - control_hidden_states, control_proj = self._gradient_checkpointing_func( - block, - control_hidden_states, - processed_encoder_hidden_states, - embedded_timestep, - temb, - image_rotary_emb, - extra_pos_emb, - attention_mask, - None, # controlnet_residual - base_hidden_states, - block_idx, - ) - else: - control_hidden_states, control_proj = block( - hidden_states=control_hidden_states, - encoder_hidden_states=processed_encoder_hidden_states, - embedded_timestep=embedded_timestep, - temb=temb, - image_rotary_emb=image_rotary_emb, - extra_pos_emb=extra_pos_emb, - attention_mask=attention_mask, - controlnet_residual=None, - latents=base_hidden_states, - block_idx=block_idx, - ) - result.append(control_proj * scale) - - if not return_dict: - return (result,) - - return CosmosControlNetOutput(control_block_samples=result) diff --git a/diffusers/models/controlnets/controlnet_flux.py b/diffusers/models/controlnets/controlnet_flux.py deleted file mode 100644 index e52465abc37c0ff36968f6ff24325dc13ad20452..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/controlnet_flux.py +++ /dev/null @@ -1,474 +0,0 @@ -# Copyright 2025 Black Forest Labs, The HuggingFace Team and The InstantX Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import ( - BaseOutput, - apply_lora_scale, - logging, -) -from ..attention import AttentionMixin -from ..controlnets.controlnet import ControlNetConditioningEmbedding, zero_module -from ..embeddings import CombinedTimestepGuidanceTextProjEmbeddings, CombinedTimestepTextProjEmbeddings, FluxPosEmbed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..transformers.transformer_flux import FluxSingleTransformerBlock, FluxTransformerBlock - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class FluxControlNetOutput(BaseOutput): - controlnet_block_samples: tuple[torch.Tensor] - controlnet_single_block_samples: tuple[torch.Tensor] - - -class FluxControlNetModel(ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin): - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - patch_size: int = 1, - in_channels: int = 64, - num_layers: int = 19, - num_single_layers: int = 38, - attention_head_dim: int = 128, - num_attention_heads: int = 24, - joint_attention_dim: int = 4096, - pooled_projection_dim: int = 768, - guidance_embeds: bool = False, - axes_dims_rope: list[int] = [16, 56, 56], - num_mode: int = None, - conditioning_embedding_channels: int = None, - ): - super().__init__() - self.out_channels = in_channels - self.inner_dim = num_attention_heads * attention_head_dim - - self.pos_embed = FluxPosEmbed(theta=10000, axes_dim=axes_dims_rope) - text_time_guidance_cls = ( - CombinedTimestepGuidanceTextProjEmbeddings if guidance_embeds else CombinedTimestepTextProjEmbeddings - ) - self.time_text_embed = text_time_guidance_cls( - embedding_dim=self.inner_dim, pooled_projection_dim=pooled_projection_dim - ) - - self.context_embedder = nn.Linear(joint_attention_dim, self.inner_dim) - self.x_embedder = torch.nn.Linear(in_channels, self.inner_dim) - - self.transformer_blocks = nn.ModuleList( - [ - FluxTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for i in range(num_layers) - ] - ) - - self.single_transformer_blocks = nn.ModuleList( - [ - FluxSingleTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for i in range(num_single_layers) - ] - ) - - # controlnet_blocks - self.controlnet_blocks = nn.ModuleList([]) - for _ in range(len(self.transformer_blocks)): - self.controlnet_blocks.append(zero_module(nn.Linear(self.inner_dim, self.inner_dim))) - - self.controlnet_single_blocks = nn.ModuleList([]) - for _ in range(len(self.single_transformer_blocks)): - self.controlnet_single_blocks.append(zero_module(nn.Linear(self.inner_dim, self.inner_dim))) - - self.union = num_mode is not None - if self.union: - self.controlnet_mode_embedder = nn.Embedding(num_mode, self.inner_dim) - - if conditioning_embedding_channels is not None: - self.input_hint_block = ControlNetConditioningEmbedding( - conditioning_embedding_channels=conditioning_embedding_channels, block_out_channels=(16, 16, 16, 16) - ) - self.controlnet_x_embedder = torch.nn.Linear(in_channels, self.inner_dim) - else: - self.input_hint_block = None - self.controlnet_x_embedder = zero_module(torch.nn.Linear(in_channels, self.inner_dim)) - - self.gradient_checkpointing = False - - @classmethod - def from_transformer( - cls, - transformer, - num_layers: int = 4, - num_single_layers: int = 10, - attention_head_dim: int = 128, - num_attention_heads: int = 24, - load_weights_from_transformer=True, - ): - config = dict(transformer.config) - config["num_layers"] = num_layers - config["num_single_layers"] = num_single_layers - config["attention_head_dim"] = attention_head_dim - config["num_attention_heads"] = num_attention_heads - - controlnet = cls.from_config(config) - - if load_weights_from_transformer: - controlnet.pos_embed.load_state_dict(transformer.pos_embed.state_dict()) - controlnet.time_text_embed.load_state_dict(transformer.time_text_embed.state_dict()) - controlnet.context_embedder.load_state_dict(transformer.context_embedder.state_dict()) - controlnet.x_embedder.load_state_dict(transformer.x_embedder.state_dict()) - controlnet.transformer_blocks.load_state_dict(transformer.transformer_blocks.state_dict(), strict=False) - controlnet.single_transformer_blocks.load_state_dict( - transformer.single_transformer_blocks.state_dict(), strict=False - ) - - controlnet.controlnet_x_embedder = zero_module(controlnet.controlnet_x_embedder) - - return controlnet - - @apply_lora_scale("joint_attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - controlnet_cond: torch.Tensor, - controlnet_mode: torch.Tensor = None, - conditioning_scale: float = 1.0, - encoder_hidden_states: torch.Tensor = None, - pooled_projections: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_ids: torch.Tensor = None, - txt_ids: torch.Tensor = None, - guidance: torch.Tensor = None, - joint_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> torch.FloatTensor | Transformer2DModelOutput: - """ - The [`FluxTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.FloatTensor` of shape `(batch size, channel, height, width)`): - Input `hidden_states`. - controlnet_cond (`torch.Tensor`): - The conditional input tensor of shape `(batch_size, sequence_length, hidden_size)`. - controlnet_mode (`torch.Tensor`): - The mode tensor of shape `(batch_size, 1)`. - conditioning_scale (`float`, defaults to `1.0`): - The scale factor for ControlNet outputs. - encoder_hidden_states (`torch.FloatTensor` of shape `(batch size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - pooled_projections (`torch.FloatTensor` of shape `(batch_size, projection_dim)`): Embeddings projected - from the embeddings of input conditions. - timestep ( `torch.LongTensor`): - Used to indicate denoising step. - img_ids (`torch.Tensor`): - Positional ids for the image tokens. - txt_ids (`torch.Tensor`): - Positional ids for the text tokens. - guidance (`torch.Tensor`, *optional*): - Guidance scale tensor used by guidance-distilled variants of the model. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - hidden_states = self.x_embedder(hidden_states) - - if self.input_hint_block is not None: - controlnet_cond = self.input_hint_block(controlnet_cond) - batch_size, channels, height_pw, width_pw = controlnet_cond.shape - height = height_pw // self.config.patch_size - width = width_pw // self.config.patch_size - controlnet_cond = controlnet_cond.reshape( - batch_size, channels, height, self.config.patch_size, width, self.config.patch_size - ) - controlnet_cond = controlnet_cond.permute(0, 2, 4, 1, 3, 5) - controlnet_cond = controlnet_cond.reshape(batch_size, height * width, -1) - # add - hidden_states = hidden_states + self.controlnet_x_embedder(controlnet_cond) - - timestep = timestep.to(hidden_states.dtype) * 1000 - if guidance is not None: - guidance = guidance.to(hidden_states.dtype) * 1000 - else: - guidance = None - temb = ( - self.time_text_embed(timestep, pooled_projections) - if guidance is None - else self.time_text_embed(timestep, guidance, pooled_projections) - ) - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - - if txt_ids.ndim == 3: - logger.warning( - "Passing `txt_ids` 3d torch.Tensor is deprecated." - "Please remove the batch dimension and pass it as a 2d torch Tensor" - ) - txt_ids = txt_ids[0] - if img_ids.ndim == 3: - logger.warning( - "Passing `img_ids` 3d torch.Tensor is deprecated." - "Please remove the batch dimension and pass it as a 2d torch Tensor" - ) - img_ids = img_ids[0] - - if self.union: - # union mode - if controlnet_mode is None: - raise ValueError("`controlnet_mode` cannot be `None` when applying ControlNet-Union") - # union mode emb - controlnet_mode_emb = self.controlnet_mode_embedder(controlnet_mode) - encoder_hidden_states = torch.cat([controlnet_mode_emb, encoder_hidden_states], dim=1) - txt_ids = torch.cat([txt_ids[:1], txt_ids], dim=0) - - ids = torch.cat((txt_ids, img_ids), dim=0) - image_rotary_emb = self.pos_embed(ids) - - block_samples = () - for index_block, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - ) - block_samples = block_samples + (hidden_states,) - - single_block_samples = () - for index_block, block in enumerate(self.single_transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - ) - single_block_samples = single_block_samples + (hidden_states,) - - # controlnet block - controlnet_block_samples = () - for block_sample, controlnet_block in zip(block_samples, self.controlnet_blocks): - block_sample = controlnet_block(block_sample) - controlnet_block_samples = controlnet_block_samples + (block_sample,) - - controlnet_single_block_samples = () - for single_block_sample, controlnet_block in zip(single_block_samples, self.controlnet_single_blocks): - single_block_sample = controlnet_block(single_block_sample) - controlnet_single_block_samples = controlnet_single_block_samples + (single_block_sample,) - - # scaling - controlnet_block_samples = [sample * conditioning_scale for sample in controlnet_block_samples] - controlnet_single_block_samples = [sample * conditioning_scale for sample in controlnet_single_block_samples] - - controlnet_block_samples = None if len(controlnet_block_samples) == 0 else controlnet_block_samples - controlnet_single_block_samples = ( - None if len(controlnet_single_block_samples) == 0 else controlnet_single_block_samples - ) - - if not return_dict: - return (controlnet_block_samples, controlnet_single_block_samples) - - return FluxControlNetOutput( - controlnet_block_samples=controlnet_block_samples, - controlnet_single_block_samples=controlnet_single_block_samples, - ) - - -class FluxMultiControlNetModel(ModelMixin): - r""" - `FluxMultiControlNetModel` wrapper class for Multi-FluxControlNetModel - - This module is a wrapper for multiple instances of the `FluxControlNetModel`. The `forward()` API is designed to be - compatible with `FluxControlNetModel`. - - Args: - controlnets (`list[FluxControlNetModel]`): - Provides additional conditioning to the unet during the denoising process. You must set multiple - `FluxControlNetModel` as a list. - """ - - def __init__(self, controlnets): - super().__init__() - self.nets = nn.ModuleList(controlnets) - - def forward( - self, - hidden_states: torch.FloatTensor, - controlnet_cond: list[torch.tensor], - controlnet_mode: list[torch.tensor], - conditioning_scale: list[float], - encoder_hidden_states: torch.Tensor = None, - pooled_projections: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_ids: torch.Tensor = None, - txt_ids: torch.Tensor = None, - guidance: torch.Tensor = None, - joint_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> FluxControlNetOutput | tuple: - r""" - Args: - hidden_states (`torch.FloatTensor` of shape `(batch size, channel, height, width)`): - Input `hidden_states`. - controlnet_cond (`list` of `torch.Tensor`): - A list of conditional input tensors, one per ControlNet. - controlnet_mode (`list` of `torch.Tensor`): - A list of mode tensors selecting the control type for each ControlNet. - conditioning_scale (`list` of `float`): - A list of scale factors applied to the ControlNet outputs. - encoder_hidden_states (`torch.FloatTensor` of shape `(batch size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - pooled_projections (`torch.FloatTensor` of shape `(batch_size, projection_dim)`): - Embeddings projected from the embeddings of input conditions. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - img_ids (`torch.Tensor`): - Positional ids for the image tokens. - txt_ids (`torch.Tensor`): - Positional ids for the text tokens. - guidance (`torch.Tensor`, *optional*): - Guidance scale tensor used by guidance-distilled variants of the model. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`FluxControlNetOutput`] instead of a plain tuple. - - Returns: - [`FluxControlNetOutput`] or `tuple`: - If `return_dict` is True, a [`FluxControlNetOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - # ControlNet-Union with multiple conditions - # only load one ControlNet for saving memories - if len(self.nets) == 1: - controlnet = self.nets[0] - - for i, (image, mode, scale) in enumerate(zip(controlnet_cond, controlnet_mode, conditioning_scale)): - block_samples, single_block_samples = controlnet( - hidden_states=hidden_states, - controlnet_cond=image, - controlnet_mode=mode[:, None], - conditioning_scale=scale, - timestep=timestep, - guidance=guidance, - pooled_projections=pooled_projections, - encoder_hidden_states=encoder_hidden_states, - txt_ids=txt_ids, - img_ids=img_ids, - joint_attention_kwargs=joint_attention_kwargs, - return_dict=return_dict, - ) - - # merge samples - if i == 0: - control_block_samples = block_samples - control_single_block_samples = single_block_samples - else: - if block_samples is not None and control_block_samples is not None: - control_block_samples = [ - control_block_sample + block_sample - for control_block_sample, block_sample in zip(control_block_samples, block_samples) - ] - if single_block_samples is not None and control_single_block_samples is not None: - control_single_block_samples = [ - control_single_block_sample + block_sample - for control_single_block_sample, block_sample in zip( - control_single_block_samples, single_block_samples - ) - ] - - # Regular Multi-ControlNets - # load all ControlNets into memories - else: - for i, (image, mode, scale, controlnet) in enumerate( - zip(controlnet_cond, controlnet_mode, conditioning_scale, self.nets) - ): - block_samples, single_block_samples = controlnet( - hidden_states=hidden_states, - controlnet_cond=image, - controlnet_mode=mode[:, None], - conditioning_scale=scale, - timestep=timestep, - guidance=guidance, - pooled_projections=pooled_projections, - encoder_hidden_states=encoder_hidden_states, - txt_ids=txt_ids, - img_ids=img_ids, - joint_attention_kwargs=joint_attention_kwargs, - return_dict=return_dict, - ) - - # merge samples - if i == 0: - control_block_samples = block_samples - control_single_block_samples = single_block_samples - else: - if block_samples is not None and control_block_samples is not None: - control_block_samples = [ - control_block_sample + block_sample - for control_block_sample, block_sample in zip(control_block_samples, block_samples) - ] - if single_block_samples is not None and control_single_block_samples is not None: - control_single_block_samples = [ - control_single_block_sample + block_sample - for control_single_block_sample, block_sample in zip( - control_single_block_samples, single_block_samples - ) - ] - - return control_block_samples, control_single_block_samples diff --git a/diffusers/models/controlnets/controlnet_hunyuan.py b/diffusers/models/controlnets/controlnet_hunyuan.py deleted file mode 100644 index 6ef92d78dd6e30c471c4c6f41b059be051a65a30..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/controlnet_hunyuan.py +++ /dev/null @@ -1,400 +0,0 @@ -# Copyright 2025 HunyuanDiT Authors, Qixun Wang and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from dataclasses import dataclass - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput, logging -from ..attention_processor import AttentionProcessor -from ..embeddings import ( - HunyuanCombinedTimestepTextSizeStyleEmbedding, - PatchEmbed, - PixArtAlphaTextProjection, -) -from ..modeling_utils import ModelMixin -from ..transformers.hunyuan_transformer_2d import HunyuanDiTBlock -from .controlnet import zero_module - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class HunyuanControlNetOutput(BaseOutput): - controlnet_block_samples: tuple[torch.Tensor] - - -class HunyuanDiT2DControlNetModel(ModelMixin, ConfigMixin): - @register_to_config - def __init__( - self, - conditioning_channels: int = 3, - num_attention_heads: int = 16, - attention_head_dim: int = 88, - in_channels: int | None = None, - patch_size: int | None = None, - activation_fn: str = "gelu-approximate", - sample_size=32, - hidden_size=1152, - transformer_num_layers: int = 40, - mlp_ratio: float = 4.0, - cross_attention_dim: int = 1024, - cross_attention_dim_t5: int = 2048, - pooled_projection_dim: int = 1024, - text_len: int = 77, - text_len_t5: int = 256, - use_style_cond_and_image_meta_size: bool = True, - ): - super().__init__() - self.num_heads = num_attention_heads - self.inner_dim = num_attention_heads * attention_head_dim - - self.text_embedder = PixArtAlphaTextProjection( - in_features=cross_attention_dim_t5, - hidden_size=cross_attention_dim_t5 * 4, - out_features=cross_attention_dim, - act_fn="silu_fp32", - ) - - self.text_embedding_padding = nn.Parameter( - torch.randn(text_len + text_len_t5, cross_attention_dim, dtype=torch.float32) - ) - - self.pos_embed = PatchEmbed( - height=sample_size, - width=sample_size, - in_channels=in_channels, - embed_dim=hidden_size, - patch_size=patch_size, - pos_embed_type=None, - ) - - self.time_extra_emb = HunyuanCombinedTimestepTextSizeStyleEmbedding( - hidden_size, - pooled_projection_dim=pooled_projection_dim, - seq_len=text_len_t5, - cross_attention_dim=cross_attention_dim_t5, - use_style_cond_and_image_meta_size=use_style_cond_and_image_meta_size, - ) - - # controlnet_blocks - self.controlnet_blocks = nn.ModuleList([]) - - # HunyuanDiT Blocks - self.blocks = nn.ModuleList( - [ - HunyuanDiTBlock( - dim=self.inner_dim, - num_attention_heads=self.config.num_attention_heads, - activation_fn=activation_fn, - ff_inner_dim=int(self.inner_dim * mlp_ratio), - cross_attention_dim=cross_attention_dim, - qk_norm=True, # See https://huggingface.co/papers/2302.05442 for details. - skip=False, # always False as it is the first half of the model - ) - for layer in range(transformer_num_layers // 2 - 1) - ] - ) - self.input_block = zero_module(nn.Linear(hidden_size, hidden_size)) - for _ in range(len(self.blocks)): - controlnet_block = nn.Linear(hidden_size, hidden_size) - controlnet_block = zero_module(controlnet_block) - self.controlnet_blocks.append(controlnet_block) - - @property - def attn_processors(self) -> dict[str, AttentionProcessor]: - r""" - Returns: - `dict` of attention processors: A dictionary containing all attention processors used in the model with - indexed by its weight name. - """ - # set recursively - processors = {} - - def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: dict[str, AttentionProcessor]): - if hasattr(module, "get_processor"): - processors[f"{name}.processor"] = module.get_processor(return_deprecated_lora=True) - - for sub_name, child in module.named_children(): - fn_recursive_add_processors(f"{name}.{sub_name}", child, processors) - - return processors - - for name, module in self.named_children(): - fn_recursive_add_processors(name, module, processors) - - return processors - - def set_attn_processor(self, processor: AttentionProcessor | dict[str, AttentionProcessor]): - r""" - Sets the attention processor to use to compute attention. - - Parameters: - processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`): - The instantiated processor class or a dictionary of processor classes that will be set as the processor - for **all** `Attention` layers. If `processor` is a dict, the key needs to define the path to the - corresponding cross attention processor. This is strongly recommended when setting trainable attention - processors. - """ - count = len(self.attn_processors.keys()) - - if isinstance(processor, dict) and len(processor) != count: - raise ValueError( - f"A dict of processors was passed, but the number of processors {len(processor)} does not match the" - f" number of attention layers: {count}. Please make sure to pass {count} processor classes." - ) - - def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor): - if hasattr(module, "set_processor"): - if not isinstance(processor, dict): - module.set_processor(processor) - else: - module.set_processor(processor.pop(f"{name}.processor")) - - for sub_name, child in module.named_children(): - fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor) - - for name, module in self.named_children(): - fn_recursive_attn_processor(name, module, processor) - - @classmethod - def from_transformer( - cls, transformer, conditioning_channels=3, transformer_num_layers=None, load_weights_from_transformer=True - ): - config = transformer.config - activation_fn = config.activation_fn - attention_head_dim = config.attention_head_dim - cross_attention_dim = config.cross_attention_dim - cross_attention_dim_t5 = config.cross_attention_dim_t5 - hidden_size = config.hidden_size - in_channels = config.in_channels - mlp_ratio = config.mlp_ratio - num_attention_heads = config.num_attention_heads - patch_size = config.patch_size - sample_size = config.sample_size - text_len = config.text_len - text_len_t5 = config.text_len_t5 - - conditioning_channels = conditioning_channels - transformer_num_layers = transformer_num_layers or config.transformer_num_layers - - controlnet = cls( - conditioning_channels=conditioning_channels, - transformer_num_layers=transformer_num_layers, - activation_fn=activation_fn, - attention_head_dim=attention_head_dim, - cross_attention_dim=cross_attention_dim, - cross_attention_dim_t5=cross_attention_dim_t5, - hidden_size=hidden_size, - in_channels=in_channels, - mlp_ratio=mlp_ratio, - num_attention_heads=num_attention_heads, - patch_size=patch_size, - sample_size=sample_size, - text_len=text_len, - text_len_t5=text_len_t5, - ) - if load_weights_from_transformer: - key = controlnet.load_state_dict(transformer.state_dict(), strict=False) - logger.warning(f"controlnet load from Hunyuan-DiT. missing_keys: {key[0]}") - return controlnet - - def forward( - self, - hidden_states, - timestep, - controlnet_cond: torch.Tensor, - conditioning_scale: float = 1.0, - encoder_hidden_states=None, - text_embedding_mask=None, - encoder_hidden_states_t5=None, - text_embedding_mask_t5=None, - image_meta_size=None, - style=None, - image_rotary_emb=None, - return_dict=True, - ): - """ - The [`HunyuanDiT2DControlNetModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch size, dim, height, width)`): - The input tensor. - timestep ( `torch.LongTensor`, *optional*): - Used to indicate denoising step. - controlnet_cond ( `torch.Tensor` ): - The conditioning input to ControlNet. - conditioning_scale ( `float` ): - Indicate the conditioning scale. - encoder_hidden_states ( `torch.Tensor` of shape `(batch size, sequence len, embed dims)`, *optional*): - Conditional embeddings for cross attention layer. This is the output of `BertModel`. - text_embedding_mask: torch.Tensor - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. This is the output - of `BertModel`. - encoder_hidden_states_t5 ( `torch.Tensor` of shape `(batch size, sequence len, embed dims)`, *optional*): - Conditional embeddings for cross attention layer. This is the output of T5 Text Encoder. - text_embedding_mask_t5: torch.Tensor - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. This is the output - of T5 Text Encoder. - image_meta_size (torch.Tensor): - Conditional embedding indicate the image sizes - style: torch.Tensor: - Conditional embedding indicate the style - image_rotary_emb (`torch.Tensor`): - The image rotary embeddings to apply on query and key tensors during attention calculation. - return_dict: bool - Whether to return a dictionary. - """ - - height, width = hidden_states.shape[-2:] - - hidden_states = self.pos_embed(hidden_states) # b,c,H,W -> b, N, C - - # 2. pre-process - hidden_states = hidden_states + self.input_block(self.pos_embed(controlnet_cond)) - - temb = self.time_extra_emb( - timestep, encoder_hidden_states_t5, image_meta_size, style, hidden_dtype=timestep.dtype - ) # [B, D] - - # text projection - batch_size, sequence_length, _ = encoder_hidden_states_t5.shape - encoder_hidden_states_t5 = self.text_embedder( - encoder_hidden_states_t5.view(-1, encoder_hidden_states_t5.shape[-1]) - ) - encoder_hidden_states_t5 = encoder_hidden_states_t5.view(batch_size, sequence_length, -1) - - encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states_t5], dim=1) - text_embedding_mask = torch.cat([text_embedding_mask, text_embedding_mask_t5], dim=-1) - text_embedding_mask = text_embedding_mask.unsqueeze(2).bool() - - encoder_hidden_states = torch.where(text_embedding_mask, encoder_hidden_states, self.text_embedding_padding) - - block_res_samples = () - for layer, block in enumerate(self.blocks): - hidden_states = block( - hidden_states, - temb=temb, - encoder_hidden_states=encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - ) # (N, L, D) - - block_res_samples = block_res_samples + (hidden_states,) - - controlnet_block_res_samples = () - for block_res_sample, controlnet_block in zip(block_res_samples, self.controlnet_blocks): - block_res_sample = controlnet_block(block_res_sample) - controlnet_block_res_samples = controlnet_block_res_samples + (block_res_sample,) - - # 6. scaling - controlnet_block_res_samples = [sample * conditioning_scale for sample in controlnet_block_res_samples] - - if not return_dict: - return (controlnet_block_res_samples,) - - return HunyuanControlNetOutput(controlnet_block_samples=controlnet_block_res_samples) - - -class HunyuanDiT2DMultiControlNetModel(ModelMixin): - r""" - `HunyuanDiT2DMultiControlNetModel` wrapper class for Multi-HunyuanDiT2DControlNetModel - - This module is a wrapper for multiple instances of the `HunyuanDiT2DControlNetModel`. The `forward()` API is - designed to be compatible with `HunyuanDiT2DControlNetModel`. - - Args: - controlnets (`list[HunyuanDiT2DControlNetModel]`): - Provides additional conditioning to the unet during the denoising process. You must set multiple - `HunyuanDiT2DControlNetModel` as a list. - """ - - def __init__(self, controlnets): - super().__init__() - self.nets = nn.ModuleList(controlnets) - - def forward( - self, - hidden_states, - timestep, - controlnet_cond: torch.Tensor, - conditioning_scale: float = 1.0, - encoder_hidden_states=None, - text_embedding_mask=None, - encoder_hidden_states_t5=None, - text_embedding_mask_t5=None, - image_meta_size=None, - style=None, - image_rotary_emb=None, - return_dict=True, - ): - """ - The [`HunyuanDiT2DControlNetModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch size, dim, height, width)`): - The input tensor. - timestep ( `torch.LongTensor`, *optional*): - Used to indicate denoising step. - controlnet_cond ( `torch.Tensor` ): - The conditioning input to ControlNet. - conditioning_scale ( `float` ): - Indicate the conditioning scale. - encoder_hidden_states ( `torch.Tensor` of shape `(batch size, sequence len, embed dims)`, *optional*): - Conditional embeddings for cross attention layer. This is the output of `BertModel`. - text_embedding_mask: torch.Tensor - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. This is the output - of `BertModel`. - encoder_hidden_states_t5 ( `torch.Tensor` of shape `(batch size, sequence len, embed dims)`, *optional*): - Conditional embeddings for cross attention layer. This is the output of T5 Text Encoder. - text_embedding_mask_t5: torch.Tensor - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. This is the output - of T5 Text Encoder. - image_meta_size (torch.Tensor): - Conditional embedding indicate the image sizes - style: torch.Tensor: - Conditional embedding indicate the style - image_rotary_emb (`torch.Tensor`): - The image rotary embeddings to apply on query and key tensors during attention calculation. - return_dict: bool - Whether to return a dictionary. - """ - for i, (image, scale, controlnet) in enumerate(zip(controlnet_cond, conditioning_scale, self.nets)): - block_samples = controlnet( - hidden_states=hidden_states, - timestep=timestep, - controlnet_cond=image, - conditioning_scale=scale, - encoder_hidden_states=encoder_hidden_states, - text_embedding_mask=text_embedding_mask, - encoder_hidden_states_t5=encoder_hidden_states_t5, - text_embedding_mask_t5=text_embedding_mask_t5, - image_meta_size=image_meta_size, - style=style, - image_rotary_emb=image_rotary_emb, - return_dict=return_dict, - ) - - # merge samples - if i == 0: - control_block_samples = block_samples - else: - control_block_samples = [ - control_block_sample + block_sample - for control_block_sample, block_sample in zip(control_block_samples[0], block_samples[0]) - ] - control_block_samples = (control_block_samples,) - - return control_block_samples diff --git a/diffusers/models/controlnets/controlnet_qwenimage.py b/diffusers/models/controlnets/controlnet_qwenimage.py deleted file mode 100644 index f721c51261e106aafeb2b1f7aa0367bea91f0d71..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/controlnet_qwenimage.py +++ /dev/null @@ -1,350 +0,0 @@ -# Copyright 2025 Black Forest Labs, The HuggingFace Team and The InstantX Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import ( - BaseOutput, - apply_lora_scale, - deprecate, - logging, -) -from ..attention import AttentionMixin -from ..cache_utils import CacheMixin -from ..controlnets.controlnet import zero_module -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..transformers.transformer_qwenimage import ( - QwenEmbedRope, - QwenImageTransformerBlock, - QwenTimestepProjEmbeddings, - RMSNorm, - compute_text_seq_len_from_mask, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class QwenImageControlNetOutput(BaseOutput): - controlnet_block_samples: tuple[torch.Tensor] - - -class QwenImageControlNetModel( - ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin -): - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - patch_size: int = 2, - in_channels: int = 64, - out_channels: int | None = 16, - num_layers: int = 60, - attention_head_dim: int = 128, - num_attention_heads: int = 24, - joint_attention_dim: int = 3584, - axes_dims_rope: tuple[int, int, int] = (16, 56, 56), - extra_condition_channels: int = 0, # for controlnet-inpainting - ): - super().__init__() - self.out_channels = out_channels or in_channels - self.inner_dim = num_attention_heads * attention_head_dim - - self.pos_embed = QwenEmbedRope(theta=10000, axes_dim=list(axes_dims_rope), scale_rope=True) - - self.time_text_embed = QwenTimestepProjEmbeddings(embedding_dim=self.inner_dim) - - self.txt_norm = RMSNorm(joint_attention_dim, eps=1e-6) - - self.img_in = nn.Linear(in_channels, self.inner_dim) - self.txt_in = nn.Linear(joint_attention_dim, self.inner_dim) - - self.transformer_blocks = nn.ModuleList( - [ - QwenImageTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for _ in range(num_layers) - ] - ) - - # controlnet_blocks - self.controlnet_blocks = nn.ModuleList([]) - for _ in range(len(self.transformer_blocks)): - self.controlnet_blocks.append(zero_module(nn.Linear(self.inner_dim, self.inner_dim))) - self.controlnet_x_embedder = zero_module( - torch.nn.Linear(in_channels + extra_condition_channels, self.inner_dim) - ) - - self.gradient_checkpointing = False - - @classmethod - def from_transformer( - cls, - transformer, - num_layers: int = 5, - attention_head_dim: int = 128, - num_attention_heads: int = 24, - load_weights_from_transformer=True, - extra_condition_channels: int = 0, - ): - config = dict(transformer.config) - config["num_layers"] = num_layers - config["attention_head_dim"] = attention_head_dim - config["num_attention_heads"] = num_attention_heads - config["extra_condition_channels"] = extra_condition_channels - - controlnet = cls.from_config(config) - - if load_weights_from_transformer: - controlnet.pos_embed.load_state_dict(transformer.pos_embed.state_dict()) - controlnet.time_text_embed.load_state_dict(transformer.time_text_embed.state_dict()) - controlnet.img_in.load_state_dict(transformer.img_in.state_dict()) - controlnet.txt_in.load_state_dict(transformer.txt_in.state_dict()) - controlnet.transformer_blocks.load_state_dict(transformer.transformer_blocks.state_dict(), strict=False) - controlnet.controlnet_x_embedder = zero_module(controlnet.controlnet_x_embedder) - - return controlnet - - @apply_lora_scale("joint_attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - controlnet_cond: torch.Tensor, - conditioning_scale: float = 1.0, - encoder_hidden_states: torch.Tensor = None, - encoder_hidden_states_mask: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_shapes: list[tuple[int, int, int]] | None = None, - txt_seq_lens: list[int] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> torch.FloatTensor | Transformer2DModelOutput: - """ - The [`QwenImageControlNetModel`] forward method. - - Args: - hidden_states (`torch.FloatTensor` of shape `(batch size, channel, height, width)`): - Input `hidden_states`. - controlnet_cond (`torch.Tensor`): - The conditional input tensor of shape `(batch_size, sequence_length, hidden_size)`. - conditioning_scale (`float`, defaults to `1.0`): - The scale factor for ControlNet outputs. - encoder_hidden_states (`torch.FloatTensor` of shape `(batch size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_hidden_states_mask (`torch.Tensor` of shape `(batch_size, text_sequence_length)`, *optional*): - Mask for the encoder hidden states. Expected to have 1.0 for valid tokens and 0.0 for padding tokens. - Used in the attention processor to prevent attending to padding tokens. The mask can have any pattern - (not just contiguous valid tokens followed by padding) since it's applied element-wise in attention. - timestep ( `torch.LongTensor`): - Used to indicate denoising step. - img_shapes (`list[tuple[int, int, int]]`, *optional*): - Image shapes for RoPE computation. - txt_seq_lens (`list[int]`, *optional*): - **Deprecated**. Not needed anymore, we use `encoder_hidden_states` instead to infer text sequence - length. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.controlnet.ControlNetOutput`] instead of a plain tuple. - - Returns: - If `return_dict` is True, a [`~models.controlnet.ControlNetOutput`] is returned, otherwise a `tuple` where - the first element is the controlnet block samples. - """ - # Handle deprecated txt_seq_lens parameter - if txt_seq_lens is not None: - deprecate( - "txt_seq_lens", - "0.39.0", - "Passing `txt_seq_lens` to `QwenImageControlNetModel.forward()` is deprecated and will be removed in " - "version 0.39.0. The text sequence length is now automatically inferred from `encoder_hidden_states` " - "and `encoder_hidden_states_mask`.", - standard_warn=False, - ) - - hidden_states = self.img_in(hidden_states) - - # add - hidden_states = hidden_states + self.controlnet_x_embedder(controlnet_cond) - - temb = self.time_text_embed(timestep, hidden_states) - - # Use the encoder_hidden_states sequence length for RoPE computation and normalize mask - text_seq_len, _, encoder_hidden_states_mask = compute_text_seq_len_from_mask( - encoder_hidden_states, encoder_hidden_states_mask - ) - - image_rotary_emb = self.pos_embed(img_shapes, max_txt_seq_len=text_seq_len, device=hidden_states.device) - - timestep = timestep.to(hidden_states.dtype) - encoder_hidden_states = self.txt_norm(encoder_hidden_states) - encoder_hidden_states = self.txt_in(encoder_hidden_states) - - block_samples = () - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - encoder_hidden_states_mask, - temb, - image_rotary_emb, - joint_attention_kwargs, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - encoder_hidden_states_mask=encoder_hidden_states_mask, - temb=temb, - image_rotary_emb=image_rotary_emb, - joint_attention_kwargs=joint_attention_kwargs, - ) - block_samples = block_samples + (hidden_states,) - - # controlnet block - controlnet_block_samples = () - for block_sample, controlnet_block in zip(block_samples, self.controlnet_blocks): - block_sample = controlnet_block(block_sample) - controlnet_block_samples = controlnet_block_samples + (block_sample,) - - # scaling - controlnet_block_samples = [sample * conditioning_scale for sample in controlnet_block_samples] - controlnet_block_samples = None if len(controlnet_block_samples) == 0 else controlnet_block_samples - - if not return_dict: - return controlnet_block_samples - - return QwenImageControlNetOutput( - controlnet_block_samples=controlnet_block_samples, - ) - - -class QwenImageMultiControlNetModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin): - r""" - `QwenImageMultiControlNetModel` wrapper class for Multi-QwenImageControlNetModel - - This module is a wrapper for multiple instances of the `QwenImageControlNetModel`. The `forward()` API is designed - to be compatible with `QwenImageControlNetModel`. - - Args: - controlnets (`list[QwenImageControlNetModel]`): - Provides additional conditioning to the unet during the denoising process. You must set multiple - `QwenImageControlNetModel` as a list. - """ - - def __init__(self, controlnets): - super().__init__() - self.nets = nn.ModuleList(controlnets) - - def forward( - self, - hidden_states: torch.FloatTensor, - controlnet_cond: list[torch.tensor], - conditioning_scale: list[float], - encoder_hidden_states: torch.Tensor = None, - encoder_hidden_states_mask: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_shapes: list[tuple[int, int, int]] | None = None, - txt_seq_lens: list[int] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> QwenImageControlNetOutput | tuple: - r""" - Args: - hidden_states (`torch.FloatTensor`): - Input `hidden_states`. - controlnet_cond (`list` of `torch.Tensor`): - A list of conditional input tensors, one per ControlNet. - conditioning_scale (`list` of `float`): - A list of scale factors applied to the ControlNet outputs. - encoder_hidden_states (`torch.Tensor`, *optional*): - Conditional embeddings (embeddings computed from the input conditions such as prompts). - encoder_hidden_states_mask (`torch.Tensor`, *optional*): - Mask for the encoder hidden states. - timestep (`torch.LongTensor`, *optional*): - Used to indicate denoising step. - img_shapes (`list` of `tuple[int, int, int]`, *optional*): - Per-sample image shapes used to construct positional encodings. - txt_seq_lens (`list` of `int`, *optional*): - Deprecated. The text sequence length is now inferred from `encoder_hidden_states` and - `encoder_hidden_states_mask`. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`QwenImageControlNetOutput`] instead of a plain tuple. - - Returns: - [`QwenImageControlNetOutput`] or `tuple`: - If `return_dict` is True, a [`QwenImageControlNetOutput`] is returned, otherwise a plain `tuple` is - returned. - """ - if txt_seq_lens is not None: - deprecate( - "txt_seq_lens", - "0.39.0", - "Passing `txt_seq_lens` to `QwenImageMultiControlNetModel.forward()` is deprecated and will be " - "removed in version 0.39.0. The text sequence length is now automatically inferred from " - "`encoder_hidden_states` and `encoder_hidden_states_mask`.", - standard_warn=False, - ) - # ControlNet-Union with multiple conditions - # only load one ControlNet for saving memories - if len(self.nets) == 1: - controlnet = self.nets[0] - - for i, (image, scale) in enumerate(zip(controlnet_cond, conditioning_scale)): - block_samples = controlnet( - hidden_states=hidden_states, - controlnet_cond=image, - conditioning_scale=scale, - encoder_hidden_states=encoder_hidden_states, - encoder_hidden_states_mask=encoder_hidden_states_mask, - timestep=timestep, - img_shapes=img_shapes, - joint_attention_kwargs=joint_attention_kwargs, - return_dict=return_dict, - ) - - # merge samples - if i == 0: - control_block_samples = block_samples - else: - if block_samples is not None and control_block_samples is not None: - control_block_samples = [ - control_block_sample + block_sample - for control_block_sample, block_sample in zip(control_block_samples, block_samples) - ] - else: - raise ValueError("QwenImageMultiControlNetModel only supports a single controlnet-union now.") - - return control_block_samples diff --git a/diffusers/models/controlnets/controlnet_sana.py b/diffusers/models/controlnets/controlnet_sana.py deleted file mode 100644 index 4b6e3010ec67c5dfbc2a41711662f99a65e8fc17..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/controlnet_sana.py +++ /dev/null @@ -1,241 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass -from typing import Any - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import BaseOutput, apply_lora_scale, logging -from ..attention import AttentionMixin -from ..embeddings import PatchEmbed, PixArtAlphaTextProjection -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormSingle, RMSNorm -from ..transformers.sana_transformer import SanaTransformerBlock -from .controlnet import zero_module - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class SanaControlNetOutput(BaseOutput): - controlnet_block_samples: tuple[torch.Tensor] - - -class SanaControlNetModel(ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin): - _supports_gradient_checkpointing = True - _no_split_modules = ["SanaTransformerBlock", "PatchEmbed"] - _skip_layerwise_casting_patterns = ["patch_embed", "norm"] - - @register_to_config - def __init__( - self, - in_channels: int = 32, - out_channels: int | None = 32, - num_attention_heads: int = 70, - attention_head_dim: int = 32, - num_layers: int = 7, - num_cross_attention_heads: int | None = 20, - cross_attention_head_dim: int | None = 112, - cross_attention_dim: int | None = 2240, - caption_channels: int = 2304, - mlp_ratio: float = 2.5, - dropout: float = 0.0, - attention_bias: bool = False, - sample_size: int = 32, - patch_size: int = 1, - norm_elementwise_affine: bool = False, - norm_eps: float = 1e-6, - interpolation_scale: int | None = None, - ) -> None: - super().__init__() - - out_channels = out_channels or in_channels - inner_dim = num_attention_heads * attention_head_dim - - # 1. Patch Embedding - self.patch_embed = PatchEmbed( - height=sample_size, - width=sample_size, - patch_size=patch_size, - in_channels=in_channels, - embed_dim=inner_dim, - interpolation_scale=interpolation_scale, - pos_embed_type="sincos" if interpolation_scale is not None else None, - ) - - # 2. Additional condition embeddings - self.time_embed = AdaLayerNormSingle(inner_dim) - - self.caption_projection = PixArtAlphaTextProjection(in_features=caption_channels, hidden_size=inner_dim) - self.caption_norm = RMSNorm(inner_dim, eps=1e-5, elementwise_affine=True) - - # 3. Transformer blocks - self.transformer_blocks = nn.ModuleList( - [ - SanaTransformerBlock( - inner_dim, - num_attention_heads, - attention_head_dim, - dropout=dropout, - num_cross_attention_heads=num_cross_attention_heads, - cross_attention_head_dim=cross_attention_head_dim, - cross_attention_dim=cross_attention_dim, - attention_bias=attention_bias, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - mlp_ratio=mlp_ratio, - ) - for _ in range(num_layers) - ] - ) - - # controlnet_blocks - self.controlnet_blocks = nn.ModuleList([]) - - self.input_block = zero_module(nn.Linear(inner_dim, inner_dim)) - for _ in range(len(self.transformer_blocks)): - controlnet_block = nn.Linear(inner_dim, inner_dim) - controlnet_block = zero_module(controlnet_block) - self.controlnet_blocks.append(controlnet_block) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: torch.LongTensor, - controlnet_cond: torch.Tensor, - conditioning_scale: float = 1.0, - encoder_attention_mask: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor, ...] | Transformer2DModelOutput: - r""" - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, channel, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - controlnet_cond (`torch.Tensor`): - The conditional input tensor for the ControlNet. - conditioning_scale (`float`, *optional*, defaults to `1.0`): - The scale factor for ControlNet outputs. - encoder_attention_mask (`torch.Tensor`, *optional*): - Attention mask applied to `encoder_hidden_states`. - attention_mask (`torch.Tensor`, *optional*): - Attention mask applied to `hidden_states`. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - [`~models.transformer_2d.Transformer2DModelOutput`] or `tuple`: - If `return_dict` is True, a [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise - a plain `tuple` is returned. - """ - # ensure attention_mask is a bias, and give it a singleton query_tokens dimension. - # we may have done this conversion already, e.g. if we came here via UNet2DConditionModel#forward. - # we can tell by counting dims; if ndim == 2: it's a mask rather than a bias. - # expects mask of shape: - # [batch, key_tokens] - # adds singleton query_tokens dimension: - # [batch, 1, key_tokens] - # this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes: - # [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn) - # [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn) - if attention_mask is not None and attention_mask.ndim == 2: - # assume that mask is expressed as: - # (1 = keep, 0 = discard) - # convert mask into a bias that can be added to attention scores: - # (keep = +0, discard = -10000.0) - attention_mask = (1 - attention_mask.to(hidden_states.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # convert encoder_attention_mask to a bias the same way we do for attention_mask - if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2: - encoder_attention_mask = (1 - encoder_attention_mask.to(hidden_states.dtype)) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - # 1. Input - batch_size, num_channels, height, width = hidden_states.shape - p = self.config.patch_size - post_patch_height, post_patch_width = height // p, width // p - - hidden_states = self.patch_embed(hidden_states) - hidden_states = hidden_states + self.input_block(self.patch_embed(controlnet_cond.to(hidden_states.dtype))) - - timestep, embedded_timestep = self.time_embed( - timestep, batch_size=batch_size, hidden_dtype=hidden_states.dtype - ) - - encoder_hidden_states = self.caption_projection(encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states.view(batch_size, -1, hidden_states.shape[-1]) - - encoder_hidden_states = self.caption_norm(encoder_hidden_states) - - # 2. Transformer blocks - block_res_samples = () - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.transformer_blocks: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - attention_mask, - encoder_hidden_states, - encoder_attention_mask, - timestep, - post_patch_height, - post_patch_width, - ) - block_res_samples = block_res_samples + (hidden_states,) - else: - for block in self.transformer_blocks: - hidden_states = block( - hidden_states, - attention_mask, - encoder_hidden_states, - encoder_attention_mask, - timestep, - post_patch_height, - post_patch_width, - ) - block_res_samples = block_res_samples + (hidden_states,) - - # 3. ControlNet blocks - controlnet_block_res_samples = () - for block_res_sample, controlnet_block in zip(block_res_samples, self.controlnet_blocks): - block_res_sample = controlnet_block(block_res_sample) - controlnet_block_res_samples = controlnet_block_res_samples + (block_res_sample,) - - controlnet_block_res_samples = [sample * conditioning_scale for sample in controlnet_block_res_samples] - - if not return_dict: - return (controlnet_block_res_samples,) - - return SanaControlNetOutput(controlnet_block_samples=controlnet_block_res_samples) diff --git a/diffusers/models/controlnets/controlnet_sd3.py b/diffusers/models/controlnets/controlnet_sd3.py deleted file mode 100644 index 1f0ca529ff16a6f78d97904c2855bff4b12d50fb..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/controlnet_sd3.py +++ /dev/null @@ -1,452 +0,0 @@ -# Copyright 2025 Stability AI, The HuggingFace Team and The InstantX Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from dataclasses import dataclass -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ..attention import AttentionMixin, JointTransformerBlock -from ..attention_processor import Attention, FusedJointAttnProcessor2_0 -from ..embeddings import CombinedTimestepTextProjEmbeddings, PatchEmbed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..transformers.transformer_sd3 import SD3SingleTransformerBlock -from .controlnet import BaseOutput, zero_module - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class SD3ControlNetOutput(BaseOutput): - controlnet_block_samples: tuple[torch.Tensor] - - -class SD3ControlNetModel(ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): - r""" - ControlNet model for [Stable Diffusion 3](https://huggingface.co/papers/2403.03206). - - Parameters: - sample_size (`int`, defaults to `128`): - The width/height of the latents. This is fixed during training since it is used to learn a number of - position embeddings. - patch_size (`int`, defaults to `2`): - Patch size to turn the input data into small patches. - in_channels (`int`, defaults to `16`): - The number of latent channels in the input. - num_layers (`int`, defaults to `18`): - The number of layers of transformer blocks to use. - attention_head_dim (`int`, defaults to `64`): - The number of channels in each head. - num_attention_heads (`int`, defaults to `18`): - The number of heads to use for multi-head attention. - joint_attention_dim (`int`, defaults to `4096`): - The embedding dimension to use for joint text-image attention. - caption_projection_dim (`int`, defaults to `1152`): - The embedding dimension of caption embeddings. - pooled_projection_dim (`int`, defaults to `2048`): - The embedding dimension of pooled text projections. - out_channels (`int`, defaults to `16`): - The number of latent channels in the output. - pos_embed_max_size (`int`, defaults to `96`): - The maximum latent height/width of positional embeddings. - extra_conditioning_channels (`int`, defaults to `0`): - The number of extra channels to use for conditioning for patch embedding. - dual_attention_layers (`tuple[int, ...]`, defaults to `()`): - The number of dual-stream transformer blocks to use. - qk_norm (`str`, *optional*, defaults to `None`): - The normalization to use for query and key in the attention layer. If `None`, no normalization is used. - pos_embed_type (`str`, defaults to `"sincos"`): - The type of positional embedding to use. Choose between `"sincos"` and `None`. - use_pos_embed (`bool`, defaults to `True`): - Whether to use positional embeddings. - force_zeros_for_pooled_projection (`bool`, defaults to `True`): - Whether to force zeros for pooled projection embeddings. This is handled in the pipelines by reading the - config value of the ControlNet model. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - sample_size: int = 128, - patch_size: int = 2, - in_channels: int = 16, - num_layers: int = 18, - attention_head_dim: int = 64, - num_attention_heads: int = 18, - joint_attention_dim: int = 4096, - caption_projection_dim: int = 1152, - pooled_projection_dim: int = 2048, - out_channels: int = 16, - pos_embed_max_size: int = 96, - extra_conditioning_channels: int = 0, - dual_attention_layers: tuple[int, ...] = (), - qk_norm: str | None = None, - pos_embed_type: str | None = "sincos", - use_pos_embed: bool = True, - force_zeros_for_pooled_projection: bool = True, - ): - super().__init__() - default_out_channels = in_channels - self.out_channels = out_channels if out_channels is not None else default_out_channels - self.inner_dim = num_attention_heads * attention_head_dim - - if use_pos_embed: - self.pos_embed = PatchEmbed( - height=sample_size, - width=sample_size, - patch_size=patch_size, - in_channels=in_channels, - embed_dim=self.inner_dim, - pos_embed_max_size=pos_embed_max_size, - pos_embed_type=pos_embed_type, - ) - else: - self.pos_embed = None - self.time_text_embed = CombinedTimestepTextProjEmbeddings( - embedding_dim=self.inner_dim, pooled_projection_dim=pooled_projection_dim - ) - if joint_attention_dim is not None: - self.context_embedder = nn.Linear(joint_attention_dim, caption_projection_dim) - - # `attention_head_dim` is doubled to account for the mixing. - # It needs to crafted when we get the actual checkpoints. - self.transformer_blocks = nn.ModuleList( - [ - JointTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - context_pre_only=False, - qk_norm=qk_norm, - use_dual_attention=True if i in dual_attention_layers else False, - ) - for i in range(num_layers) - ] - ) - else: - self.context_embedder = None - self.transformer_blocks = nn.ModuleList( - [ - SD3SingleTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for _ in range(num_layers) - ] - ) - - # controlnet_blocks - self.controlnet_blocks = nn.ModuleList([]) - for _ in range(len(self.transformer_blocks)): - controlnet_block = nn.Linear(self.inner_dim, self.inner_dim) - controlnet_block = zero_module(controlnet_block) - self.controlnet_blocks.append(controlnet_block) - pos_embed_input = PatchEmbed( - height=sample_size, - width=sample_size, - patch_size=patch_size, - in_channels=in_channels + extra_conditioning_channels, - embed_dim=self.inner_dim, - pos_embed_type=None, - ) - self.pos_embed_input = zero_module(pos_embed_input) - - self.gradient_checkpointing = False - - # Copied from diffusers.models.unets.unet_3d_condition.UNet3DConditionModel.enable_forward_chunking - def enable_forward_chunking(self, chunk_size: int | None = None, dim: int = 0) -> None: - """ - Sets the attention processor to use [feed forward - chunking](https://huggingface.co/blog/reformer#2-chunked-feed-forward-layers). - - Parameters: - chunk_size (`int`, *optional*): - The chunk size of the feed-forward layers. If not specified, will run feed-forward layer individually - over each tensor of dim=`dim`. - dim (`int`, *optional*, defaults to `0`): - The dimension over which the feed-forward computation should be chunked. Choose between dim=0 (batch) - or dim=1 (sequence length). - """ - if dim not in [0, 1]: - raise ValueError(f"Make sure to set `dim` to either 0 or 1, not {dim}") - - # By default chunk size is 1 - chunk_size = chunk_size or 1 - - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, chunk_size, dim) - - # Copied from diffusers.models.transformers.transformer_sd3.SD3Transformer2DModel.fuse_qkv_projections - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedJointAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) - - # Notes: This is for SD3.5 8b controlnet, which shares the pos_embed with the transformer - # we should have handled this in conversion script - def _get_pos_embed_from_transformer(self, transformer): - pos_embed = PatchEmbed( - height=transformer.config.sample_size, - width=transformer.config.sample_size, - patch_size=transformer.config.patch_size, - in_channels=transformer.config.in_channels, - embed_dim=transformer.inner_dim, - pos_embed_max_size=transformer.config.pos_embed_max_size, - ) - pos_embed.load_state_dict(transformer.pos_embed.state_dict(), strict=True) - return pos_embed - - @classmethod - def from_transformer( - cls, transformer, num_layers=12, num_extra_conditioning_channels=1, load_weights_from_transformer=True - ): - config = transformer.config - config["num_layers"] = num_layers or config.num_layers - config["extra_conditioning_channels"] = num_extra_conditioning_channels - controlnet = cls.from_config(config) - - if load_weights_from_transformer: - controlnet.pos_embed.load_state_dict(transformer.pos_embed.state_dict()) - controlnet.time_text_embed.load_state_dict(transformer.time_text_embed.state_dict()) - controlnet.context_embedder.load_state_dict(transformer.context_embedder.state_dict()) - controlnet.transformer_blocks.load_state_dict(transformer.transformer_blocks.state_dict(), strict=False) - - controlnet.pos_embed_input = zero_module(controlnet.pos_embed_input) - - return controlnet - - @apply_lora_scale("joint_attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - controlnet_cond: torch.Tensor, - conditioning_scale: float = 1.0, - encoder_hidden_states: torch.Tensor = None, - pooled_projections: torch.Tensor = None, - timestep: torch.LongTensor = None, - joint_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> torch.Tensor | Transformer2DModelOutput: - """ - The [`SD3Transformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch size, channel, height, width)`): - Input `hidden_states`. - controlnet_cond (`torch.Tensor`): - The conditional input tensor of shape `(batch_size, sequence_length, hidden_size)`. - conditioning_scale (`float`, defaults to `1.0`): - The scale factor for ControlNet outputs. - encoder_hidden_states (`torch.Tensor` of shape `(batch size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - pooled_projections (`torch.Tensor` of shape `(batch_size, projection_dim)`): Embeddings projected - from the embeddings of input conditions. - timestep ( `torch.LongTensor`): - Used to indicate denoising step. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - if self.pos_embed is not None and hidden_states.ndim != 4: - raise ValueError("hidden_states must be 4D when pos_embed is used") - - # SD3.5 8b controlnet does not have a `pos_embed`, - # it use the `pos_embed` from the transformer to process input before passing to controlnet - elif self.pos_embed is None and hidden_states.ndim != 3: - raise ValueError("hidden_states must be 3D when pos_embed is not used") - - if self.context_embedder is not None and encoder_hidden_states is None: - raise ValueError("encoder_hidden_states must be provided when context_embedder is used") - # SD3.5 8b controlnet does not have a `context_embedder`, it does not use `encoder_hidden_states` - elif self.context_embedder is None and encoder_hidden_states is not None: - raise ValueError("encoder_hidden_states should not be provided when context_embedder is not used") - - if self.pos_embed is not None: - hidden_states = self.pos_embed(hidden_states) # takes care of adding positional embeddings too. - - temb = self.time_text_embed(timestep, pooled_projections) - - if self.context_embedder is not None: - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - - # add - hidden_states = hidden_states + self.pos_embed_input(controlnet_cond) - - block_res_samples = () - - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - if self.context_embedder is not None: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - ) - else: - # SD3.5 8b controlnet use single transformer block, which does not use `encoder_hidden_states` - hidden_states = self._gradient_checkpointing_func(block, hidden_states, temb) - - else: - if self.context_embedder is not None: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, encoder_hidden_states=encoder_hidden_states, temb=temb - ) - else: - # SD3.5 8b controlnet use single transformer block, which does not use `encoder_hidden_states` - hidden_states = block(hidden_states, temb) - - block_res_samples = block_res_samples + (hidden_states,) - - controlnet_block_res_samples = () - for block_res_sample, controlnet_block in zip(block_res_samples, self.controlnet_blocks): - block_res_sample = controlnet_block(block_res_sample) - controlnet_block_res_samples = controlnet_block_res_samples + (block_res_sample,) - - # 6. scaling - controlnet_block_res_samples = [sample * conditioning_scale for sample in controlnet_block_res_samples] - - if not return_dict: - return (controlnet_block_res_samples,) - - return SD3ControlNetOutput(controlnet_block_samples=controlnet_block_res_samples) - - -class SD3MultiControlNetModel(ModelMixin): - r""" - `SD3ControlNetModel` wrapper class for Multi-SD3ControlNet - - This module is a wrapper for multiple instances of the `SD3ControlNetModel`. The `forward()` API is designed to be - compatible with `SD3ControlNetModel`. - - Args: - controlnets (`list[SD3ControlNetModel]`): - Provides additional conditioning to the unet during the denoising process. You must set multiple - `SD3ControlNetModel` as a list. - """ - - def __init__(self, controlnets): - super().__init__() - self.nets = nn.ModuleList(controlnets) - - def forward( - self, - hidden_states: torch.Tensor, - controlnet_cond: list[torch.tensor], - conditioning_scale: list[float], - pooled_projections: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - timestep: torch.LongTensor = None, - joint_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> SD3ControlNetOutput | tuple: - r""" - Args: - hidden_states (`torch.Tensor`): - Input `hidden_states`. - controlnet_cond (`list` of `torch.Tensor`): - A list of conditional input tensors, one per ControlNet. - conditioning_scale (`list` of `float`): - A list of scale factors applied to the ControlNet outputs. - pooled_projections (`torch.Tensor`): - Embeddings projected from the embeddings of input conditions. - encoder_hidden_states (`torch.Tensor`, *optional*): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`, *optional*): - Used to indicate denoising step. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`SD3ControlNetOutput`] instead of a plain tuple. - - Returns: - [`SD3ControlNetOutput`] or `tuple`: - If `return_dict` is True, a [`SD3ControlNetOutput`] is returned, otherwise a plain `tuple` is returned. - """ - for i, (image, scale, controlnet) in enumerate(zip(controlnet_cond, conditioning_scale, self.nets)): - block_samples = controlnet( - hidden_states=hidden_states, - timestep=timestep, - encoder_hidden_states=encoder_hidden_states, - pooled_projections=pooled_projections, - controlnet_cond=image, - conditioning_scale=scale, - joint_attention_kwargs=joint_attention_kwargs, - return_dict=return_dict, - ) - - # merge samples - if i == 0: - control_block_samples = block_samples - else: - control_block_samples = [ - control_block_sample + block_sample - for control_block_sample, block_sample in zip(control_block_samples[0], block_samples[0]) - ] - control_block_samples = (tuple(control_block_samples),) - - return control_block_samples diff --git a/diffusers/models/controlnets/controlnet_sparsectrl.py b/diffusers/models/controlnets/controlnet_sparsectrl.py deleted file mode 100644 index 55ff7cbdedc00ac5eac32c319b00c3369ada2b69..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/controlnet_sparsectrl.py +++ /dev/null @@ -1,721 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass -from typing import Any - -import torch -from torch import nn -from torch.nn import functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin -from ...utils import BaseOutput, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..attention import AttentionMixin -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - AttnAddedKVProcessor, - AttnProcessor, -) -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from ..unets.unet_2d_blocks import UNetMidBlock2DCrossAttn -from ..unets.unet_2d_condition import UNet2DConditionModel -from ..unets.unet_motion_model import CrossAttnDownBlockMotion, DownBlockMotion - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class SparseControlNetOutput(BaseOutput): - """ - The output of [`SparseControlNetModel`]. - - Args: - down_block_res_samples (`tuple[torch.Tensor]`): - A tuple of downsample activations at different resolutions for each downsampling block. Each tensor should - be of shape `(batch_size, channel * resolution, height //resolution, width // resolution)`. Output can be - used to condition the original UNet's downsampling activations. - mid_down_block_re_sample (`torch.Tensor`): - The activation of the middle block (the lowest sample resolution). Each tensor should be of shape - `(batch_size, channel * lowest_resolution, height // lowest_resolution, width // lowest_resolution)`. - Output can be used to condition the original UNet's middle block activation. - """ - - down_block_res_samples: tuple[torch.Tensor] - mid_block_res_sample: torch.Tensor - - -class SparseControlNetConditioningEmbedding(nn.Module): - def __init__( - self, - conditioning_embedding_channels: int, - conditioning_channels: int = 3, - block_out_channels: tuple[int, ...] = (16, 32, 96, 256), - ): - super().__init__() - - self.conv_in = nn.Conv2d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1) - self.blocks = nn.ModuleList([]) - - for i in range(len(block_out_channels) - 1): - channel_in = block_out_channels[i] - channel_out = block_out_channels[i + 1] - self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1)) - self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2)) - - self.conv_out = zero_module( - nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1) - ) - - def forward(self, conditioning: torch.Tensor) -> torch.Tensor: - embedding = self.conv_in(conditioning) - embedding = F.silu(embedding) - - for block in self.blocks: - embedding = block(embedding) - embedding = F.silu(embedding) - - embedding = self.conv_out(embedding) - return embedding - - -class SparseControlNetModel(ModelMixin, AttentionMixin, ConfigMixin, FromOriginalModelMixin): - """ - A SparseControlNet model as described in [SparseCtrl: Adding Sparse Controls to Text-to-Video Diffusion - Models](https://huggingface.co/papers/2311.16933). - - Args: - in_channels (`int`, defaults to 4): - The number of channels in the input sample. - conditioning_channels (`int`, defaults to 4): - The number of input channels in the controlnet conditional embedding module. If - `concat_condition_embedding` is True, the value provided here is incremented by 1. - flip_sin_to_cos (`bool`, defaults to `True`): - Whether to flip the sin to cos in the time embedding. - freq_shift (`int`, defaults to 0): - The frequency shift to apply to the time embedding. - down_block_types (`tuple[str]`, defaults to `("CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "DownBlock2D")`): - The tuple of downsample blocks to use. - only_cross_attention (`bool | tuple[bool]`, defaults to `False`): - block_out_channels (`tuple[int]`, defaults to `(320, 640, 1280, 1280)`): - The tuple of output channels for each block. - layers_per_block (`int`, defaults to 2): - The number of layers per block. - downsample_padding (`int`, defaults to 1): - The padding to use for the downsampling convolution. - mid_block_scale_factor (`float`, defaults to 1): - The scale factor to use for the mid block. - act_fn (`str`, defaults to "silu"): - The activation function to use. - norm_num_groups (`int`, *optional*, defaults to 32): - The number of groups to use for the normalization. If None, normalization and activation layers is skipped - in post-processing. - norm_eps (`float`, defaults to 1e-5): - The epsilon to use for the normalization. - cross_attention_dim (`int`, defaults to 1280): - The dimension of the cross attention features. - transformer_layers_per_block (`int` or `tuple[int]`, *optional*, defaults to 1): - The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`]. Only relevant for - [`~models.unet_2d_blocks.CrossAttnDownBlock2D`], [`~models.unet_2d_blocks.CrossAttnUpBlock2D`], - [`~models.unet_2d_blocks.UNetMidBlock2DCrossAttn`]. - transformer_layers_per_mid_block (`int` or `tuple[int]`, *optional*, defaults to 1): - The number of transformer layers to use in each layer in the middle block. - attention_head_dim (`int` or `tuple[int]`, defaults to 8): - The dimension of the attention heads. - num_attention_heads (`int` or `tuple[int]`, *optional*): - The number of heads to use for multi-head attention. - use_linear_projection (`bool`, defaults to `False`): - upcast_attention (`bool`, defaults to `False`): - resnet_time_scale_shift (`str`, defaults to `"default"`): - Time scale shift config for ResNet blocks (see `ResnetBlock2D`). Choose from `default` or `scale_shift`. - conditioning_embedding_out_channels (`tuple[int]`, defaults to `(16, 32, 96, 256)`): - The tuple of output channel for each block in the `conditioning_embedding` layer. - global_pool_conditions (`bool`, defaults to `False`): - TODO(Patrick) - unused parameter - controlnet_conditioning_channel_order (`str`, defaults to `rgb`): - motion_max_seq_length (`int`, defaults to `32`): - The maximum sequence length to use in the motion module. - motion_num_attention_heads (`int` or `tuple[int]`, defaults to `8`): - The number of heads to use in each attention layer of the motion module. - concat_conditioning_mask (`bool`, defaults to `True`): - use_simplified_condition_embedding (`bool`, defaults to `True`): - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 4, - conditioning_channels: int = 4, - flip_sin_to_cos: bool = True, - freq_shift: int = 0, - down_block_types: tuple[str, ...] = ( - "CrossAttnDownBlockMotion", - "CrossAttnDownBlockMotion", - "CrossAttnDownBlockMotion", - "DownBlockMotion", - ), - only_cross_attention: bool | tuple[bool] = False, - block_out_channels: tuple[int, ...] = (320, 640, 1280, 1280), - layers_per_block: int = 2, - downsample_padding: int = 1, - mid_block_scale_factor: float = 1, - act_fn: str = "silu", - norm_num_groups: int | None = 32, - norm_eps: float = 1e-5, - cross_attention_dim: int = 768, - transformer_layers_per_block: int | tuple[int, ...] = 1, - transformer_layers_per_mid_block: int | tuple[int] | None = None, - temporal_transformer_layers_per_block: int | tuple[int, ...] = 1, - attention_head_dim: int | tuple[int, ...] = 8, - num_attention_heads: int | tuple[int, ...] | None = None, - use_linear_projection: bool = False, - upcast_attention: bool = False, - resnet_time_scale_shift: str = "default", - conditioning_embedding_out_channels: tuple[int, ...] | None = (16, 32, 96, 256), - global_pool_conditions: bool = False, - controlnet_conditioning_channel_order: str = "rgb", - motion_max_seq_length: int = 32, - motion_num_attention_heads: int = 8, - concat_conditioning_mask: bool = True, - use_simplified_condition_embedding: bool = True, - ): - super().__init__() - self.use_simplified_condition_embedding = use_simplified_condition_embedding - - # If `num_attention_heads` is not defined (which is the case for most models) - # it will default to `attention_head_dim`. This looks weird upon first reading it and it is. - # The reason for this behavior is to correct for incorrectly named variables that were introduced - # when this library was created. The incorrect naming was only discovered much later in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131 - # Changing `attention_head_dim` to `num_attention_heads` for 40,000+ configurations is too backwards breaking - # which is why we correct for the naming here. - num_attention_heads = num_attention_heads or attention_head_dim - - # Check inputs - if len(block_out_channels) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(only_cross_attention, bool) and len(only_cross_attention) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `only_cross_attention` as `down_block_types`. `only_cross_attention`: {only_cross_attention}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}." - ) - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types) - if isinstance(temporal_transformer_layers_per_block, int): - temporal_transformer_layers_per_block = [temporal_transformer_layers_per_block] * len(down_block_types) - - # input - conv_in_kernel = 3 - conv_in_padding = (conv_in_kernel - 1) // 2 - self.conv_in = nn.Conv2d( - in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding - ) - - if concat_conditioning_mask: - conditioning_channels = conditioning_channels + 1 - - self.concat_conditioning_mask = concat_conditioning_mask - - # control net conditioning embedding - if use_simplified_condition_embedding: - self.controlnet_cond_embedding = zero_module( - nn.Conv2d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1) - ) - else: - self.controlnet_cond_embedding = SparseControlNetConditioningEmbedding( - conditioning_embedding_channels=block_out_channels[0], - block_out_channels=conditioning_embedding_out_channels, - conditioning_channels=conditioning_channels, - ) - - # time - time_embed_dim = block_out_channels[0] * 4 - self.time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos, freq_shift) - timestep_input_dim = block_out_channels[0] - - self.time_embedding = TimestepEmbedding( - timestep_input_dim, - time_embed_dim, - act_fn=act_fn, - ) - - self.down_blocks = nn.ModuleList([]) - self.controlnet_down_blocks = nn.ModuleList([]) - - if isinstance(cross_attention_dim, int): - cross_attention_dim = (cross_attention_dim,) * len(down_block_types) - - if isinstance(only_cross_attention, bool): - only_cross_attention = [only_cross_attention] * len(down_block_types) - - if isinstance(attention_head_dim, int): - attention_head_dim = (attention_head_dim,) * len(down_block_types) - - if isinstance(num_attention_heads, int): - num_attention_heads = (num_attention_heads,) * len(down_block_types) - - if isinstance(motion_num_attention_heads, int): - motion_num_attention_heads = (motion_num_attention_heads,) * len(down_block_types) - - # down - output_channel = block_out_channels[0] - - controlnet_block = nn.Conv2d(output_channel, output_channel, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_down_blocks.append(controlnet_block) - - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - if down_block_type == "CrossAttnDownBlockMotion": - down_block = CrossAttnDownBlockMotion( - in_channels=input_channel, - out_channels=output_channel, - temb_channels=time_embed_dim, - dropout=0, - num_layers=layers_per_block, - transformer_layers_per_block=transformer_layers_per_block[i], - resnet_eps=norm_eps, - resnet_time_scale_shift=resnet_time_scale_shift, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - resnet_pre_norm=True, - num_attention_heads=num_attention_heads[i], - cross_attention_dim=cross_attention_dim[i], - add_downsample=not is_final_block, - dual_cross_attention=False, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention[i], - upcast_attention=upcast_attention, - temporal_num_attention_heads=motion_num_attention_heads[i], - temporal_max_seq_length=motion_max_seq_length, - temporal_transformer_layers_per_block=temporal_transformer_layers_per_block[i], - temporal_double_self_attention=False, - ) - elif down_block_type == "DownBlockMotion": - down_block = DownBlockMotion( - in_channels=input_channel, - out_channels=output_channel, - temb_channels=time_embed_dim, - dropout=0, - num_layers=layers_per_block, - resnet_eps=norm_eps, - resnet_time_scale_shift=resnet_time_scale_shift, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - resnet_pre_norm=True, - add_downsample=not is_final_block, - temporal_num_attention_heads=motion_num_attention_heads[i], - temporal_max_seq_length=motion_max_seq_length, - temporal_transformer_layers_per_block=temporal_transformer_layers_per_block[i], - temporal_double_self_attention=False, - ) - else: - raise ValueError( - "Invalid `block_type` encountered. Must be one of `CrossAttnDownBlockMotion` or `DownBlockMotion`" - ) - - self.down_blocks.append(down_block) - - for _ in range(layers_per_block): - controlnet_block = nn.Conv2d(output_channel, output_channel, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_down_blocks.append(controlnet_block) - - if not is_final_block: - controlnet_block = nn.Conv2d(output_channel, output_channel, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_down_blocks.append(controlnet_block) - - # mid - mid_block_channels = block_out_channels[-1] - - controlnet_block = nn.Conv2d(mid_block_channels, mid_block_channels, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_mid_block = controlnet_block - - if transformer_layers_per_mid_block is None: - transformer_layers_per_mid_block = ( - transformer_layers_per_block[-1] if isinstance(transformer_layers_per_block[-1], int) else 1 - ) - - self.mid_block = UNetMidBlock2DCrossAttn( - in_channels=mid_block_channels, - temb_channels=time_embed_dim, - dropout=0, - num_layers=1, - transformer_layers_per_block=transformer_layers_per_mid_block, - resnet_eps=norm_eps, - resnet_time_scale_shift=resnet_time_scale_shift, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - resnet_pre_norm=True, - num_attention_heads=num_attention_heads[-1], - output_scale_factor=mid_block_scale_factor, - cross_attention_dim=cross_attention_dim[-1], - dual_cross_attention=False, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - attention_type="default", - ) - - @classmethod - def from_unet( - cls, - unet: UNet2DConditionModel, - controlnet_conditioning_channel_order: str = "rgb", - conditioning_embedding_out_channels: tuple[int, ...] | None = (16, 32, 96, 256), - load_weights_from_unet: bool = True, - conditioning_channels: int = 3, - ) -> "SparseControlNetModel": - r""" - Instantiate a [`SparseControlNetModel`] from [`UNet2DConditionModel`]. - - Parameters: - unet (`UNet2DConditionModel`): - The UNet model weights to copy to the [`SparseControlNetModel`]. All configuration options are also - copied where applicable. - """ - transformer_layers_per_block = ( - unet.config.transformer_layers_per_block if "transformer_layers_per_block" in unet.config else 1 - ) - down_block_types = unet.config.down_block_types - - for i in range(len(down_block_types)): - if "CrossAttn" in down_block_types[i]: - down_block_types[i] = "CrossAttnDownBlockMotion" - elif "Down" in down_block_types[i]: - down_block_types[i] = "DownBlockMotion" - else: - raise ValueError("Invalid `block_type` encountered. Must be a cross-attention or down block") - - controlnet = cls( - in_channels=unet.config.in_channels, - conditioning_channels=conditioning_channels, - flip_sin_to_cos=unet.config.flip_sin_to_cos, - freq_shift=unet.config.freq_shift, - down_block_types=unet.config.down_block_types, - only_cross_attention=unet.config.only_cross_attention, - block_out_channels=unet.config.block_out_channels, - layers_per_block=unet.config.layers_per_block, - downsample_padding=unet.config.downsample_padding, - mid_block_scale_factor=unet.config.mid_block_scale_factor, - act_fn=unet.config.act_fn, - norm_num_groups=unet.config.norm_num_groups, - norm_eps=unet.config.norm_eps, - cross_attention_dim=unet.config.cross_attention_dim, - transformer_layers_per_block=transformer_layers_per_block, - attention_head_dim=unet.config.attention_head_dim, - num_attention_heads=unet.config.num_attention_heads, - use_linear_projection=unet.config.use_linear_projection, - upcast_attention=unet.config.upcast_attention, - resnet_time_scale_shift=unet.config.resnet_time_scale_shift, - conditioning_embedding_out_channels=conditioning_embedding_out_channels, - controlnet_conditioning_channel_order=controlnet_conditioning_channel_order, - ) - - if load_weights_from_unet: - controlnet.conv_in.load_state_dict(unet.conv_in.state_dict(), strict=False) - controlnet.time_proj.load_state_dict(unet.time_proj.state_dict(), strict=False) - controlnet.time_embedding.load_state_dict(unet.time_embedding.state_dict(), strict=False) - controlnet.down_blocks.load_state_dict(unet.down_blocks.state_dict(), strict=False) - controlnet.mid_block.load_state_dict(unet.mid_block.state_dict(), strict=False) - - return controlnet - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attention_slice - def set_attention_slice(self, slice_size: str | int | list[int]) -> None: - r""" - Enable sliced attention computation. - - When this option is enabled, the attention module splits the input tensor in slices to compute attention in - several steps. This is useful for saving some memory in exchange for a small decrease in speed. - - Args: - slice_size (`str` or `int` or `list(int)`, *optional*, defaults to `"auto"`): - When `"auto"`, input to the attention heads is halved, so attention is computed in two steps. If - `"max"`, maximum amount of memory is saved by running only one slice at a time. If a number is - provided, uses as many slices as `attention_head_dim // slice_size`. In this case, `attention_head_dim` - must be a multiple of `slice_size`. - """ - sliceable_head_dims = [] - - def fn_recursive_retrieve_sliceable_dims(module: torch.nn.Module): - if hasattr(module, "set_attention_slice"): - sliceable_head_dims.append(module.sliceable_head_dim) - - for child in module.children(): - fn_recursive_retrieve_sliceable_dims(child) - - # retrieve number of attention layers - for module in self.children(): - fn_recursive_retrieve_sliceable_dims(module) - - num_sliceable_layers = len(sliceable_head_dims) - - if slice_size == "auto": - # half the attention head size is usually a good trade-off between - # speed and memory - slice_size = [dim // 2 for dim in sliceable_head_dims] - elif slice_size == "max": - # make smallest slice possible - slice_size = num_sliceable_layers * [1] - - slice_size = num_sliceable_layers * [slice_size] if not isinstance(slice_size, list) else slice_size - - if len(slice_size) != len(sliceable_head_dims): - raise ValueError( - f"You have provided {len(slice_size)}, but {self.config} has {len(sliceable_head_dims)} different" - f" attention layers. Make sure to match `len(slice_size)` to be {len(sliceable_head_dims)}." - ) - - for i in range(len(slice_size)): - size = slice_size[i] - dim = sliceable_head_dims[i] - if size is not None and size > dim: - raise ValueError(f"size {size} has to be smaller or equal to {dim}.") - - # Recursively walk through all the children. - # Any children which exposes the set_attention_slice method - # gets the message - def fn_recursive_set_attention_slice(module: torch.nn.Module, slice_size: list[int]): - if hasattr(module, "set_attention_slice"): - module.set_attention_slice(slice_size.pop()) - - for child in module.children(): - fn_recursive_set_attention_slice(child, slice_size) - - reversed_slice_size = list(reversed(slice_size)) - for module in self.children(): - fn_recursive_set_attention_slice(module, reversed_slice_size) - - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - encoder_hidden_states: torch.Tensor, - controlnet_cond: torch.Tensor, - conditioning_scale: float = 1.0, - timestep_cond: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - conditioning_mask: torch.Tensor | None = None, - guess_mode: bool = False, - return_dict: bool = True, - ) -> SparseControlNetOutput | tuple[tuple[torch.Tensor, ...], torch.Tensor]: - """ - The [`SparseControlNetModel`] forward method. - - Args: - sample (`torch.Tensor`): - The noisy input tensor. - timestep (`torch.Tensor | float | int`): - The number of timesteps to denoise an input. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states. - controlnet_cond (`torch.Tensor`): - The conditional input tensor of shape `(batch_size, sequence_length, hidden_size)`. - conditioning_scale (`float`, defaults to `1.0`): - The scale factor for ControlNet outputs. - timestep_cond (`torch.Tensor`, *optional*, defaults to `None`): - Additional conditional embeddings for timestep. If provided, the embeddings will be summed with the - timestep_embedding passed through the `self.time_embedding` layer to obtain the final timestep - embeddings. - attention_mask (`torch.Tensor`, *optional*, defaults to `None`): - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask - is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large - negative values to the attention scores corresponding to "discard" tokens. - conditioning_mask (`torch.Tensor`, *optional*, defaults to `None`): - Optional mask indicating which frames in `controlnet_cond` are valid conditioning frames. - cross_attention_kwargs (`dict[str]`, *optional*, defaults to `None`): - A kwargs dictionary that if specified is passed along to the `AttnProcessor`. - guess_mode (`bool`, defaults to `False`): - In this mode, the ControlNet encoder tries its best to recognize the input content of the input even if - you remove all prompts. A `guidance_scale` between 3.0 and 5.0 is recommended. - return_dict (`bool`, defaults to `True`): - Whether or not to return a [`~models.controlnet.ControlNetOutput`] instead of a plain tuple. - Returns: - [`~models.controlnet.ControlNetOutput`] **or** `tuple`: - If `return_dict` is `True`, a [`~models.controlnet.ControlNetOutput`] is returned, otherwise a tuple is - returned where the first element is the sample tensor. - """ - sample_batch_size, sample_channels, sample_num_frames, sample_height, sample_width = sample.shape - sample = torch.zeros_like(sample) - - # check channel order - channel_order = self.config.controlnet_conditioning_channel_order - - if channel_order == "rgb": - # in rgb order by default - ... - elif channel_order == "bgr": - controlnet_cond = torch.flip(controlnet_cond, dims=[1]) - else: - raise ValueError(f"unknown `controlnet_conditioning_channel_order`: {channel_order}") - - # prepare attention_mask - if attention_mask is not None: - attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # 1. time - timesteps = timestep - if not torch.is_tensor(timesteps): - # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can - # This would be a good case for the `match` statement (Python 3.10+) - dtype = maybe_adjust_dtype_for_device( - torch.float64 if isinstance(timestep, float) else torch.int64, sample.device - ) - timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) - elif len(timesteps.shape) == 0: - timesteps = timesteps[None].to(sample.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - timesteps = timesteps.expand(sample.shape[0]) - - t_emb = self.time_proj(timesteps) - - # timesteps does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=sample.dtype) - - emb = self.time_embedding(t_emb, timestep_cond) - emb = emb.repeat_interleave(sample_num_frames, dim=0, output_size=emb.shape[0] * sample_num_frames) - - # 2. pre-process - batch_size, channels, num_frames, height, width = sample.shape - - sample = sample.permute(0, 2, 1, 3, 4).reshape(batch_size * num_frames, channels, height, width) - sample = self.conv_in(sample) - - batch_frames, channels, height, width = sample.shape - sample = sample[:, None].reshape(sample_batch_size, sample_num_frames, channels, height, width) - - if self.concat_conditioning_mask: - controlnet_cond = torch.cat([controlnet_cond, conditioning_mask], dim=1) - - batch_size, channels, num_frames, height, width = controlnet_cond.shape - controlnet_cond = controlnet_cond.permute(0, 2, 1, 3, 4).reshape( - batch_size * num_frames, channels, height, width - ) - controlnet_cond = self.controlnet_cond_embedding(controlnet_cond) - batch_frames, channels, height, width = controlnet_cond.shape - controlnet_cond = controlnet_cond[:, None].reshape(batch_size, num_frames, channels, height, width) - - sample = sample + controlnet_cond - - batch_size, num_frames, channels, height, width = sample.shape - sample = sample.reshape(sample_batch_size * sample_num_frames, channels, height, width) - - # 3. down - down_block_res_samples = (sample,) - for downsample_block in self.down_blocks: - if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention: - sample, res_samples = downsample_block( - hidden_states=sample, - temb=emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample, res_samples = downsample_block(hidden_states=sample, temb=emb, num_frames=num_frames) - - down_block_res_samples += res_samples - - # 4. mid - if self.mid_block is not None: - if hasattr(self.mid_block, "has_cross_attention") and self.mid_block.has_cross_attention: - sample = self.mid_block( - sample, - emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample = self.mid_block(sample, emb) - - # 5. Control net blocks - controlnet_down_block_res_samples = () - - for down_block_res_sample, controlnet_block in zip(down_block_res_samples, self.controlnet_down_blocks): - down_block_res_sample = controlnet_block(down_block_res_sample) - controlnet_down_block_res_samples = controlnet_down_block_res_samples + (down_block_res_sample,) - - down_block_res_samples = controlnet_down_block_res_samples - mid_block_res_sample = self.controlnet_mid_block(sample) - - # 6. scaling - if guess_mode and not self.config.global_pool_conditions: - scales = torch.logspace(-1, 0, len(down_block_res_samples) + 1, device=sample.device) # 0.1 to 1.0 - scales = scales * conditioning_scale - down_block_res_samples = [sample * scale for sample, scale in zip(down_block_res_samples, scales)] - mid_block_res_sample = mid_block_res_sample * scales[-1] # last one - else: - down_block_res_samples = [sample * conditioning_scale for sample in down_block_res_samples] - mid_block_res_sample = mid_block_res_sample * conditioning_scale - - if self.config.global_pool_conditions: - down_block_res_samples = [ - torch.mean(sample, dim=(2, 3), keepdim=True) for sample in down_block_res_samples - ] - mid_block_res_sample = torch.mean(mid_block_res_sample, dim=(2, 3), keepdim=True) - - if not return_dict: - return (down_block_res_samples, mid_block_res_sample) - - return SparseControlNetOutput( - down_block_res_samples=down_block_res_samples, mid_block_res_sample=mid_block_res_sample - ) - - -# Copied from diffusers.models.controlnets.controlnet.zero_module -def zero_module(module: nn.Module) -> nn.Module: - for p in module.parameters(): - nn.init.zeros_(p) - return module diff --git a/diffusers/models/controlnets/controlnet_union.py b/diffusers/models/controlnets/controlnet_union.py deleted file mode 100644 index 8b3ac1c36d856418445fae3d39eeac65d683a0f5..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/controlnet_union.py +++ /dev/null @@ -1,779 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from typing import Any - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders.single_file_model import FromOriginalModelMixin -from ...utils import logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..attention import AttentionMixin -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - AttnAddedKVProcessor, - AttnProcessor, -) -from ..embeddings import TextImageTimeEmbedding, TextTimeEmbedding, TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from ..unets.unet_2d_blocks import ( - UNetMidBlock2DCrossAttn, - get_down_block, -) -from ..unets.unet_2d_condition import UNet2DConditionModel -from .controlnet import ControlNetConditioningEmbedding, ControlNetOutput, zero_module - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class QuickGELU(nn.Module): - """ - Applies GELU approximation that is fast but somewhat inaccurate. See: https://github.com/hendrycks/GELUs - """ - - def forward(self, input: torch.Tensor) -> torch.Tensor: - return input * torch.sigmoid(1.702 * input) - - -class ResidualAttentionMlp(nn.Module): - def __init__(self, d_model: int): - super().__init__() - self.c_fc = nn.Linear(d_model, d_model * 4) - self.gelu = QuickGELU() - self.c_proj = nn.Linear(d_model * 4, d_model) - - def forward(self, x: torch.Tensor): - x = self.c_fc(x) - x = self.gelu(x) - x = self.c_proj(x) - return x - - -class ResidualAttentionBlock(nn.Module): - def __init__(self, d_model: int, n_head: int, attn_mask: torch.Tensor = None): - super().__init__() - self.attn = nn.MultiheadAttention(d_model, n_head) - self.ln_1 = nn.LayerNorm(d_model) - self.mlp = ResidualAttentionMlp(d_model) - self.ln_2 = nn.LayerNorm(d_model) - self.attn_mask = attn_mask - - def attention(self, x: torch.Tensor): - self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None - return self.attn(x, x, x, need_weights=False, attn_mask=self.attn_mask)[0] - - def forward(self, x: torch.Tensor): - x = x + self.attention(self.ln_1(x)) - x = x + self.mlp(self.ln_2(x)) - return x - - -class ControlNetUnionModel(ModelMixin, AttentionMixin, ConfigMixin, FromOriginalModelMixin): - """ - A ControlNetUnion model. - - Args: - in_channels (`int`, defaults to 4): - The number of channels in the input sample. - flip_sin_to_cos (`bool`, defaults to `True`): - Whether to flip the sin to cos in the time embedding. - freq_shift (`int`, defaults to 0): - The frequency shift to apply to the time embedding. - down_block_types (`tuple[str]`, defaults to `("CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "DownBlock2D")`): - The tuple of downsample blocks to use. - only_cross_attention (`bool | tuple[bool]`, defaults to `False`): - block_out_channels (`tuple[int]`, defaults to `(320, 640, 1280, 1280)`): - The tuple of output channels for each block. - layers_per_block (`int`, defaults to 2): - The number of layers per block. - downsample_padding (`int`, defaults to 1): - The padding to use for the downsampling convolution. - mid_block_scale_factor (`float`, defaults to 1): - The scale factor to use for the mid block. - act_fn (`str`, defaults to "silu"): - The activation function to use. - norm_num_groups (`int`, *optional*, defaults to 32): - The number of groups to use for the normalization. If None, normalization and activation layers is skipped - in post-processing. - norm_eps (`float`, defaults to 1e-5): - The epsilon to use for the normalization. - cross_attention_dim (`int`, defaults to 1280): - The dimension of the cross attention features. - transformer_layers_per_block (`int` or `tuple[int]`, *optional*, defaults to 1): - The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`]. Only relevant for - [`~models.unet_2d_blocks.CrossAttnDownBlock2D`], [`~models.unet_2d_blocks.CrossAttnUpBlock2D`], - [`~models.unet_2d_blocks.UNetMidBlock2DCrossAttn`]. - encoder_hid_dim (`int`, *optional*, defaults to None): - If `encoder_hid_dim_type` is defined, `encoder_hidden_states` will be projected from `encoder_hid_dim` - dimension to `cross_attention_dim`. - encoder_hid_dim_type (`str`, *optional*, defaults to `None`): - If given, the `encoder_hidden_states` and potentially other embeddings are down-projected to text - embeddings of dimension `cross_attention` according to `encoder_hid_dim_type`. - attention_head_dim (`int | tuple[int]`, defaults to 8): - The dimension of the attention heads. - use_linear_projection (`bool`, defaults to `False`): - class_embed_type (`str`, *optional*, defaults to `None`): - The type of class embedding to use which is ultimately summed with the time embeddings. Choose from None, - `"timestep"`, `"identity"`, `"projection"`, or `"simple_projection"`. - addition_embed_type (`str`, *optional*, defaults to `None`): - Configures an optional embedding which will be summed with the time embeddings. Choose from `None` or - "text". "text" will use the `TextTimeEmbedding` layer. - num_class_embeds (`int`, *optional*, defaults to 0): - Input dimension of the learnable embedding matrix to be projected to `time_embed_dim`, when performing - class conditioning with `class_embed_type` equal to `None`. - upcast_attention (`bool`, defaults to `False`): - resnet_time_scale_shift (`str`, defaults to `"default"`): - Time scale shift config for ResNet blocks (see `ResnetBlock2D`). Choose from `default` or `scale_shift`. - projection_class_embeddings_input_dim (`int`, *optional*, defaults to `None`): - The dimension of the `class_labels` input when `class_embed_type="projection"`. Required when - `class_embed_type="projection"`. - controlnet_conditioning_channel_order (`str`, defaults to `"rgb"`): - The channel order of conditional image. Will convert to `rgb` if it's `bgr`. - conditioning_embedding_out_channels (`tuple[int]`, *optional*, defaults to `(48, 96, 192, 384)`): - The tuple of output channel for each block in the `conditioning_embedding` layer. - global_pool_conditions (`bool`, defaults to `False`): - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 4, - conditioning_channels: int = 3, - flip_sin_to_cos: bool = True, - freq_shift: int = 0, - down_block_types: tuple[str, ...] = ( - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "DownBlock2D", - ), - only_cross_attention: bool | tuple[bool] = False, - block_out_channels: tuple[int, ...] = (320, 640, 1280, 1280), - layers_per_block: int = 2, - downsample_padding: int = 1, - mid_block_scale_factor: float = 1, - act_fn: str = "silu", - norm_num_groups: int | None = 32, - norm_eps: float = 1e-5, - cross_attention_dim: int = 1280, - transformer_layers_per_block: int | tuple[int, ...] = 1, - encoder_hid_dim: int | None = None, - encoder_hid_dim_type: str | None = None, - attention_head_dim: int | tuple[int, ...] = 8, - num_attention_heads: int | tuple[int, ...] | None = None, - use_linear_projection: bool = False, - class_embed_type: str | None = None, - addition_embed_type: str | None = None, - addition_time_embed_dim: int | None = None, - num_class_embeds: int | None = None, - upcast_attention: bool = False, - resnet_time_scale_shift: str = "default", - projection_class_embeddings_input_dim: int | None = None, - controlnet_conditioning_channel_order: str = "rgb", - conditioning_embedding_out_channels: tuple[int, ...] | None = (48, 96, 192, 384), - global_pool_conditions: bool = False, - addition_embed_type_num_heads: int = 64, - num_control_type: int = 6, - num_trans_channel: int = 320, - num_trans_head: int = 8, - num_trans_layer: int = 1, - num_proj_channel: int = 320, - ): - super().__init__() - - # If `num_attention_heads` is not defined (which is the case for most models) - # it will default to `attention_head_dim`. This looks weird upon first reading it and it is. - # The reason for this behavior is to correct for incorrectly named variables that were introduced - # when this library was created. The incorrect naming was only discovered much later in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131 - # Changing `attention_head_dim` to `num_attention_heads` for 40,000+ configurations is too backwards breaking - # which is why we correct for the naming here. - num_attention_heads = num_attention_heads or attention_head_dim - - # Check inputs - if len(block_out_channels) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(only_cross_attention, bool) and len(only_cross_attention) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `only_cross_attention` as `down_block_types`. `only_cross_attention`: {only_cross_attention}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}." - ) - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types) - - # input - conv_in_kernel = 3 - conv_in_padding = (conv_in_kernel - 1) // 2 - self.conv_in = nn.Conv2d( - in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding - ) - - # time - time_embed_dim = block_out_channels[0] * 4 - self.time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos, freq_shift) - timestep_input_dim = block_out_channels[0] - self.time_embedding = TimestepEmbedding( - timestep_input_dim, - time_embed_dim, - act_fn=act_fn, - ) - - if encoder_hid_dim_type is not None: - raise ValueError(f"encoder_hid_dim_type: {encoder_hid_dim_type} must be None.") - else: - self.encoder_hid_proj = None - - # class embedding - if class_embed_type is None and num_class_embeds is not None: - self.class_embedding = nn.Embedding(num_class_embeds, time_embed_dim) - elif class_embed_type == "timestep": - self.class_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim) - elif class_embed_type == "identity": - self.class_embedding = nn.Identity(time_embed_dim, time_embed_dim) - elif class_embed_type == "projection": - if projection_class_embeddings_input_dim is None: - raise ValueError( - "`class_embed_type`: 'projection' requires `projection_class_embeddings_input_dim` be set" - ) - # The projection `class_embed_type` is the same as the timestep `class_embed_type` except - # 1. the `class_labels` inputs are not first converted to sinusoidal embeddings - # 2. it projects from an arbitrary input dimension. - # - # Note that `TimestepEmbedding` is quite general, being mainly linear layers and activations. - # When used for embedding actual timesteps, the timesteps are first converted to sinusoidal embeddings. - # As a result, `TimestepEmbedding` can be passed arbitrary vectors. - self.class_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) - else: - self.class_embedding = None - - if addition_embed_type == "text": - if encoder_hid_dim is not None: - text_time_embedding_from_dim = encoder_hid_dim - else: - text_time_embedding_from_dim = cross_attention_dim - - self.add_embedding = TextTimeEmbedding( - text_time_embedding_from_dim, time_embed_dim, num_heads=addition_embed_type_num_heads - ) - elif addition_embed_type == "text_image": - # text_embed_dim and image_embed_dim DON'T have to be `cross_attention_dim`. To not clutter the __init__ too much - # they are set to `cross_attention_dim` here as this is exactly the required dimension for the currently only use - # case when `addition_embed_type == "text_image"` (Kandinsky 2.1)` - self.add_embedding = TextImageTimeEmbedding( - text_embed_dim=cross_attention_dim, image_embed_dim=cross_attention_dim, time_embed_dim=time_embed_dim - ) - elif addition_embed_type == "text_time": - self.add_time_proj = Timesteps(addition_time_embed_dim, flip_sin_to_cos, freq_shift) - self.add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) - - elif addition_embed_type is not None: - raise ValueError(f"addition_embed_type: {addition_embed_type} must be None, 'text' or 'text_image'.") - - # control net conditioning embedding - self.controlnet_cond_embedding = ControlNetConditioningEmbedding( - conditioning_embedding_channels=block_out_channels[0], - block_out_channels=conditioning_embedding_out_channels, - conditioning_channels=conditioning_channels, - ) - - task_scale_factor = num_trans_channel**0.5 - self.task_embedding = nn.Parameter(task_scale_factor * torch.randn(num_control_type, num_trans_channel)) - self.transformer_layes = nn.ModuleList( - [ResidualAttentionBlock(num_trans_channel, num_trans_head) for _ in range(num_trans_layer)] - ) - self.spatial_ch_projs = zero_module(nn.Linear(num_trans_channel, num_proj_channel)) - self.control_type_proj = Timesteps(addition_time_embed_dim, flip_sin_to_cos, freq_shift) - self.control_add_embedding = TimestepEmbedding(addition_time_embed_dim * num_control_type, time_embed_dim) - - self.down_blocks = nn.ModuleList([]) - self.controlnet_down_blocks = nn.ModuleList([]) - - if isinstance(only_cross_attention, bool): - only_cross_attention = [only_cross_attention] * len(down_block_types) - - if isinstance(attention_head_dim, int): - attention_head_dim = (attention_head_dim,) * len(down_block_types) - - if isinstance(num_attention_heads, int): - num_attention_heads = (num_attention_heads,) * len(down_block_types) - - # down - output_channel = block_out_channels[0] - - controlnet_block = nn.Conv2d(output_channel, output_channel, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_down_blocks.append(controlnet_block) - - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - down_block = get_down_block( - down_block_type, - num_layers=layers_per_block, - transformer_layers_per_block=transformer_layers_per_block[i], - in_channels=input_channel, - out_channels=output_channel, - temb_channels=time_embed_dim, - add_downsample=not is_final_block, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads[i], - attention_head_dim=attention_head_dim[i] if attention_head_dim[i] is not None else output_channel, - downsample_padding=downsample_padding, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention[i], - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - ) - self.down_blocks.append(down_block) - - for _ in range(layers_per_block): - controlnet_block = nn.Conv2d(output_channel, output_channel, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_down_blocks.append(controlnet_block) - - if not is_final_block: - controlnet_block = nn.Conv2d(output_channel, output_channel, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_down_blocks.append(controlnet_block) - - # mid - mid_block_channel = block_out_channels[-1] - - controlnet_block = nn.Conv2d(mid_block_channel, mid_block_channel, kernel_size=1) - controlnet_block = zero_module(controlnet_block) - self.controlnet_mid_block = controlnet_block - - self.mid_block = UNetMidBlock2DCrossAttn( - transformer_layers_per_block=transformer_layers_per_block[-1], - in_channels=mid_block_channel, - temb_channels=time_embed_dim, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - output_scale_factor=mid_block_scale_factor, - resnet_time_scale_shift=resnet_time_scale_shift, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads[-1], - resnet_groups=norm_num_groups, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - ) - - @classmethod - def from_unet( - cls, - unet: UNet2DConditionModel, - controlnet_conditioning_channel_order: str = "rgb", - conditioning_embedding_out_channels: tuple[int, ...] | None = (16, 32, 96, 256), - load_weights_from_unet: bool = True, - ): - r""" - Instantiate a [`ControlNetUnionModel`] from [`UNet2DConditionModel`]. - - Parameters: - unet (`UNet2DConditionModel`): - The UNet model weights to copy to the [`ControlNetUnionModel`]. All configuration options are also - copied where applicable. - """ - transformer_layers_per_block = ( - unet.config.transformer_layers_per_block if "transformer_layers_per_block" in unet.config else 1 - ) - encoder_hid_dim = unet.config.encoder_hid_dim if "encoder_hid_dim" in unet.config else None - encoder_hid_dim_type = unet.config.encoder_hid_dim_type if "encoder_hid_dim_type" in unet.config else None - addition_embed_type = unet.config.addition_embed_type if "addition_embed_type" in unet.config else None - addition_time_embed_dim = ( - unet.config.addition_time_embed_dim if "addition_time_embed_dim" in unet.config else None - ) - - controlnet = cls( - encoder_hid_dim=encoder_hid_dim, - encoder_hid_dim_type=encoder_hid_dim_type, - addition_embed_type=addition_embed_type, - addition_time_embed_dim=addition_time_embed_dim, - transformer_layers_per_block=transformer_layers_per_block, - in_channels=unet.config.in_channels, - flip_sin_to_cos=unet.config.flip_sin_to_cos, - freq_shift=unet.config.freq_shift, - down_block_types=unet.config.down_block_types, - only_cross_attention=unet.config.only_cross_attention, - block_out_channels=unet.config.block_out_channels, - layers_per_block=unet.config.layers_per_block, - downsample_padding=unet.config.downsample_padding, - mid_block_scale_factor=unet.config.mid_block_scale_factor, - act_fn=unet.config.act_fn, - norm_num_groups=unet.config.norm_num_groups, - norm_eps=unet.config.norm_eps, - cross_attention_dim=unet.config.cross_attention_dim, - attention_head_dim=unet.config.attention_head_dim, - num_attention_heads=unet.config.num_attention_heads, - use_linear_projection=unet.config.use_linear_projection, - class_embed_type=unet.config.class_embed_type, - num_class_embeds=unet.config.num_class_embeds, - upcast_attention=unet.config.upcast_attention, - resnet_time_scale_shift=unet.config.resnet_time_scale_shift, - projection_class_embeddings_input_dim=unet.config.projection_class_embeddings_input_dim, - controlnet_conditioning_channel_order=controlnet_conditioning_channel_order, - conditioning_embedding_out_channels=conditioning_embedding_out_channels, - ) - - if load_weights_from_unet: - controlnet.conv_in.load_state_dict(unet.conv_in.state_dict()) - controlnet.time_proj.load_state_dict(unet.time_proj.state_dict()) - controlnet.time_embedding.load_state_dict(unet.time_embedding.state_dict()) - - if controlnet.class_embedding: - controlnet.class_embedding.load_state_dict(unet.class_embedding.state_dict()) - - controlnet.down_blocks.load_state_dict(unet.down_blocks.state_dict(), strict=False) - controlnet.mid_block.load_state_dict(unet.mid_block.state_dict(), strict=False) - - return controlnet - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attention_slice - def set_attention_slice(self, slice_size: str | int | list[int]) -> None: - r""" - Enable sliced attention computation. - - When this option is enabled, the attention module splits the input tensor in slices to compute attention in - several steps. This is useful for saving some memory in exchange for a small decrease in speed. - - Args: - slice_size (`str` or `int` or `list(int)`, *optional*, defaults to `"auto"`): - When `"auto"`, input to the attention heads is halved, so attention is computed in two steps. If - `"max"`, maximum amount of memory is saved by running only one slice at a time. If a number is - provided, uses as many slices as `attention_head_dim // slice_size`. In this case, `attention_head_dim` - must be a multiple of `slice_size`. - """ - sliceable_head_dims = [] - - def fn_recursive_retrieve_sliceable_dims(module: torch.nn.Module): - if hasattr(module, "set_attention_slice"): - sliceable_head_dims.append(module.sliceable_head_dim) - - for child in module.children(): - fn_recursive_retrieve_sliceable_dims(child) - - # retrieve number of attention layers - for module in self.children(): - fn_recursive_retrieve_sliceable_dims(module) - - num_sliceable_layers = len(sliceable_head_dims) - - if slice_size == "auto": - # half the attention head size is usually a good trade-off between - # speed and memory - slice_size = [dim // 2 for dim in sliceable_head_dims] - elif slice_size == "max": - # make smallest slice possible - slice_size = num_sliceable_layers * [1] - - slice_size = num_sliceable_layers * [slice_size] if not isinstance(slice_size, list) else slice_size - - if len(slice_size) != len(sliceable_head_dims): - raise ValueError( - f"You have provided {len(slice_size)}, but {self.config} has {len(sliceable_head_dims)} different" - f" attention layers. Make sure to match `len(slice_size)` to be {len(sliceable_head_dims)}." - ) - - for i in range(len(slice_size)): - size = slice_size[i] - dim = sliceable_head_dims[i] - if size is not None and size > dim: - raise ValueError(f"size {size} has to be smaller or equal to {dim}.") - - # Recursively walk through all the children. - # Any children which exposes the set_attention_slice method - # gets the message - def fn_recursive_set_attention_slice(module: torch.nn.Module, slice_size: list[int]): - if hasattr(module, "set_attention_slice"): - module.set_attention_slice(slice_size.pop()) - - for child in module.children(): - fn_recursive_set_attention_slice(child, slice_size) - - reversed_slice_size = list(reversed(slice_size)) - for module in self.children(): - fn_recursive_set_attention_slice(module, reversed_slice_size) - - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - encoder_hidden_states: torch.Tensor, - controlnet_cond: list[torch.Tensor], - control_type: torch.Tensor, - control_type_idx: list[int], - conditioning_scale: float | list[float] = 1.0, - class_labels: torch.Tensor | None = None, - timestep_cond: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - added_cond_kwargs: dict[str, torch.Tensor] | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - from_multi: bool = False, - guess_mode: bool = False, - return_dict: bool = True, - ) -> ControlNetOutput | tuple[tuple[torch.Tensor, ...], torch.Tensor]: - """ - The [`ControlNetUnionModel`] forward method. - - Args: - sample (`torch.Tensor`): - The noisy input tensor. - timestep (`torch.Tensor | float | int`): - The number of timesteps to denoise an input. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states. - controlnet_cond (`list[torch.Tensor]`): - The conditional input tensors. - control_type (`torch.Tensor`): - A tensor of shape `(batch, num_control_type)` with values `0` or `1` depending on whether the control - type is used. - control_type_idx (`list[int]`): - The indices of `control_type`. - conditioning_scale (`float`, defaults to `1.0`): - The scale factor for ControlNet outputs. - class_labels (`torch.Tensor`, *optional*, defaults to `None`): - Optional class labels for conditioning. Their embeddings will be summed with the timestep embeddings. - timestep_cond (`torch.Tensor`, *optional*, defaults to `None`): - Additional conditional embeddings for timestep. If provided, the embeddings will be summed with the - timestep_embedding passed through the `self.time_embedding` layer to obtain the final timestep - embeddings. - attention_mask (`torch.Tensor`, *optional*, defaults to `None`): - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask - is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large - negative values to the attention scores corresponding to "discard" tokens. - added_cond_kwargs (`dict`): - Additional conditions for the Stable Diffusion XL UNet. - cross_attention_kwargs (`dict[str]`, *optional*, defaults to `None`): - A kwargs dictionary that if specified is passed along to the `AttnProcessor`. - from_multi (`bool`, defaults to `False`): - Use standard scaling when called from `MultiControlNetUnionModel`. - guess_mode (`bool`, defaults to `False`): - In this mode, the ControlNet encoder tries its best to recognize the input content of the input even if - you remove all prompts. A `guidance_scale` between 3.0 and 5.0 is recommended. - return_dict (`bool`, defaults to `True`): - Whether or not to return a [`~models.controlnet.ControlNetOutput`] instead of a plain tuple. - - Returns: - [`~models.controlnet.ControlNetOutput`] **or** `tuple`: - If `return_dict` is `True`, a [`~models.controlnet.ControlNetOutput`] is returned, otherwise a tuple is - returned where the first element is the sample tensor. - """ - if isinstance(conditioning_scale, float): - conditioning_scale = [conditioning_scale] * len(controlnet_cond) - - # check channel order - channel_order = self.config.controlnet_conditioning_channel_order - - if channel_order != "rgb": - raise ValueError(f"unknown `controlnet_conditioning_channel_order`: {channel_order}") - - # prepare attention_mask - if attention_mask is not None: - attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # 1. time - timesteps = timestep - if not torch.is_tensor(timesteps): - # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can - # This would be a good case for the `match` statement (Python 3.10+) - dtype = maybe_adjust_dtype_for_device( - torch.float64 if isinstance(timestep, float) else torch.int64, sample.device - ) - timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) - elif len(timesteps.shape) == 0: - timesteps = timesteps[None].to(sample.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - timesteps = timesteps.expand(sample.shape[0]) - - t_emb = self.time_proj(timesteps) - - # timesteps does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=sample.dtype) - - emb = self.time_embedding(t_emb, timestep_cond) - aug_emb = None - - if self.class_embedding is not None: - if class_labels is None: - raise ValueError("class_labels should be provided when num_class_embeds > 0") - - if self.config.class_embed_type == "timestep": - class_labels = self.time_proj(class_labels) - - class_emb = self.class_embedding(class_labels).to(dtype=self.dtype) - emb = emb + class_emb - - if self.config.addition_embed_type is not None: - if self.config.addition_embed_type == "text": - aug_emb = self.add_embedding(encoder_hidden_states) - - elif self.config.addition_embed_type == "text_time": - if "text_embeds" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `text_embeds` to be passed in `added_cond_kwargs`" - ) - text_embeds = added_cond_kwargs.get("text_embeds") - if "time_ids" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `time_ids` to be passed in `added_cond_kwargs`" - ) - time_ids = added_cond_kwargs.get("time_ids") - time_embeds = self.add_time_proj(time_ids.flatten()) - time_embeds = time_embeds.reshape((text_embeds.shape[0], -1)) - - add_embeds = torch.concat([text_embeds, time_embeds], dim=-1) - add_embeds = add_embeds.to(emb.dtype) - aug_emb = self.add_embedding(add_embeds) - - control_embeds = self.control_type_proj(control_type.flatten()) - control_embeds = control_embeds.reshape((t_emb.shape[0], -1)) - control_embeds = control_embeds.to(emb.dtype) - control_emb = self.control_add_embedding(control_embeds) - emb = emb + control_emb - emb = emb + aug_emb if aug_emb is not None else emb - - # 2. pre-process - sample = self.conv_in(sample) - - inputs = [] - condition_list = [] - - for cond, control_idx, scale in zip(controlnet_cond, control_type_idx, conditioning_scale): - condition = self.controlnet_cond_embedding(cond) - feat_seq = torch.mean(condition, dim=(2, 3)) - feat_seq = feat_seq + self.task_embedding[control_idx] - if from_multi or len(control_type_idx) == 1: - inputs.append(feat_seq.unsqueeze(1)) - condition_list.append(condition) - else: - inputs.append(feat_seq.unsqueeze(1) * scale) - condition_list.append(condition * scale) - - condition = sample - feat_seq = torch.mean(condition, dim=(2, 3)) - inputs.append(feat_seq.unsqueeze(1)) - condition_list.append(condition) - - x = torch.cat(inputs, dim=1) - for layer in self.transformer_layes: - x = layer(x) - - controlnet_cond_fuser = sample * 0.0 - for (idx, condition), scale in zip(enumerate(condition_list[:-1]), conditioning_scale): - alpha = self.spatial_ch_projs(x[:, idx]) - alpha = alpha.unsqueeze(-1).unsqueeze(-1) - if from_multi or len(control_type_idx) == 1: - controlnet_cond_fuser += condition + alpha - else: - controlnet_cond_fuser += condition + alpha * scale - - sample = sample + controlnet_cond_fuser - - # 3. down - down_block_res_samples = (sample,) - for downsample_block in self.down_blocks: - if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention: - sample, res_samples = downsample_block( - hidden_states=sample, - temb=emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample, res_samples = downsample_block(hidden_states=sample, temb=emb) - - down_block_res_samples += res_samples - - # 4. mid - if self.mid_block is not None: - sample = self.mid_block( - sample, - emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - ) - - # 5. Control net blocks - controlnet_down_block_res_samples = () - - for down_block_res_sample, controlnet_block in zip(down_block_res_samples, self.controlnet_down_blocks): - down_block_res_sample = controlnet_block(down_block_res_sample) - controlnet_down_block_res_samples = controlnet_down_block_res_samples + (down_block_res_sample,) - - down_block_res_samples = controlnet_down_block_res_samples - - mid_block_res_sample = self.controlnet_mid_block(sample) - - # 6. scaling - if guess_mode and not self.config.global_pool_conditions: - scales = torch.logspace(-1, 0, len(down_block_res_samples) + 1, device=sample.device) # 0.1 to 1.0 - if from_multi or len(control_type_idx) == 1: - scales = scales * conditioning_scale[0] - down_block_res_samples = [sample * scale for sample, scale in zip(down_block_res_samples, scales)] - mid_block_res_sample = mid_block_res_sample * scales[-1] # last one - elif from_multi or len(control_type_idx) == 1: - down_block_res_samples = [sample * conditioning_scale[0] for sample in down_block_res_samples] - mid_block_res_sample = mid_block_res_sample * conditioning_scale[0] - - if self.config.global_pool_conditions: - down_block_res_samples = [ - torch.mean(sample, dim=(2, 3), keepdim=True) for sample in down_block_res_samples - ] - mid_block_res_sample = torch.mean(mid_block_res_sample, dim=(2, 3), keepdim=True) - - if not return_dict: - return (down_block_res_samples, mid_block_res_sample) - - return ControlNetOutput( - down_block_res_samples=down_block_res_samples, mid_block_res_sample=mid_block_res_sample - ) diff --git a/diffusers/models/controlnets/controlnet_xs.py b/diffusers/models/controlnets/controlnet_xs.py deleted file mode 100644 index a25d5d71a5b134bc91e063a24bfcb6922164eeb7..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/controlnet_xs.py +++ /dev/null @@ -1,1835 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from dataclasses import dataclass -from math import gcd -from typing import Any - -import torch -from torch import Tensor, nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput, logging -from ...utils.torch_utils import apply_freeu, maybe_adjust_dtype_for_device -from ..attention import AttentionMixin -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - Attention, - AttnAddedKVProcessor, - AttnProcessor, - FusedAttnProcessor2_0, -) -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from ..unets.unet_2d_blocks import ( - CrossAttnDownBlock2D, - CrossAttnUpBlock2D, - Downsample2D, - ResnetBlock2D, - Transformer2DModel, - UNetMidBlock2DCrossAttn, - Upsample2D, -) -from ..unets.unet_2d_condition import UNet2DConditionModel -from .controlnet import ControlNetConditioningEmbedding - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class ControlNetXSOutput(BaseOutput): - """ - The output of [`UNetControlNetXSModel`]. - - Args: - sample (`Tensor` of shape `(batch_size, num_channels, height, width)`): - The output of the `UNetControlNetXSModel`. Unlike `ControlNetOutput` this is NOT to be added to the base - model output, but is already the final output. - """ - - sample: Tensor = None - - -class DownBlockControlNetXSAdapter(nn.Module): - """Components that together with corresponding components from the base model will form a - `ControlNetXSCrossAttnDownBlock2D`""" - - def __init__( - self, - resnets: nn.ModuleList, - base_to_ctrl: nn.ModuleList, - ctrl_to_base: nn.ModuleList, - attentions: nn.ModuleList | None = None, - downsampler: nn.Conv2d | None = None, - ): - super().__init__() - self.resnets = resnets - self.base_to_ctrl = base_to_ctrl - self.ctrl_to_base = ctrl_to_base - self.attentions = attentions - self.downsamplers = downsampler - - -class MidBlockControlNetXSAdapter(nn.Module): - """Components that together with corresponding components from the base model will form a - `ControlNetXSCrossAttnMidBlock2D`""" - - def __init__(self, midblock: UNetMidBlock2DCrossAttn, base_to_ctrl: nn.ModuleList, ctrl_to_base: nn.ModuleList): - super().__init__() - self.midblock = midblock - self.base_to_ctrl = base_to_ctrl - self.ctrl_to_base = ctrl_to_base - - -class UpBlockControlNetXSAdapter(nn.Module): - """Components that together with corresponding components from the base model will form a `ControlNetXSCrossAttnUpBlock2D`""" - - def __init__(self, ctrl_to_base: nn.ModuleList): - super().__init__() - self.ctrl_to_base = ctrl_to_base - - -def get_down_block_adapter( - base_in_channels: int, - base_out_channels: int, - ctrl_in_channels: int, - ctrl_out_channels: int, - temb_channels: int, - max_norm_num_groups: int | None = 32, - has_crossattn=True, - transformer_layers_per_block: int | tuple[int] | None = 1, - num_attention_heads: int | None = 1, - cross_attention_dim: int | None = 1024, - add_downsample: bool = True, - upcast_attention: bool | None = False, - use_linear_projection: bool | None = True, -): - num_layers = 2 # only support sd + sdxl - - resnets = [] - attentions = [] - ctrl_to_base = [] - base_to_ctrl = [] - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - for i in range(num_layers): - base_in_channels = base_in_channels if i == 0 else base_out_channels - ctrl_in_channels = ctrl_in_channels if i == 0 else ctrl_out_channels - - # Before the resnet/attention application, information is concatted from base to control. - # Concat doesn't require change in number of channels - base_to_ctrl.append(make_zero_conv(base_in_channels, base_in_channels)) - - resnets.append( - ResnetBlock2D( - in_channels=ctrl_in_channels + base_in_channels, # information from base is concatted to ctrl - out_channels=ctrl_out_channels, - temb_channels=temb_channels, - groups=find_largest_factor(ctrl_in_channels + base_in_channels, max_factor=max_norm_num_groups), - groups_out=find_largest_factor(ctrl_out_channels, max_factor=max_norm_num_groups), - eps=1e-5, - ) - ) - - if has_crossattn: - attentions.append( - Transformer2DModel( - num_attention_heads, - ctrl_out_channels // num_attention_heads, - in_channels=ctrl_out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - norm_num_groups=find_largest_factor(ctrl_out_channels, max_factor=max_norm_num_groups), - ) - ) - - # After the resnet/attention application, information is added from control to base - # Addition requires change in number of channels - ctrl_to_base.append(make_zero_conv(ctrl_out_channels, base_out_channels)) - - if add_downsample: - # Before the downsampler application, information is concatted from base to control - # Concat doesn't require change in number of channels - base_to_ctrl.append(make_zero_conv(base_out_channels, base_out_channels)) - - downsamplers = Downsample2D( - ctrl_out_channels + base_out_channels, use_conv=True, out_channels=ctrl_out_channels, name="op" - ) - - # After the downsampler application, information is added from control to base - # Addition requires change in number of channels - ctrl_to_base.append(make_zero_conv(ctrl_out_channels, base_out_channels)) - else: - downsamplers = None - - down_block_components = DownBlockControlNetXSAdapter( - resnets=nn.ModuleList(resnets), - base_to_ctrl=nn.ModuleList(base_to_ctrl), - ctrl_to_base=nn.ModuleList(ctrl_to_base), - ) - - if has_crossattn: - down_block_components.attentions = nn.ModuleList(attentions) - if downsamplers is not None: - down_block_components.downsamplers = downsamplers - - return down_block_components - - -def get_mid_block_adapter( - base_channels: int, - ctrl_channels: int, - temb_channels: int | None = None, - max_norm_num_groups: int | None = 32, - transformer_layers_per_block: int = 1, - num_attention_heads: int | None = 1, - cross_attention_dim: int | None = 1024, - upcast_attention: bool = False, - use_linear_projection: bool = True, -): - # Before the midblock application, information is concatted from base to control. - # Concat doesn't require change in number of channels - base_to_ctrl = make_zero_conv(base_channels, base_channels) - - midblock = UNetMidBlock2DCrossAttn( - transformer_layers_per_block=transformer_layers_per_block, - in_channels=ctrl_channels + base_channels, - out_channels=ctrl_channels, - temb_channels=temb_channels, - # number or norm groups must divide both in_channels and out_channels - resnet_groups=find_largest_factor(gcd(ctrl_channels, ctrl_channels + base_channels), max_norm_num_groups), - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - ) - - # After the midblock application, information is added from control to base - # Addition requires change in number of channels - ctrl_to_base = make_zero_conv(ctrl_channels, base_channels) - - return MidBlockControlNetXSAdapter(base_to_ctrl=base_to_ctrl, midblock=midblock, ctrl_to_base=ctrl_to_base) - - -def get_up_block_adapter( - out_channels: int, - prev_output_channel: int, - ctrl_skip_channels: list[int], -): - ctrl_to_base = [] - num_layers = 3 # only support sd + sdxl - for i in range(num_layers): - resnet_in_channels = prev_output_channel if i == 0 else out_channels - ctrl_to_base.append(make_zero_conv(ctrl_skip_channels[i], resnet_in_channels)) - - return UpBlockControlNetXSAdapter(ctrl_to_base=nn.ModuleList(ctrl_to_base)) - - -class ControlNetXSAdapter(ModelMixin, AttentionMixin, ConfigMixin): - r""" - A `ControlNetXSAdapter` model. To use it, pass it into a `UNetControlNetXSModel` (together with a - `UNet2DConditionModel` base model). - - This model inherits from [`ModelMixin`] and [`ConfigMixin`]. Check the superclass documentation for it's generic - methods implemented for all models (such as downloading or saving). - - Like `UNetControlNetXSModel`, `ControlNetXSAdapter` is compatible with StableDiffusion and StableDiffusion-XL. It's - default parameters are compatible with StableDiffusion. - - Parameters: - conditioning_channels (`int`, defaults to 3): - Number of channels of conditioning input (e.g. an image) - conditioning_channel_order (`str`, defaults to `"rgb"`): - The channel order of conditional image. Will convert to `rgb` if it's `bgr`. - conditioning_embedding_out_channels (`tuple[int]`, defaults to `(16, 32, 96, 256)`): - The tuple of output channels for each block in the `controlnet_cond_embedding` layer. - time_embedding_mix (`float`, defaults to 1.0): - If 0, then only the control adapters's time embedding is used. If 1, then only the base unet's time - embedding is used. Otherwise, both are combined. - learn_time_embedding (`bool`, defaults to `False`): - Whether a time embedding should be learned. If yes, `UNetControlNetXSModel` will combine the time - embeddings of the base model and the control adapter. If no, `UNetControlNetXSModel` will use the base - model's time embedding. - num_attention_heads (`list[int]`, defaults to `[4]`): - The number of attention heads. - block_out_channels (`list[int]`, defaults to `[4, 8, 16, 16]`): - The tuple of output channels for each block. - base_block_out_channels (`list[int]`, defaults to `[320, 640, 1280, 1280]`): - The tuple of output channels for each block in the base unet. - cross_attention_dim (`int`, defaults to 1024): - The dimension of the cross attention features. - down_block_types (`list[str]`, defaults to `["CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "DownBlock2D"]`): - The tuple of downsample blocks to use. - sample_size (`int`, defaults to 96): - Height and width of input/output sample. - transformer_layers_per_block (`int | tuple[int]`, defaults to 1): - The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`]. Only relevant for - [`~models.unet_2d_blocks.CrossAttnDownBlock2D`], [`~models.unet_2d_blocks.UNetMidBlock2DCrossAttn`]. - upcast_attention (`bool`, defaults to `True`): - Whether the attention computation should always be upcasted. - max_norm_num_groups (`int`, defaults to 32): - Maximum number of groups in group normal. The actual number will be the largest divisor of the respective - channels, that is <= max_norm_num_groups. - """ - - @register_to_config - def __init__( - self, - conditioning_channels: int = 3, - conditioning_channel_order: str = "rgb", - conditioning_embedding_out_channels: tuple[int] = (16, 32, 96, 256), - time_embedding_mix: float = 1.0, - learn_time_embedding: bool = False, - num_attention_heads: int | tuple[int] = 4, - block_out_channels: tuple[int] = (4, 8, 16, 16), - base_block_out_channels: tuple[int] = (320, 640, 1280, 1280), - cross_attention_dim: int = 1024, - down_block_types: tuple[str] = ( - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "DownBlock2D", - ), - sample_size: int | None = 96, - transformer_layers_per_block: int | tuple[int] = 1, - upcast_attention: bool = True, - max_norm_num_groups: int = 32, - use_linear_projection: bool = True, - ): - super().__init__() - - time_embedding_input_dim = base_block_out_channels[0] - time_embedding_dim = base_block_out_channels[0] * 4 - - # Check inputs - if conditioning_channel_order not in ["rgb", "bgr"]: - raise ValueError(f"unknown `conditioning_channel_order`: {conditioning_channel_order}") - - if len(block_out_channels) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(transformer_layers_per_block, (list, tuple)): - transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types) - if not isinstance(cross_attention_dim, (list, tuple)): - cross_attention_dim = [cross_attention_dim] * len(down_block_types) - # see https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131 for why `ControlNetXSAdapter` takes `num_attention_heads` instead of `attention_head_dim` - if not isinstance(num_attention_heads, (list, tuple)): - num_attention_heads = [num_attention_heads] * len(down_block_types) - - if len(num_attention_heads) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}." - ) - - # 5 - Create conditioning hint embedding - self.controlnet_cond_embedding = ControlNetConditioningEmbedding( - conditioning_embedding_channels=block_out_channels[0], - block_out_channels=conditioning_embedding_out_channels, - conditioning_channels=conditioning_channels, - ) - - # time - if learn_time_embedding: - self.time_embedding = TimestepEmbedding(time_embedding_input_dim, time_embedding_dim) - else: - self.time_embedding = None - - self.down_blocks = nn.ModuleList([]) - self.up_connections = nn.ModuleList([]) - - # input - self.conv_in = nn.Conv2d(4, block_out_channels[0], kernel_size=3, padding=1) - self.control_to_base_for_conv_in = make_zero_conv(block_out_channels[0], base_block_out_channels[0]) - - # down - base_out_channels = base_block_out_channels[0] - ctrl_out_channels = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - base_in_channels = base_out_channels - base_out_channels = base_block_out_channels[i] - ctrl_in_channels = ctrl_out_channels - ctrl_out_channels = block_out_channels[i] - has_crossattn = "CrossAttn" in down_block_type - is_final_block = i == len(down_block_types) - 1 - - self.down_blocks.append( - get_down_block_adapter( - base_in_channels=base_in_channels, - base_out_channels=base_out_channels, - ctrl_in_channels=ctrl_in_channels, - ctrl_out_channels=ctrl_out_channels, - temb_channels=time_embedding_dim, - max_norm_num_groups=max_norm_num_groups, - has_crossattn=has_crossattn, - transformer_layers_per_block=transformer_layers_per_block[i], - num_attention_heads=num_attention_heads[i], - cross_attention_dim=cross_attention_dim[i], - add_downsample=not is_final_block, - upcast_attention=upcast_attention, - use_linear_projection=use_linear_projection, - ) - ) - - # mid - self.mid_block = get_mid_block_adapter( - base_channels=base_block_out_channels[-1], - ctrl_channels=block_out_channels[-1], - temb_channels=time_embedding_dim, - transformer_layers_per_block=transformer_layers_per_block[-1], - num_attention_heads=num_attention_heads[-1], - cross_attention_dim=cross_attention_dim[-1], - upcast_attention=upcast_attention, - use_linear_projection=use_linear_projection, - ) - - # up - # The skip connection channels are the output of the conv_in and of all the down subblocks - ctrl_skip_channels = [block_out_channels[0]] - for i, out_channels in enumerate(block_out_channels): - number_of_subblocks = ( - 3 if i < len(block_out_channels) - 1 else 2 - ) # every block has 3 subblocks, except last one, which has 2 as it has no downsampler - ctrl_skip_channels.extend([out_channels] * number_of_subblocks) - - reversed_base_block_out_channels = list(reversed(base_block_out_channels)) - - base_out_channels = reversed_base_block_out_channels[0] - for i in range(len(down_block_types)): - prev_base_output_channel = base_out_channels - base_out_channels = reversed_base_block_out_channels[i] - ctrl_skip_channels_ = [ctrl_skip_channels.pop() for _ in range(3)] - - self.up_connections.append( - get_up_block_adapter( - out_channels=base_out_channels, - prev_output_channel=prev_base_output_channel, - ctrl_skip_channels=ctrl_skip_channels_, - ) - ) - - @classmethod - def from_unet( - cls, - unet: UNet2DConditionModel, - size_ratio: float | None = None, - block_out_channels: list[int] | None = None, - num_attention_heads: list[int] | None = None, - learn_time_embedding: bool = False, - time_embedding_mix: int = 1.0, - conditioning_channels: int = 3, - conditioning_channel_order: str = "rgb", - conditioning_embedding_out_channels: tuple[int] = (16, 32, 96, 256), - ): - r""" - Instantiate a [`ControlNetXSAdapter`] from a [`UNet2DConditionModel`]. - - Parameters: - unet (`UNet2DConditionModel`): - The UNet model we want to control. The dimensions of the ControlNetXSAdapter will be adapted to it. - size_ratio (float, *optional*, defaults to `None`): - When given, block_out_channels is set to a fraction of the base model's block_out_channels. Either this - or `block_out_channels` must be given. - block_out_channels (`list[int]`, *optional*, defaults to `None`): - Down blocks output channels in control model. Either this or `size_ratio` must be given. - num_attention_heads (`list[int]`, *optional*, defaults to `None`): - The dimension of the attention heads. The naming seems a bit confusing and it is, see - https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131 for why. - learn_time_embedding (`bool`, defaults to `False`): - Whether the `ControlNetXSAdapter` should learn a time embedding. - time_embedding_mix (`float`, defaults to 1.0): - If 0, then only the control adapter's time embedding is used. If 1, then only the base unet's time - embedding is used. Otherwise, both are combined. - conditioning_channels (`int`, defaults to 3): - Number of channels of conditioning input (e.g. an image) - conditioning_channel_order (`str`, defaults to `"rgb"`): - The channel order of conditional image. Will convert to `rgb` if it's `bgr`. - conditioning_embedding_out_channels (`tuple[int]`, defaults to `(16, 32, 96, 256)`): - The tuple of output channel for each block in the `controlnet_cond_embedding` layer. - """ - - # Check input - fixed_size = block_out_channels is not None - relative_size = size_ratio is not None - if not (fixed_size ^ relative_size): - raise ValueError( - "Pass exactly one of `block_out_channels` (for absolute sizing) or `size_ratio` (for relative sizing)." - ) - - # Create model - block_out_channels = block_out_channels or [int(b * size_ratio) for b in unet.config.block_out_channels] - if num_attention_heads is None: - # The naming seems a bit confusing and it is, see https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131 for why. - num_attention_heads = unet.config.attention_head_dim - - model = cls( - conditioning_channels=conditioning_channels, - conditioning_channel_order=conditioning_channel_order, - conditioning_embedding_out_channels=conditioning_embedding_out_channels, - time_embedding_mix=time_embedding_mix, - learn_time_embedding=learn_time_embedding, - num_attention_heads=num_attention_heads, - block_out_channels=block_out_channels, - base_block_out_channels=unet.config.block_out_channels, - cross_attention_dim=unet.config.cross_attention_dim, - down_block_types=unet.config.down_block_types, - sample_size=unet.config.sample_size, - transformer_layers_per_block=unet.config.transformer_layers_per_block, - upcast_attention=unet.config.upcast_attention, - max_norm_num_groups=unet.config.norm_num_groups, - use_linear_projection=unet.config.use_linear_projection, - ) - - # ensure that the ControlNetXSAdapter is the same dtype as the UNet2DConditionModel - model.to(unet.dtype) - - return model - - def forward(self, *args, **kwargs): - raise ValueError( - "A ControlNetXSAdapter cannot be run by itself. Use it together with a UNet2DConditionModel to instantiate a UNetControlNetXSModel." - ) - - -class UNetControlNetXSModel(ModelMixin, AttentionMixin, ConfigMixin): - r""" - A UNet fused with a ControlNet-XS adapter model - - This model inherits from [`ModelMixin`] and [`ConfigMixin`]. Check the superclass documentation for it's generic - methods implemented for all models (such as downloading or saving). - - `UNetControlNetXSModel` is compatible with StableDiffusion and StableDiffusion-XL. It's default parameters are - compatible with StableDiffusion. - - It's parameters are either passed to the underlying `UNet2DConditionModel` or used exactly like in - `ControlNetXSAdapter` . See their documentation for details. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - # unet configs - sample_size: int | None = 96, - down_block_types: tuple[str] = ( - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "DownBlock2D", - ), - up_block_types: tuple[str] = ("UpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D"), - block_out_channels: tuple[int] = (320, 640, 1280, 1280), - norm_num_groups: int | None = 32, - cross_attention_dim: int | tuple[int] = 1024, - transformer_layers_per_block: int | tuple[int] = 1, - num_attention_heads: int | tuple[int] = 8, - addition_embed_type: str | None = None, - addition_time_embed_dim: int | None = None, - upcast_attention: bool = True, - use_linear_projection: bool = True, - time_cond_proj_dim: int | None = None, - projection_class_embeddings_input_dim: int | None = None, - # additional controlnet configs - time_embedding_mix: float = 1.0, - ctrl_conditioning_channels: int = 3, - ctrl_conditioning_embedding_out_channels: tuple[int] = (16, 32, 96, 256), - ctrl_conditioning_channel_order: str = "rgb", - ctrl_learn_time_embedding: bool = False, - ctrl_block_out_channels: tuple[int] = (4, 8, 16, 16), - ctrl_num_attention_heads: int | tuple[int] = 4, - ctrl_max_norm_num_groups: int = 32, - ): - super().__init__() - - if time_embedding_mix < 0 or time_embedding_mix > 1: - raise ValueError("`time_embedding_mix` needs to be between 0 and 1.") - if time_embedding_mix < 1 and not ctrl_learn_time_embedding: - raise ValueError("To use `time_embedding_mix` < 1, `ctrl_learn_time_embedding` must be `True`") - - if addition_embed_type is not None and addition_embed_type != "text_time": - raise ValueError( - "As `UNetControlNetXSModel` currently only supports StableDiffusion and StableDiffusion-XL, `addition_embed_type` must be `None` or `'text_time'`." - ) - - if not isinstance(transformer_layers_per_block, (list, tuple)): - transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types) - if not isinstance(cross_attention_dim, (list, tuple)): - cross_attention_dim = [cross_attention_dim] * len(down_block_types) - if not isinstance(num_attention_heads, (list, tuple)): - num_attention_heads = [num_attention_heads] * len(down_block_types) - if not isinstance(ctrl_num_attention_heads, (list, tuple)): - ctrl_num_attention_heads = [ctrl_num_attention_heads] * len(down_block_types) - - base_num_attention_heads = num_attention_heads - - self.in_channels = 4 - - # # Input - self.base_conv_in = nn.Conv2d(4, block_out_channels[0], kernel_size=3, padding=1) - self.controlnet_cond_embedding = ControlNetConditioningEmbedding( - conditioning_embedding_channels=ctrl_block_out_channels[0], - block_out_channels=ctrl_conditioning_embedding_out_channels, - conditioning_channels=ctrl_conditioning_channels, - ) - self.ctrl_conv_in = nn.Conv2d(4, ctrl_block_out_channels[0], kernel_size=3, padding=1) - self.control_to_base_for_conv_in = make_zero_conv(ctrl_block_out_channels[0], block_out_channels[0]) - - # # Time - time_embed_input_dim = block_out_channels[0] - time_embed_dim = block_out_channels[0] * 4 - - self.base_time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos=True, downscale_freq_shift=0) - self.base_time_embedding = TimestepEmbedding( - time_embed_input_dim, - time_embed_dim, - cond_proj_dim=time_cond_proj_dim, - ) - if ctrl_learn_time_embedding: - self.ctrl_time_embedding = TimestepEmbedding( - in_channels=time_embed_input_dim, time_embed_dim=time_embed_dim - ) - else: - self.ctrl_time_embedding = None - - if addition_embed_type is None: - self.base_add_time_proj = None - self.base_add_embedding = None - else: - self.base_add_time_proj = Timesteps(addition_time_embed_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.base_add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) - - # # Create down blocks - down_blocks = [] - base_out_channels = block_out_channels[0] - ctrl_out_channels = ctrl_block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - base_in_channels = base_out_channels - base_out_channels = block_out_channels[i] - ctrl_in_channels = ctrl_out_channels - ctrl_out_channels = ctrl_block_out_channels[i] - has_crossattn = "CrossAttn" in down_block_type - is_final_block = i == len(down_block_types) - 1 - - down_blocks.append( - ControlNetXSCrossAttnDownBlock2D( - base_in_channels=base_in_channels, - base_out_channels=base_out_channels, - ctrl_in_channels=ctrl_in_channels, - ctrl_out_channels=ctrl_out_channels, - temb_channels=time_embed_dim, - norm_num_groups=norm_num_groups, - ctrl_max_norm_num_groups=ctrl_max_norm_num_groups, - has_crossattn=has_crossattn, - transformer_layers_per_block=transformer_layers_per_block[i], - base_num_attention_heads=base_num_attention_heads[i], - ctrl_num_attention_heads=ctrl_num_attention_heads[i], - cross_attention_dim=cross_attention_dim[i], - add_downsample=not is_final_block, - upcast_attention=upcast_attention, - use_linear_projection=use_linear_projection, - ) - ) - - # # Create mid block - self.mid_block = ControlNetXSCrossAttnMidBlock2D( - base_channels=block_out_channels[-1], - ctrl_channels=ctrl_block_out_channels[-1], - temb_channels=time_embed_dim, - norm_num_groups=norm_num_groups, - ctrl_max_norm_num_groups=ctrl_max_norm_num_groups, - transformer_layers_per_block=transformer_layers_per_block[-1], - base_num_attention_heads=base_num_attention_heads[-1], - ctrl_num_attention_heads=ctrl_num_attention_heads[-1], - cross_attention_dim=cross_attention_dim[-1], - upcast_attention=upcast_attention, - use_linear_projection=use_linear_projection, - ) - - # # Create up blocks - up_blocks = [] - rev_transformer_layers_per_block = list(reversed(transformer_layers_per_block)) - rev_num_attention_heads = list(reversed(base_num_attention_heads)) - rev_cross_attention_dim = list(reversed(cross_attention_dim)) - - # The skip connection channels are the output of the conv_in and of all the down subblocks - ctrl_skip_channels = [ctrl_block_out_channels[0]] - for i, out_channels in enumerate(ctrl_block_out_channels): - number_of_subblocks = ( - 3 if i < len(ctrl_block_out_channels) - 1 else 2 - ) # every block has 3 subblocks, except last one, which has 2 as it has no downsampler - ctrl_skip_channels.extend([out_channels] * number_of_subblocks) - - reversed_block_out_channels = list(reversed(block_out_channels)) - - out_channels = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - prev_output_channel = out_channels - out_channels = reversed_block_out_channels[i] - in_channels = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)] - ctrl_skip_channels_ = [ctrl_skip_channels.pop() for _ in range(3)] - - has_crossattn = "CrossAttn" in up_block_type - is_final_block = i == len(block_out_channels) - 1 - - up_blocks.append( - ControlNetXSCrossAttnUpBlock2D( - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - ctrl_skip_channels=ctrl_skip_channels_, - temb_channels=time_embed_dim, - resolution_idx=i, - has_crossattn=has_crossattn, - transformer_layers_per_block=rev_transformer_layers_per_block[i], - num_attention_heads=rev_num_attention_heads[i], - cross_attention_dim=rev_cross_attention_dim[i], - add_upsample=not is_final_block, - upcast_attention=upcast_attention, - norm_num_groups=norm_num_groups, - use_linear_projection=use_linear_projection, - ) - ) - - self.down_blocks = nn.ModuleList(down_blocks) - self.up_blocks = nn.ModuleList(up_blocks) - - self.base_conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups) - self.base_conv_act = nn.SiLU() - self.base_conv_out = nn.Conv2d(block_out_channels[0], 4, kernel_size=3, padding=1) - - @classmethod - def from_unet( - cls, - unet: UNet2DConditionModel, - controlnet: ControlNetXSAdapter | None = None, - size_ratio: float | None = None, - ctrl_block_out_channels: list[float] | None = None, - time_embedding_mix: float | None = None, - ctrl_optional_kwargs: dict | None = None, - ): - r""" - Instantiate a [`UNetControlNetXSModel`] from a [`UNet2DConditionModel`] and an optional [`ControlNetXSAdapter`] - . - - Parameters: - unet (`UNet2DConditionModel`): - The UNet model we want to control. - controlnet (`ControlNetXSAdapter`): - The ControlNet-XS adapter with which the UNet will be fused. If none is given, a new ControlNet-XS - adapter will be created. - size_ratio (float, *optional*, defaults to `None`): - Used to construct the controlnet if none is given. See [`ControlNetXSAdapter.from_unet`] for details. - ctrl_block_out_channels (`list[int]`, *optional*, defaults to `None`): - Used to construct the controlnet if none is given. See [`ControlNetXSAdapter.from_unet`] for details, - where this parameter is called `block_out_channels`. - time_embedding_mix (`float`, *optional*, defaults to None): - Used to construct the controlnet if none is given. See [`ControlNetXSAdapter.from_unet`] for details. - ctrl_optional_kwargs (`Dict`, *optional*, defaults to `None`): - Passed to the `init` of the new controlnet if no controlnet was given. - """ - if controlnet is None: - controlnet = ControlNetXSAdapter.from_unet( - unet, size_ratio, ctrl_block_out_channels, **ctrl_optional_kwargs - ) - else: - if any( - o is not None for o in (size_ratio, ctrl_block_out_channels, time_embedding_mix, ctrl_optional_kwargs) - ): - raise ValueError( - "When a controlnet is passed, none of these parameters should be passed: size_ratio, ctrl_block_out_channels, time_embedding_mix, ctrl_optional_kwargs." - ) - - # # get params - params_for_unet = [ - "sample_size", - "down_block_types", - "up_block_types", - "block_out_channels", - "norm_num_groups", - "cross_attention_dim", - "transformer_layers_per_block", - "addition_embed_type", - "addition_time_embed_dim", - "upcast_attention", - "use_linear_projection", - "time_cond_proj_dim", - "projection_class_embeddings_input_dim", - ] - params_for_unet = {k: v for k, v in unet.config.items() if k in params_for_unet} - # The naming seems a bit confusing and it is, see https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131 for why. - params_for_unet["num_attention_heads"] = unet.config.attention_head_dim - - params_for_controlnet = [ - "conditioning_channels", - "conditioning_embedding_out_channels", - "conditioning_channel_order", - "learn_time_embedding", - "block_out_channels", - "num_attention_heads", - "max_norm_num_groups", - ] - params_for_controlnet = {"ctrl_" + k: v for k, v in controlnet.config.items() if k in params_for_controlnet} - params_for_controlnet["time_embedding_mix"] = controlnet.config.time_embedding_mix - - # # create model - model = cls.from_config({**params_for_unet, **params_for_controlnet}) - - # # load weights - # from unet - modules_from_unet = [ - "time_embedding", - "conv_in", - "conv_norm_out", - "conv_out", - ] - for m in modules_from_unet: - getattr(model, "base_" + m).load_state_dict(getattr(unet, m).state_dict()) - - optional_modules_from_unet = [ - "add_time_proj", - "add_embedding", - ] - for m in optional_modules_from_unet: - if hasattr(unet, m) and getattr(unet, m) is not None: - getattr(model, "base_" + m).load_state_dict(getattr(unet, m).state_dict()) - - # from controlnet - model.controlnet_cond_embedding.load_state_dict(controlnet.controlnet_cond_embedding.state_dict()) - model.ctrl_conv_in.load_state_dict(controlnet.conv_in.state_dict()) - if controlnet.time_embedding is not None: - model.ctrl_time_embedding.load_state_dict(controlnet.time_embedding.state_dict()) - model.control_to_base_for_conv_in.load_state_dict(controlnet.control_to_base_for_conv_in.state_dict()) - - # from both - model.down_blocks = nn.ModuleList( - ControlNetXSCrossAttnDownBlock2D.from_modules(b, c) - for b, c in zip(unet.down_blocks, controlnet.down_blocks) - ) - model.mid_block = ControlNetXSCrossAttnMidBlock2D.from_modules(unet.mid_block, controlnet.mid_block) - model.up_blocks = nn.ModuleList( - ControlNetXSCrossAttnUpBlock2D.from_modules(b, c) - for b, c in zip(unet.up_blocks, controlnet.up_connections) - ) - - # ensure that the UNetControlNetXSModel is the same dtype as the UNet2DConditionModel - model.to(unet.dtype) - - return model - - def freeze_unet_params(self) -> None: - """Freeze the weights of the parts belonging to the base UNet2DConditionModel, and leave everything else unfrozen for fine - tuning.""" - # Freeze everything - for param in self.parameters(): - param.requires_grad = True - - # Unfreeze ControlNetXSAdapter - base_parts = [ - "base_time_proj", - "base_time_embedding", - "base_add_time_proj", - "base_add_embedding", - "base_conv_in", - "base_conv_norm_out", - "base_conv_act", - "base_conv_out", - ] - base_parts = [getattr(self, part) for part in base_parts if getattr(self, part) is not None] - for part in base_parts: - for param in part.parameters(): - param.requires_grad = False - - for d in self.down_blocks: - d.freeze_base_params() - self.mid_block.freeze_base_params() - for u in self.up_blocks: - u.freeze_base_params() - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.enable_freeu - def enable_freeu(self, s1: float, s2: float, b1: float, b2: float): - r"""Enables the FreeU mechanism from https://huggingface.co/papers/2309.11497. - - The suffixes after the scaling factors represent the stage blocks where they are being applied. - - Please refer to the [official repository](https://github.com/ChenyangSi/FreeU) for combinations of values that - are known to work well for different pipelines such as Stable Diffusion v1, v2, and Stable Diffusion XL. - - Args: - s1 (`float`): - Scaling factor for stage 1 to attenuate the contributions of the skip features. This is done to - mitigate the "oversmoothing effect" in the enhanced denoising process. - s2 (`float`): - Scaling factor for stage 2 to attenuate the contributions of the skip features. This is done to - mitigate the "oversmoothing effect" in the enhanced denoising process. - b1 (`float`): Scaling factor for stage 1 to amplify the contributions of backbone features. - b2 (`float`): Scaling factor for stage 2 to amplify the contributions of backbone features. - """ - for i, upsample_block in enumerate(self.up_blocks): - setattr(upsample_block, "s1", s1) - setattr(upsample_block, "s2", s2) - setattr(upsample_block, "b1", b1) - setattr(upsample_block, "b2", b2) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.disable_freeu - def disable_freeu(self): - """Disables the FreeU mechanism.""" - freeu_keys = {"s1", "s2", "b1", "b2"} - for i, upsample_block in enumerate(self.up_blocks): - for k in freeu_keys: - if hasattr(upsample_block, k) or getattr(upsample_block, k, None) is not None: - setattr(upsample_block, k, None) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) - - def forward( - self, - sample: Tensor, - timestep: torch.Tensor | float | int, - encoder_hidden_states: torch.Tensor, - controlnet_cond: torch.Tensor | None = None, - conditioning_scale: float | None = 1.0, - class_labels: torch.Tensor | None = None, - timestep_cond: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - added_cond_kwargs: dict[str, torch.Tensor] | None = None, - return_dict: bool = True, - apply_control: bool = True, - ) -> ControlNetXSOutput | tuple: - """ - The [`ControlNetXSModel`] forward method. - - Args: - sample (`Tensor`): - The noisy input tensor. - timestep (`torch.Tensor | float | int`): - The number of timesteps to denoise an input. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states. - controlnet_cond (`Tensor`): - The conditional input tensor of shape `(batch_size, sequence_length, hidden_size)`. - conditioning_scale (`float`, defaults to `1.0`): - How much the control model affects the base model outputs. - class_labels (`torch.Tensor`, *optional*, defaults to `None`): - Optional class labels for conditioning. Their embeddings will be summed with the timestep embeddings. - timestep_cond (`torch.Tensor`, *optional*, defaults to `None`): - Additional conditional embeddings for timestep. If provided, the embeddings will be summed with the - timestep_embedding passed through the `self.time_embedding` layer to obtain the final timestep - embeddings. - attention_mask (`torch.Tensor`, *optional*, defaults to `None`): - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask - is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large - negative values to the attention scores corresponding to "discard" tokens. - cross_attention_kwargs (`dict[str]`, *optional*, defaults to `None`): - A kwargs dictionary that if specified is passed along to the `AttnProcessor`. - added_cond_kwargs (`dict`): - Additional conditions for the Stable Diffusion XL UNet. - return_dict (`bool`, defaults to `True`): - Whether or not to return a [`~models.controlnets.controlnet.ControlNetOutput`] instead of a plain - tuple. - apply_control (`bool`, defaults to `True`): - If `False`, the input is run only through the base model. - - Returns: - [`~models.controlnetxs.ControlNetXSOutput`] **or** `tuple`: - If `return_dict` is `True`, a [`~models.controlnetxs.ControlNetXSOutput`] is returned, otherwise a - tuple is returned where the first element is the sample tensor. - """ - - # check channel order - if self.config.ctrl_conditioning_channel_order == "bgr": - controlnet_cond = torch.flip(controlnet_cond, dims=[1]) - - # prepare attention_mask - if attention_mask is not None: - attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # 1. time - timesteps = timestep - if not torch.is_tensor(timesteps): - # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can - # This would be a good case for the `match` statement (Python 3.10+) - dtype = maybe_adjust_dtype_for_device( - torch.float64 if isinstance(timestep, float) else torch.int64, sample.device - ) - timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) - elif len(timesteps.shape) == 0: - timesteps = timesteps[None].to(sample.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - timesteps = timesteps.expand(sample.shape[0]) - - t_emb = self.base_time_proj(timesteps) - - # timesteps does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=sample.dtype) - - if self.config.ctrl_learn_time_embedding and apply_control: - ctrl_temb = self.ctrl_time_embedding(t_emb, timestep_cond) - base_temb = self.base_time_embedding(t_emb, timestep_cond) - interpolation_param = self.config.time_embedding_mix**0.3 - - temb = ctrl_temb * interpolation_param + base_temb * (1 - interpolation_param) - else: - temb = self.base_time_embedding(t_emb) - - # added time & text embeddings - aug_emb = None - - if self.config.addition_embed_type is None: - pass - elif self.config.addition_embed_type == "text_time": - # SDXL - style - if "text_embeds" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `text_embeds` to be passed in `added_cond_kwargs`" - ) - text_embeds = added_cond_kwargs.get("text_embeds") - if "time_ids" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `time_ids` to be passed in `added_cond_kwargs`" - ) - time_ids = added_cond_kwargs.get("time_ids") - time_embeds = self.base_add_time_proj(time_ids.flatten()) - time_embeds = time_embeds.reshape((text_embeds.shape[0], -1)) - add_embeds = torch.concat([text_embeds, time_embeds], dim=-1) - add_embeds = add_embeds.to(temb.dtype) - aug_emb = self.base_add_embedding(add_embeds) - else: - raise ValueError( - f"ControlNet-XS currently only supports StableDiffusion and StableDiffusion-XL, so addition_embed_type = {self.config.addition_embed_type} is currently not supported." - ) - - temb = temb + aug_emb if aug_emb is not None else temb - - # text embeddings - cemb = encoder_hidden_states - - # Preparation - h_ctrl = h_base = sample - hs_base, hs_ctrl = [], [] - - # Cross Control - guided_hint = self.controlnet_cond_embedding(controlnet_cond) - - # 1 - conv in & down - - h_base = self.base_conv_in(h_base) - h_ctrl = self.ctrl_conv_in(h_ctrl) - if guided_hint is not None: - h_ctrl += guided_hint - if apply_control: - h_base = h_base + self.control_to_base_for_conv_in(h_ctrl) * conditioning_scale # add ctrl -> base - - hs_base.append(h_base) - hs_ctrl.append(h_ctrl) - - for down in self.down_blocks: - h_base, h_ctrl, residual_hb, residual_hc = down( - hidden_states_base=h_base, - hidden_states_ctrl=h_ctrl, - temb=temb, - encoder_hidden_states=cemb, - conditioning_scale=conditioning_scale, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - apply_control=apply_control, - ) - hs_base.extend(residual_hb) - hs_ctrl.extend(residual_hc) - - # 2 - mid - h_base, h_ctrl = self.mid_block( - hidden_states_base=h_base, - hidden_states_ctrl=h_ctrl, - temb=temb, - encoder_hidden_states=cemb, - conditioning_scale=conditioning_scale, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - apply_control=apply_control, - ) - - # 3 - up - for up in self.up_blocks: - n_resnets = len(up.resnets) - skips_hb = hs_base[-n_resnets:] - skips_hc = hs_ctrl[-n_resnets:] - hs_base = hs_base[:-n_resnets] - hs_ctrl = hs_ctrl[:-n_resnets] - h_base = up( - hidden_states=h_base, - res_hidden_states_tuple_base=skips_hb, - res_hidden_states_tuple_ctrl=skips_hc, - temb=temb, - encoder_hidden_states=cemb, - conditioning_scale=conditioning_scale, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - apply_control=apply_control, - ) - - # 4 - conv out - h_base = self.base_conv_norm_out(h_base) - h_base = self.base_conv_act(h_base) - h_base = self.base_conv_out(h_base) - - if not return_dict: - return (h_base,) - - return ControlNetXSOutput(sample=h_base) - - -class ControlNetXSCrossAttnDownBlock2D(nn.Module): - def __init__( - self, - base_in_channels: int, - base_out_channels: int, - ctrl_in_channels: int, - ctrl_out_channels: int, - temb_channels: int, - norm_num_groups: int = 32, - ctrl_max_norm_num_groups: int = 32, - has_crossattn=True, - transformer_layers_per_block: int | tuple[int] | None = 1, - base_num_attention_heads: int | None = 1, - ctrl_num_attention_heads: int | None = 1, - cross_attention_dim: int | None = 1024, - add_downsample: bool = True, - upcast_attention: bool | None = False, - use_linear_projection: bool | None = True, - ): - super().__init__() - base_resnets = [] - base_attentions = [] - ctrl_resnets = [] - ctrl_attentions = [] - ctrl_to_base = [] - base_to_ctrl = [] - - num_layers = 2 # only support sd + sdxl - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - for i in range(num_layers): - base_in_channels = base_in_channels if i == 0 else base_out_channels - ctrl_in_channels = ctrl_in_channels if i == 0 else ctrl_out_channels - - # Before the resnet/attention application, information is concatted from base to control. - # Concat doesn't require change in number of channels - base_to_ctrl.append(make_zero_conv(base_in_channels, base_in_channels)) - - base_resnets.append( - ResnetBlock2D( - in_channels=base_in_channels, - out_channels=base_out_channels, - temb_channels=temb_channels, - groups=norm_num_groups, - ) - ) - ctrl_resnets.append( - ResnetBlock2D( - in_channels=ctrl_in_channels + base_in_channels, # information from base is concatted to ctrl - out_channels=ctrl_out_channels, - temb_channels=temb_channels, - groups=find_largest_factor( - ctrl_in_channels + base_in_channels, max_factor=ctrl_max_norm_num_groups - ), - groups_out=find_largest_factor(ctrl_out_channels, max_factor=ctrl_max_norm_num_groups), - eps=1e-5, - ) - ) - - if has_crossattn: - base_attentions.append( - Transformer2DModel( - base_num_attention_heads, - base_out_channels // base_num_attention_heads, - in_channels=base_out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - norm_num_groups=norm_num_groups, - ) - ) - ctrl_attentions.append( - Transformer2DModel( - ctrl_num_attention_heads, - ctrl_out_channels // ctrl_num_attention_heads, - in_channels=ctrl_out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - norm_num_groups=find_largest_factor(ctrl_out_channels, max_factor=ctrl_max_norm_num_groups), - ) - ) - - # After the resnet/attention application, information is added from control to base - # Addition requires change in number of channels - ctrl_to_base.append(make_zero_conv(ctrl_out_channels, base_out_channels)) - - if add_downsample: - # Before the downsampler application, information is concatted from base to control - # Concat doesn't require change in number of channels - base_to_ctrl.append(make_zero_conv(base_out_channels, base_out_channels)) - - self.base_downsamplers = Downsample2D( - base_out_channels, use_conv=True, out_channels=base_out_channels, name="op" - ) - self.ctrl_downsamplers = Downsample2D( - ctrl_out_channels + base_out_channels, use_conv=True, out_channels=ctrl_out_channels, name="op" - ) - - # After the downsampler application, information is added from control to base - # Addition requires change in number of channels - ctrl_to_base.append(make_zero_conv(ctrl_out_channels, base_out_channels)) - else: - self.base_downsamplers = None - self.ctrl_downsamplers = None - - self.base_resnets = nn.ModuleList(base_resnets) - self.ctrl_resnets = nn.ModuleList(ctrl_resnets) - self.base_attentions = nn.ModuleList(base_attentions) if has_crossattn else [None] * num_layers - self.ctrl_attentions = nn.ModuleList(ctrl_attentions) if has_crossattn else [None] * num_layers - self.base_to_ctrl = nn.ModuleList(base_to_ctrl) - self.ctrl_to_base = nn.ModuleList(ctrl_to_base) - - self.gradient_checkpointing = False - - @classmethod - def from_modules(cls, base_downblock: CrossAttnDownBlock2D, ctrl_downblock: DownBlockControlNetXSAdapter): - # get params - def get_first_cross_attention(block): - return block.attentions[0].transformer_blocks[0].attn2 - - base_in_channels = base_downblock.resnets[0].in_channels - base_out_channels = base_downblock.resnets[0].out_channels - ctrl_in_channels = ( - ctrl_downblock.resnets[0].in_channels - base_in_channels - ) # base channels are concatted to ctrl channels in init - ctrl_out_channels = ctrl_downblock.resnets[0].out_channels - temb_channels = base_downblock.resnets[0].time_emb_proj.in_features - num_groups = base_downblock.resnets[0].norm1.num_groups - ctrl_num_groups = ctrl_downblock.resnets[0].norm1.num_groups - if hasattr(base_downblock, "attentions"): - has_crossattn = True - transformer_layers_per_block = len(base_downblock.attentions[0].transformer_blocks) - base_num_attention_heads = get_first_cross_attention(base_downblock).heads - ctrl_num_attention_heads = get_first_cross_attention(ctrl_downblock).heads - cross_attention_dim = get_first_cross_attention(base_downblock).cross_attention_dim - upcast_attention = get_first_cross_attention(base_downblock).upcast_attention - use_linear_projection = base_downblock.attentions[0].use_linear_projection - else: - has_crossattn = False - transformer_layers_per_block = None - base_num_attention_heads = None - ctrl_num_attention_heads = None - cross_attention_dim = None - upcast_attention = None - use_linear_projection = None - add_downsample = base_downblock.downsamplers is not None - - # create model - model = cls( - base_in_channels=base_in_channels, - base_out_channels=base_out_channels, - ctrl_in_channels=ctrl_in_channels, - ctrl_out_channels=ctrl_out_channels, - temb_channels=temb_channels, - norm_num_groups=num_groups, - ctrl_max_norm_num_groups=ctrl_num_groups, - has_crossattn=has_crossattn, - transformer_layers_per_block=transformer_layers_per_block, - base_num_attention_heads=base_num_attention_heads, - ctrl_num_attention_heads=ctrl_num_attention_heads, - cross_attention_dim=cross_attention_dim, - add_downsample=add_downsample, - upcast_attention=upcast_attention, - use_linear_projection=use_linear_projection, - ) - - # # load weights - model.base_resnets.load_state_dict(base_downblock.resnets.state_dict()) - model.ctrl_resnets.load_state_dict(ctrl_downblock.resnets.state_dict()) - if has_crossattn: - model.base_attentions.load_state_dict(base_downblock.attentions.state_dict()) - model.ctrl_attentions.load_state_dict(ctrl_downblock.attentions.state_dict()) - if add_downsample: - model.base_downsamplers.load_state_dict(base_downblock.downsamplers[0].state_dict()) - model.ctrl_downsamplers.load_state_dict(ctrl_downblock.downsamplers.state_dict()) - model.base_to_ctrl.load_state_dict(ctrl_downblock.base_to_ctrl.state_dict()) - model.ctrl_to_base.load_state_dict(ctrl_downblock.ctrl_to_base.state_dict()) - - return model - - def freeze_base_params(self) -> None: - """Freeze the weights of the parts belonging to the base UNet2DConditionModel, and leave everything else unfrozen for fine - tuning.""" - # Unfreeze everything - for param in self.parameters(): - param.requires_grad = True - - # Freeze base part - base_parts = [self.base_resnets] - if isinstance(self.base_attentions, nn.ModuleList): # attentions can be a list of Nones - base_parts.append(self.base_attentions) - if self.base_downsamplers is not None: - base_parts.append(self.base_downsamplers) - for part in base_parts: - for param in part.parameters(): - param.requires_grad = False - - def forward( - self, - hidden_states_base: Tensor, - temb: Tensor, - encoder_hidden_states: Tensor | None = None, - hidden_states_ctrl: Tensor | None = None, - conditioning_scale: float | None = 1.0, - attention_mask: Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: Tensor | None = None, - apply_control: bool = True, - ) -> tuple[Tensor, Tensor, tuple[Tensor, ...], tuple[Tensor, ...]]: - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - h_base = hidden_states_base - h_ctrl = hidden_states_ctrl - - base_output_states = () - ctrl_output_states = () - - base_blocks = list(zip(self.base_resnets, self.base_attentions)) - ctrl_blocks = list(zip(self.ctrl_resnets, self.ctrl_attentions)) - - for (b_res, b_attn), (c_res, c_attn), b2c, c2b in zip( - base_blocks, ctrl_blocks, self.base_to_ctrl, self.ctrl_to_base - ): - # concat base -> ctrl - if apply_control: - h_ctrl = torch.cat([h_ctrl, b2c(h_base)], dim=1) - - # apply base subblock - if torch.is_grad_enabled() and self.gradient_checkpointing: - h_base = self._gradient_checkpointing_func(b_res, h_base, temb) - else: - h_base = b_res(h_base, temb) - - if b_attn is not None: - h_base = b_attn( - h_base, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - # apply ctrl subblock - if apply_control: - if torch.is_grad_enabled() and self.gradient_checkpointing: - h_ctrl = self._gradient_checkpointing_func(c_res, h_ctrl, temb) - else: - h_ctrl = c_res(h_ctrl, temb) - if c_attn is not None: - h_ctrl = c_attn( - h_ctrl, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - # add ctrl -> base - if apply_control: - h_base = h_base + c2b(h_ctrl) * conditioning_scale - - base_output_states = base_output_states + (h_base,) - ctrl_output_states = ctrl_output_states + (h_ctrl,) - - if self.base_downsamplers is not None: # if we have a base_downsampler, then also a ctrl_downsampler - b2c = self.base_to_ctrl[-1] - c2b = self.ctrl_to_base[-1] - - # concat base -> ctrl - if apply_control: - h_ctrl = torch.cat([h_ctrl, b2c(h_base)], dim=1) - # apply base subblock - h_base = self.base_downsamplers(h_base) - # apply ctrl subblock - if apply_control: - h_ctrl = self.ctrl_downsamplers(h_ctrl) - # add ctrl -> base - if apply_control: - h_base = h_base + c2b(h_ctrl) * conditioning_scale - - base_output_states = base_output_states + (h_base,) - ctrl_output_states = ctrl_output_states + (h_ctrl,) - - return h_base, h_ctrl, base_output_states, ctrl_output_states - - -class ControlNetXSCrossAttnMidBlock2D(nn.Module): - def __init__( - self, - base_channels: int, - ctrl_channels: int, - temb_channels: int | None = None, - norm_num_groups: int = 32, - ctrl_max_norm_num_groups: int = 32, - transformer_layers_per_block: int = 1, - base_num_attention_heads: int | None = 1, - ctrl_num_attention_heads: int | None = 1, - cross_attention_dim: int | None = 1024, - upcast_attention: bool = False, - use_linear_projection: bool | None = True, - ): - super().__init__() - - # Before the midblock application, information is concatted from base to control. - # Concat doesn't require change in number of channels - self.base_to_ctrl = make_zero_conv(base_channels, base_channels) - - self.base_midblock = UNetMidBlock2DCrossAttn( - transformer_layers_per_block=transformer_layers_per_block, - in_channels=base_channels, - temb_channels=temb_channels, - resnet_groups=norm_num_groups, - cross_attention_dim=cross_attention_dim, - num_attention_heads=base_num_attention_heads, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - ) - - self.ctrl_midblock = UNetMidBlock2DCrossAttn( - transformer_layers_per_block=transformer_layers_per_block, - in_channels=ctrl_channels + base_channels, - out_channels=ctrl_channels, - temb_channels=temb_channels, - # number or norm groups must divide both in_channels and out_channels - resnet_groups=find_largest_factor( - gcd(ctrl_channels, ctrl_channels + base_channels), ctrl_max_norm_num_groups - ), - cross_attention_dim=cross_attention_dim, - num_attention_heads=ctrl_num_attention_heads, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - ) - - # After the midblock application, information is added from control to base - # Addition requires change in number of channels - self.ctrl_to_base = make_zero_conv(ctrl_channels, base_channels) - - self.gradient_checkpointing = False - - @classmethod - def from_modules( - cls, - base_midblock: UNetMidBlock2DCrossAttn, - ctrl_midblock: MidBlockControlNetXSAdapter, - ): - base_to_ctrl = ctrl_midblock.base_to_ctrl - ctrl_to_base = ctrl_midblock.ctrl_to_base - ctrl_midblock = ctrl_midblock.midblock - - # get params - def get_first_cross_attention(midblock): - return midblock.attentions[0].transformer_blocks[0].attn2 - - base_channels = ctrl_to_base.out_channels - ctrl_channels = ctrl_to_base.in_channels - transformer_layers_per_block = len(base_midblock.attentions[0].transformer_blocks) - temb_channels = base_midblock.resnets[0].time_emb_proj.in_features - num_groups = base_midblock.resnets[0].norm1.num_groups - ctrl_num_groups = ctrl_midblock.resnets[0].norm1.num_groups - base_num_attention_heads = get_first_cross_attention(base_midblock).heads - ctrl_num_attention_heads = get_first_cross_attention(ctrl_midblock).heads - cross_attention_dim = get_first_cross_attention(base_midblock).cross_attention_dim - upcast_attention = get_first_cross_attention(base_midblock).upcast_attention - use_linear_projection = base_midblock.attentions[0].use_linear_projection - - # create model - model = cls( - base_channels=base_channels, - ctrl_channels=ctrl_channels, - temb_channels=temb_channels, - norm_num_groups=num_groups, - ctrl_max_norm_num_groups=ctrl_num_groups, - transformer_layers_per_block=transformer_layers_per_block, - base_num_attention_heads=base_num_attention_heads, - ctrl_num_attention_heads=ctrl_num_attention_heads, - cross_attention_dim=cross_attention_dim, - upcast_attention=upcast_attention, - use_linear_projection=use_linear_projection, - ) - - # load weights - model.base_to_ctrl.load_state_dict(base_to_ctrl.state_dict()) - model.base_midblock.load_state_dict(base_midblock.state_dict()) - model.ctrl_midblock.load_state_dict(ctrl_midblock.state_dict()) - model.ctrl_to_base.load_state_dict(ctrl_to_base.state_dict()) - - return model - - def freeze_base_params(self) -> None: - """Freeze the weights of the parts belonging to the base UNet2DConditionModel, and leave everything else unfrozen for fine - tuning.""" - # Unfreeze everything - for param in self.parameters(): - param.requires_grad = True - - # Freeze base part - for param in self.base_midblock.parameters(): - param.requires_grad = False - - def forward( - self, - hidden_states_base: Tensor, - temb: Tensor, - encoder_hidden_states: Tensor, - hidden_states_ctrl: Tensor | None = None, - conditioning_scale: float | None = 1.0, - cross_attention_kwargs: dict[str, Any] | None = None, - attention_mask: Tensor | None = None, - encoder_attention_mask: Tensor | None = None, - apply_control: bool = True, - ) -> tuple[Tensor, Tensor]: - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - h_base = hidden_states_base - h_ctrl = hidden_states_ctrl - - joint_args = { - "temb": temb, - "encoder_hidden_states": encoder_hidden_states, - "attention_mask": attention_mask, - "cross_attention_kwargs": cross_attention_kwargs, - "encoder_attention_mask": encoder_attention_mask, - } - - if apply_control: - h_ctrl = torch.cat([h_ctrl, self.base_to_ctrl(h_base)], dim=1) # concat base -> ctrl - h_base = self.base_midblock(h_base, **joint_args) # apply base mid block - if apply_control: - h_ctrl = self.ctrl_midblock(h_ctrl, **joint_args) # apply ctrl mid block - h_base = h_base + self.ctrl_to_base(h_ctrl) * conditioning_scale # add ctrl -> base - - return h_base, h_ctrl - - -class ControlNetXSCrossAttnUpBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - prev_output_channel: int, - ctrl_skip_channels: list[int], - temb_channels: int, - norm_num_groups: int = 32, - resolution_idx: int | None = None, - has_crossattn=True, - transformer_layers_per_block: int = 1, - num_attention_heads: int = 1, - cross_attention_dim: int = 1024, - add_upsample: bool = True, - upcast_attention: bool = False, - use_linear_projection: bool | None = True, - ): - super().__init__() - resnets = [] - attentions = [] - ctrl_to_base = [] - - num_layers = 3 # only support sd + sdxl - - self.has_cross_attention = has_crossattn - self.num_attention_heads = num_attention_heads - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - ctrl_to_base.append(make_zero_conv(ctrl_skip_channels[i], resnet_in_channels)) - - resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - groups=norm_num_groups, - ) - ) - - if has_crossattn: - attentions.append( - Transformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - norm_num_groups=norm_num_groups, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.attentions = nn.ModuleList(attentions) if has_crossattn else [None] * num_layers - self.ctrl_to_base = nn.ModuleList(ctrl_to_base) - - if add_upsample: - self.upsamplers = Upsample2D(out_channels, use_conv=True, out_channels=out_channels) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - @classmethod - def from_modules(cls, base_upblock: CrossAttnUpBlock2D, ctrl_upblock: UpBlockControlNetXSAdapter): - ctrl_to_base_skip_connections = ctrl_upblock.ctrl_to_base - - # get params - def get_first_cross_attention(block): - return block.attentions[0].transformer_blocks[0].attn2 - - out_channels = base_upblock.resnets[0].out_channels - in_channels = base_upblock.resnets[-1].in_channels - out_channels - prev_output_channels = base_upblock.resnets[0].in_channels - out_channels - ctrl_skip_channelss = [c.in_channels for c in ctrl_to_base_skip_connections] - temb_channels = base_upblock.resnets[0].time_emb_proj.in_features - num_groups = base_upblock.resnets[0].norm1.num_groups - resolution_idx = base_upblock.resolution_idx - if hasattr(base_upblock, "attentions"): - has_crossattn = True - transformer_layers_per_block = len(base_upblock.attentions[0].transformer_blocks) - num_attention_heads = get_first_cross_attention(base_upblock).heads - cross_attention_dim = get_first_cross_attention(base_upblock).cross_attention_dim - upcast_attention = get_first_cross_attention(base_upblock).upcast_attention - use_linear_projection = base_upblock.attentions[0].use_linear_projection - else: - has_crossattn = False - transformer_layers_per_block = None - num_attention_heads = None - cross_attention_dim = None - upcast_attention = None - use_linear_projection = None - add_upsample = base_upblock.upsamplers is not None - - # create model - model = cls( - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channels, - ctrl_skip_channels=ctrl_skip_channelss, - temb_channels=temb_channels, - norm_num_groups=num_groups, - resolution_idx=resolution_idx, - has_crossattn=has_crossattn, - transformer_layers_per_block=transformer_layers_per_block, - num_attention_heads=num_attention_heads, - cross_attention_dim=cross_attention_dim, - add_upsample=add_upsample, - upcast_attention=upcast_attention, - use_linear_projection=use_linear_projection, - ) - - # load weights - model.resnets.load_state_dict(base_upblock.resnets.state_dict()) - if has_crossattn: - model.attentions.load_state_dict(base_upblock.attentions.state_dict()) - if add_upsample: - model.upsamplers.load_state_dict(base_upblock.upsamplers[0].state_dict()) - model.ctrl_to_base.load_state_dict(ctrl_to_base_skip_connections.state_dict()) - - return model - - def freeze_base_params(self) -> None: - """Freeze the weights of the parts belonging to the base UNet2DConditionModel, and leave everything else unfrozen for fine - tuning.""" - # Unfreeze everything - for param in self.parameters(): - param.requires_grad = True - - # Freeze base part - base_parts = [self.resnets] - if isinstance(self.attentions, nn.ModuleList): # attentions can be a list of Nones - base_parts.append(self.attentions) - if self.upsamplers is not None: - base_parts.append(self.upsamplers) - for part in base_parts: - for param in part.parameters(): - param.requires_grad = False - - def forward( - self, - hidden_states: Tensor, - res_hidden_states_tuple_base: tuple[Tensor, ...], - res_hidden_states_tuple_ctrl: tuple[Tensor, ...], - temb: Tensor, - encoder_hidden_states: Tensor | None = None, - conditioning_scale: float | None = 1.0, - cross_attention_kwargs: dict[str, Any] | None = None, - attention_mask: Tensor | None = None, - upsample_size: int | None = None, - encoder_attention_mask: Tensor | None = None, - apply_control: bool = True, - ) -> Tensor: - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - is_freeu_enabled = ( - getattr(self, "s1", None) - and getattr(self, "s2", None) - and getattr(self, "b1", None) - and getattr(self, "b2", None) - ) - - def maybe_apply_freeu_to_subblock(hidden_states, res_h_base): - # FreeU: Only operate on the first two stages - if is_freeu_enabled: - return apply_freeu( - self.resolution_idx, - hidden_states, - res_h_base, - s1=self.s1, - s2=self.s2, - b1=self.b1, - b2=self.b2, - ) - else: - return hidden_states, res_h_base - - for resnet, attn, c2b, res_h_base, res_h_ctrl in zip( - self.resnets, - self.attentions, - self.ctrl_to_base, - reversed(res_hidden_states_tuple_base), - reversed(res_hidden_states_tuple_ctrl), - ): - if apply_control: - hidden_states += c2b(res_h_ctrl) * conditioning_scale - - hidden_states, res_h_base = maybe_apply_freeu_to_subblock(hidden_states, res_h_base) - hidden_states = torch.cat([hidden_states, res_h_base], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - - if attn is not None: - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - if self.upsamplers is not None: - hidden_states = self.upsamplers(hidden_states, upsample_size) - - return hidden_states - - -def make_zero_conv(in_channels, out_channels=None): - return zero_module(nn.Conv2d(in_channels, out_channels, 1, padding=0)) - - -def zero_module(module): - for p in module.parameters(): - nn.init.zeros_(p) - return module - - -def find_largest_factor(number, max_factor): - factor = max_factor - if factor >= number: - return number - while factor != 0: - residual = number % factor - if residual == 0: - return factor - factor -= 1 diff --git a/diffusers/models/controlnets/controlnet_z_image.py b/diffusers/models/controlnets/controlnet_z_image.py deleted file mode 100644 index a4800b255ef08808d66999251a2f5b27c306334c..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/controlnet_z_image.py +++ /dev/null @@ -1,862 +0,0 @@ -# Copyright 2025 Alibaba Z-Image Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import Literal - -import torch -import torch.nn as nn -import torch.nn.functional as F -from torch.nn.utils.rnn import pad_sequence - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...loaders.single_file_model import FromOriginalModelMixin -from ...models.attention_processor import Attention -from ...models.normalization import RMSNorm -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention_dispatch import dispatch_attention_fn -from ..controlnets.controlnet import zero_module -from ..modeling_utils import ModelMixin - - -ADALN_EMBED_DIM = 256 -SEQ_MULTI_OF = 32 - - -# Copied from diffusers.models.transformers.transformer_z_image.TimestepEmbedder -class TimestepEmbedder(nn.Module): - def __init__(self, out_size, mid_size=None, frequency_embedding_size=256): - super().__init__() - if mid_size is None: - mid_size = out_size - self.mlp = nn.Sequential( - nn.Linear(frequency_embedding_size, mid_size, bias=True), - nn.SiLU(), - nn.Linear(mid_size, out_size, bias=True), - ) - - self.frequency_embedding_size = frequency_embedding_size - - @staticmethod - def timestep_embedding(t, dim, max_period=10000): - with torch.amp.autocast("cuda", enabled=False): - half = dim // 2 - freqs = torch.exp( - -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32, device=t.device) / half - ) - args = t[:, None].float() * freqs[None] - embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) - if dim % 2: - embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) - return embedding - - def forward(self, t): - t_freq = self.timestep_embedding(t, self.frequency_embedding_size) - weight_dtype = self.mlp[0].weight.dtype - compute_dtype = getattr(self.mlp[0], "compute_dtype", None) - if weight_dtype.is_floating_point: - t_freq = t_freq.to(weight_dtype) - elif compute_dtype is not None: - t_freq = t_freq.to(compute_dtype) - t_emb = self.mlp(t_freq) - return t_emb - - -# Copied from diffusers.models.transformers.transformer_z_image.ZSingleStreamAttnProcessor -class ZSingleStreamAttnProcessor: - """ - Processor for Z-Image single stream attention that adapts the existing Attention class to match the behavior of the - original Z-ImageAttention module. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "ZSingleStreamAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to version 2.0 or higher." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - freqs_cis: torch.Tensor | None = None, - ) -> torch.Tensor: - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - # Apply Norms - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE - def apply_rotary_emb(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: - with torch.amp.autocast("cuda", enabled=False): - x = torch.view_as_complex(x_in.float().reshape(*x_in.shape[:-1], -1, 2)) - freqs_cis = freqs_cis.unsqueeze(2) - x_out = torch.view_as_real(x * freqs_cis).flatten(3) - return x_out.type_as(x_in) # todo - - if freqs_cis is not None: - query = apply_rotary_emb(query, freqs_cis) - key = apply_rotary_emb(key, freqs_cis) - - # Cast to correct dtype - dtype = query.dtype - query, key = query.to(dtype), key.to(dtype) - - # From [batch, seq_len] to [batch, 1, 1, seq_len] -> broadcast to [batch, heads, seq_len, seq_len] - if attention_mask is not None and attention_mask.ndim == 2: - attention_mask = attention_mask[:, None, None, :] - - # Compute joint attention - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - # Reshape back - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(dtype) - - output = attn.to_out[0](hidden_states) - if len(attn.to_out) > 1: # dropout - output = attn.to_out[1](output) - - return output - - -# Copied from diffusers.models.transformers.transformer_z_image.FeedForward -class FeedForward(nn.Module): - def __init__(self, dim: int, hidden_dim: int): - super().__init__() - self.w1 = nn.Linear(dim, hidden_dim, bias=False) - self.w2 = nn.Linear(hidden_dim, dim, bias=False) - self.w3 = nn.Linear(dim, hidden_dim, bias=False) - - def _forward_silu_gating(self, x1, x3): - return F.silu(x1) * x3 - - def forward(self, x): - return self.w2(self._forward_silu_gating(self.w1(x), self.w3(x))) - - -# Copied from diffusers.models.transformers.transformer_z_image.select_per_token -def select_per_token( - value_noisy: torch.Tensor, - value_clean: torch.Tensor, - noise_mask: torch.Tensor, - seq_len: int, -) -> torch.Tensor: - noise_mask_expanded = noise_mask.unsqueeze(-1) # (batch, seq_len, 1) - return torch.where( - noise_mask_expanded == 1, - value_noisy.unsqueeze(1).expand(-1, seq_len, -1), - value_clean.unsqueeze(1).expand(-1, seq_len, -1), - ) - - -@maybe_allow_in_graph -# Copied from diffusers.models.transformers.transformer_z_image.ZImageTransformerBlock -class ZImageTransformerBlock(nn.Module): - def __init__( - self, - layer_id: int, - dim: int, - n_heads: int, - n_kv_heads: int, - norm_eps: float, - qk_norm: bool, - modulation=True, - ): - super().__init__() - self.dim = dim - self.head_dim = dim // n_heads - - # Refactored to use diffusers Attention with custom processor - # Original Z-Image params: dim, n_heads, n_kv_heads, qk_norm - self.attention = Attention( - query_dim=dim, - cross_attention_dim=None, - dim_head=dim // n_heads, - heads=n_heads, - qk_norm="rms_norm" if qk_norm else None, - eps=1e-5, - bias=False, - out_bias=False, - processor=ZSingleStreamAttnProcessor(), - ) - - self.feed_forward = FeedForward(dim=dim, hidden_dim=int(dim / 3 * 8)) - self.layer_id = layer_id - - self.attention_norm1 = RMSNorm(dim, eps=norm_eps) - self.ffn_norm1 = RMSNorm(dim, eps=norm_eps) - - self.attention_norm2 = RMSNorm(dim, eps=norm_eps) - self.ffn_norm2 = RMSNorm(dim, eps=norm_eps) - - self.modulation = modulation - if modulation: - self.adaLN_modulation = nn.Sequential(nn.Linear(min(dim, ADALN_EMBED_DIM), 4 * dim, bias=True)) - - def forward( - self, - x: torch.Tensor, - attn_mask: torch.Tensor, - freqs_cis: torch.Tensor, - adaln_input: torch.Tensor | None = None, - noise_mask: torch.Tensor | None = None, - adaln_noisy: torch.Tensor | None = None, - adaln_clean: torch.Tensor | None = None, - ): - if self.modulation: - seq_len = x.shape[1] - - if noise_mask is not None: - # Per-token modulation: different modulation for noisy/clean tokens - mod_noisy = self.adaLN_modulation(adaln_noisy) - mod_clean = self.adaLN_modulation(adaln_clean) - - scale_msa_noisy, gate_msa_noisy, scale_mlp_noisy, gate_mlp_noisy = mod_noisy.chunk(4, dim=1) - scale_msa_clean, gate_msa_clean, scale_mlp_clean, gate_mlp_clean = mod_clean.chunk(4, dim=1) - - gate_msa_noisy, gate_mlp_noisy = gate_msa_noisy.tanh(), gate_mlp_noisy.tanh() - gate_msa_clean, gate_mlp_clean = gate_msa_clean.tanh(), gate_mlp_clean.tanh() - - scale_msa_noisy, scale_mlp_noisy = 1.0 + scale_msa_noisy, 1.0 + scale_mlp_noisy - scale_msa_clean, scale_mlp_clean = 1.0 + scale_msa_clean, 1.0 + scale_mlp_clean - - scale_msa = select_per_token(scale_msa_noisy, scale_msa_clean, noise_mask, seq_len) - scale_mlp = select_per_token(scale_mlp_noisy, scale_mlp_clean, noise_mask, seq_len) - gate_msa = select_per_token(gate_msa_noisy, gate_msa_clean, noise_mask, seq_len) - gate_mlp = select_per_token(gate_mlp_noisy, gate_mlp_clean, noise_mask, seq_len) - else: - # Global modulation: same modulation for all tokens (avoid double select) - mod = self.adaLN_modulation(adaln_input) - scale_msa, gate_msa, scale_mlp, gate_mlp = mod.unsqueeze(1).chunk(4, dim=2) - gate_msa, gate_mlp = gate_msa.tanh(), gate_mlp.tanh() - scale_msa, scale_mlp = 1.0 + scale_msa, 1.0 + scale_mlp - - # Attention block - attn_out = self.attention( - self.attention_norm1(x) * scale_msa, attention_mask=attn_mask, freqs_cis=freqs_cis - ) - x = x + gate_msa * self.attention_norm2(attn_out) - - # FFN block - x = x + gate_mlp * self.ffn_norm2(self.feed_forward(self.ffn_norm1(x) * scale_mlp)) - else: - # Attention block - attn_out = self.attention(self.attention_norm1(x), attention_mask=attn_mask, freqs_cis=freqs_cis) - x = x + self.attention_norm2(attn_out) - - # FFN block - x = x + self.ffn_norm2(self.feed_forward(self.ffn_norm1(x))) - - return x - - -# Copied from diffusers.models.transformers.transformer_z_image.RopeEmbedder -class RopeEmbedder: - def __init__( - self, - theta: float = 256.0, - axes_dims: list[int] = (16, 56, 56), - axes_lens: list[int] = (64, 128, 128), - ): - self.theta = theta - self.axes_dims = axes_dims - self.axes_lens = axes_lens - assert len(axes_dims) == len(axes_lens), "axes_dims and axes_lens must have the same length" - self.freqs_cis = None - - @staticmethod - def precompute_freqs_cis(dim: list[int], end: list[int], theta: float = 256.0): - with torch.device("cpu"): - freqs_cis = [] - for i, (d, e) in enumerate(zip(dim, end)): - freqs = 1.0 / (theta ** (torch.arange(0, d, 2, dtype=torch.float64, device="cpu") / d)) - timestep = torch.arange(e, device=freqs.device, dtype=torch.float64) - freqs = torch.outer(timestep, freqs).float() - freqs_cis_i = torch.polar(torch.ones_like(freqs), freqs).to(torch.complex64) # complex64 - freqs_cis.append(freqs_cis_i) - - return freqs_cis - - def __call__(self, ids: torch.Tensor): - assert ids.ndim == 2 - assert ids.shape[-1] == len(self.axes_dims) - device = ids.device - - if self.freqs_cis is None: - self.freqs_cis = self.precompute_freqs_cis(self.axes_dims, self.axes_lens, theta=self.theta) - self.freqs_cis = [freqs_cis.to(device) for freqs_cis in self.freqs_cis] - else: - # Ensure freqs_cis are on the same device as ids - if self.freqs_cis[0].device != device: - self.freqs_cis = [freqs_cis.to(device) for freqs_cis in self.freqs_cis] - - result = [] - for i in range(len(self.axes_dims)): - index = ids[:, i] - result.append(self.freqs_cis[i][index]) - return torch.cat(result, dim=-1) - - -@maybe_allow_in_graph -class ZImageControlTransformerBlock(nn.Module): - def __init__( - self, - layer_id: int, - dim: int, - n_heads: int, - n_kv_heads: int, - norm_eps: float, - qk_norm: bool, - modulation=True, - block_id=0, - ): - super().__init__() - self.dim = dim - self.head_dim = dim // n_heads - - # Refactored to use diffusers Attention with custom processor - # Original Z-Image params: dim, n_heads, n_kv_heads, qk_norm - self.attention = Attention( - query_dim=dim, - cross_attention_dim=None, - dim_head=dim // n_heads, - heads=n_heads, - qk_norm="rms_norm" if qk_norm else None, - eps=1e-5, - bias=False, - out_bias=False, - processor=ZSingleStreamAttnProcessor(), - ) - - self.feed_forward = FeedForward(dim=dim, hidden_dim=int(dim / 3 * 8)) - self.layer_id = layer_id - - self.attention_norm1 = RMSNorm(dim, eps=norm_eps) - self.ffn_norm1 = RMSNorm(dim, eps=norm_eps) - - self.attention_norm2 = RMSNorm(dim, eps=norm_eps) - self.ffn_norm2 = RMSNorm(dim, eps=norm_eps) - - self.modulation = modulation - if modulation: - self.adaLN_modulation = nn.Sequential(nn.Linear(min(dim, ADALN_EMBED_DIM), 4 * dim, bias=True)) - - # Control variant start - self.block_id = block_id - if block_id == 0: - self.before_proj = zero_module(nn.Linear(self.dim, self.dim)) - self.after_proj = zero_module(nn.Linear(self.dim, self.dim)) - - def forward( - self, - c: torch.Tensor, - x: torch.Tensor, - attn_mask: torch.Tensor, - freqs_cis: torch.Tensor, - adaln_input: torch.Tensor | None = None, - ): - # Control - if self.block_id == 0: - c = self.before_proj(c) + x - all_c = [] - else: - all_c = list(torch.unbind(c)) - c = all_c.pop(-1) - - # Compared to `ZImageTransformerBlock` x -> c - if self.modulation: - assert adaln_input is not None - scale_msa, gate_msa, scale_mlp, gate_mlp = self.adaLN_modulation(adaln_input).unsqueeze(1).chunk(4, dim=2) - gate_msa, gate_mlp = gate_msa.tanh(), gate_mlp.tanh() - scale_msa, scale_mlp = 1.0 + scale_msa, 1.0 + scale_mlp - - # Attention block - attn_out = self.attention( - self.attention_norm1(c) * scale_msa, attention_mask=attn_mask, freqs_cis=freqs_cis - ) - c = c + gate_msa * self.attention_norm2(attn_out) - - # FFN block - c = c + gate_mlp * self.ffn_norm2(self.feed_forward(self.ffn_norm1(c) * scale_mlp)) - else: - # Attention block - attn_out = self.attention(self.attention_norm1(c), attention_mask=attn_mask, freqs_cis=freqs_cis) - c = c + self.attention_norm2(attn_out) - - # FFN block - c = c + self.ffn_norm2(self.feed_forward(self.ffn_norm1(c))) - - # Control - c_skip = self.after_proj(c) - all_c += [c_skip, c] - c = torch.stack(all_c) - return c - - -class ZImageControlNetModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - control_layers_places: list[int] = None, - control_refiner_layers_places: list[int] = None, - control_in_dim=None, - add_control_noise_refiner: Literal["control_layers", "control_noise_refiner"] | None = None, - all_patch_size=(2,), - all_f_patch_size=(1,), - dim=3840, - n_refiner_layers=2, - n_heads=30, - n_kv_heads=30, - norm_eps=1e-5, - qk_norm=True, - ): - super().__init__() - self.control_layers_places = control_layers_places - self.control_in_dim = control_in_dim - self.control_refiner_layers_places = control_refiner_layers_places - self.add_control_noise_refiner = add_control_noise_refiner - - assert 0 in self.control_layers_places - - # control blocks - self.control_layers = nn.ModuleList( - [ - ZImageControlTransformerBlock(i, dim, n_heads, n_kv_heads, norm_eps, qk_norm, block_id=i) - for i in self.control_layers_places - ] - ) - - # control patch embeddings - all_x_embedder = {} - for patch_idx, (patch_size, f_patch_size) in enumerate(zip(all_patch_size, all_f_patch_size)): - x_embedder = nn.Linear(f_patch_size * patch_size * patch_size * self.control_in_dim, dim, bias=True) - all_x_embedder[f"{patch_size}-{f_patch_size}"] = x_embedder - - self.control_all_x_embedder = nn.ModuleDict(all_x_embedder) - if self.add_control_noise_refiner == "control_layers": - self.control_noise_refiner = None - elif self.add_control_noise_refiner == "control_noise_refiner": - self.control_noise_refiner = nn.ModuleList( - [ - ZImageControlTransformerBlock( - 1000 + layer_id, - dim, - n_heads, - n_kv_heads, - norm_eps, - qk_norm, - modulation=True, - block_id=layer_id, - ) - for layer_id in range(n_refiner_layers) - ] - ) - else: - self.control_noise_refiner = nn.ModuleList( - [ - ZImageTransformerBlock( - 1000 + layer_id, - dim, - n_heads, - n_kv_heads, - norm_eps, - qk_norm, - modulation=True, - ) - for layer_id in range(n_refiner_layers) - ] - ) - - self.t_scale: float | None = None - self.t_embedder: TimestepEmbedder | None = None - self.all_x_embedder: nn.ModuleDict | None = None - self.cap_embedder: nn.Sequential | None = None - self.rope_embedder: RopeEmbedder | None = None - self.noise_refiner: nn.ModuleList | None = None - self.context_refiner: nn.ModuleList | None = None - self.x_pad_token: nn.Parameter | None = None - self.cap_pad_token: nn.Parameter | None = None - - @classmethod - def from_transformer(cls, controlnet, transformer): - controlnet.t_scale = transformer.t_scale - controlnet.t_embedder = transformer.t_embedder - controlnet.all_x_embedder = transformer.all_x_embedder - controlnet.cap_embedder = transformer.cap_embedder - controlnet.rope_embedder = transformer.rope_embedder - controlnet.noise_refiner = transformer.noise_refiner - controlnet.context_refiner = transformer.context_refiner - controlnet.x_pad_token = transformer.x_pad_token - controlnet.cap_pad_token = transformer.cap_pad_token - return controlnet - - @staticmethod - # Copied from diffusers.models.transformers.transformer_z_image.ZImageTransformer2DModel.create_coordinate_grid - def create_coordinate_grid(size, start=None, device=None): - if start is None: - start = (0 for _ in size) - axes = [torch.arange(x0, x0 + span, dtype=torch.int32, device=device) for x0, span in zip(start, size)] - grids = torch.meshgrid(axes, indexing="ij") - return torch.stack(grids, dim=-1) - - # Copied from diffusers.models.transformers.transformer_z_image.ZImageTransformer2DModel._patchify_image - def _patchify_image(self, image: torch.Tensor, patch_size: int, f_patch_size: int): - """Patchify a single image tensor: (C, F, H, W) -> (num_patches, patch_dim).""" - pH, pW, pF = patch_size, patch_size, f_patch_size - C, F, H, W = image.size() - F_tokens, H_tokens, W_tokens = F // pF, H // pH, W // pW - image = image.view(C, F_tokens, pF, H_tokens, pH, W_tokens, pW) - image = image.permute(1, 3, 5, 2, 4, 6, 0).reshape(F_tokens * H_tokens * W_tokens, pF * pH * pW * C) - return image, (F, H, W), (F_tokens, H_tokens, W_tokens) - - # Copied from diffusers.models.transformers.transformer_z_image.ZImageTransformer2DModel._pad_with_ids - def _pad_with_ids( - self, - feat: torch.Tensor, - pos_grid_size: tuple, - pos_start: tuple, - device: torch.device, - noise_mask_val: int | None = None, - ): - """Pad feature to SEQ_MULTI_OF, create position IDs and pad mask.""" - ori_len = len(feat) - pad_len = (-ori_len) % SEQ_MULTI_OF - total_len = ori_len + pad_len - - # Pos IDs - ori_pos_ids = self.create_coordinate_grid(size=pos_grid_size, start=pos_start, device=device).flatten(0, 2) - if pad_len > 0: - pad_pos_ids = ( - self.create_coordinate_grid(size=(1, 1, 1), start=(0, 0, 0), device=device) - .flatten(0, 2) - .repeat(pad_len, 1) - ) - pos_ids = torch.cat([ori_pos_ids, pad_pos_ids], dim=0) - padded_feat = torch.cat([feat, feat[-1:].repeat(pad_len, 1)], dim=0) - pad_mask = torch.cat( - [ - torch.zeros(ori_len, dtype=torch.bool, device=device), - torch.ones(pad_len, dtype=torch.bool, device=device), - ] - ) - else: - pos_ids = ori_pos_ids - padded_feat = feat - pad_mask = torch.zeros(ori_len, dtype=torch.bool, device=device) - - noise_mask = [noise_mask_val] * total_len if noise_mask_val is not None else None # token level - return padded_feat, pos_ids, pad_mask, total_len, noise_mask - - # Copied from diffusers.models.transformers.transformer_z_image.ZImageTransformer2DModel.patchify_and_embed - def patchify_and_embed( - self, all_image: list[torch.Tensor], all_cap_feats: list[torch.Tensor], patch_size: int, f_patch_size: int - ): - """Patchify for basic mode: single image per batch item.""" - device = all_image[0].device - all_img_out, all_img_size, all_img_pos_ids, all_img_pad_mask = [], [], [], [] - all_cap_out, all_cap_pos_ids, all_cap_pad_mask = [], [], [] - - for image, cap_feat in zip(all_image, all_cap_feats): - # Caption - cap_out, cap_pos_ids, cap_pad_mask, cap_len, _ = self._pad_with_ids( - cap_feat, (len(cap_feat) + (-len(cap_feat)) % SEQ_MULTI_OF, 1, 1), (1, 0, 0), device - ) - all_cap_out.append(cap_out) - all_cap_pos_ids.append(cap_pos_ids) - all_cap_pad_mask.append(cap_pad_mask) - - # Image - img_patches, size, (F_t, H_t, W_t) = self._patchify_image(image, patch_size, f_patch_size) - img_out, img_pos_ids, img_pad_mask, _, _ = self._pad_with_ids( - img_patches, (F_t, H_t, W_t), (cap_len + 1, 0, 0), device - ) - all_img_out.append(img_out) - all_img_size.append(size) - all_img_pos_ids.append(img_pos_ids) - all_img_pad_mask.append(img_pad_mask) - - return ( - all_img_out, - all_cap_out, - all_img_size, - all_img_pos_ids, - all_cap_pos_ids, - all_img_pad_mask, - all_cap_pad_mask, - ) - - def patchify( - self, - all_image: list[torch.Tensor], - patch_size: int, - f_patch_size: int, - ): - pH = pW = patch_size - pF = f_patch_size - all_image_out = [] - - for i, image in enumerate(all_image): - ### Process Image - C, F, H, W = image.size() - F_tokens, H_tokens, W_tokens = F // pF, H // pH, W // pW - - image = image.view(C, F_tokens, pF, H_tokens, pH, W_tokens, pW) - # "c f pf h ph w pw -> (f h w) (pf ph pw c)" - image = image.permute(1, 3, 5, 2, 4, 6, 0).reshape(F_tokens * H_tokens * W_tokens, pF * pH * pW * C) - - image_ori_len = len(image) - image_padding_len = (-image_ori_len) % SEQ_MULTI_OF - - # padded feature - image_padded_feat = torch.cat([image, image[-1:].repeat(image_padding_len, 1)], dim=0) - all_image_out.append(image_padded_feat) - - return all_image_out - - def forward( - self, - x: list[torch.Tensor], - t, - cap_feats: list[torch.Tensor], - control_context: list[torch.Tensor], - conditioning_scale: float = 1.0, - patch_size=2, - f_patch_size=1, - ): - r""" - Args: - x (`list` of `torch.Tensor`): - A list of input image latents, one tensor per sample in the batch. - t (`torch.Tensor`): - Timestep tensor used to indicate the denoising step. - cap_feats (`list` of `torch.Tensor`): - A list of caption (text) feature tensors, one per sample. - control_context (`list` of `torch.Tensor`): - A list of control conditioning feature tensors, one per sample. - conditioning_scale (`float`, *optional*, defaults to `1.0`): - The scale factor for ControlNet outputs. - patch_size (`int`, *optional*, defaults to `2`): - Spatial patch size used to tokenize the latent. - f_patch_size (`int`, *optional*, defaults to `1`): - Temporal (frame) patch size used to tokenize the latent. - """ - if ( - self.t_scale is None - or self.t_embedder is None - or self.all_x_embedder is None - or self.cap_embedder is None - or self.rope_embedder is None - or self.noise_refiner is None - or self.context_refiner is None - or self.x_pad_token is None - or self.cap_pad_token is None - ): - raise ValueError( - "Required modules are `None`, use `from_transformer` to share required modules from `transformer`." - ) - - assert patch_size in self.config.all_patch_size - assert f_patch_size in self.config.all_f_patch_size - - bsz = len(x) - device = x[0].device - t = t * self.t_scale - t = self.t_embedder(t) - - ( - x, - cap_feats, - x_size, - x_pos_ids, - cap_pos_ids, - x_inner_pad_mask, - cap_inner_pad_mask, - ) = self.patchify_and_embed(x, cap_feats, patch_size, f_patch_size) - - x_item_seqlens = [len(_) for _ in x] - assert all(_ % SEQ_MULTI_OF == 0 for _ in x_item_seqlens) - x_max_item_seqlen = max(x_item_seqlens) - - control_context = self.patchify(control_context, patch_size, f_patch_size) - control_context = torch.cat(control_context, dim=0) - control_context = self.control_all_x_embedder[f"{patch_size}-{f_patch_size}"](control_context) - - control_context[torch.cat(x_inner_pad_mask)] = self.x_pad_token - control_context = list(control_context.split(x_item_seqlens, dim=0)) - - control_context = pad_sequence(control_context, batch_first=True, padding_value=0.0) - - # x embed & refine - x = torch.cat(x, dim=0) - x = self.all_x_embedder[f"{patch_size}-{f_patch_size}"](x) - - # Match t_embedder output dtype to x for layerwise casting compatibility - adaln_input = t.type_as(x) - x[torch.cat(x_inner_pad_mask)] = self.x_pad_token - x = list(x.split(x_item_seqlens, dim=0)) - x_freqs_cis = list(self.rope_embedder(torch.cat(x_pos_ids, dim=0)).split([len(_) for _ in x_pos_ids], dim=0)) - - x = pad_sequence(x, batch_first=True, padding_value=0.0) - x_freqs_cis = pad_sequence(x_freqs_cis, batch_first=True, padding_value=0.0) - # Clarify the length matches to satisfy Dynamo due to "Symbolic Shape Inference" to avoid compilation errors - x_freqs_cis = x_freqs_cis[:, : x.shape[1]] - - x_attn_mask = torch.zeros((bsz, x_max_item_seqlen), dtype=torch.bool, device=device) - for i, seq_len in enumerate(x_item_seqlens): - x_attn_mask[i, :seq_len] = 1 - - if self.add_control_noise_refiner is not None: - if self.add_control_noise_refiner == "control_layers": - layers = self.control_layers - elif self.add_control_noise_refiner == "control_noise_refiner": - layers = self.control_noise_refiner - else: - raise ValueError(f"Unsupported `add_control_noise_refiner` type: {self.add_control_noise_refiner}.") - for layer in layers: - if torch.is_grad_enabled() and self.gradient_checkpointing: - control_context = self._gradient_checkpointing_func( - layer, control_context, x, x_attn_mask, x_freqs_cis, adaln_input - ) - else: - control_context = layer(control_context, x, x_attn_mask, x_freqs_cis, adaln_input) - - hints = torch.unbind(control_context)[:-1] - control_context = torch.unbind(control_context)[-1] - noise_refiner_block_samples = { - layer_idx: hints[idx] * conditioning_scale - for idx, layer_idx in enumerate(self.control_refiner_layers_places) - } - else: - noise_refiner_block_samples = None - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for layer_idx, layer in enumerate(self.noise_refiner): - x = self._gradient_checkpointing_func(layer, x, x_attn_mask, x_freqs_cis, adaln_input) - if noise_refiner_block_samples is not None: - if layer_idx in noise_refiner_block_samples: - x = x + noise_refiner_block_samples[layer_idx] - else: - for layer_idx, layer in enumerate(self.noise_refiner): - x = layer(x, x_attn_mask, x_freqs_cis, adaln_input) - if noise_refiner_block_samples is not None: - if layer_idx in noise_refiner_block_samples: - x = x + noise_refiner_block_samples[layer_idx] - - # cap embed & refine - cap_item_seqlens = [len(_) for _ in cap_feats] - cap_max_item_seqlen = max(cap_item_seqlens) - - cap_feats = torch.cat(cap_feats, dim=0) - cap_feats = self.cap_embedder(cap_feats) - cap_feats[torch.cat(cap_inner_pad_mask)] = self.cap_pad_token - cap_feats = list(cap_feats.split(cap_item_seqlens, dim=0)) - cap_freqs_cis = list( - self.rope_embedder(torch.cat(cap_pos_ids, dim=0)).split([len(_) for _ in cap_pos_ids], dim=0) - ) - - cap_feats = pad_sequence(cap_feats, batch_first=True, padding_value=0.0) - cap_freqs_cis = pad_sequence(cap_freqs_cis, batch_first=True, padding_value=0.0) - # Clarify the length matches to satisfy Dynamo due to "Symbolic Shape Inference" to avoid compilation errors - cap_freqs_cis = cap_freqs_cis[:, : cap_feats.shape[1]] - - cap_attn_mask = torch.zeros((bsz, cap_max_item_seqlen), dtype=torch.bool, device=device) - for i, seq_len in enumerate(cap_item_seqlens): - cap_attn_mask[i, :seq_len] = 1 - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for layer in self.context_refiner: - cap_feats = self._gradient_checkpointing_func(layer, cap_feats, cap_attn_mask, cap_freqs_cis) - else: - for layer in self.context_refiner: - cap_feats = layer(cap_feats, cap_attn_mask, cap_freqs_cis) - - # unified - unified = [] - unified_freqs_cis = [] - for i in range(bsz): - x_len = x_item_seqlens[i] - cap_len = cap_item_seqlens[i] - unified.append(torch.cat([x[i][:x_len], cap_feats[i][:cap_len]])) - unified_freqs_cis.append(torch.cat([x_freqs_cis[i][:x_len], cap_freqs_cis[i][:cap_len]])) - unified_item_seqlens = [a + b for a, b in zip(cap_item_seqlens, x_item_seqlens)] - assert unified_item_seqlens == [len(_) for _ in unified] - unified_max_item_seqlen = max(unified_item_seqlens) - - unified = pad_sequence(unified, batch_first=True, padding_value=0.0) - unified_freqs_cis = pad_sequence(unified_freqs_cis, batch_first=True, padding_value=0.0) - unified_attn_mask = torch.zeros((bsz, unified_max_item_seqlen), dtype=torch.bool, device=device) - for i, seq_len in enumerate(unified_item_seqlens): - unified_attn_mask[i, :seq_len] = 1 - - ## ControlNet start - if not self.add_control_noise_refiner: - if torch.is_grad_enabled() and self.gradient_checkpointing: - for layer in self.control_noise_refiner: - control_context = self._gradient_checkpointing_func( - layer, control_context, x_attn_mask, x_freqs_cis, adaln_input - ) - else: - for layer in self.control_noise_refiner: - control_context = layer(control_context, x_attn_mask, x_freqs_cis, adaln_input) - - # unified - control_context_unified = [] - for i in range(bsz): - x_len = x_item_seqlens[i] - cap_len = cap_item_seqlens[i] - control_context_unified.append(torch.cat([control_context[i][:x_len], cap_feats[i][:cap_len]])) - control_context_unified = pad_sequence(control_context_unified, batch_first=True, padding_value=0.0) - - for layer in self.control_layers: - if torch.is_grad_enabled() and self.gradient_checkpointing: - control_context_unified = self._gradient_checkpointing_func( - layer, control_context_unified, unified, unified_attn_mask, unified_freqs_cis, adaln_input - ) - else: - control_context_unified = layer( - control_context_unified, unified, unified_attn_mask, unified_freqs_cis, adaln_input - ) - - hints = torch.unbind(control_context_unified)[:-1] - controlnet_block_samples = { - layer_idx: hints[idx] * conditioning_scale for idx, layer_idx in enumerate(self.control_layers_places) - } - return controlnet_block_samples diff --git a/diffusers/models/controlnets/multicontrolnet.py b/diffusers/models/controlnets/multicontrolnet.py deleted file mode 100644 index 41586950a56be8de698a060b5318383588ce746b..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/multicontrolnet.py +++ /dev/null @@ -1,214 +0,0 @@ -import os -from typing import Any, Callable - -import torch -from torch import nn - -from ...utils import logging -from ..controlnets.controlnet import ControlNetModel, ControlNetOutput -from ..modeling_utils import ModelMixin - - -logger = logging.get_logger(__name__) - - -class MultiControlNetModel(ModelMixin): - r""" - Multiple `ControlNetModel` wrapper class for Multi-ControlNet - - This module is a wrapper for multiple instances of the `ControlNetModel`. The `forward()` API is designed to be - compatible with `ControlNetModel`. - - Args: - controlnets (`list[ControlNetModel]`): - Provides additional conditioning to the unet during the denoising process. You must set multiple - `ControlNetModel` as a list. - """ - - def __init__(self, controlnets: list[ControlNetModel] | tuple[ControlNetModel]): - super().__init__() - self.nets = nn.ModuleList(controlnets) - - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - encoder_hidden_states: torch.Tensor, - controlnet_cond: list[torch.tensor], - conditioning_scale: list[float], - class_labels: torch.Tensor | None = None, - timestep_cond: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - added_cond_kwargs: dict[str, torch.Tensor] | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - guess_mode: bool = False, - return_dict: bool = True, - ) -> ControlNetOutput | tuple: - r""" - Args: - sample (`torch.Tensor`): - The noisy input tensor. - timestep (`torch.Tensor`, `float`, or `int`): - The number of timesteps to denoise an input. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states. - controlnet_cond (`list` of `torch.Tensor`): - A list of conditional input tensors, one per ControlNet. - conditioning_scale (`list` of `float`): - A list of scale factors applied to the ControlNet outputs. - class_labels (`torch.Tensor`, *optional*): - Optional class labels for conditioning. - timestep_cond (`torch.Tensor`, *optional*): - Additional conditional embeddings for timestep. - attention_mask (`torch.Tensor`, *optional*): - Attention mask applied to `encoder_hidden_states`. - added_cond_kwargs (`dict`, *optional*): - Additional conditions for the Stable Diffusion XL UNet. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttnProcessor`. - guess_mode (`bool`, *optional*, defaults to `False`): - In this mode, the ControlNet encoder tries its best to recognize the input content even if you remove - all prompts. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`ControlNetOutput`] instead of a plain tuple. - - Returns: - [`~models.controlnets.controlnet.ControlNetOutput`] or `tuple`: - If `return_dict` is True, a [`~models.controlnets.controlnet.ControlNetOutput`] is returned, otherwise - a plain `tuple` is returned. - """ - for i, (image, scale, controlnet) in enumerate(zip(controlnet_cond, conditioning_scale, self.nets)): - down_samples, mid_sample = controlnet( - sample=sample, - timestep=timestep, - encoder_hidden_states=encoder_hidden_states, - controlnet_cond=image, - conditioning_scale=scale, - class_labels=class_labels, - timestep_cond=timestep_cond, - attention_mask=attention_mask, - added_cond_kwargs=added_cond_kwargs, - cross_attention_kwargs=cross_attention_kwargs, - guess_mode=guess_mode, - return_dict=return_dict, - ) - - # merge samples - if i == 0: - down_block_res_samples, mid_block_res_sample = down_samples, mid_sample - else: - down_block_res_samples = [ - samples_prev + samples_curr - for samples_prev, samples_curr in zip(down_block_res_samples, down_samples) - ] - mid_block_res_sample += mid_sample - - return down_block_res_samples, mid_block_res_sample - - def save_pretrained( - self, - save_directory: str | os.PathLike, - is_main_process: bool = True, - save_function: Callable = None, - safe_serialization: bool = True, - variant: str | None = None, - ): - """ - Save a model and its configuration file to a directory, so that it can be re-loaded using the - `[`~models.controlnets.multicontrolnet.MultiControlNetModel.from_pretrained`]` class method. - - Arguments: - save_directory (`str` or `os.PathLike`): - Directory to which to save. Will be created if it doesn't exist. - is_main_process (`bool`, *optional*, defaults to `True`): - Whether the process calling this is the main process or not. Useful when in distributed training like - TPUs and need to call this function on all processes. In this case, set `is_main_process=True` only on - the main process to avoid race conditions. - save_function (`Callable`): - The function to use to save the state dictionary. Useful on distributed training like TPUs when one - need to replace `torch.save` by another method. Can be configured with the environment variable - `DIFFUSERS_SAVE_MODE`. - safe_serialization (`bool`, *optional*, defaults to `True`): - Whether to save the model using `safetensors` or the traditional PyTorch way (that uses `pickle`). - variant (`str`, *optional*): - If specified, weights are saved in the format pytorch_model..bin. - """ - for idx, controlnet in enumerate(self.nets): - suffix = "" if idx == 0 else f"_{idx}" - controlnet.save_pretrained( - save_directory + suffix, - is_main_process=is_main_process, - save_function=save_function, - safe_serialization=safe_serialization, - variant=variant, - ) - - @classmethod - def from_pretrained(cls, pretrained_model_path: str | os.PathLike | None, **kwargs): - r""" - Instantiate a pretrained MultiControlNet model from multiple pre-trained controlnet models. - - The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated). To train - the model, you should first set it back in training mode with `model.train()`. - - The warning *Weights from XXX not initialized from pretrained model* means that the weights of XXX do not come - pretrained with the rest of the model. It is up to you to train those weights with a downstream fine-tuning - task. - - The warning *Weights from XXX not used in YYY* means that the layer XXX is not used by YYY, therefore those - weights are discarded. - - Parameters: - pretrained_model_path (`os.PathLike`): - A path to a *directory* containing model weights saved using - [`~models.controlnets.multicontrolnet.MultiControlNetModel.save_pretrained`], e.g., - `./my_model_directory/controlnet`. - dtype (`torch.dtype`, *optional*): - Override the default `torch.dtype` and load the model under this dtype. - output_loading_info(`bool`, *optional*, defaults to `False`): - Whether or not to also return a dictionary containing missing keys, unexpected keys and error messages. - device_map (`str` or `dict[str, int | str | torch.device]`, *optional*): - A map that specifies where each submodule should go. It doesn't need to be refined to each - parameter/buffer name, once a given module name is inside, every submodule of it will be sent to the - same device. - - To have Accelerate compute the most optimized `device_map` automatically, set `device_map="auto"`. For - more information about each option see [designing a device - map](https://hf.co/docs/accelerate/main/en/usage_guides/big_modeling#designing-a-device-map). - max_memory (`Dict`, *optional*): - A dictionary device identifier to maximum memory. Will default to the maximum memory available for each - GPU and the available CPU RAM if unset. - low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 else `False`): - Speed up model loading by not initializing the weights and only loading the pre-trained weights. This - also tries to not use more than 1x model size in CPU memory (including peak memory) while loading the - model. This is only supported when torch version >= 1.9.0. If you are using an older version of torch, - setting this argument to `True` will raise an error. - variant (`str`, *optional*): - If specified load weights from `variant` filename, *e.g.* pytorch_model..bin. - use_safetensors (`bool`, *optional*, defaults to `None`): - If set to `None`, the `safetensors` weights will be downloaded if they're available **and** if the - `safetensors` library is installed. If set to `True`, the model will be forcibly loaded from - `safetensors` weights. If set to `False`, loading will *not* use `safetensors`. - """ - idx = 0 - controlnets = [] - - # load controlnet and append to list until no controlnet directory exists anymore - # first controlnet has to be saved under `./mydirectory/controlnet` to be compliant with `DiffusionPipeline.from_prertained` - # second, third, ... controlnets have to be saved under `./mydirectory/controlnet_1`, `./mydirectory/controlnet_2`, ... - model_path_to_load = pretrained_model_path - while os.path.isdir(model_path_to_load): - controlnet = ControlNetModel.from_pretrained(model_path_to_load, **kwargs) - controlnets.append(controlnet) - - idx += 1 - model_path_to_load = pretrained_model_path + f"_{idx}" - - logger.info(f"{len(controlnets)} controlnets loaded from {pretrained_model_path}.") - - if len(controlnets) == 0: - raise ValueError( - f"No ControlNets found under {os.path.dirname(pretrained_model_path)}. Expected at least {pretrained_model_path + '_0'}." - ) - - return cls(controlnets) diff --git a/diffusers/models/controlnets/multicontrolnet_union.py b/diffusers/models/controlnets/multicontrolnet_union.py deleted file mode 100644 index 7dd8f12eb037f7eb99feaa7dcf5811b11688d750..0000000000000000000000000000000000000000 --- a/diffusers/models/controlnets/multicontrolnet_union.py +++ /dev/null @@ -1,231 +0,0 @@ -import os -from typing import Any, Callable - -import torch -from torch import nn - -from ...utils import logging -from ..controlnets.controlnet import ControlNetOutput -from ..controlnets.controlnet_union import ControlNetUnionModel -from ..modeling_utils import ModelMixin - - -logger = logging.get_logger(__name__) - - -class MultiControlNetUnionModel(ModelMixin): - r""" - Multiple `ControlNetUnionModel` wrapper class for Multi-ControlNet-Union. - - This module is a wrapper for multiple instances of the `ControlNetUnionModel`. The `forward()` API is designed to - be compatible with `ControlNetUnionModel`. - - Args: - controlnets (`list[ControlNetUnionModel]`): - Provides additional conditioning to the unet during the denoising process. You must set multiple - `ControlNetUnionModel` as a list. - """ - - def __init__(self, controlnets: list[ControlNetUnionModel] | tuple[ControlNetUnionModel]): - super().__init__() - self.nets = nn.ModuleList(controlnets) - - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - encoder_hidden_states: torch.Tensor, - controlnet_cond: list[torch.tensor], - control_type: list[torch.Tensor], - control_type_idx: list[list[int]], - conditioning_scale: list[float], - class_labels: torch.Tensor | None = None, - timestep_cond: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - added_cond_kwargs: dict[str, torch.Tensor] | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - guess_mode: bool = False, - return_dict: bool = True, - ) -> ControlNetOutput | tuple: - r""" - Args: - sample (`torch.Tensor`): - The noisy input tensor. - timestep (`torch.Tensor`, `float`, or `int`): - The number of timesteps to denoise an input. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states. - controlnet_cond (`list` of `torch.Tensor`): - A list of conditional input tensors, one per ControlNet. - control_type (`list` of `torch.Tensor`): - A list of control type tensors, one per ControlNet, indicating the active control types. - control_type_idx (`list` of `list` of `int`): - Per-ControlNet list of control type indices corresponding to `controlnet_cond`. - conditioning_scale (`list` of `float`): - A list of scale factors applied to the ControlNet outputs. - class_labels (`torch.Tensor`, *optional*): - Optional class labels for conditioning. - timestep_cond (`torch.Tensor`, *optional*): - Additional conditional embeddings for timestep. - attention_mask (`torch.Tensor`, *optional*): - Attention mask applied to `encoder_hidden_states`. - added_cond_kwargs (`dict`, *optional*): - Additional conditions for the Stable Diffusion XL UNet. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttnProcessor`. - guess_mode (`bool`, *optional*, defaults to `False`): - In this mode, the ControlNet encoder tries its best to recognize the input content even if you remove - all prompts. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`ControlNetOutput`] instead of a plain tuple. - - Returns: - [`~models.controlnets.controlnet.ControlNetOutput`] or `tuple`: - If `return_dict` is True, a [`~models.controlnets.controlnet.ControlNetOutput`] is returned, otherwise - a plain `tuple` is returned. - """ - down_block_res_samples, mid_block_res_sample = None, None - for i, (image, ctype, ctype_idx, scale, controlnet) in enumerate( - zip(controlnet_cond, control_type, control_type_idx, conditioning_scale, self.nets) - ): - if scale == 0.0: - continue - down_samples, mid_sample = controlnet( - sample=sample, - timestep=timestep, - encoder_hidden_states=encoder_hidden_states, - controlnet_cond=image, - control_type=ctype, - control_type_idx=ctype_idx, - conditioning_scale=scale, - class_labels=class_labels, - timestep_cond=timestep_cond, - attention_mask=attention_mask, - added_cond_kwargs=added_cond_kwargs, - cross_attention_kwargs=cross_attention_kwargs, - from_multi=True, - guess_mode=guess_mode, - return_dict=return_dict, - ) - - # merge samples - if down_block_res_samples is None and mid_block_res_sample is None: - down_block_res_samples, mid_block_res_sample = down_samples, mid_sample - else: - down_block_res_samples = [ - samples_prev + samples_curr - for samples_prev, samples_curr in zip(down_block_res_samples, down_samples) - ] - mid_block_res_sample += mid_sample - - return down_block_res_samples, mid_block_res_sample - - # Copied from diffusers.models.controlnets.multicontrolnet.MultiControlNetModel.save_pretrained with ControlNet->ControlNetUnion - def save_pretrained( - self, - save_directory: str | os.PathLike, - is_main_process: bool = True, - save_function: Callable = None, - safe_serialization: bool = True, - variant: str | None = None, - ): - """ - Save a model and its configuration file to a directory, so that it can be re-loaded using the - `[`~models.controlnets.multicontrolnet.MultiControlNetUnionModel.from_pretrained`]` class method. - - Arguments: - save_directory (`str` or `os.PathLike`): - Directory to which to save. Will be created if it doesn't exist. - is_main_process (`bool`, *optional*, defaults to `True`): - Whether the process calling this is the main process or not. Useful when in distributed training like - TPUs and need to call this function on all processes. In this case, set `is_main_process=True` only on - the main process to avoid race conditions. - save_function (`Callable`): - The function to use to save the state dictionary. Useful on distributed training like TPUs when one - need to replace `torch.save` by another method. Can be configured with the environment variable - `DIFFUSERS_SAVE_MODE`. - safe_serialization (`bool`, *optional*, defaults to `True`): - Whether to save the model using `safetensors` or the traditional PyTorch way (that uses `pickle`). - variant (`str`, *optional*): - If specified, weights are saved in the format pytorch_model..bin. - """ - for idx, controlnet in enumerate(self.nets): - suffix = "" if idx == 0 else f"_{idx}" - controlnet.save_pretrained( - save_directory + suffix, - is_main_process=is_main_process, - save_function=save_function, - safe_serialization=safe_serialization, - variant=variant, - ) - - @classmethod - # Copied from diffusers.models.controlnets.multicontrolnet.MultiControlNetModel.from_pretrained with ControlNet->ControlNetUnion - def from_pretrained(cls, pretrained_model_path: str | os.PathLike | None, **kwargs): - r""" - Instantiate a pretrained MultiControlNetUnion model from multiple pre-trained controlnet models. - - The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated). To train - the model, you should first set it back in training mode with `model.train()`. - - The warning *Weights from XXX not initialized from pretrained model* means that the weights of XXX do not come - pretrained with the rest of the model. It is up to you to train those weights with a downstream fine-tuning - task. - - The warning *Weights from XXX not used in YYY* means that the layer XXX is not used by YYY, therefore those - weights are discarded. - - Parameters: - pretrained_model_path (`os.PathLike`): - A path to a *directory* containing model weights saved using - [`~models.controlnets.multicontrolnet.MultiControlNetUnionModel.save_pretrained`], e.g., - `./my_model_directory/controlnet`. - dtype (`torch.dtype`, *optional*): - Override the default `torch.dtype` and load the model under this dtype. - output_loading_info(`bool`, *optional*, defaults to `False`): - Whether or not to also return a dictionary containing missing keys, unexpected keys and error messages. - device_map (`str` or `dict[str, int | str | torch.device]`, *optional*): - A map that specifies where each submodule should go. It doesn't need to be refined to each - parameter/buffer name, once a given module name is inside, every submodule of it will be sent to the - same device. - - To have Accelerate compute the most optimized `device_map` automatically, set `device_map="auto"`. For - more information about each option see [designing a device - map](https://hf.co/docs/accelerate/main/en/usage_guides/big_modeling#designing-a-device-map). - max_memory (`Dict`, *optional*): - A dictionary device identifier to maximum memory. Will default to the maximum memory available for each - GPU and the available CPU RAM if unset. - low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 else `False`): - Speed up model loading by not initializing the weights and only loading the pre-trained weights. This - also tries to not use more than 1x model size in CPU memory (including peak memory) while loading the - model. This is only supported when torch version >= 1.9.0. If you are using an older version of torch, - setting this argument to `True` will raise an error. - variant (`str`, *optional*): - If specified load weights from `variant` filename, *e.g.* pytorch_model..bin. - use_safetensors (`bool`, *optional*, defaults to `None`): - If set to `None`, the `safetensors` weights will be downloaded if they're available **and** if the - `safetensors` library is installed. If set to `True`, the model will be forcibly loaded from - `safetensors` weights. If set to `False`, loading will *not* use `safetensors`. - """ - idx = 0 - controlnets = [] - - # load controlnet and append to list until no controlnet directory exists anymore - # first controlnet has to be saved under `./mydirectory/controlnet` to be compliant with `DiffusionPipeline.from_prertained` - # second, third, ... controlnets have to be saved under `./mydirectory/controlnet_1`, `./mydirectory/controlnet_2`, ... - model_path_to_load = pretrained_model_path - while os.path.isdir(model_path_to_load): - controlnet = ControlNetUnionModel.from_pretrained(model_path_to_load, **kwargs) - controlnets.append(controlnet) - - idx += 1 - model_path_to_load = pretrained_model_path + f"_{idx}" - - logger.info(f"{len(controlnets)} controlnets loaded from {pretrained_model_path}.") - - if len(controlnets) == 0: - raise ValueError( - f"No ControlNetUnions found under {os.path.dirname(pretrained_model_path)}. Expected at least {pretrained_model_path + '_0'}." - ) - - return cls(controlnets) diff --git a/diffusers/models/downsampling.py b/diffusers/models/downsampling.py deleted file mode 100644 index 6ae1b647a1914ec5875bb3f48658847cedf67c92..0000000000000000000000000000000000000000 --- a/diffusers/models/downsampling.py +++ /dev/null @@ -1,399 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ..utils import deprecate -from .normalization import RMSNorm -from .upsampling import upfirdn2d_native - - -class Downsample1D(nn.Module): - """A 1D downsampling layer with an optional convolution. - - Parameters: - channels (`int`): - number of channels in the inputs and outputs. - use_conv (`bool`, default `False`): - option to use a convolution. - out_channels (`int`, optional): - number of output channels. Defaults to `channels`. - padding (`int`, default `1`): - padding for the convolution. - name (`str`, default `conv`): - name of the downsampling 1D layer. - """ - - def __init__( - self, - channels: int, - use_conv: bool = False, - out_channels: int | None = None, - padding: int = 1, - name: str = "conv", - ): - super().__init__() - self.channels = channels - self.out_channels = out_channels or channels - self.use_conv = use_conv - self.padding = padding - stride = 2 - self.name = name - - if use_conv: - self.conv = nn.Conv1d(self.channels, self.out_channels, 3, stride=stride, padding=padding) - else: - assert self.channels == self.out_channels - self.conv = nn.AvgPool1d(kernel_size=stride, stride=stride) - - def forward(self, inputs: torch.Tensor) -> torch.Tensor: - assert inputs.shape[1] == self.channels - return self.conv(inputs) - - -class Downsample2D(nn.Module): - """A 2D downsampling layer with an optional convolution. - - Parameters: - channels (`int`): - number of channels in the inputs and outputs. - use_conv (`bool`, default `False`): - option to use a convolution. - out_channels (`int`, optional): - number of output channels. Defaults to `channels`. - padding (`int`, default `1`): - padding for the convolution. - name (`str`, default `conv`): - name of the downsampling 2D layer. - """ - - def __init__( - self, - channels: int, - use_conv: bool = False, - out_channels: int | None = None, - padding: int = 1, - name: str = "conv", - kernel_size=3, - norm_type=None, - eps=None, - elementwise_affine=None, - bias=True, - ): - super().__init__() - self.channels = channels - self.out_channels = out_channels or channels - self.use_conv = use_conv - self.padding = padding - stride = 2 - self.name = name - - if norm_type == "ln_norm": - self.norm = nn.LayerNorm(channels, eps, elementwise_affine) - elif norm_type == "rms_norm": - self.norm = RMSNorm(channels, eps, elementwise_affine) - elif norm_type is None: - self.norm = None - else: - raise ValueError(f"unknown norm_type: {norm_type}") - - if use_conv: - conv = nn.Conv2d( - self.channels, self.out_channels, kernel_size=kernel_size, stride=stride, padding=padding, bias=bias - ) - else: - assert self.channels == self.out_channels - conv = nn.AvgPool2d(kernel_size=stride, stride=stride) - - # TODO(Suraj, Patrick) - clean up after weight dicts are correctly renamed - if name == "conv": - self.Conv2d_0 = conv - self.conv = conv - elif name == "Conv2d_0": - self.conv = conv - else: - self.conv = conv - - def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - assert hidden_states.shape[1] == self.channels - - if self.norm is not None: - hidden_states = self.norm(hidden_states.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) - - if self.use_conv and self.padding == 0: - pad = (0, 1, 0, 1) - hidden_states = F.pad(hidden_states, pad, mode="constant", value=0) - - assert hidden_states.shape[1] == self.channels - - hidden_states = self.conv(hidden_states) - - return hidden_states - - -class FirDownsample2D(nn.Module): - """A 2D FIR downsampling layer with an optional convolution. - - Parameters: - channels (`int`): - number of channels in the inputs and outputs. - use_conv (`bool`, default `False`): - option to use a convolution. - out_channels (`int`, optional): - number of output channels. Defaults to `channels`. - fir_kernel (`tuple`, default `(1, 3, 3, 1)`): - kernel for the FIR filter. - """ - - def __init__( - self, - channels: int | None = None, - out_channels: int | None = None, - use_conv: bool = False, - fir_kernel: tuple[int, int, int, int] = (1, 3, 3, 1), - ): - super().__init__() - out_channels = out_channels if out_channels else channels - if use_conv: - self.Conv2d_0 = nn.Conv2d(channels, out_channels, kernel_size=3, stride=1, padding=1) - self.fir_kernel = fir_kernel - self.use_conv = use_conv - self.out_channels = out_channels - - def _downsample_2d( - self, - hidden_states: torch.Tensor, - weight: torch.Tensor | None = None, - kernel: torch.Tensor | None = None, - factor: int = 2, - gain: float = 1, - ) -> torch.Tensor: - """Fused `Conv2d()` followed by `downsample_2d()`. - Padding is performed only once at the beginning, not between the operations. The fused op is considerably more - efficient than performing the same calculation using standard TensorFlow ops. It supports gradients of - arbitrary order. - - Args: - hidden_states (`torch.Tensor`): - Input tensor of the shape `[N, C, H, W]` or `[N, H, W, C]`. - weight (`torch.Tensor`, *optional*): - Weight tensor of the shape `[filterH, filterW, inChannels, outChannels]`. Grouped convolution can be - performed by `inChannels = x.shape[0] // numGroups`. - kernel (`torch.Tensor`, *optional*): - FIR filter of the shape `[firH, firW]` or `[firN]` (separable). The default is `[1] * factor`, which - corresponds to average pooling. - factor (`int`, *optional*, default to `2`): - Integer downsampling factor. - gain (`float`, *optional*, default to `1.0`): - Scaling factor for signal magnitude. - - Returns: - output (`torch.Tensor`): - Tensor of the shape `[N, C, H // factor, W // factor]` or `[N, H // factor, W // factor, C]`, and same - datatype as `x`. - """ - - assert isinstance(factor, int) and factor >= 1 - if kernel is None: - kernel = [1] * factor - - # setup kernel - kernel = torch.tensor(kernel, dtype=torch.float32) - if kernel.ndim == 1: - kernel = torch.outer(kernel, kernel) - kernel /= torch.sum(kernel) - - kernel = kernel * gain - - if self.use_conv: - _, _, convH, convW = weight.shape - pad_value = (kernel.shape[0] - factor) + (convW - 1) - stride_value = [factor, factor] - upfirdn_input = upfirdn2d_native( - hidden_states, - kernel.to(device=hidden_states.device, dtype=hidden_states.dtype), - pad=((pad_value + 1) // 2, pad_value // 2), - ) - output = F.conv2d(upfirdn_input, weight, stride=stride_value, padding=0) - else: - pad_value = kernel.shape[0] - factor - output = upfirdn2d_native( - hidden_states, - kernel.to(device=hidden_states.device, dtype=hidden_states.dtype), - down=factor, - pad=((pad_value + 1) // 2, pad_value // 2), - ) - - return output - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.use_conv: - downsample_input = self._downsample_2d(hidden_states, weight=self.Conv2d_0.weight, kernel=self.fir_kernel) - hidden_states = downsample_input + self.Conv2d_0.bias.reshape(1, -1, 1, 1) - else: - hidden_states = self._downsample_2d(hidden_states, kernel=self.fir_kernel, factor=2) - - return hidden_states - - -# downsample/upsample layer used in k-upscaler, might be able to use FirDownsample2D/DirUpsample2D instead -class KDownsample2D(nn.Module): - r"""A 2D K-downsampling layer. - - Parameters: - pad_mode (`str`, *optional*, default to `"reflect"`): the padding mode to use. - """ - - def __init__(self, pad_mode: str = "reflect"): - super().__init__() - self.pad_mode = pad_mode - kernel_1d = torch.tensor([[1 / 8, 3 / 8, 3 / 8, 1 / 8]]) - self.pad = kernel_1d.shape[1] // 2 - 1 - self.register_buffer("kernel", kernel_1d.T @ kernel_1d, persistent=False) - - def forward(self, inputs: torch.Tensor) -> torch.Tensor: - inputs = F.pad(inputs, (self.pad,) * 4, self.pad_mode) - weight = inputs.new_zeros( - [ - inputs.shape[1], - inputs.shape[1], - self.kernel.shape[0], - self.kernel.shape[1], - ] - ) - indices = torch.arange(inputs.shape[1], device=inputs.device) - kernel = self.kernel.to(weight)[None, :].expand(inputs.shape[1], -1, -1) - weight[indices, indices] = kernel - return F.conv2d(inputs, weight, stride=2) - - -class CogVideoXDownsample3D(nn.Module): - # Todo: Wait for paper release. - r""" - A 3D Downsampling layer using in [CogVideoX]() by Tsinghua University & ZhipuAI - - Args: - in_channels (`int`): - Number of channels in the input image. - out_channels (`int`): - Number of channels produced by the convolution. - kernel_size (`int`, defaults to `3`): - Size of the convolving kernel. - stride (`int`, defaults to `2`): - Stride of the convolution. - padding (`int`, defaults to `0`): - Padding added to all four sides of the input. - compress_time (`bool`, defaults to `False`): - Whether or not to compress the time dimension. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int = 3, - stride: int = 2, - padding: int = 0, - compress_time: bool = False, - ): - super().__init__() - - self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding) - self.compress_time = compress_time - - def forward(self, x: torch.Tensor) -> torch.Tensor: - if self.compress_time: - batch_size, channels, frames, height, width = x.shape - - # (batch_size, channels, frames, height, width) -> (batch_size, height, width, channels, frames) -> (batch_size * height * width, channels, frames) - x = x.permute(0, 3, 4, 1, 2).reshape(batch_size * height * width, channels, frames) - - if x.shape[-1] % 2 == 1: - x_first, x_rest = x[..., 0], x[..., 1:] - if x_rest.shape[-1] > 0: - # (batch_size * height * width, channels, frames - 1) -> (batch_size * height * width, channels, (frames - 1) // 2) - x_rest = F.avg_pool1d(x_rest, kernel_size=2, stride=2) - - x = torch.cat([x_first[..., None], x_rest], dim=-1) - # (batch_size * height * width, channels, (frames // 2) + 1) -> (batch_size, height, width, channels, (frames // 2) + 1) -> (batch_size, channels, (frames // 2) + 1, height, width) - x = x.reshape(batch_size, height, width, channels, x.shape[-1]).permute(0, 3, 4, 1, 2) - else: - # (batch_size * height * width, channels, frames) -> (batch_size * height * width, channels, frames // 2) - x = F.avg_pool1d(x, kernel_size=2, stride=2) - # (batch_size * height * width, channels, frames // 2) -> (batch_size, height, width, channels, frames // 2) -> (batch_size, channels, frames // 2, height, width) - x = x.reshape(batch_size, height, width, channels, x.shape[-1]).permute(0, 3, 4, 1, 2) - - # Pad the tensor - pad = (0, 1, 0, 1) - x = F.pad(x, pad, mode="constant", value=0) - batch_size, channels, frames, height, width = x.shape - # (batch_size, channels, frames, height, width) -> (batch_size, frames, channels, height, width) -> (batch_size * frames, channels, height, width) - x = x.permute(0, 2, 1, 3, 4).reshape(batch_size * frames, channels, height, width) - x = self.conv(x) - # (batch_size * frames, channels, height, width) -> (batch_size, frames, channels, height, width) -> (batch_size, channels, frames, height, width) - x = x.reshape(batch_size, frames, x.shape[1], x.shape[2], x.shape[3]).permute(0, 2, 1, 3, 4) - return x - - -def downsample_2d( - hidden_states: torch.Tensor, - kernel: torch.Tensor | None = None, - factor: int = 2, - gain: float = 1, -) -> torch.Tensor: - r"""Downsample2D a batch of 2D images with the given filter. - Accepts a batch of 2D images of the shape `[N, C, H, W]` or `[N, H, W, C]` and downsamples each image with the - given filter. The filter is normalized so that if the input pixels are constant, they will be scaled by the - specified `gain`. Pixels outside the image are assumed to be zero, and the filter is padded with zeros so that its - shape is a multiple of the downsampling factor. - - Args: - hidden_states (`torch.Tensor`) - Input tensor of the shape `[N, C, H, W]` or `[N, H, W, C]`. - kernel (`torch.Tensor`, *optional*): - FIR filter of the shape `[firH, firW]` or `[firN]` (separable). The default is `[1] * factor`, which - corresponds to average pooling. - factor (`int`, *optional*, default to `2`): - Integer downsampling factor. - gain (`float`, *optional*, default to `1.0`): - Scaling factor for signal magnitude. - - Returns: - output (`torch.Tensor`): - Tensor of the shape `[N, C, H // factor, W // factor]` - """ - - assert isinstance(factor, int) and factor >= 1 - if kernel is None: - kernel = [1] * factor - - kernel = torch.tensor(kernel, dtype=torch.float32) - if kernel.ndim == 1: - kernel = torch.outer(kernel, kernel) - kernel /= torch.sum(kernel) - - kernel = kernel * gain - pad_value = kernel.shape[0] - factor - output = upfirdn2d_native( - hidden_states, - kernel.to(device=hidden_states.device, dtype=hidden_states.dtype), - down=factor, - pad=((pad_value + 1) // 2, pad_value // 2), - ) - return output diff --git a/diffusers/models/embeddings.py b/diffusers/models/embeddings.py deleted file mode 100644 index 888ae58100ee8b92f111de7ff6ac72a2d81d97e8..0000000000000000000000000000000000000000 --- a/diffusers/models/embeddings.py +++ /dev/null @@ -1,2621 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import math - -import numpy as np -import torch -import torch.nn.functional as F -from torch import nn - -from ..utils import deprecate -from ..utils.torch_utils import maybe_adjust_dtype_for_device -from .activations import FP32SiLU, get_activation -from .attention_processor import Attention - - -def get_timestep_embedding( - timesteps: torch.Tensor, - embedding_dim: int, - flip_sin_to_cos: bool = False, - downscale_freq_shift: float = 1, - scale: float = 1, - max_period: int = 10000, -) -> torch.Tensor: - """ - This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings. - - Args - timesteps (torch.Tensor): - a 1-D Tensor of N indices, one per batch element. These may be fractional. - embedding_dim (int): - the dimension of the output. - flip_sin_to_cos (bool): - Whether the embedding order should be `cos, sin` (if True) or `sin, cos` (if False) - downscale_freq_shift (float): - Controls the delta between frequencies between dimensions - scale (float): - Scaling factor applied to the embeddings. - max_period (int): - Controls the maximum frequency of the embeddings - Returns - torch.Tensor: an [N x dim] Tensor of positional embeddings. - """ - assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array" - - half_dim = embedding_dim // 2 - exponent = -math.log(max_period) * torch.arange( - start=0, end=half_dim, dtype=torch.float32, device=timesteps.device - ) - exponent = exponent / (half_dim - downscale_freq_shift) - - emb = torch.exp(exponent) - emb = timesteps[:, None].float() * emb[None, :] - - # scale embeddings - emb = scale * emb - - # concat sine and cosine embeddings - emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) - - # flip sine and cosine embeddings - if flip_sin_to_cos: - emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) - - # zero pad - if embedding_dim % 2 == 1: - emb = torch.nn.functional.pad(emb, (0, 1, 0, 0)) - return emb - - -def get_3d_sincos_pos_embed( - embed_dim: int, - spatial_size: int | tuple[int, int], - temporal_size: int, - spatial_interpolation_scale: float = 1.0, - temporal_interpolation_scale: float = 1.0, - device: torch.device | None = None, - output_type: str = "np", -) -> torch.Tensor: - r""" - Creates 3D sinusoidal positional embeddings. - - Args: - embed_dim (`int`): - The embedding dimension of inputs. It must be divisible by 16. - spatial_size (`int` or `tuple[int, int]`): - The spatial dimension of positional embeddings. If an integer is provided, the same size is applied to both - spatial dimensions (height and width). - temporal_size (`int`): - The temporal dimension of positional embeddings (number of frames). - spatial_interpolation_scale (`float`, defaults to 1.0): - Scale factor for spatial grid interpolation. - temporal_interpolation_scale (`float`, defaults to 1.0): - Scale factor for temporal grid interpolation. - - Returns: - `torch.Tensor`: - The 3D sinusoidal positional embeddings of shape `[temporal_size, spatial_size[0] * spatial_size[1], - embed_dim]`. - """ - if output_type == "np": - return _get_3d_sincos_pos_embed_np( - embed_dim=embed_dim, - spatial_size=spatial_size, - temporal_size=temporal_size, - spatial_interpolation_scale=spatial_interpolation_scale, - temporal_interpolation_scale=temporal_interpolation_scale, - ) - if embed_dim % 4 != 0: - raise ValueError("`embed_dim` must be divisible by 4") - if isinstance(spatial_size, int): - spatial_size = (spatial_size, spatial_size) - - embed_dim_spatial = 3 * embed_dim // 4 - embed_dim_temporal = embed_dim // 4 - - # 1. Spatial - grid_h = torch.arange(spatial_size[1], device=device, dtype=torch.float32) / spatial_interpolation_scale - grid_w = torch.arange(spatial_size[0], device=device, dtype=torch.float32) / spatial_interpolation_scale - grid = torch.meshgrid(grid_w, grid_h, indexing="xy") # here w goes first - grid = torch.stack(grid, dim=0) - - grid = grid.reshape([2, 1, spatial_size[1], spatial_size[0]]) - pos_embed_spatial = get_2d_sincos_pos_embed_from_grid(embed_dim_spatial, grid, output_type="pt") - - # 2. Temporal - grid_t = torch.arange(temporal_size, device=device, dtype=torch.float32) / temporal_interpolation_scale - pos_embed_temporal = get_1d_sincos_pos_embed_from_grid(embed_dim_temporal, grid_t, output_type="pt") - - # 3. Concat - pos_embed_spatial = pos_embed_spatial[None, :, :] - pos_embed_spatial = pos_embed_spatial.repeat_interleave( - temporal_size, dim=0, output_size=pos_embed_spatial.shape[0] * temporal_size - ) # [T, H*W, D // 4 * 3] - - pos_embed_temporal = pos_embed_temporal[:, None, :] - pos_embed_temporal = pos_embed_temporal.repeat_interleave( - spatial_size[0] * spatial_size[1], dim=1 - ) # [T, H*W, D // 4] - - pos_embed = torch.concat([pos_embed_temporal, pos_embed_spatial], dim=-1) # [T, H*W, D] - return pos_embed - - -def _get_3d_sincos_pos_embed_np( - embed_dim: int, - spatial_size: int | tuple[int, int], - temporal_size: int, - spatial_interpolation_scale: float = 1.0, - temporal_interpolation_scale: float = 1.0, -) -> np.ndarray: - r""" - Creates 3D sinusoidal positional embeddings. - - Args: - embed_dim (`int`): - The embedding dimension of inputs. It must be divisible by 16. - spatial_size (`int` or `tuple[int, int]`): - The spatial dimension of positional embeddings. If an integer is provided, the same size is applied to both - spatial dimensions (height and width). - temporal_size (`int`): - The temporal dimension of positional embeddings (number of frames). - spatial_interpolation_scale (`float`, defaults to 1.0): - Scale factor for spatial grid interpolation. - temporal_interpolation_scale (`float`, defaults to 1.0): - Scale factor for temporal grid interpolation. - - Returns: - `np.ndarray`: - The 3D sinusoidal positional embeddings of shape `[temporal_size, spatial_size[0] * spatial_size[1], - embed_dim]`. - """ - deprecation_message = ( - "`get_3d_sincos_pos_embed` uses `torch` and supports `device`." - " `from_numpy` is no longer required." - " Pass `output_type='pt' to use the new version now." - ) - deprecate("output_type=='np'", "0.33.0", deprecation_message, standard_warn=False) - if embed_dim % 4 != 0: - raise ValueError("`embed_dim` must be divisible by 4") - if isinstance(spatial_size, int): - spatial_size = (spatial_size, spatial_size) - - embed_dim_spatial = 3 * embed_dim // 4 - embed_dim_temporal = embed_dim // 4 - - # 1. Spatial - grid_h = np.arange(spatial_size[1], dtype=np.float32) / spatial_interpolation_scale - grid_w = np.arange(spatial_size[0], dtype=np.float32) / spatial_interpolation_scale - grid = np.meshgrid(grid_w, grid_h) # here w goes first - grid = np.stack(grid, axis=0) - - grid = grid.reshape([2, 1, spatial_size[1], spatial_size[0]]) - pos_embed_spatial = get_2d_sincos_pos_embed_from_grid(embed_dim_spatial, grid) - - # 2. Temporal - grid_t = np.arange(temporal_size, dtype=np.float32) / temporal_interpolation_scale - pos_embed_temporal = get_1d_sincos_pos_embed_from_grid(embed_dim_temporal, grid_t) - - # 3. Concat - pos_embed_spatial = pos_embed_spatial[np.newaxis, :, :] - pos_embed_spatial = np.repeat(pos_embed_spatial, temporal_size, axis=0) # [T, H*W, D // 4 * 3] - - pos_embed_temporal = pos_embed_temporal[:, np.newaxis, :] - pos_embed_temporal = np.repeat(pos_embed_temporal, spatial_size[0] * spatial_size[1], axis=1) # [T, H*W, D // 4] - - pos_embed = np.concatenate([pos_embed_temporal, pos_embed_spatial], axis=-1) # [T, H*W, D] - return pos_embed - - -def get_2d_sincos_pos_embed( - embed_dim, - grid_size, - cls_token=False, - extra_tokens=0, - interpolation_scale=1.0, - base_size=16, - device: torch.device | None = None, - output_type: str = "np", -): - """ - Creates 2D sinusoidal positional embeddings. - - Args: - embed_dim (`int`): - The embedding dimension. - grid_size (`int`): - The size of the grid height and width. - cls_token (`bool`, defaults to `False`): - Whether or not to add a classification token. - extra_tokens (`int`, defaults to `0`): - The number of extra tokens to add. - interpolation_scale (`float`, defaults to `1.0`): - The scale of the interpolation. - - Returns: - pos_embed (`torch.Tensor`): - Shape is either `[grid_size * grid_size, embed_dim]` if not using cls_token, or `[1 + grid_size*grid_size, - embed_dim]` if using cls_token - """ - if output_type == "np": - deprecation_message = ( - "`get_2d_sincos_pos_embed` uses `torch` and supports `device`." - " `from_numpy` is no longer required." - " Pass `output_type='pt' to use the new version now." - ) - deprecate("output_type=='np'", "0.33.0", deprecation_message, standard_warn=False) - return get_2d_sincos_pos_embed_np( - embed_dim=embed_dim, - grid_size=grid_size, - cls_token=cls_token, - extra_tokens=extra_tokens, - interpolation_scale=interpolation_scale, - base_size=base_size, - ) - if isinstance(grid_size, int): - grid_size = (grid_size, grid_size) - - grid_h = ( - torch.arange(grid_size[0], device=device, dtype=torch.float32) - / (grid_size[0] / base_size) - / interpolation_scale - ) - grid_w = ( - torch.arange(grid_size[1], device=device, dtype=torch.float32) - / (grid_size[1] / base_size) - / interpolation_scale - ) - grid = torch.meshgrid(grid_w, grid_h, indexing="xy") # here w goes first - grid = torch.stack(grid, dim=0) - - grid = grid.reshape([2, 1, grid_size[1], grid_size[0]]) - pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid, output_type=output_type) - if cls_token and extra_tokens > 0: - pos_embed = torch.concat([torch.zeros([extra_tokens, embed_dim]), pos_embed], dim=0) - return pos_embed - - -def get_2d_sincos_pos_embed_from_grid(embed_dim, grid, output_type="np"): - r""" - This function generates 2D sinusoidal positional embeddings from a grid. - - Args: - embed_dim (`int`): The embedding dimension. - grid (`torch.Tensor`): Grid of positions with shape `(H * W,)`. - - Returns: - `torch.Tensor`: The 2D sinusoidal positional embeddings with shape `(H * W, embed_dim)` - """ - if output_type == "np": - deprecation_message = ( - "`get_2d_sincos_pos_embed_from_grid` uses `torch` and supports `device`." - " `from_numpy` is no longer required." - " Pass `output_type='pt' to use the new version now." - ) - deprecate("output_type=='np'", "0.33.0", deprecation_message, standard_warn=False) - return get_2d_sincos_pos_embed_from_grid_np( - embed_dim=embed_dim, - grid=grid, - ) - if embed_dim % 2 != 0: - raise ValueError("embed_dim must be divisible by 2") - - # use half of dimensions to encode grid_h - emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0], output_type=output_type) # (H*W, D/2) - emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1], output_type=output_type) # (H*W, D/2) - - emb = torch.concat([emb_h, emb_w], dim=1) # (H*W, D) - return emb - - -def get_1d_sincos_pos_embed_from_grid(embed_dim, pos, output_type="np", flip_sin_to_cos=False, dtype=None): - """ - This function generates 1D positional embeddings from a grid. - - Args: - embed_dim (`int`): The embedding dimension `D` - pos (`torch.Tensor`): 1D tensor of positions with shape `(M,)` - output_type (`str`, *optional*, defaults to `"np"`): Output type. Use `"pt"` for PyTorch tensors. - flip_sin_to_cos (`bool`, *optional*, defaults to `False`): Whether to flip sine and cosine embeddings. - dtype (`torch.dtype`, *optional*): Data type for frequency calculations. If `None`, defaults to - `torch.float32` on MPS devices (which don't support `torch.float64`) and `torch.float64` on other devices. - - Returns: - `torch.Tensor`: Sinusoidal positional embeddings of shape `(M, D)`. - """ - if output_type == "np": - deprecation_message = ( - "`get_1d_sincos_pos_embed_from_grid` uses `torch` and supports `device`." - " `from_numpy` is no longer required." - " Pass `output_type='pt' to use the new version now." - ) - deprecate("output_type=='np'", "0.34.0", deprecation_message, standard_warn=False) - return get_1d_sincos_pos_embed_from_grid_np(embed_dim=embed_dim, pos=pos) - if embed_dim % 2 != 0: - raise ValueError("embed_dim must be divisible by 2") - - # Auto-detect appropriate dtype if not specified - if dtype is None: - dtype = maybe_adjust_dtype_for_device(torch.float64, pos.device) - - omega = torch.arange(embed_dim // 2, device=pos.device, dtype=dtype) - omega /= embed_dim / 2.0 - omega = 1.0 / 10000**omega # (D/2,) - - pos = pos.reshape(-1) # (M,) - out = torch.outer(pos, omega) # (M, D/2), outer product - - emb_sin = torch.sin(out) # (M, D/2) - emb_cos = torch.cos(out) # (M, D/2) - - emb = torch.concat([emb_sin, emb_cos], dim=1) # (M, D) - - # flip sine and cosine embeddings - if flip_sin_to_cos: - emb = torch.cat([emb[:, embed_dim // 2 :], emb[:, : embed_dim // 2]], dim=1) - - return emb - - -def get_2d_sincos_pos_embed_np( - embed_dim, grid_size, cls_token=False, extra_tokens=0, interpolation_scale=1.0, base_size=16 -): - """ - Creates 2D sinusoidal positional embeddings. - - Args: - embed_dim (`int`): - The embedding dimension. - grid_size (`int`): - The size of the grid height and width. - cls_token (`bool`, defaults to `False`): - Whether or not to add a classification token. - extra_tokens (`int`, defaults to `0`): - The number of extra tokens to add. - interpolation_scale (`float`, defaults to `1.0`): - The scale of the interpolation. - - Returns: - pos_embed (`np.ndarray`): - Shape is either `[grid_size * grid_size, embed_dim]` if not using cls_token, or `[1 + grid_size*grid_size, - embed_dim]` if using cls_token - """ - if isinstance(grid_size, int): - grid_size = (grid_size, grid_size) - - grid_h = np.arange(grid_size[0], dtype=np.float32) / (grid_size[0] / base_size) / interpolation_scale - grid_w = np.arange(grid_size[1], dtype=np.float32) / (grid_size[1] / base_size) / interpolation_scale - grid = np.meshgrid(grid_w, grid_h) # here w goes first - grid = np.stack(grid, axis=0) - - grid = grid.reshape([2, 1, grid_size[1], grid_size[0]]) - pos_embed = get_2d_sincos_pos_embed_from_grid_np(embed_dim, grid) - if cls_token and extra_tokens > 0: - pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0) - return pos_embed - - -def get_2d_sincos_pos_embed_from_grid_np(embed_dim, grid): - r""" - This function generates 2D sinusoidal positional embeddings from a grid. - - Args: - embed_dim (`int`): The embedding dimension. - grid (`np.ndarray`): Grid of positions with shape `(H * W,)`. - - Returns: - `np.ndarray`: The 2D sinusoidal positional embeddings with shape `(H * W, embed_dim)` - """ - if embed_dim % 2 != 0: - raise ValueError("embed_dim must be divisible by 2") - - # use half of dimensions to encode grid_h - emb_h = get_1d_sincos_pos_embed_from_grid_np(embed_dim // 2, grid[0]) # (H*W, D/2) - emb_w = get_1d_sincos_pos_embed_from_grid_np(embed_dim // 2, grid[1]) # (H*W, D/2) - - emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D) - return emb - - -def get_1d_sincos_pos_embed_from_grid_np(embed_dim, pos): - """ - This function generates 1D positional embeddings from a grid. - - Args: - embed_dim (`int`): The embedding dimension `D` - pos (`numpy.ndarray`): 1D tensor of positions with shape `(M,)` - - Returns: - `numpy.ndarray`: Sinusoidal positional embeddings of shape `(M, D)`. - """ - if embed_dim % 2 != 0: - raise ValueError("embed_dim must be divisible by 2") - - omega = np.arange(embed_dim // 2, dtype=np.float64) - omega /= embed_dim / 2.0 - omega = 1.0 / 10000**omega # (D/2,) - - pos = pos.reshape(-1) # (M,) - out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product - - emb_sin = np.sin(out) # (M, D/2) - emb_cos = np.cos(out) # (M, D/2) - - emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D) - return emb - - -class PatchEmbed(nn.Module): - """ - 2D Image to Patch Embedding with support for SD3 cropping. - - Args: - height (`int`, defaults to `224`): The height of the image. - width (`int`, defaults to `224`): The width of the image. - patch_size (`int`, defaults to `16`): The size of the patches. - in_channels (`int`, defaults to `3`): The number of input channels. - embed_dim (`int`, defaults to `768`): The output dimension of the embedding. - layer_norm (`bool`, defaults to `False`): Whether or not to use layer normalization. - flatten (`bool`, defaults to `True`): Whether or not to flatten the output. - bias (`bool`, defaults to `True`): Whether or not to use bias. - interpolation_scale (`float`, defaults to `1`): The scale of the interpolation. - pos_embed_type (`str`, defaults to `"sincos"`): The type of positional embedding. - pos_embed_max_size (`int`, defaults to `None`): The maximum size of the positional embedding. - """ - - def __init__( - self, - height=224, - width=224, - patch_size=16, - in_channels=3, - embed_dim=768, - layer_norm=False, - flatten=True, - bias=True, - interpolation_scale=1, - pos_embed_type="sincos", - pos_embed_max_size=None, # For SD3 cropping - ): - super().__init__() - - num_patches = (height // patch_size) * (width // patch_size) - self.flatten = flatten - self.layer_norm = layer_norm - self.pos_embed_max_size = pos_embed_max_size - - self.proj = nn.Conv2d( - in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias - ) - if layer_norm: - self.norm = nn.LayerNorm(embed_dim, elementwise_affine=False, eps=1e-6) - else: - self.norm = None - - self.patch_size = patch_size - self.height, self.width = height // patch_size, width // patch_size - self.base_size = height // patch_size - self.interpolation_scale = interpolation_scale - - # Calculate positional embeddings based on max size or default - if pos_embed_max_size: - grid_size = pos_embed_max_size - else: - grid_size = int(num_patches**0.5) - - if pos_embed_type is None: - self.pos_embed = None - elif pos_embed_type == "sincos": - pos_embed = get_2d_sincos_pos_embed( - embed_dim, - grid_size, - base_size=self.base_size, - interpolation_scale=self.interpolation_scale, - output_type="pt", - ) - persistent = True if pos_embed_max_size else False - self.register_buffer("pos_embed", pos_embed.float().unsqueeze(0), persistent=persistent) - else: - raise ValueError(f"Unsupported pos_embed_type: {pos_embed_type}") - - def cropped_pos_embed(self, height, width): - """Crops positional embeddings for SD3 compatibility.""" - if self.pos_embed_max_size is None: - raise ValueError("`pos_embed_max_size` must be set for cropping.") - - height = height // self.patch_size - width = width // self.patch_size - if height > self.pos_embed_max_size: - raise ValueError( - f"Height ({height}) cannot be greater than `pos_embed_max_size`: {self.pos_embed_max_size}." - ) - if width > self.pos_embed_max_size: - raise ValueError( - f"Width ({width}) cannot be greater than `pos_embed_max_size`: {self.pos_embed_max_size}." - ) - - top = (self.pos_embed_max_size - height) // 2 - left = (self.pos_embed_max_size - width) // 2 - spatial_pos_embed = self.pos_embed.reshape(1, self.pos_embed_max_size, self.pos_embed_max_size, -1) - spatial_pos_embed = spatial_pos_embed[:, top : top + height, left : left + width, :] - spatial_pos_embed = spatial_pos_embed.reshape(1, -1, spatial_pos_embed.shape[-1]) - return spatial_pos_embed - - def forward(self, latent): - if self.pos_embed_max_size is not None: - height, width = latent.shape[-2:] - else: - height, width = latent.shape[-2] // self.patch_size, latent.shape[-1] // self.patch_size - latent = self.proj(latent) - if self.flatten: - latent = latent.flatten(2).transpose(1, 2) # BCHW -> BNC - if self.layer_norm: - latent = self.norm(latent) - if self.pos_embed is None: - return latent.to(latent.dtype) - # Interpolate or crop positional embeddings as needed - if self.pos_embed_max_size: - pos_embed = self.cropped_pos_embed(height, width) - else: - if self.height != height or self.width != width: - pos_embed = get_2d_sincos_pos_embed( - embed_dim=self.pos_embed.shape[-1], - grid_size=(height, width), - base_size=self.base_size, - interpolation_scale=self.interpolation_scale, - device=latent.device, - output_type="pt", - ) - pos_embed = pos_embed.float().unsqueeze(0) - else: - pos_embed = self.pos_embed - - return (latent + pos_embed).to(latent.dtype) - - -class LuminaPatchEmbed(nn.Module): - """ - 2D Image to Patch Embedding with support for Lumina-T2X - - Args: - patch_size (`int`, defaults to `2`): The size of the patches. - in_channels (`int`, defaults to `4`): The number of input channels. - embed_dim (`int`, defaults to `768`): The output dimension of the embedding. - bias (`bool`, defaults to `True`): Whether or not to use bias. - """ - - def __init__(self, patch_size=2, in_channels=4, embed_dim=768, bias=True): - super().__init__() - self.patch_size = patch_size - self.proj = nn.Linear( - in_features=patch_size * patch_size * in_channels, - out_features=embed_dim, - bias=bias, - ) - - def forward(self, x, freqs_cis): - """ - Patchifies and embeds the input tensor(s). - - Args: - x (list[torch.Tensor] | torch.Tensor): The input tensor(s) to be patchified and embedded. - - Returns: - tuple[torch.Tensor, torch.Tensor, list[tuple[int, int]], torch.Tensor]: A tuple containing the patchified - and embedded tensor(s), the mask indicating the valid patches, the original image size(s), and the - frequency tensor(s). - """ - freqs_cis = freqs_cis.to(x[0].device) - patch_height = patch_width = self.patch_size - batch_size, channel, height, width = x.size() - height_tokens, width_tokens = height // patch_height, width // patch_width - - x = x.view(batch_size, channel, height_tokens, patch_height, width_tokens, patch_width).permute( - 0, 2, 4, 1, 3, 5 - ) - x = x.flatten(3) - x = self.proj(x) - x = x.flatten(1, 2) - - mask = torch.ones(x.shape[0], x.shape[1], dtype=torch.int32, device=x.device) - - return ( - x, - mask, - [(height, width)] * batch_size, - freqs_cis[:height_tokens, :width_tokens].flatten(0, 1).unsqueeze(0), - ) - - -class CogVideoXPatchEmbed(nn.Module): - def __init__( - self, - patch_size: int = 2, - patch_size_t: int | None = None, - in_channels: int = 16, - embed_dim: int = 1920, - text_embed_dim: int = 4096, - bias: bool = True, - sample_width: int = 90, - sample_height: int = 60, - sample_frames: int = 49, - temporal_compression_ratio: int = 4, - max_text_seq_length: int = 226, - spatial_interpolation_scale: float = 1.875, - temporal_interpolation_scale: float = 1.0, - use_positional_embeddings: bool = True, - use_learned_positional_embeddings: bool = True, - ) -> None: - super().__init__() - - self.patch_size = patch_size - self.patch_size_t = patch_size_t - self.embed_dim = embed_dim - self.sample_height = sample_height - self.sample_width = sample_width - self.sample_frames = sample_frames - self.temporal_compression_ratio = temporal_compression_ratio - self.max_text_seq_length = max_text_seq_length - self.spatial_interpolation_scale = spatial_interpolation_scale - self.temporal_interpolation_scale = temporal_interpolation_scale - self.use_positional_embeddings = use_positional_embeddings - self.use_learned_positional_embeddings = use_learned_positional_embeddings - - if patch_size_t is None: - # CogVideoX 1.0 checkpoints - self.proj = nn.Conv2d( - in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias - ) - else: - # CogVideoX 1.5 checkpoints - self.proj = nn.Linear(in_channels * patch_size * patch_size * patch_size_t, embed_dim) - - self.text_proj = nn.Linear(text_embed_dim, embed_dim) - - if use_positional_embeddings or use_learned_positional_embeddings: - persistent = use_learned_positional_embeddings - pos_embedding = self._get_positional_embeddings(sample_height, sample_width, sample_frames) - self.register_buffer("pos_embedding", pos_embedding, persistent=persistent) - - def _get_positional_embeddings( - self, sample_height: int, sample_width: int, sample_frames: int, device: torch.device | None = None - ) -> torch.Tensor: - post_patch_height = sample_height // self.patch_size - post_patch_width = sample_width // self.patch_size - post_time_compression_frames = (sample_frames - 1) // self.temporal_compression_ratio + 1 - num_patches = post_patch_height * post_patch_width * post_time_compression_frames - - pos_embedding = get_3d_sincos_pos_embed( - self.embed_dim, - (post_patch_width, post_patch_height), - post_time_compression_frames, - self.spatial_interpolation_scale, - self.temporal_interpolation_scale, - device=device, - output_type="pt", - ) - pos_embedding = pos_embedding.flatten(0, 1) - joint_pos_embedding = pos_embedding.new_zeros( - 1, self.max_text_seq_length + num_patches, self.embed_dim, requires_grad=False - ) - joint_pos_embedding.data[:, self.max_text_seq_length :].copy_(pos_embedding) - - return joint_pos_embedding - - def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor): - r""" - Args: - text_embeds (`torch.Tensor`): - Input text embeddings. Expected shape: (batch_size, seq_length, embedding_dim). - image_embeds (`torch.Tensor`): - Input image embeddings. Expected shape: (batch_size, num_frames, channels, height, width). - """ - text_embeds = self.text_proj(text_embeds) - - batch_size, num_frames, channels, height, width = image_embeds.shape - - if self.patch_size_t is None: - image_embeds = image_embeds.reshape(-1, channels, height, width) - image_embeds = self.proj(image_embeds) - image_embeds = image_embeds.view(batch_size, num_frames, *image_embeds.shape[1:]) - image_embeds = image_embeds.flatten(3).transpose(2, 3) # [batch, num_frames, height x width, channels] - image_embeds = image_embeds.flatten(1, 2) # [batch, num_frames x height x width, channels] - else: - p = self.patch_size - p_t = self.patch_size_t - - image_embeds = image_embeds.permute(0, 1, 3, 4, 2) - image_embeds = image_embeds.reshape( - batch_size, num_frames // p_t, p_t, height // p, p, width // p, p, channels - ) - image_embeds = image_embeds.permute(0, 1, 3, 5, 7, 2, 4, 6).flatten(4, 7).flatten(1, 3) - image_embeds = self.proj(image_embeds) - - embeds = torch.cat( - [text_embeds, image_embeds], dim=1 - ).contiguous() # [batch, seq_length + num_frames x height x width, channels] - - if self.use_positional_embeddings or self.use_learned_positional_embeddings: - if self.use_learned_positional_embeddings and (self.sample_width != width or self.sample_height != height): - raise ValueError( - "It is currently not possible to generate videos at a different resolution that the defaults. This should only be the case with 'THUDM/CogVideoX-5b-I2V'." - "If you think this is incorrect, please open an issue at https://github.com/huggingface/diffusers/issues." - ) - - pre_time_compression_frames = (num_frames - 1) * self.temporal_compression_ratio + 1 - - if ( - self.sample_height != height - or self.sample_width != width - or self.sample_frames != pre_time_compression_frames - ): - pos_embedding = self._get_positional_embeddings( - height, width, pre_time_compression_frames, device=embeds.device - ) - else: - pos_embedding = self.pos_embedding - - pos_embedding = pos_embedding.to(dtype=embeds.dtype) - embeds = embeds + pos_embedding - - return embeds - - -class CogView3PlusPatchEmbed(nn.Module): - def __init__( - self, - in_channels: int = 16, - hidden_size: int = 2560, - patch_size: int = 2, - text_hidden_size: int = 4096, - pos_embed_max_size: int = 128, - ): - super().__init__() - self.in_channels = in_channels - self.hidden_size = hidden_size - self.patch_size = patch_size - self.text_hidden_size = text_hidden_size - self.pos_embed_max_size = pos_embed_max_size - # Linear projection for image patches - self.proj = nn.Linear(in_channels * patch_size**2, hidden_size) - - # Linear projection for text embeddings - self.text_proj = nn.Linear(text_hidden_size, hidden_size) - - pos_embed = get_2d_sincos_pos_embed( - hidden_size, pos_embed_max_size, base_size=pos_embed_max_size, output_type="pt" - ) - pos_embed = pos_embed.reshape(pos_embed_max_size, pos_embed_max_size, hidden_size) - self.register_buffer("pos_embed", pos_embed.float(), persistent=False) - - def forward(self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, channel, height, width = hidden_states.shape - - if height % self.patch_size != 0 or width % self.patch_size != 0: - raise ValueError("Height and width must be divisible by patch size") - - height = height // self.patch_size - width = width // self.patch_size - hidden_states = hidden_states.view(batch_size, channel, height, self.patch_size, width, self.patch_size) - hidden_states = hidden_states.permute(0, 2, 4, 1, 3, 5).contiguous() - hidden_states = hidden_states.view(batch_size, height * width, channel * self.patch_size * self.patch_size) - - # Project the patches - hidden_states = self.proj(hidden_states) - encoder_hidden_states = self.text_proj(encoder_hidden_states) - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - # Calculate text_length - text_length = encoder_hidden_states.shape[1] - - image_pos_embed = self.pos_embed[:height, :width].reshape(height * width, -1) - text_pos_embed = torch.zeros( - (text_length, self.hidden_size), dtype=image_pos_embed.dtype, device=image_pos_embed.device - ) - pos_embed = torch.cat([text_pos_embed, image_pos_embed], dim=0)[None, ...] - - return (hidden_states + pos_embed).to(hidden_states.dtype) - - -def get_3d_rotary_pos_embed( - embed_dim, - crops_coords, - grid_size, - temporal_size, - theta: int = 10000, - use_real: bool = True, - grid_type: str = "linspace", - max_size: tuple[int, int] | None = None, - device: torch.device | None = None, -) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: - """ - RoPE for video tokens with 3D structure. - - Args: - embed_dim: (`int`): - The embedding dimension size, corresponding to hidden_size_head. - crops_coords (`tuple[int]`): - The top-left and bottom-right coordinates of the crop. - grid_size (`tuple[int]`): - The grid size of the spatial positional embedding (height, width). - temporal_size (`int`): - The size of the temporal dimension. - theta (`float`): - Scaling factor for frequency computation. - grid_type (`str`): - Whether to use "linspace" or "slice" to compute grids. - - Returns: - `torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`. - """ - if use_real is not True: - raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed") - - if grid_type == "linspace": - start, stop = crops_coords - grid_size_h, grid_size_w = grid_size - grid_h = torch.linspace( - start[0], stop[0] * (grid_size_h - 1) / grid_size_h, grid_size_h, device=device, dtype=torch.float32 - ) - grid_w = torch.linspace( - start[1], stop[1] * (grid_size_w - 1) / grid_size_w, grid_size_w, device=device, dtype=torch.float32 - ) - grid_t = torch.arange(temporal_size, device=device, dtype=torch.float32) - grid_t = torch.linspace( - 0, temporal_size * (temporal_size - 1) / temporal_size, temporal_size, device=device, dtype=torch.float32 - ) - elif grid_type == "slice": - max_h, max_w = max_size - grid_size_h, grid_size_w = grid_size - grid_h = torch.arange(max_h, device=device, dtype=torch.float32) - grid_w = torch.arange(max_w, device=device, dtype=torch.float32) - grid_t = torch.arange(temporal_size, device=device, dtype=torch.float32) - else: - raise ValueError("Invalid value passed for `grid_type`.") - - # Compute dimensions for each axis - dim_t = embed_dim // 4 - dim_h = embed_dim // 8 * 3 - dim_w = embed_dim // 8 * 3 - - # Temporal frequencies - freqs_t = get_1d_rotary_pos_embed(dim_t, grid_t, theta=theta, use_real=True) - # Spatial frequencies for height and width - freqs_h = get_1d_rotary_pos_embed(dim_h, grid_h, theta=theta, use_real=True) - freqs_w = get_1d_rotary_pos_embed(dim_w, grid_w, theta=theta, use_real=True) - - # BroadCast and concatenate temporal and spaial frequencie (height and width) into a 3d tensor - def combine_time_height_width(freqs_t, freqs_h, freqs_w): - freqs_t = freqs_t[:, None, None, :].expand( - -1, grid_size_h, grid_size_w, -1 - ) # temporal_size, grid_size_h, grid_size_w, dim_t - freqs_h = freqs_h[None, :, None, :].expand( - temporal_size, -1, grid_size_w, -1 - ) # temporal_size, grid_size_h, grid_size_2, dim_h - freqs_w = freqs_w[None, None, :, :].expand( - temporal_size, grid_size_h, -1, -1 - ) # temporal_size, grid_size_h, grid_size_2, dim_w - - freqs = torch.cat( - [freqs_t, freqs_h, freqs_w], dim=-1 - ) # temporal_size, grid_size_h, grid_size_w, (dim_t + dim_h + dim_w) - freqs = freqs.view( - temporal_size * grid_size_h * grid_size_w, -1 - ) # (temporal_size * grid_size_h * grid_size_w), (dim_t + dim_h + dim_w) - return freqs - - t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t - h_cos, h_sin = freqs_h # both h_cos and h_sin has shape: grid_size_h, dim_h - w_cos, w_sin = freqs_w # both w_cos and w_sin has shape: grid_size_w, dim_w - - if grid_type == "slice": - t_cos, t_sin = t_cos[:temporal_size], t_sin[:temporal_size] - h_cos, h_sin = h_cos[:grid_size_h], h_sin[:grid_size_h] - w_cos, w_sin = w_cos[:grid_size_w], w_sin[:grid_size_w] - - cos = combine_time_height_width(t_cos, h_cos, w_cos) - sin = combine_time_height_width(t_sin, h_sin, w_sin) - return cos, sin - - -def get_3d_rotary_pos_embed_allegro( - embed_dim, - crops_coords, - grid_size, - temporal_size, - interpolation_scale: tuple[float, float, float] = (1.0, 1.0, 1.0), - theta: int = 10000, - device: torch.device | None = None, -) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: - # TODO(aryan): docs - start, stop = crops_coords - grid_size_h, grid_size_w = grid_size - interpolation_scale_t, interpolation_scale_h, interpolation_scale_w = interpolation_scale - grid_t = torch.linspace( - 0, temporal_size * (temporal_size - 1) / temporal_size, temporal_size, device=device, dtype=torch.float32 - ) - grid_h = torch.linspace( - start[0], stop[0] * (grid_size_h - 1) / grid_size_h, grid_size_h, device=device, dtype=torch.float32 - ) - grid_w = torch.linspace( - start[1], stop[1] * (grid_size_w - 1) / grid_size_w, grid_size_w, device=device, dtype=torch.float32 - ) - - # Compute dimensions for each axis - dim_t = embed_dim // 3 - dim_h = embed_dim // 3 - dim_w = embed_dim // 3 - - # Temporal frequencies - freqs_t = get_1d_rotary_pos_embed( - dim_t, grid_t / interpolation_scale_t, theta=theta, use_real=True, repeat_interleave_real=False - ) - # Spatial frequencies for height and width - freqs_h = get_1d_rotary_pos_embed( - dim_h, grid_h / interpolation_scale_h, theta=theta, use_real=True, repeat_interleave_real=False - ) - freqs_w = get_1d_rotary_pos_embed( - dim_w, grid_w / interpolation_scale_w, theta=theta, use_real=True, repeat_interleave_real=False - ) - - return freqs_t, freqs_h, freqs_w, grid_t, grid_h, grid_w - - -def get_2d_rotary_pos_embed( - embed_dim, crops_coords, grid_size, use_real=True, device: torch.device | None = None, output_type: str = "np" -): - """ - RoPE for image tokens with 2d structure. - - Args: - embed_dim: (`int`): - The embedding dimension size - crops_coords (`tuple[int]`) - The top-left and bottom-right coordinates of the crop. - grid_size (`tuple[int]`): - The grid size of the positional embedding. - use_real (`bool`): - If True, return real part and imaginary part separately. Otherwise, return complex numbers. - device: (`torch.device`, **optional**): - The device used to create tensors. - - Returns: - `torch.Tensor`: positional embedding with shape `( grid_size * grid_size, embed_dim/2)`. - """ - if output_type == "np": - deprecation_message = ( - "`get_2d_sincos_pos_embed` uses `torch` and supports `device`." - " `from_numpy` is no longer required." - " Pass `output_type='pt' to use the new version now." - ) - deprecate("output_type=='np'", "0.33.0", deprecation_message, standard_warn=False) - return _get_2d_rotary_pos_embed_np( - embed_dim=embed_dim, - crops_coords=crops_coords, - grid_size=grid_size, - use_real=use_real, - ) - start, stop = crops_coords - # scale end by (steps−1)/steps matches np.linspace(..., endpoint=False) - grid_h = torch.linspace( - start[0], stop[0] * (grid_size[0] - 1) / grid_size[0], grid_size[0], device=device, dtype=torch.float32 - ) - grid_w = torch.linspace( - start[1], stop[1] * (grid_size[1] - 1) / grid_size[1], grid_size[1], device=device, dtype=torch.float32 - ) - grid = torch.meshgrid(grid_w, grid_h, indexing="xy") - grid = torch.stack(grid, dim=0) # [2, W, H] - - grid = grid.reshape([2, 1, *grid.shape[1:]]) - pos_embed = get_2d_rotary_pos_embed_from_grid(embed_dim, grid, use_real=use_real) - return pos_embed - - -def _get_2d_rotary_pos_embed_np(embed_dim, crops_coords, grid_size, use_real=True): - """ - RoPE for image tokens with 2d structure. - - Args: - embed_dim: (`int`): - The embedding dimension size - crops_coords (`tuple[int]`) - The top-left and bottom-right coordinates of the crop. - grid_size (`tuple[int]`): - The grid size of the positional embedding. - use_real (`bool`): - If True, return real part and imaginary part separately. Otherwise, return complex numbers. - - Returns: - `torch.Tensor`: positional embedding with shape `( grid_size * grid_size, embed_dim/2)`. - """ - start, stop = crops_coords - grid_h = np.linspace(start[0], stop[0], grid_size[0], endpoint=False, dtype=np.float32) - grid_w = np.linspace(start[1], stop[1], grid_size[1], endpoint=False, dtype=np.float32) - grid = np.meshgrid(grid_w, grid_h) # here w goes first - grid = np.stack(grid, axis=0) # [2, W, H] - - grid = grid.reshape([2, 1, *grid.shape[1:]]) - pos_embed = get_2d_rotary_pos_embed_from_grid(embed_dim, grid, use_real=use_real) - return pos_embed - - -def get_2d_rotary_pos_embed_from_grid(embed_dim, grid, use_real=False): - """ - Get 2D RoPE from grid. - - Args: - embed_dim: (`int`): - The embedding dimension size, corresponding to hidden_size_head. - grid (`np.ndarray`): - The grid of the positional embedding. - use_real (`bool`): - If True, return real part and imaginary part separately. Otherwise, return complex numbers. - - Returns: - `torch.Tensor`: positional embedding with shape `( grid_size * grid_size, embed_dim/2)`. - """ - assert embed_dim % 4 == 0 - - # use half of dimensions to encode grid_h - emb_h = get_1d_rotary_pos_embed( - embed_dim // 2, grid[0].reshape(-1), use_real=use_real - ) # (H*W, D/2) if use_real else (H*W, D/4) - emb_w = get_1d_rotary_pos_embed( - embed_dim // 2, grid[1].reshape(-1), use_real=use_real - ) # (H*W, D/2) if use_real else (H*W, D/4) - - if use_real: - cos = torch.cat([emb_h[0], emb_w[0]], dim=1) # (H*W, D) - sin = torch.cat([emb_h[1], emb_w[1]], dim=1) # (H*W, D) - return cos, sin - else: - emb = torch.cat([emb_h, emb_w], dim=1) # (H*W, D/2) - return emb - - -def get_2d_rotary_pos_embed_lumina(embed_dim, len_h, len_w, linear_factor=1.0, ntk_factor=1.0): - """ - Get 2D RoPE from grid. - - Args: - embed_dim: (`int`): - The embedding dimension size, corresponding to hidden_size_head. - grid (`np.ndarray`): - The grid of the positional embedding. - linear_factor (`float`): - The linear factor of the positional embedding, which is used to scale the positional embedding in the linear - layer. - ntk_factor (`float`): - The ntk factor of the positional embedding, which is used to scale the positional embedding in the ntk layer. - - Returns: - `torch.Tensor`: positional embedding with shape `( grid_size * grid_size, embed_dim/2)`. - """ - assert embed_dim % 4 == 0 - - emb_h = get_1d_rotary_pos_embed( - embed_dim // 2, len_h, linear_factor=linear_factor, ntk_factor=ntk_factor - ) # (H, D/4) - emb_w = get_1d_rotary_pos_embed( - embed_dim // 2, len_w, linear_factor=linear_factor, ntk_factor=ntk_factor - ) # (W, D/4) - emb_h = emb_h.view(len_h, 1, embed_dim // 4, 1).repeat(1, len_w, 1, 1) # (H, W, D/4, 1) - emb_w = emb_w.view(1, len_w, embed_dim // 4, 1).repeat(len_h, 1, 1, 1) # (H, W, D/4, 1) - - emb = torch.cat([emb_h, emb_w], dim=-1).flatten(2) # (H, W, D/2) - return emb - - -def get_1d_rotary_pos_embed( - dim: int, - pos: np.ndarray | int, - theta: float = 10000.0, - use_real=False, - linear_factor=1.0, - ntk_factor=1.0, - repeat_interleave_real=True, - freqs_dtype=torch.float32, # torch.float32, torch.float64 (flux) -): - """ - Precompute the frequency tensor for complex exponentials (cis) with given dimensions. - - This function calculates a frequency tensor with complex exponentials using the given dimension 'dim' and the end - index 'end'. The 'theta' parameter scales the frequencies. The returned tensor contains complex values in complex64 - data type. - - Args: - dim (`int`): Dimension of the frequency tensor. - pos (`np.ndarray` or `int`): Position indices for the frequency tensor. [S] or scalar - theta (`float`, *optional*, defaults to 10000.0): - Scaling factor for frequency computation. Defaults to 10000.0. - use_real (`bool`, *optional*): - If True, return real part and imaginary part separately. Otherwise, return complex numbers. - linear_factor (`float`, *optional*, defaults to 1.0): - Scaling factor for the context extrapolation. Defaults to 1.0. - ntk_factor (`float`, *optional*, defaults to 1.0): - Scaling factor for the NTK-Aware RoPE. Defaults to 1.0. - repeat_interleave_real (`bool`, *optional*, defaults to `True`): - If `True` and `use_real`, real part and imaginary part are each interleaved with themselves to reach `dim`. - Otherwise, they are concateanted with themselves. - freqs_dtype (`torch.float32` or `torch.float64`, *optional*, defaults to `torch.float32`): - the dtype of the frequency tensor. - Returns: - `torch.Tensor`: Precomputed frequency tensor with complex exponentials. [S, D/2] - """ - assert dim % 2 == 0 - - if isinstance(pos, int): - pos = torch.arange(pos) - if isinstance(pos, np.ndarray): - pos = torch.from_numpy(pos) # type: ignore # [S] - - theta = theta * ntk_factor - freqs = ( - 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device) / dim)) / linear_factor - ) # [D/2] - freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2] - is_npu = freqs.device.type == "npu" - if is_npu: - freqs = freqs.float() - if use_real and repeat_interleave_real: - # flux, hunyuan-dit, cogvideox - freqs_cos = freqs.cos().repeat_interleave(2, dim=1, output_size=freqs.shape[1] * 2).float() # [S, D] - freqs_sin = freqs.sin().repeat_interleave(2, dim=1, output_size=freqs.shape[1] * 2).float() # [S, D] - return freqs_cos, freqs_sin - elif use_real: - # stable audio, allegro - freqs_cos = torch.cat([freqs.cos(), freqs.cos()], dim=-1).float() # [S, D] - freqs_sin = torch.cat([freqs.sin(), freqs.sin()], dim=-1).float() # [S, D] - return freqs_cos, freqs_sin - else: - # lumina - freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2] - return freqs_cis - - -def apply_rotary_emb( - x: torch.Tensor, - freqs_cis: torch.Tensor | tuple[torch.Tensor], - use_real: bool = True, - use_real_unbind_dim: int = -1, - sequence_dim: int = 2, -) -> tuple[torch.Tensor, torch.Tensor]: - """ - Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings - to the given query or key 'x' tensors using the provided frequency tensor 'freqs_cis'. The input tensors are - reshaped as complex numbers, and the frequency tensor is reshaped for broadcasting compatibility. The resulting - tensors contain rotary embeddings and are returned as real tensors. - - Args: - x (`torch.Tensor`): - Query or key tensor to apply rotary embeddings. [B, H, S, D] xk (torch.Tensor): Key tensor to apply - freqs_cis (`tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],) - - Returns: - tuple[torch.Tensor, torch.Tensor]: tuple of modified query tensor and key tensor with rotary embeddings. - """ - if use_real: - cos, sin = freqs_cis # [S, D] - if sequence_dim == 2: - cos = cos[None, None, :, :] - sin = sin[None, None, :, :] - elif sequence_dim == 1: - cos = cos[None, :, None, :] - sin = sin[None, :, None, :] - else: - raise ValueError(f"`sequence_dim={sequence_dim}` but should be 1 or 2.") - - cos, sin = cos.to(x.device), sin.to(x.device) - - if use_real_unbind_dim == -1: - # Used for flux, cogvideox, hunyuan-dit - x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, H, S, D//2] - x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3) - elif use_real_unbind_dim == -2: - # Used for Stable Audio, OmniGen, CogView4 and Cosmos - x_real, x_imag = x.reshape(*x.shape[:-1], 2, -1).unbind(-2) # [B, H, S, D//2] - x_rotated = torch.cat([-x_imag, x_real], dim=-1) - else: - raise ValueError(f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2.") - - out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype) - - return out - else: - # used for lumina - x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2)) - freqs_cis = freqs_cis.unsqueeze(2) - x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3) - - return x_out.type_as(x) - - -def apply_rotary_emb_allegro(x: torch.Tensor, freqs_cis, positions): - # TODO(aryan): rewrite - def apply_1d_rope(tokens, pos, cos, sin): - cos = F.embedding(pos, cos)[:, None, :, :] - sin = F.embedding(pos, sin)[:, None, :, :] - x1, x2 = tokens[..., : tokens.shape[-1] // 2], tokens[..., tokens.shape[-1] // 2 :] - tokens_rotated = torch.cat((-x2, x1), dim=-1) - return (tokens.float() * cos + tokens_rotated.float() * sin).to(tokens.dtype) - - (t_cos, t_sin), (h_cos, h_sin), (w_cos, w_sin) = freqs_cis - t, h, w = x.chunk(3, dim=-1) - t = apply_1d_rope(t, positions[0], t_cos, t_sin) - h = apply_1d_rope(h, positions[1], h_cos, h_sin) - w = apply_1d_rope(w, positions[2], w_cos, w_sin) - x = torch.cat([t, h, w], dim=-1) - return x - - -class TimestepEmbedding(nn.Module): - def __init__( - self, - in_channels: int, - time_embed_dim: int, - act_fn: str = "silu", - out_dim: int = None, - post_act_fn: str | None = None, - cond_proj_dim=None, - sample_proj_bias=True, - ): - super().__init__() - - self.linear_1 = nn.Linear(in_channels, time_embed_dim, sample_proj_bias) - - if cond_proj_dim is not None: - self.cond_proj = nn.Linear(cond_proj_dim, in_channels, bias=False) - else: - self.cond_proj = None - - self.act = get_activation(act_fn) - - if out_dim is not None: - time_embed_dim_out = out_dim - else: - time_embed_dim_out = time_embed_dim - self.linear_2 = nn.Linear(time_embed_dim, time_embed_dim_out, sample_proj_bias) - - if post_act_fn is None: - self.post_act = None - else: - self.post_act = get_activation(post_act_fn) - - def forward(self, sample, condition=None): - if condition is not None: - sample = sample + self.cond_proj(condition) - sample = self.linear_1(sample) - - if self.act is not None: - sample = self.act(sample) - - sample = self.linear_2(sample) - - if self.post_act is not None: - sample = self.post_act(sample) - return sample - - -class Timesteps(nn.Module): - def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float, scale: int = 1): - super().__init__() - self.num_channels = num_channels - self.flip_sin_to_cos = flip_sin_to_cos - self.downscale_freq_shift = downscale_freq_shift - self.scale = scale - - def forward(self, timesteps: torch.Tensor) -> torch.Tensor: - t_emb = get_timestep_embedding( - timesteps, - self.num_channels, - flip_sin_to_cos=self.flip_sin_to_cos, - downscale_freq_shift=self.downscale_freq_shift, - scale=self.scale, - ) - return t_emb - - -class GaussianFourierProjection(nn.Module): - """Gaussian Fourier embeddings for noise levels.""" - - def __init__( - self, embedding_size: int = 256, scale: float = 1.0, set_W_to_weight=True, log=True, flip_sin_to_cos=False - ): - super().__init__() - self.weight = nn.Parameter(torch.randn(embedding_size) * scale, requires_grad=False) - self.log = log - self.flip_sin_to_cos = flip_sin_to_cos - - if set_W_to_weight: - # to delete later - del self.weight - self.W = nn.Parameter(torch.randn(embedding_size) * scale, requires_grad=False) - self.weight = self.W - del self.W - - def forward(self, x): - if self.log: - x = torch.log(x) - - x_proj = x[:, None] * self.weight[None, :] * 2 * np.pi - - if self.flip_sin_to_cos: - out = torch.cat([torch.cos(x_proj), torch.sin(x_proj)], dim=-1) - else: - out = torch.cat([torch.sin(x_proj), torch.cos(x_proj)], dim=-1) - return out - - -class SinusoidalPositionalEmbedding(nn.Module): - """Apply positional information to a sequence of embeddings. - - Takes in a sequence of embeddings with shape (batch_size, seq_length, embed_dim) and adds positional embeddings to - them - - Args: - embed_dim: (int): Dimension of the positional embedding. - max_seq_length: Maximum sequence length to apply positional embeddings - - """ - - def __init__(self, embed_dim: int, max_seq_length: int = 32): - super().__init__() - position = torch.arange(max_seq_length).unsqueeze(1) - div_term = torch.exp(torch.arange(0, embed_dim, 2) * (-math.log(10000.0) / embed_dim)) - pe = torch.zeros(1, max_seq_length, embed_dim) - pe[0, :, 0::2] = torch.sin(position * div_term) - pe[0, :, 1::2] = torch.cos(position * div_term) - self.register_buffer("pe", pe) - - def forward(self, x): - _, seq_length, _ = x.shape - x = x + self.pe[:, :seq_length] - return x - - -class ImagePositionalEmbeddings(nn.Module): - """ - Converts latent image classes into vector embeddings. Sums the vector embeddings with positional embeddings for the - height and width of the latent space. - - For more details, see figure 10 of the dall-e paper: https://huggingface.co/papers/2102.12092 - - For VQ-diffusion: - - Output vector embeddings are used as input for the transformer. - - Note that the vector embeddings for the transformer are different than the vector embeddings from the VQVAE. - - Args: - num_embed (`int`): - Number of embeddings for the latent pixels embeddings. - height (`int`): - Height of the latent image i.e. the number of height embeddings. - width (`int`): - Width of the latent image i.e. the number of width embeddings. - embed_dim (`int`): - Dimension of the produced vector embeddings. Used for the latent pixel, height, and width embeddings. - """ - - def __init__( - self, - num_embed: int, - height: int, - width: int, - embed_dim: int, - ): - super().__init__() - - self.height = height - self.width = width - self.num_embed = num_embed - self.embed_dim = embed_dim - - self.emb = nn.Embedding(self.num_embed, embed_dim) - self.height_emb = nn.Embedding(self.height, embed_dim) - self.width_emb = nn.Embedding(self.width, embed_dim) - - def forward(self, index): - emb = self.emb(index) - - height_emb = self.height_emb(torch.arange(self.height, device=index.device).view(1, self.height)) - - # 1 x H x D -> 1 x H x 1 x D - height_emb = height_emb.unsqueeze(2) - - width_emb = self.width_emb(torch.arange(self.width, device=index.device).view(1, self.width)) - - # 1 x W x D -> 1 x 1 x W x D - width_emb = width_emb.unsqueeze(1) - - pos_emb = height_emb + width_emb - - # 1 x H x W x D -> 1 x L xD - pos_emb = pos_emb.view(1, self.height * self.width, -1) - - emb = emb + pos_emb[:, : emb.shape[1], :] - - return emb - - -class LabelEmbedding(nn.Module): - """ - Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance. - - Args: - num_classes (`int`): The number of classes. - hidden_size (`int`): The size of the vector embeddings. - dropout_prob (`float`): The probability of dropping a label. - """ - - def __init__(self, num_classes, hidden_size, dropout_prob): - super().__init__() - use_cfg_embedding = dropout_prob > 0 - self.embedding_table = nn.Embedding(num_classes + use_cfg_embedding, hidden_size) - self.num_classes = num_classes - self.dropout_prob = dropout_prob - - def token_drop(self, labels, force_drop_ids=None): - """ - Drops labels to enable classifier-free guidance. - """ - if force_drop_ids is None: - drop_ids = torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob - else: - drop_ids = torch.tensor(force_drop_ids == 1) - labels = torch.where(drop_ids, self.num_classes, labels) - return labels - - def forward(self, labels: torch.LongTensor, force_drop_ids=None): - use_dropout = self.dropout_prob > 0 - if (self.training and use_dropout) or (force_drop_ids is not None): - labels = self.token_drop(labels, force_drop_ids) - embeddings = self.embedding_table(labels) - return embeddings - - -class TextImageProjection(nn.Module): - def __init__( - self, - text_embed_dim: int = 1024, - image_embed_dim: int = 768, - cross_attention_dim: int = 768, - num_image_text_embeds: int = 10, - ): - super().__init__() - - self.num_image_text_embeds = num_image_text_embeds - self.image_embeds = nn.Linear(image_embed_dim, self.num_image_text_embeds * cross_attention_dim) - self.text_proj = nn.Linear(text_embed_dim, cross_attention_dim) - - def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor): - batch_size = text_embeds.shape[0] - - # image - image_text_embeds = self.image_embeds(image_embeds) - image_text_embeds = image_text_embeds.reshape(batch_size, self.num_image_text_embeds, -1) - - # text - text_embeds = self.text_proj(text_embeds) - - return torch.cat([image_text_embeds, text_embeds], dim=1) - - -class ImageProjection(nn.Module): - def __init__( - self, - image_embed_dim: int = 768, - cross_attention_dim: int = 768, - num_image_text_embeds: int = 32, - ): - super().__init__() - - self.num_image_text_embeds = num_image_text_embeds - self.image_embeds = nn.Linear(image_embed_dim, self.num_image_text_embeds * cross_attention_dim) - self.norm = nn.LayerNorm(cross_attention_dim) - - def forward(self, image_embeds: torch.Tensor): - batch_size = image_embeds.shape[0] - - # image - image_embeds = self.image_embeds(image_embeds.to(self.image_embeds.weight.dtype)) - image_embeds = image_embeds.reshape(batch_size, self.num_image_text_embeds, -1) - image_embeds = self.norm(image_embeds) - return image_embeds - - -class IPAdapterFullImageProjection(nn.Module): - def __init__(self, image_embed_dim=1024, cross_attention_dim=1024): - super().__init__() - from .attention import FeedForward - - self.ff = FeedForward(image_embed_dim, cross_attention_dim, mult=1, activation_fn="gelu") - self.norm = nn.LayerNorm(cross_attention_dim) - - def forward(self, image_embeds: torch.Tensor): - return self.norm(self.ff(image_embeds)) - - -class IPAdapterFaceIDImageProjection(nn.Module): - def __init__(self, image_embed_dim=1024, cross_attention_dim=1024, mult=1, num_tokens=1): - super().__init__() - from .attention import FeedForward - - self.num_tokens = num_tokens - self.cross_attention_dim = cross_attention_dim - self.ff = FeedForward(image_embed_dim, cross_attention_dim * num_tokens, mult=mult, activation_fn="gelu") - self.norm = nn.LayerNorm(cross_attention_dim) - - def forward(self, image_embeds: torch.Tensor): - x = self.ff(image_embeds) - x = x.reshape(-1, self.num_tokens, self.cross_attention_dim) - return self.norm(x) - - -class CombinedTimestepLabelEmbeddings(nn.Module): - def __init__(self, num_classes, embedding_dim, class_dropout_prob=0.1): - super().__init__() - - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=1) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - self.class_embedder = LabelEmbedding(num_classes, embedding_dim, class_dropout_prob) - - def forward(self, timestep, class_labels, hidden_dtype=None): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) - - class_labels = self.class_embedder(class_labels) # (N, D) - - conditioning = timesteps_emb + class_labels # (N, D) - - return conditioning - - -class CombinedTimestepTextProjEmbeddings(nn.Module): - def __init__(self, embedding_dim, pooled_projection_dim): - super().__init__() - - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - self.text_embedder = PixArtAlphaTextProjection(pooled_projection_dim, embedding_dim, act_fn="silu") - - def forward(self, timestep, pooled_projection): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D) - - pooled_projections = self.text_embedder(pooled_projection) - - conditioning = timesteps_emb + pooled_projections - - return conditioning - - -class CombinedTimestepGuidanceTextProjEmbeddings(nn.Module): - def __init__(self, embedding_dim, pooled_projection_dim): - super().__init__() - - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - self.guidance_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - self.text_embedder = PixArtAlphaTextProjection(pooled_projection_dim, embedding_dim, act_fn="silu") - - def forward(self, timestep, guidance, pooled_projection): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D) - - guidance_proj = self.time_proj(guidance) - guidance_emb = self.guidance_embedder(guidance_proj.to(dtype=pooled_projection.dtype)) # (N, D) - - time_guidance_emb = timesteps_emb + guidance_emb - - pooled_projections = self.text_embedder(pooled_projection) - conditioning = time_guidance_emb + pooled_projections - - return conditioning - - -class CogView3CombinedTimestepSizeEmbeddings(nn.Module): - def __init__(self, embedding_dim: int, condition_dim: int, pooled_projection_dim: int, timesteps_dim: int = 256): - super().__init__() - - self.time_proj = Timesteps(num_channels=timesteps_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.condition_proj = Timesteps(num_channels=condition_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=timesteps_dim, time_embed_dim=embedding_dim) - self.condition_embedder = PixArtAlphaTextProjection(pooled_projection_dim, embedding_dim, act_fn="silu") - - def forward( - self, - timestep: torch.Tensor, - original_size: torch.Tensor, - target_size: torch.Tensor, - crop_coords: torch.Tensor, - hidden_dtype: torch.dtype, - ) -> torch.Tensor: - timesteps_proj = self.time_proj(timestep) - - original_size_proj = self.condition_proj(original_size.flatten()).view(original_size.size(0), -1) - crop_coords_proj = self.condition_proj(crop_coords.flatten()).view(crop_coords.size(0), -1) - target_size_proj = self.condition_proj(target_size.flatten()).view(target_size.size(0), -1) - - # (B, 3 * condition_dim) - condition_proj = torch.cat([original_size_proj, crop_coords_proj, target_size_proj], dim=1) - - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (B, embedding_dim) - condition_emb = self.condition_embedder(condition_proj.to(dtype=hidden_dtype)) # (B, embedding_dim) - - conditioning = timesteps_emb + condition_emb - return conditioning - - -class HunyuanDiTAttentionPool(nn.Module): - # Copied from https://github.com/Tencent/HunyuanDiT/blob/cb709308d92e6c7e8d59d0dff41b74d35088db6a/hydit/modules/poolers.py#L6 - - def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: int = None): - super().__init__() - self.positional_embedding = nn.Parameter(torch.randn(spacial_dim + 1, embed_dim) / embed_dim**0.5) - self.k_proj = nn.Linear(embed_dim, embed_dim) - self.q_proj = nn.Linear(embed_dim, embed_dim) - self.v_proj = nn.Linear(embed_dim, embed_dim) - self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim) - self.num_heads = num_heads - - def forward(self, x): - x = x.permute(1, 0, 2) # NLC -> LNC - x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (L+1)NC - x = x + self.positional_embedding[:, None, :].to(x.dtype) # (L+1)NC - x, _ = F.multi_head_attention_forward( - query=x[:1], - key=x, - value=x, - embed_dim_to_check=x.shape[-1], - num_heads=self.num_heads, - q_proj_weight=self.q_proj.weight, - k_proj_weight=self.k_proj.weight, - v_proj_weight=self.v_proj.weight, - in_proj_weight=None, - in_proj_bias=torch.cat([self.q_proj.bias, self.k_proj.bias, self.v_proj.bias]), - bias_k=None, - bias_v=None, - add_zero_attn=False, - dropout_p=0, - out_proj_weight=self.c_proj.weight, - out_proj_bias=self.c_proj.bias, - use_separate_proj_weight=True, - training=self.training, - need_weights=False, - ) - return x.squeeze(0) - - -class HunyuanCombinedTimestepTextSizeStyleEmbedding(nn.Module): - def __init__( - self, - embedding_dim, - pooled_projection_dim=1024, - seq_len=256, - cross_attention_dim=2048, - use_style_cond_and_image_meta_size=True, - ): - super().__init__() - - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - self.size_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - - self.pooler = HunyuanDiTAttentionPool( - seq_len, cross_attention_dim, num_heads=8, output_dim=pooled_projection_dim - ) - - # Here we use a default learned embedder layer for future extension. - self.use_style_cond_and_image_meta_size = use_style_cond_and_image_meta_size - if use_style_cond_and_image_meta_size: - self.style_embedder = nn.Embedding(1, embedding_dim) - extra_in_dim = 256 * 6 + embedding_dim + pooled_projection_dim - else: - extra_in_dim = pooled_projection_dim - - self.extra_embedder = PixArtAlphaTextProjection( - in_features=extra_in_dim, - hidden_size=embedding_dim * 4, - out_features=embedding_dim, - act_fn="silu_fp32", - ) - - def forward(self, timestep, encoder_hidden_states, image_meta_size, style, hidden_dtype=None): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, 256) - - # extra condition1: text - pooled_projections = self.pooler(encoder_hidden_states) # (N, 1024) - - if self.use_style_cond_and_image_meta_size: - # extra condition2: image meta size embedding - image_meta_size = self.size_proj(image_meta_size.view(-1)) - image_meta_size = image_meta_size.to(dtype=hidden_dtype) - image_meta_size = image_meta_size.view(-1, 6 * 256) # (N, 1536) - - # extra condition3: style embedding - style_embedding = self.style_embedder(style) # (N, embedding_dim) - - # Concatenate all extra vectors - extra_cond = torch.cat([pooled_projections, image_meta_size, style_embedding], dim=1) - else: - extra_cond = torch.cat([pooled_projections], dim=1) - - conditioning = timesteps_emb + self.extra_embedder(extra_cond) # [B, D] - - return conditioning - - -class LuminaCombinedTimestepCaptionEmbedding(nn.Module): - def __init__(self, hidden_size=4096, cross_attention_dim=2048, frequency_embedding_size=256): - super().__init__() - self.time_proj = Timesteps( - num_channels=frequency_embedding_size, flip_sin_to_cos=True, downscale_freq_shift=0.0 - ) - - self.timestep_embedder = TimestepEmbedding(in_channels=frequency_embedding_size, time_embed_dim=hidden_size) - - self.caption_embedder = nn.Sequential( - nn.LayerNorm(cross_attention_dim), - nn.Linear( - cross_attention_dim, - hidden_size, - bias=True, - ), - ) - - def forward(self, timestep, caption_feat, caption_mask): - # timestep embedding: - time_freq = self.time_proj(timestep) - time_embed = self.timestep_embedder(time_freq.to(dtype=caption_feat.dtype)) - - # caption condition embedding: - caption_mask_float = caption_mask.float().unsqueeze(-1) - caption_feats_pool = (caption_feat * caption_mask_float).sum(dim=1) / caption_mask_float.sum(dim=1) - caption_feats_pool = caption_feats_pool.to(caption_feat) - caption_embed = self.caption_embedder(caption_feats_pool) - - conditioning = time_embed + caption_embed - - return conditioning - - -class MochiCombinedTimestepCaptionEmbedding(nn.Module): - def __init__( - self, - embedding_dim: int, - pooled_projection_dim: int, - text_embed_dim: int, - time_embed_dim: int = 256, - num_attention_heads: int = 8, - ) -> None: - super().__init__() - - self.time_proj = Timesteps(num_channels=time_embed_dim, flip_sin_to_cos=True, downscale_freq_shift=0.0) - self.timestep_embedder = TimestepEmbedding(in_channels=time_embed_dim, time_embed_dim=embedding_dim) - self.pooler = MochiAttentionPool( - num_attention_heads=num_attention_heads, embed_dim=text_embed_dim, output_dim=embedding_dim - ) - self.caption_proj = nn.Linear(text_embed_dim, pooled_projection_dim) - - def forward( - self, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_attention_mask: torch.Tensor, - hidden_dtype: torch.dtype | None = None, - ): - time_proj = self.time_proj(timestep) - time_emb = self.timestep_embedder(time_proj.to(dtype=hidden_dtype)) - - pooled_projections = self.pooler(encoder_hidden_states, encoder_attention_mask) - caption_proj = self.caption_proj(encoder_hidden_states) - - conditioning = time_emb + pooled_projections - return conditioning, caption_proj - - -class TextTimeEmbedding(nn.Module): - def __init__(self, encoder_dim: int, time_embed_dim: int, num_heads: int = 64): - super().__init__() - self.norm1 = nn.LayerNorm(encoder_dim) - self.pool = AttentionPooling(num_heads, encoder_dim) - self.proj = nn.Linear(encoder_dim, time_embed_dim) - self.norm2 = nn.LayerNorm(time_embed_dim) - - def forward(self, hidden_states): - hidden_states = self.norm1(hidden_states) - hidden_states = self.pool(hidden_states) - hidden_states = self.proj(hidden_states) - hidden_states = self.norm2(hidden_states) - return hidden_states - - -class TextImageTimeEmbedding(nn.Module): - def __init__(self, text_embed_dim: int = 768, image_embed_dim: int = 768, time_embed_dim: int = 1536): - super().__init__() - self.text_proj = nn.Linear(text_embed_dim, time_embed_dim) - self.text_norm = nn.LayerNorm(time_embed_dim) - self.image_proj = nn.Linear(image_embed_dim, time_embed_dim) - - def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor): - # text - time_text_embeds = self.text_proj(text_embeds) - time_text_embeds = self.text_norm(time_text_embeds) - - # image - time_image_embeds = self.image_proj(image_embeds) - - return time_image_embeds + time_text_embeds - - -class ImageTimeEmbedding(nn.Module): - def __init__(self, image_embed_dim: int = 768, time_embed_dim: int = 1536): - super().__init__() - self.image_proj = nn.Linear(image_embed_dim, time_embed_dim) - self.image_norm = nn.LayerNorm(time_embed_dim) - - def forward(self, image_embeds: torch.Tensor): - # image - time_image_embeds = self.image_proj(image_embeds) - time_image_embeds = self.image_norm(time_image_embeds) - return time_image_embeds - - -class ImageHintTimeEmbedding(nn.Module): - def __init__(self, image_embed_dim: int = 768, time_embed_dim: int = 1536): - super().__init__() - self.image_proj = nn.Linear(image_embed_dim, time_embed_dim) - self.image_norm = nn.LayerNorm(time_embed_dim) - self.input_hint_block = nn.Sequential( - nn.Conv2d(3, 16, 3, padding=1), - nn.SiLU(), - nn.Conv2d(16, 16, 3, padding=1), - nn.SiLU(), - nn.Conv2d(16, 32, 3, padding=1, stride=2), - nn.SiLU(), - nn.Conv2d(32, 32, 3, padding=1), - nn.SiLU(), - nn.Conv2d(32, 96, 3, padding=1, stride=2), - nn.SiLU(), - nn.Conv2d(96, 96, 3, padding=1), - nn.SiLU(), - nn.Conv2d(96, 256, 3, padding=1, stride=2), - nn.SiLU(), - nn.Conv2d(256, 4, 3, padding=1), - ) - - def forward(self, image_embeds: torch.Tensor, hint: torch.Tensor): - # image - time_image_embeds = self.image_proj(image_embeds) - time_image_embeds = self.image_norm(time_image_embeds) - hint = self.input_hint_block(hint) - return time_image_embeds, hint - - -class AttentionPooling(nn.Module): - # Copied from https://github.com/deep-floyd/IF/blob/2f91391f27dd3c468bf174be5805b4cc92980c0b/deepfloyd_if/model/nn.py#L54 - - def __init__(self, num_heads, embed_dim, dtype=None): - super().__init__() - self.dtype = dtype - self.positional_embedding = nn.Parameter(torch.randn(1, embed_dim) / embed_dim**0.5) - self.k_proj = nn.Linear(embed_dim, embed_dim, dtype=self.dtype) - self.q_proj = nn.Linear(embed_dim, embed_dim, dtype=self.dtype) - self.v_proj = nn.Linear(embed_dim, embed_dim, dtype=self.dtype) - self.num_heads = num_heads - self.dim_per_head = embed_dim // self.num_heads - - def forward(self, x): - bs, length, width = x.size() - - def shape(x): - # (bs, length, width) --> (bs, length, n_heads, dim_per_head) - x = x.view(bs, -1, self.num_heads, self.dim_per_head) - # (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head) - x = x.transpose(1, 2) - # (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head) - x = x.reshape(bs * self.num_heads, -1, self.dim_per_head) - # (bs*n_heads, length, dim_per_head) --> (bs*n_heads, dim_per_head, length) - x = x.transpose(1, 2) - return x - - class_token = x.mean(dim=1, keepdim=True) + self.positional_embedding.to(x.dtype) - x = torch.cat([class_token, x], dim=1) # (bs, length+1, width) - - # (bs*n_heads, class_token_length, dim_per_head) - q = shape(self.q_proj(class_token)) - # (bs*n_heads, length+class_token_length, dim_per_head) - k = shape(self.k_proj(x)) - v = shape(self.v_proj(x)) - - # (bs*n_heads, class_token_length, length+class_token_length): - scale = 1 / math.sqrt(math.sqrt(self.dim_per_head)) - weight = torch.einsum("bct,bcs->bts", q * scale, k * scale) # More stable with f16 than dividing afterwards - weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) - - # (bs*n_heads, dim_per_head, class_token_length) - a = torch.einsum("bts,bcs->bct", weight, v) - - # (bs, length+1, width) - a = a.reshape(bs, -1, 1).transpose(1, 2) - - return a[:, 0, :] # cls_token - - -class MochiAttentionPool(nn.Module): - def __init__( - self, - num_attention_heads: int, - embed_dim: int, - output_dim: int | None = None, - ) -> None: - super().__init__() - - self.output_dim = output_dim or embed_dim - self.num_attention_heads = num_attention_heads - - self.to_kv = nn.Linear(embed_dim, 2 * embed_dim) - self.to_q = nn.Linear(embed_dim, embed_dim) - self.to_out = nn.Linear(embed_dim, self.output_dim) - - @staticmethod - def pool_tokens(x: torch.Tensor, mask: torch.Tensor, *, keepdim=False) -> torch.Tensor: - """ - Pool tokens in x using mask. - - NOTE: We assume x does not require gradients. - - Args: - x: (B, L, D) tensor of tokens. - mask: (B, L) boolean tensor indicating which tokens are not padding. - - Returns: - pooled: (B, D) tensor of pooled tokens. - """ - assert x.size(1) == mask.size(1) # Expected mask to have same length as tokens. - assert x.size(0) == mask.size(0) # Expected mask to have same batch size as tokens. - mask = mask[:, :, None].to(dtype=x.dtype) - mask = mask / mask.sum(dim=1, keepdim=True).clamp(min=1) - pooled = (x * mask).sum(dim=1, keepdim=keepdim) - return pooled - - def forward(self, x: torch.Tensor, mask: torch.BoolTensor) -> torch.Tensor: - r""" - Args: - x (`torch.Tensor`): - Tensor of shape `(B, S, D)` of input tokens. - mask (`torch.Tensor`): - Boolean ensor of shape `(B, S)` indicating which tokens are not padding. - - Returns: - `torch.Tensor`: - `(B, D)` tensor of pooled tokens. - """ - D = x.size(2) - - # Construct attention mask, shape: (B, 1, num_queries=1, num_keys=1+L). - attn_mask = mask[:, None, None, :].bool() # (B, 1, 1, L). - attn_mask = F.pad(attn_mask, (1, 0), value=True) # (B, 1, 1, 1+L). - - # Average non-padding token features. These will be used as the query. - x_pool = self.pool_tokens(x, mask, keepdim=True) # (B, 1, D) - - # Concat pooled features to input sequence. - x = torch.cat([x_pool, x], dim=1) # (B, L+1, D) - - # Compute queries, keys, values. Only the mean token is used to create a query. - kv = self.to_kv(x) # (B, L+1, 2 * D) - q = self.to_q(x[:, 0]) # (B, D) - - # Extract heads. - head_dim = D // self.num_attention_heads - kv = kv.unflatten(2, (2, self.num_attention_heads, head_dim)) # (B, 1+L, 2, H, head_dim) - kv = kv.transpose(1, 3) # (B, H, 2, 1+L, head_dim) - k, v = kv.unbind(2) # (B, H, 1+L, head_dim) - q = q.unflatten(1, (self.num_attention_heads, head_dim)) # (B, H, head_dim) - q = q.unsqueeze(2) # (B, H, 1, head_dim) - - # Compute attention. - x = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, dropout_p=0.0) # (B, H, 1, head_dim) - - # Concatenate heads and run output. - x = x.squeeze(2).flatten(1, 2) # (B, D = H * head_dim) - x = self.to_out(x) - return x - - -def get_fourier_embeds_from_boundingbox(embed_dim, box): - """ - Args: - embed_dim: int - box: a 3-D tensor [B x N x 4] representing the bounding boxes for GLIGEN pipeline - Returns: - [B x N x embed_dim] tensor of positional embeddings - """ - - batch_size, num_boxes = box.shape[:2] - - emb = 100 ** (torch.arange(embed_dim) / embed_dim) - emb = emb[None, None, None].to(device=box.device, dtype=box.dtype) - emb = emb * box.unsqueeze(-1) - - emb = torch.stack((emb.sin(), emb.cos()), dim=-1) - emb = emb.permute(0, 1, 3, 4, 2).reshape(batch_size, num_boxes, embed_dim * 2 * 4) - - return emb - - -class GLIGENTextBoundingboxProjection(nn.Module): - def __init__(self, positive_len, out_dim, feature_type="text-only", fourier_freqs=8): - super().__init__() - self.positive_len = positive_len - self.out_dim = out_dim - - self.fourier_embedder_dim = fourier_freqs - self.position_dim = fourier_freqs * 2 * 4 # 2: sin/cos, 4: xyxy - - if isinstance(out_dim, tuple): - out_dim = out_dim[0] - - if feature_type == "text-only": - self.linears = nn.Sequential( - nn.Linear(self.positive_len + self.position_dim, 512), - nn.SiLU(), - nn.Linear(512, 512), - nn.SiLU(), - nn.Linear(512, out_dim), - ) - self.null_positive_feature = torch.nn.Parameter(torch.zeros([self.positive_len])) - - elif feature_type == "text-image": - self.linears_text = nn.Sequential( - nn.Linear(self.positive_len + self.position_dim, 512), - nn.SiLU(), - nn.Linear(512, 512), - nn.SiLU(), - nn.Linear(512, out_dim), - ) - self.linears_image = nn.Sequential( - nn.Linear(self.positive_len + self.position_dim, 512), - nn.SiLU(), - nn.Linear(512, 512), - nn.SiLU(), - nn.Linear(512, out_dim), - ) - self.null_text_feature = torch.nn.Parameter(torch.zeros([self.positive_len])) - self.null_image_feature = torch.nn.Parameter(torch.zeros([self.positive_len])) - - self.null_position_feature = torch.nn.Parameter(torch.zeros([self.position_dim])) - - def forward( - self, - boxes, - masks, - positive_embeddings=None, - phrases_masks=None, - image_masks=None, - phrases_embeddings=None, - image_embeddings=None, - ): - masks = masks.unsqueeze(-1) - - # embedding position (it may includes padding as placeholder) - xyxy_embedding = get_fourier_embeds_from_boundingbox(self.fourier_embedder_dim, boxes) # B*N*4 -> B*N*C - - # learnable null embedding - xyxy_null = self.null_position_feature.view(1, 1, -1) - - # replace padding with learnable null embedding - xyxy_embedding = xyxy_embedding * masks + (1 - masks) * xyxy_null - - # positionet with text only information - if positive_embeddings is not None: - # learnable null embedding - positive_null = self.null_positive_feature.view(1, 1, -1) - - # replace padding with learnable null embedding - positive_embeddings = positive_embeddings * masks + (1 - masks) * positive_null - - objs = self.linears(torch.cat([positive_embeddings, xyxy_embedding], dim=-1)) - - # positionet with text and image information - else: - phrases_masks = phrases_masks.unsqueeze(-1) - image_masks = image_masks.unsqueeze(-1) - - # learnable null embedding - text_null = self.null_text_feature.view(1, 1, -1) - image_null = self.null_image_feature.view(1, 1, -1) - - # replace padding with learnable null embedding - phrases_embeddings = phrases_embeddings * phrases_masks + (1 - phrases_masks) * text_null - image_embeddings = image_embeddings * image_masks + (1 - image_masks) * image_null - - objs_text = self.linears_text(torch.cat([phrases_embeddings, xyxy_embedding], dim=-1)) - objs_image = self.linears_image(torch.cat([image_embeddings, xyxy_embedding], dim=-1)) - objs = torch.cat([objs_text, objs_image], dim=1) - - return objs - - -class PixArtAlphaCombinedTimestepSizeEmbeddings(nn.Module): - """ - For PixArt-Alpha. - - Reference: - https://github.com/PixArt-alpha/PixArt-alpha/blob/0f55e922376d8b797edd44d25d0e7464b260dcab/diffusion/model/nets/PixArtMS.py#L164C9-L168C29 - """ - - def __init__(self, embedding_dim, size_emb_dim, use_additional_conditions: bool = False): - super().__init__() - - self.outdim = size_emb_dim - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - self.use_additional_conditions = use_additional_conditions - if use_additional_conditions: - self.additional_condition_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.resolution_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=size_emb_dim) - self.aspect_ratio_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=size_emb_dim) - - def forward(self, timestep, resolution, aspect_ratio, batch_size, hidden_dtype): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) - - if self.use_additional_conditions: - resolution_emb = self.additional_condition_proj(resolution.flatten()).to(hidden_dtype) - resolution_emb = self.resolution_embedder(resolution_emb).reshape(batch_size, -1) - aspect_ratio_emb = self.additional_condition_proj(aspect_ratio.flatten()).to(hidden_dtype) - aspect_ratio_emb = self.aspect_ratio_embedder(aspect_ratio_emb).reshape(batch_size, -1) - conditioning = timesteps_emb + torch.cat([resolution_emb, aspect_ratio_emb], dim=1) - else: - conditioning = timesteps_emb - - return conditioning - - -class PixArtAlphaTextProjection(nn.Module): - """ - Projects caption embeddings. Also handles dropout for classifier-free guidance. - - Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py - """ - - def __init__(self, in_features, hidden_size, out_features=None, act_fn="gelu_tanh"): - super().__init__() - if out_features is None: - out_features = hidden_size - self.linear_1 = nn.Linear(in_features=in_features, out_features=hidden_size, bias=True) - if act_fn == "gelu_tanh": - self.act_1 = nn.GELU(approximate="tanh") - elif act_fn == "silu": - self.act_1 = nn.SiLU() - elif act_fn == "silu_fp32": - self.act_1 = FP32SiLU() - else: - raise ValueError(f"Unknown activation function: {act_fn}") - self.linear_2 = nn.Linear(in_features=hidden_size, out_features=out_features, bias=True) - - def forward(self, caption): - hidden_states = self.linear_1(caption) - hidden_states = self.act_1(hidden_states) - hidden_states = self.linear_2(hidden_states) - return hidden_states - - -class IPAdapterPlusImageProjectionBlock(nn.Module): - def __init__( - self, - embed_dims: int = 768, - dim_head: int = 64, - heads: int = 16, - ffn_ratio: float = 4, - ) -> None: - super().__init__() - from .attention import FeedForward - - self.ln0 = nn.LayerNorm(embed_dims) - self.ln1 = nn.LayerNorm(embed_dims) - self.attn = Attention( - query_dim=embed_dims, - dim_head=dim_head, - heads=heads, - out_bias=False, - ) - self.ff = nn.Sequential( - nn.LayerNorm(embed_dims), - FeedForward(embed_dims, embed_dims, activation_fn="gelu", mult=ffn_ratio, bias=False), - ) - - def forward(self, x, latents, residual): - encoder_hidden_states = self.ln0(x) - latents = self.ln1(latents) - encoder_hidden_states = torch.cat([encoder_hidden_states, latents], dim=-2) - latents = self.attn(latents, encoder_hidden_states) + residual - latents = self.ff(latents) + latents - return latents - - -class IPAdapterPlusImageProjection(nn.Module): - """Resampler of IP-Adapter Plus. - - Args: - embed_dims (int): The feature dimension. Defaults to 768. output_dims (int): The number of output channels, - that is the same - number of the channels in the `unet.config.cross_attention_dim`. Defaults to 1024. - hidden_dims (int): - The number of hidden channels. Defaults to 1280. depth (int): The number of blocks. Defaults - to 8. dim_head (int): The number of head channels. Defaults to 64. heads (int): Parallel attention heads. - Defaults to 16. num_queries (int): - The number of queries. Defaults to 8. ffn_ratio (float): The expansion ratio - of feedforward network hidden - layer channels. Defaults to 4. - """ - - def __init__( - self, - embed_dims: int = 768, - output_dims: int = 1024, - hidden_dims: int = 1280, - depth: int = 4, - dim_head: int = 64, - heads: int = 16, - num_queries: int = 8, - ffn_ratio: float = 4, - ) -> None: - super().__init__() - self.latents = nn.Parameter(torch.randn(1, num_queries, hidden_dims) / hidden_dims**0.5) - - self.proj_in = nn.Linear(embed_dims, hidden_dims) - - self.proj_out = nn.Linear(hidden_dims, output_dims) - self.norm_out = nn.LayerNorm(output_dims) - - self.layers = nn.ModuleList( - [IPAdapterPlusImageProjectionBlock(hidden_dims, dim_head, heads, ffn_ratio) for _ in range(depth)] - ) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - """Forward pass. - - Args: - x (torch.Tensor): Input Tensor. - Returns: - torch.Tensor: Output Tensor. - """ - latents = self.latents.repeat(x.size(0), 1, 1) - - x = self.proj_in(x) - - for block in self.layers: - residual = latents - latents = block(x, latents, residual) - - latents = self.proj_out(latents) - return self.norm_out(latents) - - -class IPAdapterFaceIDPlusImageProjection(nn.Module): - """FacePerceiverResampler of IP-Adapter Plus. - - Args: - embed_dims (int): The feature dimension. Defaults to 768. output_dims (int): The number of output channels, - that is the same - number of the channels in the `unet.config.cross_attention_dim`. Defaults to 1024. - hidden_dims (int): - The number of hidden channels. Defaults to 1280. depth (int): The number of blocks. Defaults - to 8. dim_head (int): The number of head channels. Defaults to 64. heads (int): Parallel attention heads. - Defaults to 16. num_tokens (int): Number of tokens num_queries (int): The number of queries. Defaults to 8. - ffn_ratio (float): The expansion ratio of feedforward network hidden - layer channels. Defaults to 4. - ffproj_ratio (float): The expansion ratio of feedforward network hidden - layer channels (for ID embeddings). Defaults to 4. - """ - - def __init__( - self, - embed_dims: int = 768, - output_dims: int = 768, - hidden_dims: int = 1280, - id_embeddings_dim: int = 512, - depth: int = 4, - dim_head: int = 64, - heads: int = 16, - num_tokens: int = 4, - num_queries: int = 8, - ffn_ratio: float = 4, - ffproj_ratio: int = 2, - ) -> None: - super().__init__() - from .attention import FeedForward - - self.num_tokens = num_tokens - self.embed_dim = embed_dims - self.clip_embeds = None - self.shortcut = False - self.shortcut_scale = 1.0 - - self.proj = FeedForward(id_embeddings_dim, embed_dims * num_tokens, activation_fn="gelu", mult=ffproj_ratio) - self.norm = nn.LayerNorm(embed_dims) - - self.proj_in = nn.Linear(hidden_dims, embed_dims) - - self.proj_out = nn.Linear(embed_dims, output_dims) - self.norm_out = nn.LayerNorm(output_dims) - - self.layers = nn.ModuleList( - [IPAdapterPlusImageProjectionBlock(embed_dims, dim_head, heads, ffn_ratio) for _ in range(depth)] - ) - - def forward(self, id_embeds: torch.Tensor) -> torch.Tensor: - """Forward pass. - - Args: - id_embeds (torch.Tensor): Input Tensor (ID embeds). - Returns: - torch.Tensor: Output Tensor. - """ - id_embeds = id_embeds.to(self.clip_embeds.dtype) - id_embeds = self.proj(id_embeds) - id_embeds = id_embeds.reshape(-1, self.num_tokens, self.embed_dim) - id_embeds = self.norm(id_embeds) - latents = id_embeds - - clip_embeds = self.proj_in(self.clip_embeds) - x = clip_embeds.reshape(-1, clip_embeds.shape[2], clip_embeds.shape[3]) - - for block in self.layers: - residual = latents - latents = block(x, latents, residual) - - latents = self.proj_out(latents) - out = self.norm_out(latents) - if self.shortcut: - out = id_embeds + self.shortcut_scale * out - return out - - -class IPAdapterTimeImageProjectionBlock(nn.Module): - """Block for IPAdapterTimeImageProjection. - - Args: - hidden_dim (`int`, defaults to 1280): - The number of hidden channels. - dim_head (`int`, defaults to 64): - The number of head channels. - heads (`int`, defaults to 20): - Parallel attention heads. - ffn_ratio (`int`, defaults to 4): - The expansion ratio of feedforward network hidden layer channels. - """ - - def __init__( - self, - hidden_dim: int = 1280, - dim_head: int = 64, - heads: int = 20, - ffn_ratio: int = 4, - ) -> None: - super().__init__() - from .attention import FeedForward - - self.ln0 = nn.LayerNorm(hidden_dim) - self.ln1 = nn.LayerNorm(hidden_dim) - self.attn = Attention( - query_dim=hidden_dim, - cross_attention_dim=hidden_dim, - dim_head=dim_head, - heads=heads, - bias=False, - out_bias=False, - ) - self.ff = FeedForward(hidden_dim, hidden_dim, activation_fn="gelu", mult=ffn_ratio, bias=False) - - # AdaLayerNorm - self.adaln_silu = nn.SiLU() - self.adaln_proj = nn.Linear(hidden_dim, 4 * hidden_dim) - self.adaln_norm = nn.LayerNorm(hidden_dim) - - # Set attention scale and fuse KV - self.attn.scale = 1 / math.sqrt(math.sqrt(dim_head)) - self.attn.fuse_projections() - self.attn.to_k = None - self.attn.to_v = None - - def forward(self, x: torch.Tensor, latents: torch.Tensor, timestep_emb: torch.Tensor) -> torch.Tensor: - """Forward pass. - - Args: - x (`torch.Tensor`): - Image features. - latents (`torch.Tensor`): - Latent features. - timestep_emb (`torch.Tensor`): - Timestep embedding. - - Returns: - `torch.Tensor`: Output latent features. - """ - - # Shift and scale for AdaLayerNorm - emb = self.adaln_proj(self.adaln_silu(timestep_emb)) - shift_msa, scale_msa, shift_mlp, scale_mlp = emb.chunk(4, dim=1) - - # Fused Attention - residual = latents - x = self.ln0(x) - latents = self.ln1(latents) * (1 + scale_msa[:, None]) + shift_msa[:, None] - - batch_size = latents.shape[0] - - query = self.attn.to_q(latents) - kv_input = torch.cat((x, latents), dim=-2) - key, value = self.attn.to_kv(kv_input).chunk(2, dim=-1) - - inner_dim = key.shape[-1] - head_dim = inner_dim // self.attn.heads - - query = query.view(batch_size, -1, self.attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, self.attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, self.attn.heads, head_dim).transpose(1, 2) - - weight = (query * self.attn.scale) @ (key * self.attn.scale).transpose(-2, -1) - weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) - latents = weight @ value - - latents = latents.transpose(1, 2).reshape(batch_size, -1, self.attn.heads * head_dim) - latents = self.attn.to_out[0](latents) - latents = self.attn.to_out[1](latents) - latents = latents + residual - - ## FeedForward - residual = latents - latents = self.adaln_norm(latents) * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - return self.ff(latents) + residual - - -# Modified from https://github.com/mlfoundations/open_flamingo/blob/main/open_flamingo/src/helpers.py -class IPAdapterTimeImageProjection(nn.Module): - """Resampler of SD3 IP-Adapter with timestep embedding. - - Args: - embed_dim (`int`, defaults to 1152): - The feature dimension. - output_dim (`int`, defaults to 2432): - The number of output channels. - hidden_dim (`int`, defaults to 1280): - The number of hidden channels. - depth (`int`, defaults to 4): - The number of blocks. - dim_head (`int`, defaults to 64): - The number of head channels. - heads (`int`, defaults to 20): - Parallel attention heads. - num_queries (`int`, defaults to 64): - The number of queries. - ffn_ratio (`int`, defaults to 4): - The expansion ratio of feedforward network hidden layer channels. - timestep_in_dim (`int`, defaults to 320): - The number of input channels for timestep embedding. - timestep_flip_sin_to_cos (`bool`, defaults to True): - Flip the timestep embedding order to `cos, sin` (if True) or `sin, cos` (if False). - timestep_freq_shift (`int`, defaults to 0): - Controls the timestep delta between frequencies between dimensions. - """ - - def __init__( - self, - embed_dim: int = 1152, - output_dim: int = 2432, - hidden_dim: int = 1280, - depth: int = 4, - dim_head: int = 64, - heads: int = 20, - num_queries: int = 64, - ffn_ratio: int = 4, - timestep_in_dim: int = 320, - timestep_flip_sin_to_cos: bool = True, - timestep_freq_shift: int = 0, - ) -> None: - super().__init__() - self.latents = nn.Parameter(torch.randn(1, num_queries, hidden_dim) / hidden_dim**0.5) - self.proj_in = nn.Linear(embed_dim, hidden_dim) - self.proj_out = nn.Linear(hidden_dim, output_dim) - self.norm_out = nn.LayerNorm(output_dim) - self.layers = nn.ModuleList( - [IPAdapterTimeImageProjectionBlock(hidden_dim, dim_head, heads, ffn_ratio) for _ in range(depth)] - ) - self.time_proj = Timesteps(timestep_in_dim, timestep_flip_sin_to_cos, timestep_freq_shift) - self.time_embedding = TimestepEmbedding(timestep_in_dim, hidden_dim, act_fn="silu") - - def forward(self, x: torch.Tensor, timestep: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - """Forward pass. - - Args: - x (`torch.Tensor`): - Image features. - timestep (`torch.Tensor`): - Timestep in denoising process. - Returns: - `tuple`[`torch.Tensor`, `torch.Tensor`]: The pair (latents, timestep_emb). - """ - timestep_emb = self.time_proj(timestep).to(dtype=x.dtype) - timestep_emb = self.time_embedding(timestep_emb) - - latents = self.latents.repeat(x.size(0), 1, 1) - - x = self.proj_in(x) - x = x + timestep_emb[:, None] - - for block in self.layers: - latents = block(x, latents, timestep_emb) - - latents = self.proj_out(latents) - latents = self.norm_out(latents) - - return latents, timestep_emb - - -class MultiIPAdapterImageProjection(nn.Module): - def __init__(self, IPAdapterImageProjectionLayers: list[nn.Module] | tuple[nn.Module]): - super().__init__() - self.image_projection_layers = nn.ModuleList(IPAdapterImageProjectionLayers) - - @property - def num_ip_adapters(self) -> int: - """Number of IP-Adapters loaded.""" - return len(self.image_projection_layers) - - def forward(self, image_embeds: list[torch.Tensor]): - projected_image_embeds = [] - - # currently, we accept `image_embeds` as - # 1. a tensor (deprecated) with shape [batch_size, embed_dim] or [batch_size, sequence_length, embed_dim] - # 2. list of `n` tensors where `n` is number of ip-adapters, each tensor can hae shape [batch_size, num_images, embed_dim] or [batch_size, num_images, sequence_length, embed_dim] - if not isinstance(image_embeds, list): - deprecation_message = ( - "You have passed a tensor as `image_embeds`.This is deprecated and will be removed in a future release." - " Please make sure to update your script to pass `image_embeds` as a list of tensors to suppress this warning." - ) - deprecate("image_embeds not a list", "1.0.0", deprecation_message, standard_warn=False) - image_embeds = [image_embeds.unsqueeze(1)] - - if len(image_embeds) != len(self.image_projection_layers): - raise ValueError( - f"image_embeds must have the same length as image_projection_layers, got {len(image_embeds)} and {len(self.image_projection_layers)}" - ) - - for image_embed, image_projection_layer in zip(image_embeds, self.image_projection_layers): - batch_size, num_images = image_embed.shape[0], image_embed.shape[1] - image_embed = image_embed.reshape((batch_size * num_images,) + image_embed.shape[2:]) - image_embed = image_projection_layer(image_embed) - image_embed = image_embed.reshape((batch_size, num_images) + image_embed.shape[1:]) - - projected_image_embeds.append(image_embed) - - return projected_image_embeds - - -class FluxPosEmbed(nn.Module): - def __new__(cls, *args, **kwargs): - deprecation_message = "Importing and using `FluxPosEmbed` from `diffusers.models.embeddings` is deprecated. Please import it from `diffusers.models.transformers.transformer_flux`." - deprecate("FluxPosEmbed", "1.0.0", deprecation_message) - - from .transformers.transformer_flux import FluxPosEmbed - - return FluxPosEmbed(*args, **kwargs) diff --git a/diffusers/models/lora.py b/diffusers/models/lora.py deleted file mode 100644 index 72e285832737c54ad6e38a6f5cb620011091bf7e..0000000000000000000000000000000000000000 --- a/diffusers/models/lora.py +++ /dev/null @@ -1,455 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -# IMPORTANT: # -################################################################### -# ----------------------------------------------------------------# -# This file is deprecated and will be removed soon # -# (as soon as PEFT will become a required dependency for LoRA) # -# ----------------------------------------------------------------# -################################################################### - -import torch -import torch.nn.functional as F -from torch import nn - -from ..utils import deprecate, logging -from ..utils.import_utils import is_transformers_available - - -if is_transformers_available(): - from transformers import CLIPTextModel, CLIPTextModelWithProjection - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def text_encoder_attn_modules(text_encoder: nn.Module): - attn_modules = [] - - if isinstance(text_encoder, (CLIPTextModel, CLIPTextModelWithProjection)): - for i, layer in enumerate(text_encoder.text_model.encoder.layers): - name = f"text_model.encoder.layers.{i}.self_attn" - mod = layer.self_attn - attn_modules.append((name, mod)) - else: - raise ValueError(f"do not know how to get attention modules for: {text_encoder.__class__.__name__}") - - return attn_modules - - -def text_encoder_mlp_modules(text_encoder: nn.Module): - mlp_modules = [] - - if isinstance(text_encoder, (CLIPTextModel, CLIPTextModelWithProjection)): - for i, layer in enumerate(text_encoder.text_model.encoder.layers): - mlp_mod = layer.mlp - name = f"text_model.encoder.layers.{i}.mlp" - mlp_modules.append((name, mlp_mod)) - else: - raise ValueError(f"do not know how to get mlp modules for: {text_encoder.__class__.__name__}") - - return mlp_modules - - -def adjust_lora_scale_text_encoder(text_encoder, lora_scale: float = 1.0): - for _, attn_module in text_encoder_attn_modules(text_encoder): - if isinstance(attn_module.q_proj, PatchedLoraProjection): - attn_module.q_proj.lora_scale = lora_scale - attn_module.k_proj.lora_scale = lora_scale - attn_module.v_proj.lora_scale = lora_scale - attn_module.out_proj.lora_scale = lora_scale - - for _, mlp_module in text_encoder_mlp_modules(text_encoder): - if isinstance(mlp_module.fc1, PatchedLoraProjection): - mlp_module.fc1.lora_scale = lora_scale - mlp_module.fc2.lora_scale = lora_scale - - -class PatchedLoraProjection(torch.nn.Module): - def __init__(self, regular_linear_layer, lora_scale=1, network_alpha=None, rank=4, dtype=None): - deprecation_message = "Use of `PatchedLoraProjection` is deprecated. Please switch to PEFT backend by installing PEFT: `pip install peft`." - deprecate("PatchedLoraProjection", "1.0.0", deprecation_message) - - super().__init__() - from ..models.lora import LoRALinearLayer - - self.regular_linear_layer = regular_linear_layer - - device = self.regular_linear_layer.weight.device - - if dtype is None: - dtype = self.regular_linear_layer.weight.dtype - - self.lora_linear_layer = LoRALinearLayer( - self.regular_linear_layer.in_features, - self.regular_linear_layer.out_features, - network_alpha=network_alpha, - device=device, - dtype=dtype, - rank=rank, - ) - - self.lora_scale = lora_scale - - # overwrite PyTorch's `state_dict` to be sure that only the 'regular_linear_layer' weights are saved - # when saving the whole text encoder model and when LoRA is unloaded or fused - def state_dict(self, *args, destination=None, prefix="", keep_vars=False): - if self.lora_linear_layer is None: - return self.regular_linear_layer.state_dict( - *args, destination=destination, prefix=prefix, keep_vars=keep_vars - ) - - return super().state_dict(*args, destination=destination, prefix=prefix, keep_vars=keep_vars) - - def _fuse_lora(self, lora_scale=1.0, safe_fusing=False): - if self.lora_linear_layer is None: - return - - dtype, device = self.regular_linear_layer.weight.data.dtype, self.regular_linear_layer.weight.data.device - - w_orig = self.regular_linear_layer.weight.data.float() - w_up = self.lora_linear_layer.up.weight.data.float() - w_down = self.lora_linear_layer.down.weight.data.float() - - if self.lora_linear_layer.network_alpha is not None: - w_up = w_up * self.lora_linear_layer.network_alpha / self.lora_linear_layer.rank - - fused_weight = w_orig + (lora_scale * torch.bmm(w_up[None, :], w_down[None, :])[0]) - - if safe_fusing and torch.isnan(fused_weight).any().item(): - raise ValueError( - "This LoRA weight seems to be broken. " - f"Encountered NaN values when trying to fuse LoRA weights for {self}." - "LoRA weights will not be fused." - ) - - self.regular_linear_layer.weight.data = fused_weight.to(device=device, dtype=dtype) - - # we can drop the lora layer now - self.lora_linear_layer = None - - # offload the up and down matrices to CPU to not blow the memory - self.w_up = w_up.cpu() - self.w_down = w_down.cpu() - self.lora_scale = lora_scale - - def _unfuse_lora(self): - if not (getattr(self, "w_up", None) is not None and getattr(self, "w_down", None) is not None): - return - - fused_weight = self.regular_linear_layer.weight.data - dtype, device = fused_weight.dtype, fused_weight.device - - w_up = self.w_up.to(device=device).float() - w_down = self.w_down.to(device).float() - - unfused_weight = fused_weight.float() - (self.lora_scale * torch.bmm(w_up[None, :], w_down[None, :])[0]) - self.regular_linear_layer.weight.data = unfused_weight.to(device=device, dtype=dtype) - - self.w_up = None - self.w_down = None - - def forward(self, input): - if self.lora_scale is None: - self.lora_scale = 1.0 - if self.lora_linear_layer is None: - return self.regular_linear_layer(input) - return self.regular_linear_layer(input) + (self.lora_scale * self.lora_linear_layer(input)) - - -class LoRALinearLayer(nn.Module): - r""" - A linear layer that is used with LoRA. - - Parameters: - in_features (`int`): - Number of input features. - out_features (`int`): - Number of output features. - rank (`int`, `optional`, defaults to 4): - The rank of the LoRA layer. - network_alpha (`float`, `optional`, defaults to `None`): - The value of the network alpha used for stable learning and preventing underflow. This value has the same - meaning as the `--network_alpha` option in the kohya-ss trainer script. See - https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning - device (`torch.device`, `optional`, defaults to `None`): - The device to use for the layer's weights. - dtype (`torch.dtype`, `optional`, defaults to `None`): - The dtype to use for the layer's weights. - """ - - def __init__( - self, - in_features: int, - out_features: int, - rank: int = 4, - network_alpha: float | None = None, - device: torch.device | str | None = None, - dtype: torch.dtype | None = None, - ): - super().__init__() - - deprecation_message = "Use of `LoRALinearLayer` is deprecated. Please switch to PEFT backend by installing PEFT: `pip install peft`." - deprecate("LoRALinearLayer", "1.0.0", deprecation_message) - - self.down = nn.Linear(in_features, rank, bias=False, device=device, dtype=dtype) - self.up = nn.Linear(rank, out_features, bias=False, device=device, dtype=dtype) - # This value has the same meaning as the `--network_alpha` option in the kohya-ss trainer script. - # See https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning - self.network_alpha = network_alpha - self.rank = rank - self.out_features = out_features - self.in_features = in_features - - nn.init.normal_(self.down.weight, std=1 / rank) - nn.init.zeros_(self.up.weight) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - orig_dtype = hidden_states.dtype - dtype = self.down.weight.dtype - - down_hidden_states = self.down(hidden_states.to(dtype)) - up_hidden_states = self.up(down_hidden_states) - - if self.network_alpha is not None: - up_hidden_states *= self.network_alpha / self.rank - - return up_hidden_states.to(orig_dtype) - - -class LoRAConv2dLayer(nn.Module): - r""" - A convolutional layer that is used with LoRA. - - Parameters: - in_features (`int`): - Number of input features. - out_features (`int`): - Number of output features. - rank (`int`, `optional`, defaults to 4): - The rank of the LoRA layer. - kernel_size (`int` or `tuple` of two `int`, `optional`, defaults to 1): - The kernel size of the convolution. - stride (`int` or `tuple` of two `int`, `optional`, defaults to 1): - The stride of the convolution. - padding (`int` or `tuple` of two `int` or `str`, `optional`, defaults to 0): - The padding of the convolution. - network_alpha (`float`, `optional`, defaults to `None`): - The value of the network alpha used for stable learning and preventing underflow. This value has the same - meaning as the `--network_alpha` option in the kohya-ss trainer script. See - https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning - """ - - def __init__( - self, - in_features: int, - out_features: int, - rank: int = 4, - kernel_size: int | tuple[int, int] = (1, 1), - stride: int | tuple[int, int] = (1, 1), - padding: int | tuple[int, int] | str = 0, - network_alpha: float | None = None, - ): - super().__init__() - - deprecation_message = "Use of `LoRAConv2dLayer` is deprecated. Please switch to PEFT backend by installing PEFT: `pip install peft`." - deprecate("LoRAConv2dLayer", "1.0.0", deprecation_message) - - self.down = nn.Conv2d(in_features, rank, kernel_size=kernel_size, stride=stride, padding=padding, bias=False) - # according to the official kohya_ss trainer kernel_size are always fixed for the up layer - # # see: https://github.com/bmaltais/kohya_ss/blob/2accb1305979ba62f5077a23aabac23b4c37e935/networks/lora_diffusers.py#L129 - self.up = nn.Conv2d(rank, out_features, kernel_size=(1, 1), stride=(1, 1), bias=False) - - # This value has the same meaning as the `--network_alpha` option in the kohya-ss trainer script. - # See https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning - self.network_alpha = network_alpha - self.rank = rank - - nn.init.normal_(self.down.weight, std=1 / rank) - nn.init.zeros_(self.up.weight) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - orig_dtype = hidden_states.dtype - dtype = self.down.weight.dtype - - down_hidden_states = self.down(hidden_states.to(dtype)) - up_hidden_states = self.up(down_hidden_states) - - if self.network_alpha is not None: - up_hidden_states *= self.network_alpha / self.rank - - return up_hidden_states.to(orig_dtype) - - -class LoRACompatibleConv(nn.Conv2d): - """ - A convolutional layer that can be used with LoRA. - """ - - def __init__(self, *args, lora_layer: LoRAConv2dLayer | None = None, **kwargs): - deprecation_message = "Use of `LoRACompatibleConv` is deprecated. Please switch to PEFT backend by installing PEFT: `pip install peft`." - deprecate("LoRACompatibleConv", "1.0.0", deprecation_message) - - super().__init__(*args, **kwargs) - self.lora_layer = lora_layer - - def set_lora_layer(self, lora_layer: LoRAConv2dLayer | None): - deprecation_message = "Use of `set_lora_layer()` is deprecated. Please switch to PEFT backend by installing PEFT: `pip install peft`." - deprecate("set_lora_layer", "1.0.0", deprecation_message) - - self.lora_layer = lora_layer - - def _fuse_lora(self, lora_scale: float = 1.0, safe_fusing: bool = False): - if self.lora_layer is None: - return - - dtype, device = self.weight.data.dtype, self.weight.data.device - - w_orig = self.weight.data.float() - w_up = self.lora_layer.up.weight.data.float() - w_down = self.lora_layer.down.weight.data.float() - - if self.lora_layer.network_alpha is not None: - w_up = w_up * self.lora_layer.network_alpha / self.lora_layer.rank - - fusion = torch.mm(w_up.flatten(start_dim=1), w_down.flatten(start_dim=1)) - fusion = fusion.reshape((w_orig.shape)) - fused_weight = w_orig + (lora_scale * fusion) - - if safe_fusing and torch.isnan(fused_weight).any().item(): - raise ValueError( - "This LoRA weight seems to be broken. " - f"Encountered NaN values when trying to fuse LoRA weights for {self}." - "LoRA weights will not be fused." - ) - - self.weight.data = fused_weight.to(device=device, dtype=dtype) - - # we can drop the lora layer now - self.lora_layer = None - - # offload the up and down matrices to CPU to not blow the memory - self.w_up = w_up.cpu() - self.w_down = w_down.cpu() - self._lora_scale = lora_scale - - def _unfuse_lora(self): - if not (getattr(self, "w_up", None) is not None and getattr(self, "w_down", None) is not None): - return - - fused_weight = self.weight.data - dtype, device = fused_weight.data.dtype, fused_weight.data.device - - self.w_up = self.w_up.to(device=device).float() - self.w_down = self.w_down.to(device).float() - - fusion = torch.mm(self.w_up.flatten(start_dim=1), self.w_down.flatten(start_dim=1)) - fusion = fusion.reshape((fused_weight.shape)) - unfused_weight = fused_weight.float() - (self._lora_scale * fusion) - self.weight.data = unfused_weight.to(device=device, dtype=dtype) - - self.w_up = None - self.w_down = None - - def forward(self, hidden_states: torch.Tensor, scale: float = 1.0) -> torch.Tensor: - if self.padding_mode != "zeros": - hidden_states = F.pad(hidden_states, self._reversed_padding_repeated_twice, mode=self.padding_mode) - padding = (0, 0) - else: - padding = self.padding - - original_outputs = F.conv2d( - hidden_states, self.weight, self.bias, self.stride, padding, self.dilation, self.groups - ) - - if self.lora_layer is None: - return original_outputs - else: - return original_outputs + (scale * self.lora_layer(hidden_states)) - - -class LoRACompatibleLinear(nn.Linear): - """ - A Linear layer that can be used with LoRA. - """ - - def __init__(self, *args, lora_layer: LoRALinearLayer | None = None, **kwargs): - deprecation_message = "Use of `LoRACompatibleLinear` is deprecated. Please switch to PEFT backend by installing PEFT: `pip install peft`." - deprecate("LoRACompatibleLinear", "1.0.0", deprecation_message) - - super().__init__(*args, **kwargs) - self.lora_layer = lora_layer - - def set_lora_layer(self, lora_layer: LoRALinearLayer | None): - deprecation_message = "Use of `set_lora_layer()` is deprecated. Please switch to PEFT backend by installing PEFT: `pip install peft`." - deprecate("set_lora_layer", "1.0.0", deprecation_message) - self.lora_layer = lora_layer - - def _fuse_lora(self, lora_scale: float = 1.0, safe_fusing: bool = False): - if self.lora_layer is None: - return - - dtype, device = self.weight.data.dtype, self.weight.data.device - - w_orig = self.weight.data.float() - w_up = self.lora_layer.up.weight.data.float() - w_down = self.lora_layer.down.weight.data.float() - - if self.lora_layer.network_alpha is not None: - w_up = w_up * self.lora_layer.network_alpha / self.lora_layer.rank - - fused_weight = w_orig + (lora_scale * torch.bmm(w_up[None, :], w_down[None, :])[0]) - - if safe_fusing and torch.isnan(fused_weight).any().item(): - raise ValueError( - "This LoRA weight seems to be broken. " - f"Encountered NaN values when trying to fuse LoRA weights for {self}." - "LoRA weights will not be fused." - ) - - self.weight.data = fused_weight.to(device=device, dtype=dtype) - - # we can drop the lora layer now - self.lora_layer = None - - # offload the up and down matrices to CPU to not blow the memory - self.w_up = w_up.cpu() - self.w_down = w_down.cpu() - self._lora_scale = lora_scale - - def _unfuse_lora(self): - if not (getattr(self, "w_up", None) is not None and getattr(self, "w_down", None) is not None): - return - - fused_weight = self.weight.data - dtype, device = fused_weight.dtype, fused_weight.device - - w_up = self.w_up.to(device=device).float() - w_down = self.w_down.to(device).float() - - unfused_weight = fused_weight.float() - (self._lora_scale * torch.bmm(w_up[None, :], w_down[None, :])[0]) - self.weight.data = unfused_weight.to(device=device, dtype=dtype) - - self.w_up = None - self.w_down = None - - def forward(self, hidden_states: torch.Tensor, scale: float = 1.0) -> torch.Tensor: - if self.lora_layer is None: - out = super().forward(hidden_states) - return out - else: - out = super().forward(hidden_states) + (scale * self.lora_layer(hidden_states)) - return out diff --git a/diffusers/models/model_loading_utils.py b/diffusers/models/model_loading_utils.py deleted file mode 100644 index abbde8082bb5b1d3d13b61e0082ac0f4bd9d797a..0000000000000000000000000000000000000000 --- a/diffusers/models/model_loading_utils.py +++ /dev/null @@ -1,761 +0,0 @@ -# coding=utf-8 -# Copyright 2025 The HuggingFace Inc. team. -# Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import functools -import importlib -import inspect -import os -from array import array -from collections import OrderedDict, defaultdict -from concurrent.futures import ThreadPoolExecutor, as_completed -from pathlib import Path -from zipfile import is_zipfile - -import safetensors -import torch -from huggingface_hub import DDUFEntry -from huggingface_hub.utils import EntryNotFoundError - -from ..quantizers import DiffusersQuantizer -from ..utils import ( - DEFAULT_HF_PARALLEL_LOADING_WORKERS, - GGUF_FILE_EXTENSION, - SAFE_WEIGHTS_INDEX_NAME, - SAFETENSORS_FILE_EXTENSION, - WEIGHTS_INDEX_NAME, - _add_variant, - _get_model_file, - deprecate, - is_accelerate_available, - is_accelerate_version, - is_gguf_available, - is_torch_available, - is_torch_version, - logging, -) -from ..utils.distributed_utils import is_torch_dist_rank_zero - - -logger = logging.get_logger(__name__) - -_CLASS_REMAPPING_DICT = { - "Transformer2DModel": { - "ada_norm_zero": "DiTTransformer2DModel", - "ada_norm_single": "PixArtTransformer2DModel", - } -} - - -if is_accelerate_available(): - from accelerate import infer_auto_device_map - from accelerate.utils import get_balanced_memory, get_max_memory, offload_weight, set_module_tensor_to_device - - -# Adapted from `transformers` (see modeling_utils.py) -def _determine_device_map( - model: torch.nn.Module, device_map, max_memory, torch_dtype, keep_in_fp32_modules=[], hf_quantizer=None -): - if isinstance(device_map, str): - special_dtypes = {} - if hf_quantizer is not None: - special_dtypes.update(hf_quantizer.get_special_dtypes_update(model, torch_dtype)) - special_dtypes.update( - { - name: torch.float32 - for name, _ in model.named_parameters() - if any(m in name for m in keep_in_fp32_modules) - } - ) - - target_dtype = torch_dtype - if hf_quantizer is not None: - target_dtype = hf_quantizer.adjust_target_dtype(target_dtype) - - no_split_modules = model._get_no_split_modules(device_map) - device_map_kwargs = {"no_split_module_classes": no_split_modules} - - if "special_dtypes" in inspect.signature(infer_auto_device_map).parameters: - device_map_kwargs["special_dtypes"] = special_dtypes - elif len(special_dtypes) > 0: - logger.warning( - "This model has some weights that should be kept in higher precision, you need to upgrade " - "`accelerate` to properly deal with them (`pip install --upgrade accelerate`)." - ) - - if device_map != "sequential": - max_memory = get_balanced_memory( - model, - dtype=torch_dtype, - low_zero=(device_map == "balanced_low_0"), - max_memory=max_memory, - **device_map_kwargs, - ) - else: - max_memory = get_max_memory(max_memory) - - if hf_quantizer is not None: - max_memory = hf_quantizer.adjust_max_memory(max_memory) - - device_map_kwargs["max_memory"] = max_memory - device_map = infer_auto_device_map(model, dtype=target_dtype, **device_map_kwargs) - - return device_map - - -def _fetch_remapped_cls_from_config(config, old_class): - previous_class_name = old_class.__name__ - remapped_class_name = _CLASS_REMAPPING_DICT.get(previous_class_name).get(config["norm_type"], None) - - # Details: - # https://github.com/huggingface/diffusers/pull/7647#discussion_r1621344818 - if remapped_class_name: - # load diffusers library to import compatible and original scheduler - diffusers_library = importlib.import_module(__name__.split(".")[0]) - remapped_class = getattr(diffusers_library, remapped_class_name) - logger.info( - f"Changing class object to be of `{remapped_class_name}` type from `{previous_class_name}` type." - f"This is because `{previous_class_name}` is scheduled to be deprecated in a future version. Note that this" - " DOESN'T affect the final results." - ) - return remapped_class - else: - return old_class - - -def _determine_param_device(param_name: str, device_map: dict[str, int | str | torch.device] | None): - """ - Find the device of param_name from the device_map. - """ - if device_map is None: - return "cpu" - else: - module_name = param_name - # find next higher level module that is defined in device_map: - # bert.lm_head.weight -> bert.lm_head -> bert -> '' - while len(module_name) > 0 and module_name not in device_map: - module_name = ".".join(module_name.split(".")[:-1]) - if module_name == "" and "" not in device_map: - raise ValueError(f"{param_name} doesn't have any device set.") - return device_map[module_name] - - -def load_state_dict( - checkpoint_file: str | os.PathLike, - dduf_entries: dict[str, DDUFEntry] | None = None, - disable_mmap: bool = False, - map_location: str | torch.device = "cpu", -): - """ - Reads a checkpoint file, returning properly formatted errors if they arise. - """ - # TODO: maybe refactor a bit this part where we pass a dict here - if isinstance(checkpoint_file, dict): - return checkpoint_file - try: - file_extension = os.path.basename(checkpoint_file).split(".")[-1] - if file_extension == SAFETENSORS_FILE_EXTENSION: - if dduf_entries: - # tensors are loaded on cpu - with dduf_entries[checkpoint_file].as_mmap() as mm: - return safetensors.torch.load(mm) - if disable_mmap: - return safetensors.torch.load(open(checkpoint_file, "rb").read()) - else: - return safetensors.torch.load_file(checkpoint_file, device=map_location) - elif file_extension == GGUF_FILE_EXTENSION: - return load_gguf_checkpoint(checkpoint_file) - else: - extra_args = {} - weights_only_kwarg = {"weights_only": True} if is_torch_version(">=", "1.13") else {} - # mmap can only be used with files serialized with zipfile-based format. - if ( - isinstance(checkpoint_file, str) - and map_location != "meta" - and is_torch_version(">=", "2.1.0") - and is_zipfile(checkpoint_file) - and not disable_mmap - ): - extra_args = {"mmap": True} - return torch.load(checkpoint_file, map_location=map_location, **weights_only_kwarg, **extra_args) - except Exception as e: - try: - with open(checkpoint_file) as f: - if f.read().startswith("version"): - raise OSError( - "You seem to have cloned a repository without having git-lfs installed. Please install " - "git-lfs and run `git lfs install` followed by `git lfs pull` in the folder " - "you cloned." - ) - else: - raise ValueError( - f"Unable to locate the file {checkpoint_file} which is necessary to load this pretrained " - "model. Make sure you have saved the model properly." - ) from e - except (UnicodeDecodeError, ValueError): - raise OSError( - f"Unable to load weights from checkpoint file for '{checkpoint_file}' at '{checkpoint_file}'. " - ) - - -def load_model_dict_into_meta( - model, - state_dict: OrderedDict, - dtype: str | torch.dtype | None = None, - model_name_or_path: str | None = None, - hf_quantizer: DiffusersQuantizer | None = None, - keep_in_fp32_modules: list | None = None, - device_map: dict[str, int | str | torch.device] | None = None, - unexpected_keys: list[str] | None = None, - offload_folder: str | os.PathLike | None = None, - offload_index: dict | None = None, - state_dict_index: dict | None = None, - state_dict_folder: str | os.PathLike | None = None, -) -> list[str]: - """ - This is somewhat similar to `_load_state_dict_into_model`, but deals with a model that has some or all of its - params on a `meta` device. It replaces the model params with the data from the `state_dict` - """ - - is_quantized = hf_quantizer is not None - empty_state_dict = model.state_dict() - - for param_name, param in state_dict.items(): - if param_name not in empty_state_dict: - continue - - set_module_kwargs = {} - # We convert floating dtypes to the `dtype` passed. We also want to keep the buffers/params - # in int/uint/bool and not cast them. - # TODO: revisit cases when param.dtype == torch.float8_e4m3fn - if dtype is not None and torch.is_floating_point(param): - if keep_in_fp32_modules is not None and any( - module_to_keep_in_fp32 in param_name.split(".") for module_to_keep_in_fp32 in keep_in_fp32_modules - ): - param = param.to(torch.float32) - set_module_kwargs["dtype"] = torch.float32 - # For quantizers have save weights using torch.float8_e4m3fn - elif hf_quantizer is not None and param.dtype == getattr(torch, "float8_e4m3fn", None): - pass - else: - param = param.to(dtype) - set_module_kwargs["dtype"] = dtype - - if is_accelerate_version(">", "1.8.1"): - set_module_kwargs["non_blocking"] = True - set_module_kwargs["clear_cache"] = False - - # For compatibility with PyTorch load_state_dict which converts state dict dtype to existing dtype in model, and which - # uses `param.copy_(input_param)` that preserves the contiguity of the parameter in the model. - # Reference: https://github.com/pytorch/pytorch/blob/db79ceb110f6646523019a59bbd7b838f43d4a86/torch/nn/modules/module.py#L2040C29-L2040C29 - old_param = model - splits = param_name.split(".") - for split in splits: - old_param = getattr(old_param, split) - - if not isinstance(old_param, (torch.nn.Parameter, torch.Tensor)): - old_param = None - - if old_param is not None: - if dtype is None: - param = param.to(old_param.dtype) - - if old_param.is_contiguous(): - param = param.contiguous() - - param_device = _determine_param_device(param_name, device_map) - - # bnb params are flattened. - # gguf quants have a different shape based on the type of quantization applied - if empty_state_dict[param_name].shape != param.shape: - if ( - is_quantized - and hf_quantizer.pre_quantized - and hf_quantizer.check_if_quantized_param( - model, param, param_name, state_dict, param_device=param_device - ) - ): - hf_quantizer.check_quantized_param_shape(param_name, empty_state_dict[param_name], param) - else: - model_name_or_path_str = f"{model_name_or_path} " if model_name_or_path is not None else "" - raise ValueError( - f"Cannot load {model_name_or_path_str} because {param_name} expected shape {empty_state_dict[param_name].shape}, but got {param.shape}. If you want to instead overwrite randomly initialized weights, please make sure to pass both `low_cpu_mem_usage=False` and `ignore_mismatched_sizes=True`. For more information, see also: https://github.com/huggingface/diffusers/issues/1619#issuecomment-1345604389 as an example." - ) - if param_device == "disk": - offload_index = offload_weight(param, param_name, offload_folder, offload_index) - elif param_device == "cpu" and state_dict_index is not None: - state_dict_index = offload_weight(param, param_name, state_dict_folder, state_dict_index) - elif is_quantized and ( - hf_quantizer.check_if_quantized_param(model, param, param_name, state_dict, param_device=param_device) - ): - hf_quantizer.create_quantized_param( - model, param, param_name, param_device, state_dict, unexpected_keys, dtype=dtype - ) - else: - set_module_tensor_to_device(model, param_name, param_device, value=param, **set_module_kwargs) - - return offload_index, state_dict_index - - -def check_support_param_buffer_assignment(model_to_load, state_dict, start_prefix=""): - """ - Checks if `model_to_load` supports param buffer assignment (such as when loading in empty weights) by first - checking if the model explicitly disables it, then by ensuring that the state dict keys are a subset of the model's - parameters. - - """ - if model_to_load.device.type == "meta": - return False - - if len([key for key in state_dict if key.startswith(start_prefix)]) == 0: - return False - - # Some models explicitly do not support param buffer assignment - if not getattr(model_to_load, "_supports_param_buffer_assignment", True): - logger.debug( - f"{model_to_load.__class__.__name__} does not support param buffer assignment, loading will be slower" - ) - return False - - # If the model does, the incoming `state_dict` and the `model_to_load` must be the same dtype - first_key = next(iter(model_to_load.state_dict().keys())) - if start_prefix + first_key in state_dict: - return state_dict[start_prefix + first_key].dtype == model_to_load.state_dict()[first_key].dtype - - return False - - -def _load_shard_file( - shard_file, - model, - model_state_dict, - device_map=None, - dtype=None, - hf_quantizer=None, - keep_in_fp32_modules=None, - dduf_entries=None, - loaded_keys=None, - unexpected_keys=None, - offload_index=None, - offload_folder=None, - state_dict_index=None, - state_dict_folder=None, - ignore_mismatched_sizes=False, - low_cpu_mem_usage=False, - disable_mmap=False, -): - state_dict = load_state_dict(shard_file, dduf_entries=dduf_entries, disable_mmap=disable_mmap) - if hf_quantizer is not None: - state_dict = hf_quantizer.maybe_update_state_dict(state_dict) - - mismatched_keys = _find_mismatched_keys( - state_dict, - model_state_dict, - loaded_keys, - ignore_mismatched_sizes, - ) - error_msgs = [] - if low_cpu_mem_usage: - offload_index, state_dict_index = load_model_dict_into_meta( - model, - state_dict, - device_map=device_map, - dtype=dtype, - hf_quantizer=hf_quantizer, - keep_in_fp32_modules=keep_in_fp32_modules, - unexpected_keys=unexpected_keys, - offload_folder=offload_folder, - offload_index=offload_index, - state_dict_index=state_dict_index, - state_dict_folder=state_dict_folder, - ) - else: - assign_to_params_buffers = check_support_param_buffer_assignment(model, state_dict) - - error_msgs += _load_state_dict_into_model(model, state_dict, assign_to_params_buffers) - return offload_index, state_dict_index, mismatched_keys, error_msgs - - -def _load_shard_files_with_threadpool( - shard_files, - model, - model_state_dict, - device_map=None, - dtype=None, - hf_quantizer=None, - keep_in_fp32_modules=None, - dduf_entries=None, - loaded_keys=None, - unexpected_keys=None, - offload_index=None, - offload_folder=None, - state_dict_index=None, - state_dict_folder=None, - ignore_mismatched_sizes=False, - low_cpu_mem_usage=False, - disable_mmap=False, -): - # Do not spawn anymore workers than you need - num_workers = min(len(shard_files), DEFAULT_HF_PARALLEL_LOADING_WORKERS) - - logger.info(f"Loading model weights in parallel with {num_workers} workers...") - - error_msgs = [] - mismatched_keys = [] - - load_one = functools.partial( - _load_shard_file, - model=model, - model_state_dict=model_state_dict, - device_map=device_map, - dtype=dtype, - hf_quantizer=hf_quantizer, - keep_in_fp32_modules=keep_in_fp32_modules, - dduf_entries=dduf_entries, - loaded_keys=loaded_keys, - unexpected_keys=unexpected_keys, - offload_index=offload_index, - offload_folder=offload_folder, - state_dict_index=state_dict_index, - state_dict_folder=state_dict_folder, - ignore_mismatched_sizes=ignore_mismatched_sizes, - low_cpu_mem_usage=low_cpu_mem_usage, - disable_mmap=disable_mmap, - ) - - tqdm_kwargs = {"total": len(shard_files), "desc": "Loading checkpoint shards"} - if not is_torch_dist_rank_zero(): - tqdm_kwargs["disable"] = True - - with ThreadPoolExecutor(max_workers=num_workers) as executor: - with logging.tqdm(**tqdm_kwargs) as pbar: - futures = [executor.submit(load_one, shard_file) for shard_file in shard_files] - for future in as_completed(futures): - result = future.result() - offload_index, state_dict_index, _mismatched_keys, _error_msgs = result - error_msgs += _error_msgs - mismatched_keys += _mismatched_keys - pbar.update(1) - - return offload_index, state_dict_index, mismatched_keys, error_msgs - - -def _find_mismatched_keys( - state_dict, - model_state_dict, - loaded_keys, - ignore_mismatched_sizes, -): - mismatched_keys = [] - if ignore_mismatched_sizes: - for checkpoint_key in loaded_keys: - model_key = checkpoint_key - # If the checkpoint is sharded, we may not have the key here. - if checkpoint_key not in state_dict: - continue - - if model_key in model_state_dict and state_dict[checkpoint_key].shape != model_state_dict[model_key].shape: - mismatched_keys.append( - (checkpoint_key, state_dict[checkpoint_key].shape, model_state_dict[model_key].shape) - ) - del state_dict[checkpoint_key] - return mismatched_keys - - -def _load_state_dict_into_model( - model_to_load, state_dict: OrderedDict, assign_to_params_buffers: bool = False -) -> list[str]: - # Convert old format to new format if needed from a PyTorch state_dict - # copy state_dict so _load_from_state_dict can modify it - state_dict = state_dict.copy() - error_msgs = [] - - # PyTorch's `_load_from_state_dict` does not copy parameters in a module's descendants - # so we need to apply the function recursively. - def load(module: torch.nn.Module, prefix: str = "", assign_to_params_buffers: bool = False): - local_metadata = {} - local_metadata["assign_to_params_buffers"] = assign_to_params_buffers - if assign_to_params_buffers and not is_torch_version(">=", "2.1"): - logger.info("You need to have torch>=2.1 in order to load the model with assign_to_params_buffers=True") - args = (state_dict, prefix, local_metadata, True, [], [], error_msgs) - module._load_from_state_dict(*args) - - for name, child in module._modules.items(): - if child is not None: - load(child, prefix + name + ".", assign_to_params_buffers) - - load(model_to_load, assign_to_params_buffers=assign_to_params_buffers) - - return error_msgs - - -def _fetch_index_file( - is_local, - pretrained_model_name_or_path, - subfolder, - use_safetensors, - cache_dir, - variant, - force_download, - proxies, - local_files_only, - token, - revision, - user_agent, - commit_hash, - dduf_entries: dict[str, DDUFEntry] | None = None, -): - if is_local: - index_file = Path( - pretrained_model_name_or_path, - subfolder or "", - _add_variant(SAFE_WEIGHTS_INDEX_NAME if use_safetensors else WEIGHTS_INDEX_NAME, variant), - ) - else: - index_file_in_repo = Path( - subfolder or "", - _add_variant(SAFE_WEIGHTS_INDEX_NAME if use_safetensors else WEIGHTS_INDEX_NAME, variant), - ).as_posix() - try: - index_file = _get_model_file( - pretrained_model_name_or_path, - weights_name=index_file_in_repo, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=None, - user_agent=user_agent, - commit_hash=commit_hash, - dduf_entries=dduf_entries, - ) - if not dduf_entries: - index_file = Path(index_file) - except (EntryNotFoundError, EnvironmentError): - index_file = None - - return index_file - - -def _fetch_index_file_legacy( - is_local, - pretrained_model_name_or_path, - subfolder, - use_safetensors, - cache_dir, - variant, - force_download, - proxies, - local_files_only, - token, - revision, - user_agent, - commit_hash, - dduf_entries: dict[str, DDUFEntry] | None = None, -): - if is_local: - index_file = Path( - pretrained_model_name_or_path, - subfolder or "", - SAFE_WEIGHTS_INDEX_NAME if use_safetensors else WEIGHTS_INDEX_NAME, - ).as_posix() - splits = index_file.split(".") - split_index = -3 if ".cache" in index_file else -2 - splits = splits[:-split_index] + [variant] + splits[-split_index:] - index_file = ".".join(splits) - if os.path.exists(index_file): - deprecation_message = f"This serialization format is now deprecated to standardize the serialization format between `transformers` and `diffusers`. We recommend you to remove the existing files associated with the current variant ({variant}) and re-obtain them by running a `save_pretrained()`." - deprecate("legacy_sharded_ckpts_with_variant", "1.0.0", deprecation_message, standard_warn=False) - index_file = Path(index_file) - else: - index_file = None - else: - if variant is not None: - index_file_in_repo = Path( - subfolder or "", - SAFE_WEIGHTS_INDEX_NAME if use_safetensors else WEIGHTS_INDEX_NAME, - ).as_posix() - splits = index_file_in_repo.split(".") - split_index = -2 - splits = splits[:-split_index] + [variant] + splits[-split_index:] - index_file_in_repo = ".".join(splits) - try: - index_file = _get_model_file( - pretrained_model_name_or_path, - weights_name=index_file_in_repo, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=None, - user_agent=user_agent, - commit_hash=commit_hash, - dduf_entries=dduf_entries, - ) - index_file = Path(index_file) - deprecation_message = f"This serialization format is now deprecated to standardize the serialization format between `transformers` and `diffusers`. We recommend you to remove the existing files associated with the current variant ({variant}) and re-obtain them by running a `save_pretrained()`." - deprecate("legacy_sharded_ckpts_with_variant", "1.0.0", deprecation_message, standard_warn=False) - except (EntryNotFoundError, EnvironmentError): - index_file = None - - return index_file - - -def _gguf_parse_value(_value, data_type): - if not isinstance(data_type, list): - data_type = [data_type] - if len(data_type) == 1: - data_type = data_type[0] - array_data_type = None - else: - if data_type[0] != 9: - raise ValueError("Received multiple types, therefore expected the first type to indicate an array.") - data_type, array_data_type = data_type - - if data_type in [0, 1, 2, 3, 4, 5, 10, 11]: - _value = int(_value[0]) - elif data_type in [6, 12]: - _value = float(_value[0]) - elif data_type in [7]: - _value = bool(_value[0]) - elif data_type in [8]: - _value = array("B", list(_value)).tobytes().decode() - elif data_type in [9]: - _value = _gguf_parse_value(_value, array_data_type) - return _value - - -def load_gguf_checkpoint(gguf_checkpoint_path, return_tensors=False): - """ - Load a GGUF file and return a dictionary of parsed parameters containing tensors, the parsed tokenizer and config - attributes. - - Args: - gguf_checkpoint_path (`str`): - The path the to GGUF file to load - return_tensors (`bool`, defaults to `True`): - Whether to read the tensors from the file and return them. Not doing so is faster and only loads the - metadata in memory. - """ - - if is_gguf_available() and is_torch_available(): - import gguf - from gguf import GGUFReader - - from ..quantizers.gguf.utils import SUPPORTED_GGUF_QUANT_TYPES, GGUFParameter - else: - logger.error( - "Loading a GGUF checkpoint in PyTorch, requires both PyTorch and GGUF>=0.10.0 to be installed. Please see " - "https://pytorch.org/ and https://github.com/ggerganov/llama.cpp/tree/master/gguf-py for installation instructions." - ) - raise ImportError("Please install torch and gguf>=0.10.0 to load a GGUF checkpoint in PyTorch.") - - reader = GGUFReader(gguf_checkpoint_path) - - parsed_parameters = {} - for tensor in reader.tensors: - name = tensor.name - quant_type = tensor.tensor_type - - # if the tensor is a torch supported dtype do not use GGUFParameter - is_gguf_quant = quant_type not in [gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16] - if is_gguf_quant and quant_type not in SUPPORTED_GGUF_QUANT_TYPES: - _supported_quants_str = "\n".join([str(type) for type in SUPPORTED_GGUF_QUANT_TYPES]) - raise ValueError( - ( - f"{name} has a quantization type: {str(quant_type)} which is unsupported." - "\n\nCurrently the following quantization types are supported: \n\n" - f"{_supported_quants_str}" - "\n\nTo request support for this quantization type please open an issue here: https://github.com/huggingface/diffusers" - ) - ) - - weights = torch.from_numpy(tensor.data.copy()) - parsed_parameters[name] = GGUFParameter(weights, quant_type=quant_type) if is_gguf_quant else weights - - return parsed_parameters - - -def _find_mismatched_keys(state_dict, model_state_dict, loaded_keys, ignore_mismatched_sizes): - mismatched_keys = [] - if not ignore_mismatched_sizes: - return mismatched_keys - for checkpoint_key in loaded_keys: - model_key = checkpoint_key - # If the checkpoint is sharded, we may not have the key here. - if checkpoint_key not in state_dict: - continue - - if model_key in model_state_dict and state_dict[checkpoint_key].shape != model_state_dict[model_key].shape: - mismatched_keys.append( - (checkpoint_key, state_dict[checkpoint_key].shape, model_state_dict[model_key].shape) - ) - del state_dict[checkpoint_key] - return mismatched_keys - - -def _expand_device_map(device_map, param_names): - """ - Expand a device map to return the correspondence parameter name to device. - """ - new_device_map = {} - for module, device in device_map.items(): - new_device_map.update( - {p: device for p in param_names if p == module or p.startswith(f"{module}.") or module == ""} - ) - return new_device_map - - -# Adapted from: https://github.com/huggingface/transformers/blob/0687d481e2c71544501ef9cb3eef795a6e79b1de/src/transformers/modeling_utils.py#L5859 -def _caching_allocator_warmup( - model, expanded_device_map: dict[str, torch.device], dtype: torch.dtype, hf_quantizer: DiffusersQuantizer | None -) -> None: - """ - This function warm-ups the caching allocator based on the size of the model tensors that will reside on each - device. It allows to have one large call to Malloc, instead of recursively calling it later when loading the model, - which is actually the loading speed bottleneck. Calling this function allows to cut the model loading time by a - very large margin. - """ - factor = 2 if hf_quantizer is None else hf_quantizer.get_cuda_warm_up_factor() - - # Keep only accelerator devices - accelerator_device_map = { - param: torch.device(device) - for param, device in expanded_device_map.items() - if str(device) not in ["cpu", "disk"] - } - if not accelerator_device_map: - return - - elements_per_device = defaultdict(int) - for param_name, device in accelerator_device_map.items(): - try: - p = model.get_parameter(param_name) - except AttributeError: - try: - p = model.get_buffer(param_name) - except AttributeError: - raise AttributeError(f"Parameter or buffer with name={param_name} not found in model") - # TODO: account for TP when needed. - elements_per_device[device] += p.numel() - - # This will kick off the caching allocator to avoid having to Malloc afterwards - for device, elem_count in elements_per_device.items(): - warmup_elems = max(1, elem_count // factor) - _ = torch.empty(warmup_elems, dtype=dtype, device=device, requires_grad=False) diff --git a/diffusers/models/modeling_outputs.py b/diffusers/models/modeling_outputs.py deleted file mode 100644 index 0120a34d9052fe6b499b519ba4366e7a088f7910..0000000000000000000000000000000000000000 --- a/diffusers/models/modeling_outputs.py +++ /dev/null @@ -1,31 +0,0 @@ -from dataclasses import dataclass - -from ..utils import BaseOutput - - -@dataclass -class AutoencoderKLOutput(BaseOutput): - """ - Output of AutoencoderKL encoding method. - - Args: - latent_dist (`DiagonalGaussianDistribution`): - Encoded outputs of `Encoder` represented as the mean and logvar of `DiagonalGaussianDistribution`. - `DiagonalGaussianDistribution` allows for sampling latents from the distribution. - """ - - latent_dist: "DiagonalGaussianDistribution" # noqa: F821 - - -@dataclass -class Transformer2DModelOutput(BaseOutput): - """ - The output of [`Transformer2DModel`]. - - Args: - sample (`torch.Tensor` of shape `(batch_size, num_channels, height, width)` or `(batch size, num_vector_embeds - 1, num_latent_pixels)` if [`Transformer2DModel`] is discrete): - The hidden states output conditioned on the `encoder_hidden_states` input. If discrete, returns probability - distributions for the unnoised latent pixels. - """ - - sample: "torch.Tensor" # noqa: F821 diff --git a/diffusers/models/modeling_utils.py b/diffusers/models/modeling_utils.py deleted file mode 100644 index 61dfc3133fbd702d69a4d055eeea08b2ee5049ee..0000000000000000000000000000000000000000 --- a/diffusers/models/modeling_utils.py +++ /dev/null @@ -1,2138 +0,0 @@ -# coding=utf-8 -# Copyright 2025 The HuggingFace Inc. team. -# Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import copy -import functools -import inspect -import itertools -import json -import os -import re -import shutil -import tempfile -from collections import OrderedDict -from contextlib import ExitStack, contextmanager -from functools import wraps -from pathlib import Path -from typing import Any, Callable, ContextManager, Type - -import safetensors -import torch -import torch.utils.checkpoint -from huggingface_hub import DDUFEntry, create_repo, split_torch_state_dict_into_shards -from huggingface_hub.utils import validate_hf_hub_args -from torch import Tensor, nn -from typing_extensions import Self - -from .. import __version__ -from ..quantizers import DiffusersAutoQuantizer, DiffusersQuantizer -from ..quantizers.quantization_config import QuantizationMethod -from ..utils import ( - CONFIG_NAME, - FLASHPACK_WEIGHTS_NAME, - HF_ENABLE_PARALLEL_LOADING, - SAFE_WEIGHTS_INDEX_NAME, - SAFETENSORS_WEIGHTS_NAME, - WEIGHTS_INDEX_NAME, - WEIGHTS_NAME, - _add_variant, - _get_checkpoint_shard_files, - _get_model_file, - deprecate, - is_accelerate_available, - is_bitsandbytes_available, - is_bitsandbytes_version, - is_flashpack_available, - is_peft_available, - is_torch_version, - logging, -) -from ..utils.distributed_utils import is_torch_dist_rank_zero -from ..utils.hub_utils import PushToHubMixin, load_or_create_model_card, populate_model_card -from ..utils.torch_utils import empty_device_cache -from ._modeling_parallel import ContextParallelConfig, ContextParallelModelPlan, ParallelConfig -from .model_loading_utils import ( - _caching_allocator_warmup, - _determine_device_map, - _expand_device_map, - _fetch_index_file, - _fetch_index_file_legacy, - _load_shard_file, - _load_shard_files_with_threadpool, - load_state_dict, -) - - -class ContextManagers: - """ - Wrapper for `contextlib.ExitStack` which enters a collection of context managers. Adaptation of `ContextManagers` - in the `fastcore` library. - """ - - def __init__(self, context_managers: list[ContextManager]): - self.context_managers = context_managers - self.stack = ExitStack() - - def __enter__(self): - for context_manager in self.context_managers: - self.stack.enter_context(context_manager) - - def __exit__(self, *args, **kwargs): - self.stack.__exit__(*args, **kwargs) - - -logger = logging.get_logger(__name__) - -_REGEX_SHARD = re.compile(r"(.*?)-\d{5}-of-\d{5}") - -# The `user_agent` dict is flattened into a single `user-agent` HTTP header. Serializing an -# unbounded `quantization_config` into it can exceed server header size limits, so we only -# attach the serialized config for telemetry when it stays under this many characters. -_MAX_QUANT_CONFIG_USER_AGENT_CHARS = 2048 - -TORCH_INIT_FUNCTIONS = { - "uniform_": nn.init.uniform_, - "normal_": nn.init.normal_, - "trunc_normal_": nn.init.trunc_normal_, - "constant_": nn.init.constant_, - "xavier_uniform_": nn.init.xavier_uniform_, - "xavier_normal_": nn.init.xavier_normal_, - "kaiming_uniform_": nn.init.kaiming_uniform_, - "kaiming_normal_": nn.init.kaiming_normal_, - "uniform": nn.init.uniform, - "normal": nn.init.normal, - "xavier_uniform": nn.init.xavier_uniform, - "xavier_normal": nn.init.xavier_normal, - "kaiming_uniform": nn.init.kaiming_uniform, - "kaiming_normal": nn.init.kaiming_normal, -} - -if is_torch_version(">=", "1.9.0"): - _LOW_CPU_MEM_USAGE_DEFAULT = True -else: - _LOW_CPU_MEM_USAGE_DEFAULT = False - - -if is_accelerate_available(): - import accelerate - from accelerate import dispatch_model - from accelerate.utils import load_offloaded_weights, save_offload_index - - -def get_parameter_device(parameter: torch.nn.Module) -> torch.device: - from ..hooks.group_offloading import _get_group_onload_device - - try: - # Try to get the onload device from the group offloading hook - return _get_group_onload_device(parameter) - except ValueError: - pass - - try: - # If the onload device is not available due to no group offloading hooks, try to get the device - # from the first parameter or buffer - parameters_and_buffers = itertools.chain(parameter.parameters(), parameter.buffers()) - return next(parameters_and_buffers).device - except StopIteration: - # For torch.nn.DataParallel compatibility in PyTorch 1.5 - - def find_tensor_attributes(module: torch.nn.Module) -> list[tuple[str, Tensor]]: - tuples = [(k, v) for k, v in module.__dict__.items() if torch.is_tensor(v)] - return tuples - - gen = parameter._named_members(get_members_fn=find_tensor_attributes) - first_tuple = next(gen) - return first_tuple[1].device - - -def get_parameter_dtype(parameter: torch.nn.Module) -> torch.dtype: - """ - Returns the first found floating dtype in parameters if there is one, otherwise returns the last dtype it found. - """ - # 1. Check if we have attached any dtype modifying hooks (eg. layerwise casting) - if isinstance(parameter, nn.Module): - for name, submodule in parameter.named_modules(): - if not hasattr(submodule, "_diffusers_hook"): - continue - registry = submodule._diffusers_hook - hook = registry.get_hook("layerwise_casting") - if hook is not None: - return hook.compute_dtype - - # 2. If no dtype modifying hooks are attached, return the dtype of the first floating point parameter/buffer - last_dtype = None - - for name, param in parameter.named_parameters(): - last_dtype = param.dtype - if ( - hasattr(parameter, "_keep_in_fp32_modules") - and parameter._keep_in_fp32_modules - and any(m in name for m in parameter._keep_in_fp32_modules) - ): - continue - - if param.is_floating_point(): - return param.dtype - - for buffer in parameter.buffers(): - last_dtype = buffer.dtype - if buffer.is_floating_point(): - return buffer.dtype - - if last_dtype is not None: - # if no floating dtype was found return whatever the first dtype is - return last_dtype - - # For nn.DataParallel compatibility in PyTorch > 1.5 - def find_tensor_attributes(module: nn.Module) -> list[tuple[str, Tensor]]: - tuples = [(k, v) for k, v in module.__dict__.items() if torch.is_tensor(v)] - return tuples - - gen = parameter._named_members(get_members_fn=find_tensor_attributes) - last_tuple = None - for tuple in gen: - last_tuple = tuple - if tuple[1].is_floating_point(): - return tuple[1].dtype - - if last_tuple is not None: - # fallback to the last dtype - return last_tuple[1].dtype - - -@contextmanager -def no_init_weights(): - """ - Context manager to globally disable weight initialization to speed up loading large models. To do that, all the - torch.nn.init function are all replaced with skip. - """ - - def _skip_init(*args, **kwargs): - pass - - for name, init_func in TORCH_INIT_FUNCTIONS.items(): - setattr(torch.nn.init, name, _skip_init) - try: - yield - finally: - # Restore the original initialization functions - for name, init_func in TORCH_INIT_FUNCTIONS.items(): - setattr(torch.nn.init, name, init_func) - - -class ModelMixin(torch.nn.Module, PushToHubMixin): - r""" - Base class for all models. - - [`ModelMixin`] takes care of storing the model configuration and provides methods for loading, downloading and - saving models. - - - **config_name** ([`str`]) -- Filename to save a model to when calling [`~models.ModelMixin.save_pretrained`]. - """ - - config_name = CONFIG_NAME - _automatically_saved_args = ["_diffusers_version", "_class_name", "_name_or_path"] - _supports_gradient_checkpointing = False - _keys_to_ignore_on_load_unexpected = None - _no_split_modules = None - _keep_in_fp32_modules = None - _skip_layerwise_casting_patterns = None - _supports_group_offloading = True - _repeated_blocks = [] - _parallel_config = None - _cp_plan = None - _skip_keys = None - - def __init__(self): - super().__init__() - - self._gradient_checkpointing_func = None - - def __getattr__(self, name: str) -> Any: - """The only reason we overwrite `getattr` here is to gracefully deprecate accessing - config attributes directly. See https://github.com/huggingface/diffusers/pull/3129 We need to overwrite - __getattr__ here in addition so that we don't trigger `torch.nn.Module`'s __getattr__': - https://pytorch.org/docs/stable/_modules/torch/nn/modules/module.html#Module - """ - - is_in_config = "_internal_dict" in self.__dict__ and hasattr(self.__dict__["_internal_dict"], name) - is_attribute = name in self.__dict__ - - if is_in_config and not is_attribute: - 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. 'unet.config.{name}'." - deprecate("direct config name access", "1.0.0", deprecation_message, standard_warn=False, stacklevel=3) - return self._internal_dict[name] - - # call PyTorch's https://pytorch.org/docs/stable/_modules/torch/nn/modules/module.html#Module - return super().__getattr__(name) - - @property - def is_gradient_checkpointing(self) -> bool: - """ - Whether gradient checkpointing is activated for this model or not. - """ - return any(hasattr(m, "gradient_checkpointing") and m.gradient_checkpointing for m in self.modules()) - - def enable_gradient_checkpointing(self, gradient_checkpointing_func: Callable | None = None) -> None: - """ - Activates gradient checkpointing for the current model (may be referred to as *activation checkpointing* or - *checkpoint activations* in other frameworks). - - Args: - gradient_checkpointing_func (`Callable`, *optional*): - The function to use for gradient checkpointing. If `None`, the default PyTorch checkpointing function - is used (`torch.utils.checkpoint.checkpoint`). - """ - if not self._supports_gradient_checkpointing: - raise ValueError( - f"{self.__class__.__name__} does not support gradient checkpointing. Please make sure to set the boolean attribute " - f"`_supports_gradient_checkpointing` to `True` in the class definition." - ) - - if gradient_checkpointing_func is None: - - def _gradient_checkpointing_func(module, *args): - ckpt_kwargs = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} - return torch.utils.checkpoint.checkpoint( - module.__call__, - *args, - **ckpt_kwargs, - ) - - gradient_checkpointing_func = _gradient_checkpointing_func - - self._set_gradient_checkpointing(enable=True, gradient_checkpointing_func=gradient_checkpointing_func) - - def disable_gradient_checkpointing(self) -> None: - """ - Deactivates gradient checkpointing for the current model (may be referred to as *activation checkpointing* or - *checkpoint activations* in other frameworks). - """ - if self._supports_gradient_checkpointing: - self._set_gradient_checkpointing(enable=False) - - def set_use_npu_flash_attention(self, valid: bool) -> None: - r""" - Set the switch for the npu flash attention. - """ - - def fn_recursive_set_npu_flash_attention(module: torch.nn.Module): - if hasattr(module, "set_use_npu_flash_attention"): - module.set_use_npu_flash_attention(valid) - - for child in module.children(): - fn_recursive_set_npu_flash_attention(child) - - for module in self.children(): - if isinstance(module, torch.nn.Module): - fn_recursive_set_npu_flash_attention(module) - - def enable_npu_flash_attention(self) -> None: - r""" - Enable npu flash attention from torch_npu - - """ - self.set_use_npu_flash_attention(True) - - def disable_npu_flash_attention(self) -> None: - r""" - disable npu flash attention from torch_npu - - """ - self.set_use_npu_flash_attention(False) - - def set_use_xla_flash_attention( - self, use_xla_flash_attention: bool, partition_spec: Callable | None = None, **kwargs - ) -> None: - # Recursively walk through all the children. - # Any children which exposes the set_use_xla_flash_attention method - # gets the message - def fn_recursive_set_flash_attention(module: torch.nn.Module): - if hasattr(module, "set_use_xla_flash_attention"): - module.set_use_xla_flash_attention(use_xla_flash_attention, partition_spec, **kwargs) - - for child in module.children(): - fn_recursive_set_flash_attention(child) - - for module in self.children(): - if isinstance(module, torch.nn.Module): - fn_recursive_set_flash_attention(module) - - def enable_xla_flash_attention(self, partition_spec: Callable | None = None, **kwargs): - r""" - Enable the flash attention pallals kernel for torch_xla. - """ - self.set_use_xla_flash_attention(True, partition_spec, **kwargs) - - def disable_xla_flash_attention(self): - r""" - Disable the flash attention pallals kernel for torch_xla. - """ - self.set_use_xla_flash_attention(False) - - def set_use_memory_efficient_attention_xformers(self, valid: bool, attention_op: Callable | None = None) -> None: - # Recursively walk through all the children. - # Any children which exposes the set_use_memory_efficient_attention_xformers method - # gets the message - def fn_recursive_set_mem_eff(module: torch.nn.Module): - if hasattr(module, "set_use_memory_efficient_attention_xformers"): - module.set_use_memory_efficient_attention_xformers(valid, attention_op) - - for child in module.children(): - fn_recursive_set_mem_eff(child) - - for module in self.children(): - if isinstance(module, torch.nn.Module): - fn_recursive_set_mem_eff(module) - - def enable_xformers_memory_efficient_attention(self, attention_op: Callable | None = None) -> None: - r""" - Enable memory efficient attention from [xFormers](https://facebookresearch.github.io/xformers/). - - When this option is enabled, you should observe lower GPU memory usage and a potential speed up during - inference. Speed up during training is not guaranteed. - - > [!WARNING] > ⚠️ When memory efficient attention and sliced attention are both enabled, memory efficient - attention takes > precedent. - - Parameters: - attention_op (`Callable`, *optional*): - Override the default `None` operator for use as `op` argument to the - [`memory_efficient_attention()`](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.memory_efficient_attention) - function of xFormers. - - Examples: - - ```py - >>> import torch - >>> from diffusers import UNet2DConditionModel - >>> from xformers.ops import MemoryEfficientAttentionFlashAttentionOp - - >>> model = UNet2DConditionModel.from_pretrained( - ... "stabilityai/stable-diffusion-2-1", subfolder="unet", torch_dtype=torch.float16 - ... ) - >>> model = model.to("cuda") - >>> model.enable_xformers_memory_efficient_attention(attention_op=MemoryEfficientAttentionFlashAttentionOp) - ``` - """ - self.set_use_memory_efficient_attention_xformers(True, attention_op) - - def disable_xformers_memory_efficient_attention(self) -> None: - r""" - Disable memory efficient attention from [xFormers](https://facebookresearch.github.io/xformers/). - """ - self.set_use_memory_efficient_attention_xformers(False) - - def enable_layerwise_casting( - self, - storage_dtype: torch.dtype = torch.float8_e4m3fn, - compute_dtype: torch.dtype | None = None, - skip_modules_pattern: tuple[str, ...] | None = None, - skip_modules_classes: tuple[Type[torch.nn.Module], ...] | None = None, - non_blocking: bool = False, - ) -> None: - r""" - Activates layerwise casting for the current model. - - Layerwise casting is a technique that casts the model weights to a lower precision dtype for storage but - upcasts them on-the-fly to a higher precision dtype for computation. This process can significantly reduce the - memory footprint from model weights, but may lead to some quality degradation in the outputs. Most degradations - are negligible, mostly stemming from weight casting in normalization and modulation layers. - - By default, most models in diffusers set the `_skip_layerwise_casting_patterns` attribute to ignore patch - embedding, positional embedding and normalization layers. This is because these layers are most likely - precision-critical for quality. If you wish to change this behavior, you can set the - `_skip_layerwise_casting_patterns` attribute to `None`, or call - [`~hooks.layerwise_casting.apply_layerwise_casting`] with custom arguments. - - Example: - Using [`~models.ModelMixin.enable_layerwise_casting`]: - - ```python - >>> from diffusers import CogVideoXTransformer3DModel - - >>> transformer = CogVideoXTransformer3DModel.from_pretrained( - ... "THUDM/CogVideoX-5b", subfolder="transformer", torch_dtype=torch.bfloat16 - ... ) - - >>> # Enable layerwise casting via the model, which ignores certain modules by default - >>> transformer.enable_layerwise_casting(storage_dtype=torch.float8_e4m3fn, compute_dtype=torch.bfloat16) - ``` - - Args: - storage_dtype (`torch.dtype`): - The dtype to which the model should be cast for storage. - compute_dtype (`torch.dtype`): - The dtype to which the model weights should be cast during the forward pass. - skip_modules_pattern (`tuple[str, ...]`, *optional*): - A list of patterns to match the names of the modules to skip during the layerwise casting process. If - set to `None`, default skip patterns are used to ignore certain internal layers of modules and PEFT - layers. - skip_modules_classes (`tuple[Type[torch.nn.Module], ...]`, *optional*): - A list of module classes to skip during the layerwise casting process. - non_blocking (`bool`, *optional*, defaults to `False`): - If `True`, the weight casting operations are non-blocking. - """ - from ..hooks import apply_layerwise_casting - - user_provided_patterns = True - if skip_modules_pattern is None: - from ..hooks.layerwise_casting import DEFAULT_SKIP_MODULES_PATTERN - - skip_modules_pattern = DEFAULT_SKIP_MODULES_PATTERN - user_provided_patterns = False - if self._keep_in_fp32_modules is not None: - skip_modules_pattern += tuple(self._keep_in_fp32_modules) - if self._skip_layerwise_casting_patterns is not None: - skip_modules_pattern += tuple(self._skip_layerwise_casting_patterns) - skip_modules_pattern = tuple(set(skip_modules_pattern)) - - if is_peft_available() and not user_provided_patterns: - # By default, we want to skip all peft layers because they have a very low memory footprint. - # If users want to apply layerwise casting on peft layers as well, they can utilize the - # `~diffusers.hooks.layerwise_casting.apply_layerwise_casting` function which provides - # them with more flexibility and control. - - from peft.tuners.loha.layer import LoHaLayer - from peft.tuners.lokr.layer import LoKrLayer - from peft.tuners.lora.layer import LoraLayer - - for layer in (LoHaLayer, LoKrLayer, LoraLayer): - skip_modules_pattern += tuple(layer.adapter_layer_names) - - if compute_dtype is None: - logger.info("`compute_dtype` not provided when enabling layerwise casting. Using dtype of the model.") - compute_dtype = self.dtype - - apply_layerwise_casting( - self, storage_dtype, compute_dtype, skip_modules_pattern, skip_modules_classes, non_blocking - ) - - def enable_group_offload( - self, - onload_device: torch.device, - offload_device: torch.device = torch.device("cpu"), - offload_type: str = "block_level", - num_blocks_per_group: int | None = None, - non_blocking: bool = False, - use_stream: bool = False, - record_stream: bool = False, - low_cpu_mem_usage=False, - offload_to_disk_path: str | None = None, - block_modules: str | None = None, - exclude_kwargs: str | None = None, - ) -> None: - r""" - Activates group offloading for the current model. - - See [`~hooks.group_offloading.apply_group_offloading`] for more information. - - Example: - - ```python - >>> from diffusers import CogVideoXTransformer3DModel - - >>> transformer = CogVideoXTransformer3DModel.from_pretrained( - ... "THUDM/CogVideoX-5b", subfolder="transformer", torch_dtype=torch.bfloat16 - ... ) - - >>> transformer.enable_group_offload( - ... onload_device=torch.device("cuda"), - ... offload_device=torch.device("cpu"), - ... offload_type="leaf_level", - ... use_stream=True, - ... ) - ``` - """ - from ..hooks import apply_group_offloading - - if getattr(self, "enable_tiling", None) is not None and getattr(self, "use_tiling", False) and use_stream: - msg = ( - "Applying group offloading on autoencoders, with CUDA streams, may not work as expected if the first " - "forward pass is executed with tiling enabled. Please make sure to either:\n" - "1. Run a forward pass with small input shapes.\n" - "2. Or, run a forward pass with tiling disabled (can still use small dummy inputs)." - ) - logger.warning(msg) - if not self._supports_group_offloading: - raise ValueError( - f"{self.__class__.__name__} does not support group offloading. Please make sure to set the boolean attribute " - f"`_supports_group_offloading` to `True` in the class definition. If you believe this is a mistake, please " - f"open an issue at https://github.com/huggingface/diffusers/issues." - ) - - apply_group_offloading( - module=self, - onload_device=onload_device, - offload_device=offload_device, - offload_type=offload_type, - num_blocks_per_group=num_blocks_per_group, - non_blocking=non_blocking, - use_stream=use_stream, - record_stream=record_stream, - low_cpu_mem_usage=low_cpu_mem_usage, - offload_to_disk_path=offload_to_disk_path, - block_modules=block_modules, - exclude_kwargs=exclude_kwargs, - ) - - def set_attention_backend(self, backend: str) -> None: - """ - Set the attention backend for the model. - - Args: - backend (`str`): - The name of the backend to set. Must be one of the available backends defined in - `AttentionBackendName`. Available backends can be found in - `diffusers.attention_dispatch.AttentionBackendName`. Defaults to torch native scaled dot product - attention as backend. - """ - from .attention import AttentionModuleMixin - from .attention_dispatch import ( - AttentionBackendName, - _AttentionBackendRegistry, - _check_attention_backend_requirements, - _maybe_download_kernel_for_backend, - ) - - # TODO: the following will not be required when everything is refactored to AttentionModuleMixin - from .attention_processor import Attention, MochiAttention - - logger.warning("Attention backends are an experimental feature and the API may be subject to change.") - attention_classes = (Attention, MochiAttention, AttentionModuleMixin) - - parallel_config_set = False - for module in self.modules(): - if not isinstance(module, attention_classes): - continue - processor = module.processor - if getattr(processor, "_parallel_config", None) is not None: - parallel_config_set = True - break - - backend = backend.lower() - available_backends = {x.value for x in AttentionBackendName.__members__.values()} - if backend not in available_backends: - raise ValueError(f"`{backend=}` must be one of the following: " + ", ".join(available_backends)) - - backend = AttentionBackendName(backend) - if parallel_config_set and not _AttentionBackendRegistry._is_context_parallel_available(backend): - compatible_backends = sorted(_AttentionBackendRegistry._supports_context_parallel) - raise ValueError( - f"Context parallelism is enabled but current attention backend '{backend.value}' " - f"does not support context parallelism. " - f"Please set a compatible attention backend: {compatible_backends} using `model.set_attention_backend()`." - ) - - _check_attention_backend_requirements(backend) - _maybe_download_kernel_for_backend(backend) - - for module in self.modules(): - if not isinstance(module, attention_classes): - continue - processor = module.processor - if processor is None or not hasattr(processor, "_attention_backend"): - continue - processor._attention_backend = backend - - # Important to set the active backend so that it propagates gracefully throughout. - _AttentionBackendRegistry.set_active_backend(backend) - - def reset_attention_backend(self) -> None: - """ - Resets the attention backend for the model. Following calls to `forward` will use the environment default, if - set, or the torch native scaled dot product attention. - """ - from .attention import AttentionModuleMixin - from .attention_processor import Attention, MochiAttention - - logger.warning("Attention backends are an experimental feature and the API may be subject to change.") - - attention_classes = (Attention, MochiAttention, AttentionModuleMixin) - for module in self.modules(): - if not isinstance(module, attention_classes): - continue - processor = module.processor - if processor is None or not hasattr(processor, "_attention_backend"): - continue - processor._attention_backend = None - - def save_pretrained( - self, - save_directory: str | os.PathLike, - is_main_process: bool = True, - save_function: Callable | None = None, - safe_serialization: bool = True, - variant: str | None = None, - max_shard_size: int | str = "10GB", - push_to_hub: bool = False, - use_flashpack: bool = False, - **kwargs, - ): - """ - Save a model and its configuration file to a directory so that it can be reloaded using the - [`~models.ModelMixin.from_pretrained`] class method. - - Arguments: - save_directory (`str` or `os.PathLike`): - Directory to save a model and its configuration file to. Will be created if it doesn't exist. - is_main_process (`bool`, *optional*, defaults to `True`): - Whether the process calling this is the main process or not. Useful during distributed training and you - need to call this function on all processes. In this case, set `is_main_process=True` only on the main - process to avoid race conditions. - save_function (`Callable`): - The function to use to save the state dictionary. Useful during distributed training when you need to - replace `torch.save` with another method. Can be configured with the environment variable - `DIFFUSERS_SAVE_MODE`. - safe_serialization (`bool`, *optional*, defaults to `True`): - Whether to save the model using `safetensors` or the traditional PyTorch way with `pickle`. - variant (`str`, *optional*): - If specified, weights are saved in the format `pytorch_model..bin`. - max_shard_size (`int` or `str`, defaults to `"10GB"`): - The maximum size for a checkpoint before being sharded. Checkpoints shard will then be each of size - lower than this size. If expressed as a string, needs to be digits followed by a unit (like `"5GB"`). - If expressed as an integer, the unit is bytes. Note that this limit will be decreased after a certain - period of time (starting from Oct 2024) to allow users to upgrade to the latest version of `diffusers`. - This is to establish a common default size for this argument across different libraries in the Hugging - Face ecosystem (`transformers`, and `accelerate`, for example). - push_to_hub (`bool`, *optional*, defaults to `False`): - Whether or not to push your model to the Hugging Face Hub after saving it. You can specify the - repository you want to push to with `repo_id` (will default to the name of `save_directory` in your - namespace). - kwargs (`dict[str, Any]`, *optional*): - Additional keyword arguments passed along to the [`~utils.PushToHubMixin.push_to_hub`] method. - """ - if os.path.isfile(save_directory): - logger.error(f"Provided path ({save_directory}) should be a directory, not a file") - return - - hf_quantizer = getattr(self, "hf_quantizer", None) - if hf_quantizer is not None: - quantization_serializable = ( - hf_quantizer is not None - and isinstance(hf_quantizer, DiffusersQuantizer) - and hf_quantizer.is_serializable - ) - if safe_serialization and quantization_serializable: - quantization_serializable = ( - quantization_serializable and hf_quantizer.supports_safetensors_serialization - ) - if not quantization_serializable: - raise ValueError( - f"The model is quantized with {hf_quantizer.quantization_config.quant_method} and is not serializable - check out the warnings from" - " the logger on the traceback to understand the reason why the quantized model is not serializable." - ) - - weights_name = WEIGHTS_NAME - if use_flashpack: - weights_name = FLASHPACK_WEIGHTS_NAME - elif safe_serialization: - weights_name = SAFETENSORS_WEIGHTS_NAME - - weights_name = _add_variant(weights_name, variant) - weights_name_pattern = weights_name.replace(".bin", "{suffix}.bin").replace( - ".safetensors", "{suffix}.safetensors" - ) - - os.makedirs(save_directory, exist_ok=True) - - if push_to_hub: - commit_message = kwargs.pop("commit_message", None) - private = kwargs.pop("private", None) - create_pr = kwargs.pop("create_pr", False) - token = kwargs.pop("token", None) - repo_id = kwargs.pop("repo_id", save_directory.split(os.path.sep)[-1]) - repo_id = create_repo(repo_id, exist_ok=True, private=private, token=token).repo_id - - # Only save the model itself if we are using distributed training - model_to_save = self - - # Attach architecture to the config - # Save the config - if is_main_process: - model_to_save.save_config(save_directory) - - # Save the model - state_dict = model_to_save.state_dict() - quantization_metadata = {} - if hf_quantizer is not None: - state_dict, quantization_metadata = hf_quantizer.get_state_dict_and_metadata( - state_dict, safe_serialization=safe_serialization - ) - - if use_flashpack: - if is_flashpack_available(): - import flashpack - else: - logger.error( - "Saving a FlashPack checkpoint in PyTorch, requires both PyTorch and flashpack to be installed. Please see " - "https://pytorch.org/ and https://github.com/fal-ai/flashpack for installation instructions." - ) - raise ImportError("Please install torch and flashpack to save a FlashPack checkpoint in PyTorch.") - - flashpack.serialization.pack_to_file( - state_dict_or_model=state_dict, - destination_path=os.path.join(save_directory, weights_name), - target_dtype=self.dtype, - ) - else: - # Save the model - state_dict_split = split_torch_state_dict_into_shards( - state_dict, max_shard_size=max_shard_size, filename_pattern=weights_name_pattern - ) - - # Clean the folder from a previous save - if is_main_process: - for filename in os.listdir(save_directory): - if filename in state_dict_split.filename_to_tensors.keys(): - continue - full_filename = os.path.join(save_directory, filename) - if not os.path.isfile(full_filename): - continue - weights_without_ext = weights_name_pattern.replace(".bin", "").replace(".safetensors", "") - weights_without_ext = weights_without_ext.replace("{suffix}", "") - filename_without_ext = filename.replace(".bin", "").replace(".safetensors", "") - # make sure that file to be deleted matches format of sharded file, e.g. pytorch_model-00001-of-00005 - if ( - filename.startswith(weights_without_ext) - and _REGEX_SHARD.fullmatch(filename_without_ext) is not None - ): - os.remove(full_filename) - - for filename, tensors in state_dict_split.filename_to_tensors.items(): - shard = {tensor: state_dict[tensor].contiguous() for tensor in tensors} - filepath = os.path.join(save_directory, filename) - if safe_serialization: - metadata = {"format": "pt"} - if quantization_metadata: - metadata.update(quantization_metadata) - metadata = {k: str(v) if not isinstance(v, str) else v for k, v in metadata.items()} - # At some point we will need to deal better with save_function (used for TPU and other distributed - # joyfulness), but for now this enough. - safetensors.torch.save_file(shard, filepath, metadata=metadata) - else: - torch.save(shard, filepath) - - if state_dict_split.is_sharded: - metadata = dict(state_dict_split.metadata) - if quantization_metadata: - metadata.update(quantization_metadata) - index = { - "metadata": metadata, - "weight_map": state_dict_split.tensor_to_filename, - } - save_index_file = SAFE_WEIGHTS_INDEX_NAME if safe_serialization else WEIGHTS_INDEX_NAME - save_index_file = os.path.join(save_directory, _add_variant(save_index_file, variant)) - # Save the index as well - with open(save_index_file, "w", encoding="utf-8") as f: - content = json.dumps(index, indent=2, sort_keys=True) + "\n" - f.write(content) - logger.info( - f"The model is bigger than the maximum size per checkpoint ({max_shard_size}) and is going to be " - f"split in {len(state_dict_split.filename_to_tensors)} checkpoint shards. You can find where each parameters has been saved in the " - f"index located at {save_index_file}." - ) - else: - path_to_weights = os.path.join(save_directory, weights_name) - logger.info(f"Model weights saved in {path_to_weights}") - - if push_to_hub: - # Create a new empty model card and eventually tag it - model_card = load_or_create_model_card(repo_id, token=token) - model_card = populate_model_card(model_card) - model_card.save(Path(save_directory, "README.md").as_posix()) - - self._upload_folder( - save_directory, - repo_id, - token=token, - commit_message=commit_message, - create_pr=create_pr, - ) - - def dequantize(self): - """ - Potentially dequantize the model in case it has been quantized by a quantization method that support - dequantization. - """ - hf_quantizer = getattr(self, "hf_quantizer", None) - - if hf_quantizer is None: - raise ValueError("You need to first quantize your model in order to dequantize it") - - return hf_quantizer.dequantize(self) - - @classmethod - @validate_hf_hub_args - def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None, **kwargs) -> Self: - r""" - Instantiate a pretrained PyTorch model from a pretrained model configuration. - - The model is set in evaluation mode - `model.eval()` - by default, and dropout modules are deactivated. To - train the model, set it back in training mode with `model.train()`. - - Parameters: - pretrained_model_name_or_path (`str` or `os.PathLike`, *optional*): - Can be either: - - - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on - the Hub. - - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved - with [`~ModelMixin.save_pretrained`]. - - cache_dir (`str | os.PathLike`, *optional*): - Path to a directory where a downloaded pretrained model configuration is cached if the standard cache - is not used. - dtype (`torch.dtype`, *optional*): - Override the default `torch.dtype` and load the model with another dtype. - force_download (`bool`, *optional*, defaults to `False`): - Whether or not to force the (re-)download of the model weights and configuration files, overriding the - cached versions if they exist. - proxies (`dict[str, str]`, *optional*): - A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', - 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. - output_loading_info (`bool`, *optional*, defaults to `False`): - Whether or not to also return a dictionary containing missing keys, unexpected keys and error messages. - local_files_only(`bool`, *optional*, defaults to `False`): - Whether to only load local model weights and configuration files or not. If set to `True`, the model - won't be downloaded from the Hub. - token (`str` or *bool*, *optional*): - The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from - `diffusers-cli login` (stored in `~/.huggingface`) is used. - revision (`str`, *optional*, defaults to `"main"`): - The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier - allowed by Git. - subfolder (`str`, *optional*, defaults to `""`): - The subfolder location of a model file within a larger model repository on the Hub or locally. - mirror (`str`, *optional*): - Mirror source to resolve accessibility issues if you're downloading a model in China. We do not - guarantee the timeliness or safety of the source, and you should refer to the mirror site for more - information. - device_map (`int | str | torch.device` or `dict[str, int | str | torch.device]`, *optional*): - A map that specifies where each submodule should go. It doesn't need to be defined for each - parameter/buffer name; once a given module name is inside, every submodule of it will be sent to the - same device. Defaults to `None`, meaning that the model will be loaded on CPU. - - Examples: - - ```py - >>> from diffusers import AutoModel - >>> import torch - - >>> # This works. - >>> model = AutoModel.from_pretrained( - ... "stabilityai/stable-diffusion-xl-base-1.0", subfolder="unet", device_map="cuda" - ... ) - >>> # This also works (integer accelerator device ID). - >>> model = AutoModel.from_pretrained( - ... "stabilityai/stable-diffusion-xl-base-1.0", subfolder="unet", device_map=0 - ... ) - >>> # Specifying a supported offloading strategy like "auto" also works. - >>> model = AutoModel.from_pretrained( - ... "stabilityai/stable-diffusion-xl-base-1.0", subfolder="unet", device_map="auto" - ... ) - >>> # Specifying a dictionary as `device_map` also works. - >>> model = AutoModel.from_pretrained( - ... "stabilityai/stable-diffusion-xl-base-1.0", - ... subfolder="unet", - ... device_map={"": torch.device("cuda")}, - ... ) - ``` - - Set `device_map="auto"` to have 🤗 Accelerate automatically compute the most optimized `device_map`. For - more information about each option see [designing a device - map](https://huggingface.co/docs/accelerate/en/concept_guides/big_model_inference#the-devicemap). You - can also refer to the [Diffusers-specific - documentation](https://huggingface.co/docs/diffusers/main/en/training/distributed_inference#model-sharding) - for more concrete examples. - max_memory (`Dict`, *optional*): - A dictionary device identifier for the maximum memory. Will default to the maximum memory available for - each GPU and the available CPU RAM if unset. - offload_folder (`str` or `os.PathLike`, *optional*): - The path to offload weights if `device_map` contains the value `"disk"`. - offload_state_dict (`bool`, *optional*): - If `True`, temporarily offloads the CPU state dict to the hard drive to avoid running out of CPU RAM if - the weight of the CPU state dict + the biggest shard of the checkpoint does not fit. Defaults to `True` - when there is some disk offload. - low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 else `False`): - Speed up model loading only loading the pretrained weights and not initializing the weights. This also - tries to not use more than 1x model size in CPU memory (including peak memory) while loading the model. - Only supported for PyTorch >= 1.9.0. If you are using an older version of PyTorch, setting this - argument to `True` will raise an error. - variant (`str`, *optional*): - Load weights from a specified `variant` filename such as `"fp16"` or `"ema"`. - use_safetensors (`bool`, *optional*, defaults to `None`): - If set to `None`, the `safetensors` weights are downloaded if they're available **and** if the - `safetensors` library is installed. If set to `True`, the model is forcibly loaded from `safetensors` - weights. If set to `False`, `safetensors` weights are not loaded. - disable_mmap ('bool', *optional*, defaults to 'False'): - Whether to disable mmap when loading a Safetensors model. This option can perform better when the model - is on a network mount or hard drive, which may not handle the seeky-ness of mmap very well. - use_flashpack (`bool`, *optional*, defaults to `False`): - If set to `True`, the model is loaded from `flashpack` weights. - flashpack_kwargs(`dict[str, Any]`, *optional*, defaults to `{}`): - Kwargs passed to - [`flashpack.deserialization.assign_from_file`](https://github.com/fal-ai/flashpack/blob/f1aa91c5cd9532a3dbf5bcc707ab9b01c274b76c/src/flashpack/deserialization.py#L408-L422) - - - > [!TIP] > To use private or [gated models](https://huggingface.co/docs/hub/models-gated#gated-models), log-in - with `hf > auth login`. You can also activate the special > - ["offline-mode"](https://huggingface.co/diffusers/installation.html#offline-mode) to use this method in a > - firewalled environment. - - Example: - - ```py - from diffusers import UNet2DConditionModel - - unet = UNet2DConditionModel.from_pretrained("stable-diffusion-v1-5/stable-diffusion-v1-5", subfolder="unet") - ``` - - If you get the error message below, you need to finetune the weights for your downstream task: - - ```bash - Some weights of UNet2DConditionModel were not initialized from the model checkpoint at stable-diffusion-v1-5/stable-diffusion-v1-5 and are newly initialized because the shapes did not match: - - conv_in.weight: found shape torch.Size([320, 4, 3, 3]) in the checkpoint and torch.Size([320, 9, 3, 3]) in the model instantiated - You should probably TRAIN this model on a down-stream task to be able to use it for predictions and inference. - ``` - """ - cache_dir = kwargs.pop("cache_dir", None) - ignore_mismatched_sizes = kwargs.pop("ignore_mismatched_sizes", False) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - output_loading_info = kwargs.pop("output_loading_info", False) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - torch_dtype = kwargs.pop("torch_dtype", None) - dtype = kwargs.pop("dtype", None) - torch_dtype = dtype if dtype is not None else torch_dtype - subfolder = kwargs.pop("subfolder", None) - device_map = kwargs.pop("device_map", None) - max_memory = kwargs.pop("max_memory", None) - offload_folder = kwargs.pop("offload_folder", None) - offload_state_dict = kwargs.pop("offload_state_dict", None) - low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT) - variant = kwargs.pop("variant", None) - use_safetensors = kwargs.pop("use_safetensors", None) - quantization_config = kwargs.pop("quantization_config", None) - dduf_entries: dict[str, DDUFEntry] | None = kwargs.pop("dduf_entries", None) - disable_mmap = kwargs.pop("disable_mmap", False) - parallel_config: ParallelConfig | ContextParallelConfig | None = kwargs.pop("parallel_config", None) - use_flashpack = kwargs.pop("use_flashpack", False) - flashpack_kwargs = kwargs.pop("flashpack_kwargs", {}) - - is_parallel_loading_enabled = HF_ENABLE_PARALLEL_LOADING - if is_parallel_loading_enabled and not low_cpu_mem_usage: - raise NotImplementedError("Parallel loading is not supported when not using `low_cpu_mem_usage`.") - - if torch_dtype is not None and not isinstance(torch_dtype, torch.dtype): - torch_dtype = torch.float32 - logger.warning( - f"Passed `torch_dtype` {torch_dtype} is not a `torch.dtype`. Defaulting to `torch.float32`." - ) - - allow_pickle = False - if use_safetensors is None: - use_safetensors = True - allow_pickle = True - - if low_cpu_mem_usage and not is_accelerate_available(): - low_cpu_mem_usage = False - logger.warning( - "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the" - " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install" - " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip" - " install accelerate\n```\n." - ) - - if device_map is not None and not is_accelerate_available(): - raise NotImplementedError( - "Loading and dispatching requires `accelerate`. Please make sure to install accelerate or set" - " `device_map=None`. You can install accelerate with `pip install accelerate`." - ) - - # Check if we can handle device_map and dispatching the weights - if device_map is not None and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Loading and dispatching requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `device_map=None`." - ) - - if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"): - raise NotImplementedError( - "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set" - " `low_cpu_mem_usage=False`." - ) - - if low_cpu_mem_usage is False and device_map is not None: - raise ValueError( - f"You cannot set `low_cpu_mem_usage` to `False` while using device_map={device_map} for loading and" - " dispatching. Please make sure to set `low_cpu_mem_usage=True`." - ) - - # change device_map into a map if we passed an int, a str or a torch.device - if isinstance(device_map, torch.device): - device_map = {"": device_map} - elif isinstance(device_map, str) and device_map not in ["auto", "balanced", "balanced_low_0", "sequential"]: - try: - device_map = {"": torch.device(device_map)} - except RuntimeError: - raise ValueError( - "When passing device_map as a string, the value needs to be a device name (e.g. cpu, cuda:0) or " - f"'auto', 'balanced', 'balanced_low_0', 'sequential' but found {device_map}." - ) - elif isinstance(device_map, int): - if device_map < 0: - raise ValueError( - "You can't pass device_map as a negative int. If you want to put the model on the cpu, pass device_map = 'cpu' " - ) - else: - device_map = {"": device_map} - - if device_map is not None: - if low_cpu_mem_usage is None: - low_cpu_mem_usage = True - elif not low_cpu_mem_usage: - raise ValueError("Passing along a `device_map` requires `low_cpu_mem_usage=True`") - - if low_cpu_mem_usage: - if device_map is not None and not is_torch_version(">=", "1.10"): - # The max memory utils require PyTorch >= 1.10 to have torch.cuda.mem_get_info. - raise ValueError("`low_cpu_mem_usage` and `device_map` require PyTorch >= 1.10.") - - user_agent = { - "diffusers": __version__, - "file_type": "model", - "framework": "pytorch", - "model_class": str(cls.__name__), - } - unused_kwargs = {} - - # Load config if we don't provide a configuration - config_path = pretrained_model_name_or_path - - # load config - config, unused_kwargs, commit_hash = cls.load_config( - config_path, - cache_dir=cache_dir, - return_unused_kwargs=True, - return_commit_hash=True, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - dduf_entries=dduf_entries, - **kwargs, - ) - # no in-place modification of the original config. - config = copy.deepcopy(config) - - # determine initial quantization config. - ####################################### - pre_quantized = "quantization_config" in config and config["quantization_config"] is not None - if pre_quantized or quantization_config is not None: - if pre_quantized: - config["quantization_config"] = DiffusersAutoQuantizer.merge_quantization_configs( - config["quantization_config"], quantization_config - ) - else: - config["quantization_config"] = quantization_config - hf_quantizer = DiffusersAutoQuantizer.from_config( - config["quantization_config"], pre_quantized=pre_quantized - ) - else: - hf_quantizer = None - - if hf_quantizer is not None: - hf_quantizer.validate_environment(torch_dtype=torch_dtype, device_map=device_map) - torch_dtype = hf_quantizer.update_torch_dtype(torch_dtype) - device_map = hf_quantizer.update_device_map(device_map) - - # In order to ensure popular quantization methods are supported. Can be disabled with `disable_telemetry` - user_agent["quant"] = hf_quantizer.quantization_config.quant_method.value - # Attach the full serialized config for telemetry, but skip it when it is large enough to - # risk exceeding HTTP header size limits (see `_MAX_QUANT_CONFIG_USER_AGENT_CHARS`). - serialized_quant_config = json.dumps(hf_quantizer.quantization_config.to_dict(), sort_keys=True) - if len(serialized_quant_config) <= _MAX_QUANT_CONFIG_USER_AGENT_CHARS: - user_agent["quant_config"] = serialized_quant_config - - # Force-set to `True` for more mem efficiency - if low_cpu_mem_usage is None: - low_cpu_mem_usage = True - logger.info("Set `low_cpu_mem_usage` to True as `hf_quantizer` is not None.") - elif not low_cpu_mem_usage: - raise ValueError("`low_cpu_mem_usage` cannot be False or None when using quantization.") - - # Check if `_keep_in_fp32_modules` is not None - use_keep_in_fp32_modules = cls._keep_in_fp32_modules is not None and ( - hf_quantizer is None or getattr(hf_quantizer, "use_keep_in_fp32_modules", False) - ) - - if use_keep_in_fp32_modules: - keep_in_fp32_modules = cls._keep_in_fp32_modules - if not isinstance(keep_in_fp32_modules, list): - keep_in_fp32_modules = [keep_in_fp32_modules] - - if low_cpu_mem_usage is None: - low_cpu_mem_usage = True - logger.info("Set `low_cpu_mem_usage` to True as `_keep_in_fp32_modules` is not None.") - elif not low_cpu_mem_usage: - raise ValueError("`low_cpu_mem_usage` cannot be False when `keep_in_fp32_modules` is True.") - else: - keep_in_fp32_modules = [] - - is_sharded = False - resolved_model_file = None - - # Determine if we're loading from a directory of sharded checkpoints. - sharded_metadata = None - index_file = None - is_local = os.path.isdir(pretrained_model_name_or_path) - index_file_kwargs = { - "is_local": is_local, - "pretrained_model_name_or_path": pretrained_model_name_or_path, - "subfolder": subfolder or "", - "use_safetensors": use_safetensors, - "cache_dir": cache_dir, - "variant": variant, - "force_download": force_download, - "proxies": proxies, - "local_files_only": local_files_only, - "token": token, - "revision": revision, - "user_agent": user_agent, - "commit_hash": commit_hash, - "dduf_entries": dduf_entries, - } - index_file = _fetch_index_file(**index_file_kwargs) - # In case the index file was not found we still have to consider the legacy format. - # this becomes applicable when the variant is not None. - if variant is not None and (index_file is None or not os.path.exists(index_file)): - index_file = _fetch_index_file_legacy(**index_file_kwargs) - if index_file is not None and (dduf_entries or index_file.is_file()): - is_sharded = True - - # load model - # in the case it is sharded, we have already the index - if is_sharded: - resolved_model_file, sharded_metadata = _get_checkpoint_shard_files( - pretrained_model_name_or_path, - index_file, - cache_dir=cache_dir, - proxies=proxies, - local_files_only=local_files_only, - token=token, - user_agent=user_agent, - revision=revision, - subfolder=subfolder or "", - dduf_entries=dduf_entries, - ) - else: - if use_flashpack: - weights_name = FLASHPACK_WEIGHTS_NAME - elif use_safetensors: - weights_name = _add_variant(SAFETENSORS_WEIGHTS_NAME, variant) - else: - weights_name = None - if weights_name is not None: - try: - resolved_model_file = _get_model_file( - pretrained_model_name_or_path, - weights_name=weights_name, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - commit_hash=commit_hash, - dduf_entries=dduf_entries, - ) - - except IOError as e: - logger.error(f"An error occurred while trying to fetch {pretrained_model_name_or_path}: {e}") - if not allow_pickle: - raise - logger.warning( - "Defaulting to unsafe serialization. Pass `allow_pickle=False` to raise an error instead." - ) - - if resolved_model_file is None and not is_sharded: - resolved_model_file = _get_model_file( - pretrained_model_name_or_path, - weights_name=_add_variant(WEIGHTS_NAME, variant), - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - commit_hash=commit_hash, - dduf_entries=dduf_entries, - ) - - if not isinstance(resolved_model_file, list): - resolved_model_file = [resolved_model_file] - - # set dtype to instantiate the model under: - # 1. If torch_dtype is not None, we use that dtype - # 2. If torch_dtype is float8, we don't use _set_default_torch_dtype and we downcast after loading the model - dtype_orig = None - if torch_dtype is not None and not torch_dtype == getattr(torch, "float8_e4m3fn", None): - if not isinstance(torch_dtype, torch.dtype): - raise ValueError( - f"{torch_dtype} needs to be of type `torch.dtype`, e.g. `torch.float16`, but is {type(torch_dtype)}." - ) - dtype_orig = cls._set_default_torch_dtype(torch_dtype) - - init_contexts = [no_init_weights()] - - if low_cpu_mem_usage: - init_contexts.append(accelerate.init_empty_weights()) - - with ContextManagers(init_contexts): - model = cls.from_config(config, **unused_kwargs) - - if use_flashpack: - if is_flashpack_available(): - import flashpack - else: - logger.error( - "Loading a FlashPack checkpoint in PyTorch, requires both PyTorch and flashpack to be installed. Please see " - "https://pytorch.org/ and https://github.com/fal-ai/flashpack for installation instructions." - ) - raise ImportError("Please install torch and flashpack to load a FlashPack checkpoint in PyTorch.") - - if device_map is None: - logger.warning( - "`device_map` has not been provided for FlashPack, model will be on `cpu` - provide `device_map` to fully utilize " - "the benefit of FlashPack." - ) - flashpack_device = torch.device("cpu") - else: - device = device_map[""] - if isinstance(device, str) and device in ["auto", "balanced", "balanced_low_0", "sequential"]: - raise ValueError( - "FlashPack `device_map` should not be one of `auto`, `balanced`, `balanced_low_0`, `sequential`. Use a specific device instead, e.g., `device_map='cuda'` or `device_map='cuda:0'" - ) - flashpack_device = torch.device(device) if not isinstance(device, torch.device) else device - - flashpack.mixin.assign_from_file( - model=model, - path=resolved_model_file[0], - device=flashpack_device, - **flashpack_kwargs, - ) - if dtype_orig is not None: - torch.set_default_dtype(dtype_orig) - if output_loading_info: - logger.warning("`output_loading_info` is not supported with FlashPack.") - return model, {} - - return model - - if dtype_orig is not None: - torch.set_default_dtype(dtype_orig) - - state_dict = None - if not is_sharded: - # Time to load the checkpoint - state_dict = load_state_dict(resolved_model_file[0], disable_mmap=disable_mmap, dduf_entries=dduf_entries) - # We only fix it for non sharded checkpoints as we don't need it yet for sharded one. - model._fix_state_dict_keys_on_load(state_dict) - - if is_sharded: - loaded_keys = sharded_metadata["all_checkpoint_keys"] - else: - loaded_keys = list(state_dict.keys()) - - checkpoint_files = resolved_model_file - if hf_quantizer is not None: - loaded_keys = hf_quantizer.maybe_update_loaded_keys(loaded_keys, checkpoint_files) - - if hf_quantizer is not None: - hf_quantizer.preprocess_model( - model=model, - device_map=device_map, - keep_in_fp32_modules=keep_in_fp32_modules, - ) - - if hf_quantizer is not None and not hf_quantizer.supports_parallel_loading: - is_parallel_loading_enabled = False - - # Now that the model is loaded, we can determine the device_map - device_map = _determine_device_map( - model, device_map, max_memory, torch_dtype, keep_in_fp32_modules, hf_quantizer - ) - if hf_quantizer is not None: - hf_quantizer.validate_environment(device_map=device_map) - - ( - model, - missing_keys, - unexpected_keys, - mismatched_keys, - offload_index, - error_msgs, - ) = cls._load_pretrained_model( - model, - state_dict, - resolved_model_file, - pretrained_model_name_or_path, - loaded_keys, - ignore_mismatched_sizes=ignore_mismatched_sizes, - low_cpu_mem_usage=low_cpu_mem_usage, - device_map=device_map, - offload_folder=offload_folder, - offload_state_dict=offload_state_dict, - dtype=torch_dtype, - hf_quantizer=hf_quantizer, - keep_in_fp32_modules=keep_in_fp32_modules, - dduf_entries=dduf_entries, - is_parallel_loading_enabled=is_parallel_loading_enabled, - disable_mmap=disable_mmap, - ) - loading_info = { - "missing_keys": missing_keys, - "unexpected_keys": unexpected_keys, - "mismatched_keys": mismatched_keys, - "error_msgs": error_msgs, - } - - # Dispatch model with hooks on all devices if necessary - if device_map is not None: - device_map_kwargs = { - "device_map": device_map, - "offload_dir": offload_folder, - "offload_index": offload_index, - } - dispatch_model(model, **device_map_kwargs) - - if hf_quantizer is not None: - hf_quantizer.postprocess_model(model) - model.hf_quantizer = hf_quantizer - - if ( - torch_dtype is not None - and torch_dtype == getattr(torch, "float8_e4m3fn", None) - and hf_quantizer is None - and not use_keep_in_fp32_modules - ): - model = model.to(torch_dtype) - - if hf_quantizer is not None: - # We also make sure to purge `_pre_quantization_dtype` when we serialize - # the model config because `_pre_quantization_dtype` is `torch.dtype`, not JSON serializable. - model.register_to_config(_name_or_path=pretrained_model_name_or_path, _pre_quantization_dtype=torch_dtype) - else: - model.register_to_config(_name_or_path=pretrained_model_name_or_path) - - # Set model in evaluation mode to deactivate DropOut modules by default - model.eval() - - if parallel_config is not None: - model.enable_parallelism(config=parallel_config) - - if output_loading_info: - return model, loading_info - - return model - - # Adapted from `transformers`. - @wraps(torch.nn.Module.cuda) - def cuda(self, *args, **kwargs): - from ..hooks.group_offloading import _is_group_offload_enabled - - # Checks if the model has been loaded in 4-bit or 8-bit with BNB - if getattr(self, "quantization_method", None) == QuantizationMethod.BITS_AND_BYTES: - if getattr(self, "is_loaded_in_8bit", False) and is_bitsandbytes_version("<", "0.48.0"): - raise ValueError( - "Calling `cuda()` is not supported for `8-bit` quantized models with the installed version of bitsandbytes. " - f"The current device is `{self.device}`. If you intended to move the model, please install bitsandbytes >= 0.48.0." - ) - elif getattr(self, "is_loaded_in_4bit", False) and is_bitsandbytes_version("<", "0.43.2"): - raise ValueError( - "Calling `cuda()` is not supported for `4-bit` quantized models with the installed version of bitsandbytes. " - f"The current device is `{self.device}`. If you intended to move the model, please install bitsandbytes >= 0.43.2." - ) - - # Checks if group offloading is enabled - if _is_group_offload_enabled(self): - logger.warning( - f"The module '{self.__class__.__name__}' is group offloaded and moving it using `.cuda()` is not supported." - ) - return self - - return super().cuda(*args, **kwargs) - - # Adapted from `transformers`. - @wraps(torch.nn.Module.to) - def to(self, *args, **kwargs): - from ..hooks.group_offloading import _is_group_offload_enabled - - fp32_modules = self._keep_in_fp32_modules or [] - - device_arg_or_kwarg_present = any(isinstance(arg, torch.device) for arg in args) or "device" in kwargs - dtype_present_in_args = "dtype" in kwargs - - # Try converting arguments to torch.device in case they are passed as strings - for arg in args: - if not isinstance(arg, str): - continue - try: - torch.device(arg) - device_arg_or_kwarg_present = True - except RuntimeError: - pass - - if not dtype_present_in_args: - for arg in args: - if isinstance(arg, torch.dtype): - dtype_present_in_args = True - break - - if dtype_present_in_args and fp32_modules is not None: - logger.warning( - f"There are modules in {self.__class__.__name__} that should be kept in float32: {fp32_modules}. Casting directly with `to()` can lead to inconsistent results; set `torch_dtype` in `from_pretrained()` instead to keep these modules in float32." - ) - - if getattr(self, "is_quantized", False): - if dtype_present_in_args: - raise ValueError( - "Casting a quantized model to a new `dtype` is unsupported. To set the dtype of unquantized layers, please " - "use the `torch_dtype` argument when loading the model using `from_pretrained` or `from_single_file`" - ) - - if getattr(self, "quantization_method", None) == QuantizationMethod.BITS_AND_BYTES: - if getattr(self, "is_loaded_in_8bit", False) and is_bitsandbytes_version("<", "0.48.0"): - raise ValueError( - "Calling `to()` is not supported for `8-bit` quantized models with the installed version of bitsandbytes. " - f"The current device is `{self.device}`. If you intended to move the model, please install bitsandbytes >= 0.48.0." - ) - elif getattr(self, "is_loaded_in_4bit", False) and is_bitsandbytes_version("<", "0.43.2"): - raise ValueError( - "Calling `to()` is not supported for `4-bit` quantized models with the installed version of bitsandbytes. " - f"The current device is `{self.device}`. If you intended to move the model, please install bitsandbytes >= 0.43.2." - ) - if _is_group_offload_enabled(self) and device_arg_or_kwarg_present: - logger.warning( - f"The module '{self.__class__.__name__}' is group offloaded and moving it using `.to()` is not supported." - ) - return self - - return super().to(*args, **kwargs) - - # Taken from `transformers`. - def half(self, *args): - # Checks if the model is quantized - if getattr(self, "is_quantized", False): - raise ValueError( - "`.half()` is not supported for quantized model. Please use the model as it is, since the" - " model has already been cast to the correct `dtype`." - ) - else: - return super().half(*args) - - # Taken from `transformers`. - def float(self, *args): - # Checks if the model is quantized - if getattr(self, "is_quantized", False): - raise ValueError( - "`.float()` is not supported for quantized model. Please use the model as it is, since the" - " model has already been cast to the correct `dtype`." - ) - else: - return super().float(*args) - - def compile_repeated_blocks(self, *args, **kwargs): - """ - Compiles *only* the frequently repeated sub-modules of a model (e.g. the Transformer layers) instead of - compiling the entire model. This technique—often called **regional compilation** (see the PyTorch recipe - https://docs.pytorch.org/tutorials/recipes/regional_compilation.html) can reduce end-to-end compile time - substantially, while preserving the runtime speed-ups you would expect from a full `torch.compile`. - - The set of sub-modules to compile is discovered by the presence of **`_repeated_blocks`** attribute in the - model definition. Define this attribute on your model subclass as a list/tuple of class names (strings). Every - module whose class name matches will be compiled. - - Once discovered, each matching sub-module is compiled by calling `submodule.compile(*args, **kwargs)`. Any - positional or keyword arguments you supply to `compile_repeated_blocks` are forwarded verbatim to - `torch.compile`. - """ - repeated_blocks = getattr(self, "_repeated_blocks", None) - - if not repeated_blocks: - raise ValueError( - "`_repeated_blocks` attribute is empty. " - f"Set `_repeated_blocks` for the class `{self.__class__.__name__}` to benefit from faster compilation. " - ) - has_compiled_region = False - for submod in self.modules(): - if submod.__class__.__name__ in repeated_blocks: - submod.compile(*args, **kwargs) - has_compiled_region = True - - if not has_compiled_region: - raise ValueError( - f"Regional compilation failed because {repeated_blocks} classes are not found in the model. " - ) - - def enable_parallelism( - self, - *, - config: ParallelConfig | ContextParallelConfig, - cp_plan: dict[str, ContextParallelModelPlan] | None = None, - ): - logger.warning( - "`enable_parallelism` is an experimental feature. The API may change in the future and breaking changes may be introduced at any time without warning." - ) - - if not torch.distributed.is_available() and not torch.distributed.is_initialized(): - raise RuntimeError( - "torch.distributed must be available and initialized before calling `enable_parallelism`." - ) - - from ..hooks.context_parallel import apply_context_parallel - from .attention import AttentionModuleMixin - from .attention_dispatch import AttentionBackendName, _AttentionBackendRegistry - from .attention_processor import Attention, MochiAttention - - if isinstance(config, ContextParallelConfig): - config = ParallelConfig(context_parallel_config=config) - - rank = torch.distributed.get_rank() - world_size = torch.distributed.get_world_size() - device_type = torch._C._get_accelerator().type - device_module = torch.get_device_module(device_type) - device = torch.device(device_type, rank % device_module.device_count()) - - attention_classes = (Attention, MochiAttention, AttentionModuleMixin) - - if config.context_parallel_config is not None: - for module in self.modules(): - if not isinstance(module, attention_classes): - continue - - processor = module.processor - if processor is None or not hasattr(processor, "_attention_backend"): - continue - - attention_backend = processor._attention_backend - if attention_backend is None: - attention_backend, _ = _AttentionBackendRegistry.get_active_backend() - else: - attention_backend = AttentionBackendName(attention_backend) - - if not _AttentionBackendRegistry._is_context_parallel_available(attention_backend): - compatible_backends = sorted(_AttentionBackendRegistry._supports_context_parallel) - raise ValueError( - f"Context parallelism is enabled but the attention processor '{processor.__class__.__name__}' " - f"is using backend '{attention_backend.value}' which does not support context parallelism. " - f"Please set a compatible attention backend: {compatible_backends} using `model.set_attention_backend()` before " - f"calling `model.enable_parallelism()`." - ) - - # All modules use the same attention processor and backend. We don't need to - # iterate over all modules after checking the first processor - break - - mesh = None - if config.context_parallel_config is not None: - cp_config = config.context_parallel_config - mesh = cp_config.mesh or torch.distributed.device_mesh.init_device_mesh( - device_type=device_type, - mesh_shape=cp_config.mesh_shape, - mesh_dim_names=cp_config.mesh_dim_names, - ) - - config.setup(rank, world_size, device, mesh=mesh) - self._parallel_config = config - - for module in self.modules(): - if not isinstance(module, attention_classes): - continue - processor = module.processor - if processor is None or not hasattr(processor, "_parallel_config"): - continue - processor._parallel_config = config - - if config.context_parallel_config is not None: - if cp_plan is None and self._cp_plan is None: - raise ValueError( - "`cp_plan` must be provided either as an argument or set in the model's `_cp_plan` attribute." - ) - cp_plan = cp_plan if cp_plan is not None else self._cp_plan - apply_context_parallel(self, config.context_parallel_config, cp_plan) - - @classmethod - def _load_pretrained_model( - cls, - model, - state_dict: OrderedDict, - resolved_model_file: list[str], - pretrained_model_name_or_path: str | os.PathLike, - loaded_keys: list[str], - ignore_mismatched_sizes: bool = False, - assign_to_params_buffers: bool = False, - hf_quantizer: DiffusersQuantizer | None = None, - low_cpu_mem_usage: bool = True, - dtype: str | torch.dtype | None = None, - keep_in_fp32_modules: list[str] | None = None, - device_map: str | int | torch.device | dict[str, str | int | torch.device] = None, - offload_state_dict: bool | None = None, - offload_folder: str | os.PathLike | None = None, - dduf_entries: dict[str, DDUFEntry] | None = None, - is_parallel_loading_enabled: bool | None = False, - disable_mmap: bool = False, - ): - model_state_dict = model.state_dict() - expected_keys = list(model_state_dict.keys()) - missing_keys = list(set(expected_keys) - set(loaded_keys)) - if hf_quantizer is not None: - missing_keys = hf_quantizer.update_missing_keys(model, missing_keys, prefix="") - unexpected_keys = list(set(loaded_keys) - set(expected_keys)) - # Some models may have keys that are not in the state by design, removing them before needlessly warning - # the user. - if cls._keys_to_ignore_on_load_unexpected is not None: - for pat in cls._keys_to_ignore_on_load_unexpected: - unexpected_keys = [k for k in unexpected_keys if re.search(pat, k) is None] - - mismatched_keys = [] - error_msgs = [] - - # Deal with offload - if device_map is not None and "disk" in device_map.values(): - if offload_folder is None: - raise ValueError( - "The current `device_map` had weights offloaded to the disk. Please provide an `offload_folder`" - " for them. Alternatively, make sure you have `safetensors` installed if the model you are using" - " offers the weights in this format." - ) - else: - os.makedirs(offload_folder, exist_ok=True) - if offload_state_dict is None: - offload_state_dict = True - - # If a device map has been used, we can speedup the load time by warming up the device caching allocator. - # If we don't warmup, each tensor allocation on device calls to the allocator for memory (effectively, a - # lot of individual calls to device malloc). We can, however, preallocate the memory required by the - # tensors using their expected shape and not performing any initialization of the memory (empty data). - # When the actual device allocations happen, the allocator already has a pool of unused device memory - # that it can re-use for faster loading of the model. - if device_map is not None: - expanded_device_map = _expand_device_map(device_map, expected_keys) - _caching_allocator_warmup(model, expanded_device_map, dtype, hf_quantizer) - - offload_index = {} if device_map is not None and "disk" in device_map.values() else None - state_dict_folder, state_dict_index = None, None - if offload_state_dict: - state_dict_folder = tempfile.mkdtemp() - state_dict_index = {} - - if state_dict is not None: - # load_state_dict will manage the case where we pass a dict instead of a file - # if state dict is not None, it means that we don't need to read the files from resolved_model_file also - resolved_model_file = [state_dict] - - # Prepare the loading function sharing the attributes shared between them. - load_fn = functools.partial( - _load_shard_files_with_threadpool if is_parallel_loading_enabled else _load_shard_file, - model=model, - model_state_dict=model_state_dict, - device_map=device_map, - dtype=dtype, - hf_quantizer=hf_quantizer, - keep_in_fp32_modules=keep_in_fp32_modules, - dduf_entries=dduf_entries, - loaded_keys=loaded_keys, - unexpected_keys=unexpected_keys, - offload_index=offload_index, - offload_folder=offload_folder, - state_dict_index=state_dict_index, - state_dict_folder=state_dict_folder, - ignore_mismatched_sizes=ignore_mismatched_sizes, - low_cpu_mem_usage=low_cpu_mem_usage, - disable_mmap=disable_mmap, - ) - - if is_parallel_loading_enabled: - offload_index, state_dict_index, _mismatched_keys, _error_msgs = load_fn(resolved_model_file) - error_msgs += _error_msgs - mismatched_keys += _mismatched_keys - else: - shard_files = resolved_model_file - if len(resolved_model_file) > 1: - shard_tqdm_kwargs = {"desc": "Loading checkpoint shards"} - if not is_torch_dist_rank_zero(): - shard_tqdm_kwargs["disable"] = True - shard_files = logging.tqdm(resolved_model_file, **shard_tqdm_kwargs) - - for shard_file in shard_files: - offload_index, state_dict_index, _mismatched_keys, _error_msgs = load_fn(shard_file) - error_msgs += _error_msgs - mismatched_keys += _mismatched_keys - - empty_device_cache() - - if offload_index is not None and len(offload_index) > 0: - save_offload_index(offload_index, offload_folder) - offload_index = None - - if offload_state_dict: - load_offloaded_weights(model, state_dict_index, state_dict_folder) - shutil.rmtree(state_dict_folder) - - if len(error_msgs) > 0: - error_msg = "\n\t".join(error_msgs) - if "size mismatch" in error_msg: - error_msg += ( - "\n\tYou may consider adding `ignore_mismatched_sizes=True` in the model `from_pretrained` method." - ) - raise RuntimeError(f"Error(s) in loading state_dict for {model.__class__.__name__}:\n\t{error_msg}") - - if len(unexpected_keys) > 0: - logger.warning( - f"Some weights of the model checkpoint at {pretrained_model_name_or_path} were not used when initializing {cls.__name__}: \n {[', '.join(unexpected_keys)]}" - ) - else: - logger.info(f"All model checkpoint weights were used when initializing {model.__class__.__name__}.\n") - - if len(missing_keys) > 0: - logger.warning( - f"Some weights of {model.__class__.__name__} were not initialized from the model checkpoint at" - f" {pretrained_model_name_or_path} and are newly initialized: {missing_keys}\nYou should probably" - " TRAIN this model on a down-stream task to be able to use it for predictions and inference." - ) - elif len(mismatched_keys) == 0: - logger.info( - f"All the weights of {model.__class__.__name__} were initialized from the model checkpoint at" - f" {pretrained_model_name_or_path}.\nIf your task is similar to the task the model of the" - f" checkpoint was trained on, you can already use {model.__class__.__name__} for predictions" - " without further training." - ) - if len(mismatched_keys) > 0: - mismatched_warning = "\n".join( - [ - f"- {key}: found shape {shape1} in the checkpoint and {shape2} in the model instantiated" - for key, shape1, shape2 in mismatched_keys - ] - ) - logger.warning( - f"Some weights of {model.__class__.__name__} were not initialized from the model checkpoint at" - f" {pretrained_model_name_or_path} and are newly initialized because the shapes did not" - f" match:\n{mismatched_warning}\nYou should probably TRAIN this model on a down-stream task to be" - " able to use it for predictions and inference." - ) - - return model, missing_keys, unexpected_keys, mismatched_keys, offload_index, error_msgs - - @classmethod - def _get_signature_keys(cls, obj): - parameters = inspect.signature(obj.__init__).parameters - required_parameters = {k: v for k, v in parameters.items() if v.default == inspect._empty} - optional_parameters = set({k for k, v in parameters.items() if v.default != inspect._empty}) - expected_modules = set(required_parameters.keys()) - {"self"} - - return expected_modules, optional_parameters - - # Adapted from `transformers` modeling_utils.py - def _get_no_split_modules(self, device_map: str): - """ - Get the modules of the model that should not be split when using device_map. We iterate through the modules to - get the underlying `_no_split_modules`. - - Args: - device_map (`str`): - The device map value. Options are ["auto", "balanced", "balanced_low_0", "sequential"] - - Returns: - `list[str]`: list of modules that should not be split - """ - _no_split_modules = set() - modules_to_check = [self] - while len(modules_to_check) > 0: - module = modules_to_check.pop(-1) - # if the module does not appear in _no_split_modules, we also check the children - if module.__class__.__name__ not in _no_split_modules: - if isinstance(module, ModelMixin): - if module._no_split_modules is None: - raise ValueError( - f"{module.__class__.__name__} does not support `device_map='{device_map}'`. To implement support, the model " - "class needs to implement the `_no_split_modules` attribute." - ) - else: - _no_split_modules = _no_split_modules | set(module._no_split_modules) - modules_to_check += list(module.children()) - return list(_no_split_modules) - - @classmethod - def _set_default_torch_dtype(cls, dtype: torch.dtype) -> torch.dtype: - """ - Change the default dtype and return the previous one. This is needed when wanting to instantiate the model - under specific dtype. - - Args: - dtype (`torch.dtype`): - a floating dtype to set to. - - Returns: - `torch.dtype`: the original `dtype` that can be used to restore `torch.set_default_dtype(dtype)` if it was - modified. If it wasn't, returns `None`. - - Note `set_default_dtype` currently only works with floating-point types and asserts if for example, - `torch.int64` is passed. So if a non-float `dtype` is passed this functions will throw an exception. - """ - if not dtype.is_floating_point: - raise ValueError( - f"Can't instantiate {cls.__name__} model under dtype={dtype} since it is not a floating point dtype" - ) - - logger.info(f"Instantiating {cls.__name__} model under default dtype {dtype}.") - dtype_orig = torch.get_default_dtype() - torch.set_default_dtype(dtype) - return dtype_orig - - @property - def device(self) -> torch.device: - """ - `torch.device`: The device on which the module is (assuming that all the module parameters are on the same - device). - """ - return get_parameter_device(self) - - @property - def dtype(self) -> torch.dtype: - """ - `torch.dtype`: The dtype of the module (assuming that all the module parameters have the same dtype). - """ - return get_parameter_dtype(self) - - def num_parameters(self, only_trainable: bool = False, exclude_embeddings: bool = False) -> int: - """ - Get number of (trainable or non-embedding) parameters in the module. - - Args: - only_trainable (`bool`, *optional*, defaults to `False`): - Whether or not to return only the number of trainable parameters. - exclude_embeddings (`bool`, *optional*, defaults to `False`): - Whether or not to return only the number of non-embedding parameters. - - Returns: - `int`: The number of parameters. - - Example: - - ```py - from diffusers import UNet2DConditionModel - - model_id = "stable-diffusion-v1-5/stable-diffusion-v1-5" - unet = UNet2DConditionModel.from_pretrained(model_id, subfolder="unet") - unet.num_parameters(only_trainable=True) - 859520964 - ``` - """ - is_loaded_in_4bit = getattr(self, "is_loaded_in_4bit", False) - - if is_loaded_in_4bit: - if is_bitsandbytes_available(): - import bitsandbytes as bnb - else: - raise ValueError( - "bitsandbytes is not installed but it seems that the model has been loaded in 4bit precision, something went wrong" - " make sure to install bitsandbytes with `pip install bitsandbytes`. You also need a GPU. " - ) - - if exclude_embeddings: - embedding_param_names = [ - f"{name}.weight" for name, module_type in self.named_modules() if isinstance(module_type, nn.Embedding) - ] - total_parameters = [ - parameter for name, parameter in self.named_parameters() if name not in embedding_param_names - ] - else: - total_parameters = list(self.parameters()) - - total_numel = [] - - for param in total_parameters: - if param.requires_grad or not only_trainable: - # For 4bit models, we need to multiply the number of parameters by 2 as half of the parameters are - # used for the 4bit quantization (uint8 tensors are stored) - if is_loaded_in_4bit and isinstance(param, bnb.nn.Params4bit): - if hasattr(param, "element_size"): - num_bytes = param.element_size() - elif hasattr(param, "quant_storage"): - num_bytes = param.quant_storage.itemsize - else: - num_bytes = 1 - total_numel.append(param.numel() * 2 * num_bytes) - else: - total_numel.append(param.numel()) - - return sum(total_numel) - - def get_memory_footprint(self, return_buffers=True): - r""" - Get the memory footprint of a model. This will return the memory footprint of the current model in bytes. - Useful to benchmark the memory footprint of the current model and design some tests. Solution inspired from the - PyTorch discussions: https://discuss.pytorch.org/t/gpu-memory-that-model-uses/56822/2 - - Arguments: - return_buffers (`bool`, *optional*, defaults to `True`): - Whether to return the size of the buffer tensors in the computation of the memory footprint. Buffers - are tensors that do not require gradients and not registered as parameters. E.g. mean and std in batch - norm layers. Please see: https://discuss.pytorch.org/t/what-pytorch-means-by-buffers/120266/2 - """ - mem = sum([param.nelement() * param.element_size() for param in self.parameters()]) - if return_buffers: - mem_bufs = sum([buf.nelement() * buf.element_size() for buf in self.buffers()]) - mem = mem + mem_bufs - return mem - - def _set_gradient_checkpointing( - self, enable: bool = True, gradient_checkpointing_func: Callable = torch.utils.checkpoint.checkpoint - ) -> None: - is_gradient_checkpointing_set = False - - for name, module in self.named_modules(): - if hasattr(module, "gradient_checkpointing"): - logger.debug(f"Setting `gradient_checkpointing={enable}` for '{name}'") - module._gradient_checkpointing_func = gradient_checkpointing_func - module.gradient_checkpointing = enable - is_gradient_checkpointing_set = True - - if not is_gradient_checkpointing_set: - raise ValueError( - f"The module {self.__class__.__name__} does not support gradient checkpointing. Please make sure to " - f"use a module that supports gradient checkpointing by creating a boolean attribute `gradient_checkpointing`." - ) - - def _fix_state_dict_keys_on_load(self, state_dict: OrderedDict) -> None: - """ - This function fix the state dict of the model to take into account some changes that were made in the model - architecture: - - deprecated attention blocks (happened before we introduced sharded checkpoint, - so this is why we apply this method only when loading non sharded checkpoints for now) - """ - deprecated_attention_block_paths = [] - - def recursive_find_attn_block(name, module): - if hasattr(module, "_from_deprecated_attn_block") and module._from_deprecated_attn_block: - deprecated_attention_block_paths.append(name) - - for sub_name, sub_module in module.named_children(): - sub_name = sub_name if name == "" else f"{name}.{sub_name}" - recursive_find_attn_block(sub_name, sub_module) - - recursive_find_attn_block("", self) - - # NOTE: we have to check if the deprecated parameters are in the state dict - # because it is possible we are loading from a state dict that was already - # converted - - for path in deprecated_attention_block_paths: - # group_norm path stays the same - - # query -> to_q - if f"{path}.query.weight" in state_dict: - state_dict[f"{path}.to_q.weight"] = state_dict.pop(f"{path}.query.weight") - if f"{path}.query.bias" in state_dict: - state_dict[f"{path}.to_q.bias"] = state_dict.pop(f"{path}.query.bias") - - # key -> to_k - if f"{path}.key.weight" in state_dict: - state_dict[f"{path}.to_k.weight"] = state_dict.pop(f"{path}.key.weight") - if f"{path}.key.bias" in state_dict: - state_dict[f"{path}.to_k.bias"] = state_dict.pop(f"{path}.key.bias") - - # value -> to_v - if f"{path}.value.weight" in state_dict: - state_dict[f"{path}.to_v.weight"] = state_dict.pop(f"{path}.value.weight") - if f"{path}.value.bias" in state_dict: - state_dict[f"{path}.to_v.bias"] = state_dict.pop(f"{path}.value.bias") - - # proj_attn -> to_out.0 - if f"{path}.proj_attn.weight" in state_dict: - state_dict[f"{path}.to_out.0.weight"] = state_dict.pop(f"{path}.proj_attn.weight") - if f"{path}.proj_attn.bias" in state_dict: - state_dict[f"{path}.to_out.0.bias"] = state_dict.pop(f"{path}.proj_attn.bias") - return state_dict - - -class LegacyModelMixin(ModelMixin): - r""" - A subclass of `ModelMixin` to resolve class mapping from legacy classes (like `Transformer2DModel`) to more - pipeline-specific classes (like `DiTTransformer2DModel`). - """ - - @classmethod - @validate_hf_hub_args - def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None, **kwargs): - # To prevent dependency import problem. - from .model_loading_utils import _fetch_remapped_cls_from_config - - # Create a copy of the kwargs so that we don't mess with the keyword arguments in the downstream calls. - kwargs_copy = kwargs.copy() - - cache_dir = kwargs.pop("cache_dir", None) - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - local_files_only = kwargs.pop("local_files_only", None) - token = kwargs.pop("token", None) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - - # Load config if we don't provide a configuration - config_path = pretrained_model_name_or_path - - user_agent = { - "diffusers": __version__, - "file_type": "model", - "framework": "pytorch", - } - - # load config - config, _, _ = cls.load_config( - config_path, - cache_dir=cache_dir, - return_unused_kwargs=True, - return_commit_hash=True, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - **kwargs, - ) - # resolve remapping - remapped_class = _fetch_remapped_cls_from_config(config, cls) - - if remapped_class is cls: - return super(LegacyModelMixin, remapped_class).from_pretrained( - pretrained_model_name_or_path, **kwargs_copy - ) - else: - return remapped_class.from_pretrained(pretrained_model_name_or_path, **kwargs_copy) diff --git a/diffusers/models/normalization.py b/diffusers/models/normalization.py deleted file mode 100644 index 84ffb67bfd6ac23147e3a7f08416374e5089b1d6..0000000000000000000000000000000000000000 --- a/diffusers/models/normalization.py +++ /dev/null @@ -1,647 +0,0 @@ -# coding=utf-8 -# Copyright 2025 HuggingFace Inc. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import numbers - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ..utils import is_torch_npu_available, is_torch_version -from .activations import get_activation -from .embeddings import CombinedTimestepLabelEmbeddings, PixArtAlphaCombinedTimestepSizeEmbeddings - - -class AdaLayerNorm(nn.Module): - r""" - Norm layer modified to incorporate timestep embeddings. - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - num_embeddings (`int`, *optional*): The size of the embeddings dictionary. - output_dim (`int`, *optional*): - norm_elementwise_affine (`bool`, defaults to `False): - norm_eps (`bool`, defaults to `False`): - chunk_dim (`int`, defaults to `0`): - """ - - def __init__( - self, - embedding_dim: int, - num_embeddings: int | None = None, - output_dim: int | None = None, - norm_elementwise_affine: bool = False, - norm_eps: float = 1e-5, - chunk_dim: int = 0, - ): - super().__init__() - - self.chunk_dim = chunk_dim - output_dim = output_dim or embedding_dim * 2 - - if num_embeddings is not None: - self.emb = nn.Embedding(num_embeddings, embedding_dim) - else: - self.emb = None - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, output_dim) - self.norm = nn.LayerNorm(output_dim // 2, norm_eps, norm_elementwise_affine) - - def forward( - self, x: torch.Tensor, timestep: torch.Tensor | None = None, temb: torch.Tensor | None = None - ) -> torch.Tensor: - if self.emb is not None: - temb = self.emb(timestep) - - temb = self.linear(self.silu(temb)) - - if self.chunk_dim == 1: - # This is a bit weird why we have the order of "shift, scale" here and "scale, shift" in the - # other if-branch. This branch is specific to CogVideoX and OmniGen for now. - shift, scale = temb.chunk(2, dim=1) - shift = shift[:, None, :] - scale = scale[:, None, :] - else: - scale, shift = temb.chunk(2, dim=0) - - x = self.norm(x) * (1 + scale) + shift - return x - - -class FP32LayerNorm(nn.LayerNorm): - def forward(self, inputs: torch.Tensor) -> torch.Tensor: - origin_dtype = inputs.dtype - return F.layer_norm( - inputs.float(), - self.normalized_shape, - self.weight.float() if self.weight is not None else None, - self.bias.float() if self.bias is not None else None, - self.eps, - ).to(origin_dtype) - - -class SD35AdaLayerNormZeroX(nn.Module): - r""" - Norm layer adaptive layer norm zero (AdaLN-Zero). - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - num_embeddings (`int`): The size of the embeddings dictionary. - """ - - def __init__(self, embedding_dim: int, norm_type: str = "layer_norm", bias: bool = True) -> None: - super().__init__() - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, 9 * embedding_dim, bias=bias) - if norm_type == "layer_norm": - self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) - else: - raise ValueError(f"Unsupported `norm_type` ({norm_type}) provided. Supported ones are: 'layer_norm'.") - - def forward( - self, - hidden_states: torch.Tensor, - emb: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, ...]: - emb = self.linear(self.silu(emb)) - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp, shift_msa2, scale_msa2, gate_msa2 = emb.chunk( - 9, dim=1 - ) - norm_hidden_states = self.norm(hidden_states) - hidden_states = norm_hidden_states * (1 + scale_msa[:, None]) + shift_msa[:, None] - norm_hidden_states2 = norm_hidden_states * (1 + scale_msa2[:, None]) + shift_msa2[:, None] - return hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp, norm_hidden_states2, gate_msa2 - - -class AdaLayerNormZero(nn.Module): - r""" - Norm layer adaptive layer norm zero (adaLN-Zero). - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - num_embeddings (`int`): The size of the embeddings dictionary. - """ - - def __init__(self, embedding_dim: int, num_embeddings: int | None = None, norm_type="layer_norm", bias=True): - super().__init__() - if num_embeddings is not None: - self.emb = CombinedTimestepLabelEmbeddings(num_embeddings, embedding_dim) - else: - self.emb = None - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=bias) - if norm_type == "layer_norm": - self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) - elif norm_type == "fp32_layer_norm": - self.norm = FP32LayerNorm(embedding_dim, elementwise_affine=False, bias=False) - else: - raise ValueError( - f"Unsupported `norm_type` ({norm_type}) provided. Supported ones are: 'layer_norm', 'fp32_layer_norm'." - ) - - def forward( - self, - x: torch.Tensor, - timestep: torch.Tensor | None = None, - class_labels: torch.LongTensor | None = None, - hidden_dtype: torch.dtype | None = None, - emb: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - if self.emb is not None: - emb = self.emb(timestep, class_labels, hidden_dtype=hidden_dtype) - emb = self.linear(self.silu(emb)) - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = emb.chunk(6, dim=1) - x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None] - return x, gate_msa, shift_mlp, scale_mlp, gate_mlp - - -class AdaLayerNormZeroSingle(nn.Module): - r""" - Norm layer adaptive layer norm zero (adaLN-Zero). - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - num_embeddings (`int`): The size of the embeddings dictionary. - """ - - def __init__(self, embedding_dim: int, norm_type="layer_norm", bias=True): - super().__init__() - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, 3 * embedding_dim, bias=bias) - if norm_type == "layer_norm": - self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) - else: - raise ValueError( - f"Unsupported `norm_type` ({norm_type}) provided. Supported ones are: 'layer_norm', 'fp32_layer_norm'." - ) - - def forward( - self, - x: torch.Tensor, - emb: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - emb = self.linear(self.silu(emb)) - shift_msa, scale_msa, gate_msa = emb.chunk(3, dim=1) - x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None] - return x, gate_msa - - -class LuminaRMSNormZero(nn.Module): - """ - Norm layer adaptive RMS normalization zero. - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - """ - - def __init__(self, embedding_dim: int, norm_eps: float, norm_elementwise_affine: bool): - super().__init__() - self.silu = nn.SiLU() - self.linear = nn.Linear( - min(embedding_dim, 1024), - 4 * embedding_dim, - bias=True, - ) - self.norm = RMSNorm(embedding_dim, eps=norm_eps) - - def forward( - self, - x: torch.Tensor, - emb: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - emb = self.linear(self.silu(emb)) - scale_msa, gate_msa, scale_mlp, gate_mlp = emb.chunk(4, dim=1) - x = self.norm(x) * (1 + scale_msa[:, None]) - - return x, gate_msa, scale_mlp, gate_mlp - - -class AdaLayerNormSingle(nn.Module): - r""" - Norm layer adaptive layer norm single (adaLN-single). - - As proposed in PixArt-Alpha (see: https://huggingface.co/papers/2310.00426; Section 2.3). - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - use_additional_conditions (`bool`): To use additional conditions for normalization or not. - """ - - def __init__(self, embedding_dim: int, use_additional_conditions: bool = False): - super().__init__() - - self.emb = PixArtAlphaCombinedTimestepSizeEmbeddings( - embedding_dim, size_emb_dim=embedding_dim // 3, use_additional_conditions=use_additional_conditions - ) - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=True) - - def forward( - self, - timestep: torch.Tensor, - added_cond_kwargs: dict[str, torch.Tensor] | None = None, - batch_size: int | None = None, - hidden_dtype: torch.dtype | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - # No modulation happening here. - added_cond_kwargs = added_cond_kwargs or {"resolution": None, "aspect_ratio": None} - embedded_timestep = self.emb(timestep, **added_cond_kwargs, batch_size=batch_size, hidden_dtype=hidden_dtype) - return self.linear(self.silu(embedded_timestep)), embedded_timestep - - -class AdaGroupNorm(nn.Module): - r""" - GroupNorm layer modified to incorporate timestep embeddings. - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - num_embeddings (`int`): The size of the embeddings dictionary. - num_groups (`int`): The number of groups to separate the channels into. - act_fn (`str`, *optional*, defaults to `None`): The activation function to use. - eps (`float`, *optional*, defaults to `1e-5`): The epsilon value to use for numerical stability. - """ - - def __init__( - self, embedding_dim: int, out_dim: int, num_groups: int, act_fn: str | None = None, eps: float = 1e-5 - ): - super().__init__() - self.num_groups = num_groups - self.eps = eps - - if act_fn is None: - self.act = None - else: - self.act = get_activation(act_fn) - - self.linear = nn.Linear(embedding_dim, out_dim * 2) - - def forward(self, x: torch.Tensor, emb: torch.Tensor) -> torch.Tensor: - if self.act: - emb = self.act(emb) - emb = self.linear(emb) - emb = emb[:, :, None, None] - scale, shift = emb.chunk(2, dim=1) - - x = F.group_norm(x, self.num_groups, eps=self.eps) - x = x * (1 + scale) + shift - return x - - -class AdaLayerNormContinuous(nn.Module): - r""" - Adaptive normalization layer with a norm layer (layer_norm or rms_norm). - - Args: - embedding_dim (`int`): Embedding dimension to use during projection. - conditioning_embedding_dim (`int`): Dimension of the input condition. - elementwise_affine (`bool`, defaults to `True`): - Boolean flag to denote if affine transformation should be applied. - eps (`float`, defaults to 1e-5): Epsilon factor. - bias (`bias`, defaults to `True`): Boolean flag to denote if bias should be use. - norm_type (`str`, defaults to `"layer_norm"`): - Normalization layer to use. Values supported: "layer_norm", "rms_norm". - """ - - def __init__( - self, - embedding_dim: int, - conditioning_embedding_dim: int, - # NOTE: It is a bit weird that the norm layer can be configured to have scale and shift parameters - # because the output is immediately scaled and shifted by the projected conditioning embeddings. - # Note that AdaLayerNorm does not let the norm layer have scale and shift parameters. - # However, this is how it was implemented in the original code, and it's rather likely you should - # set `elementwise_affine` to False. - elementwise_affine=True, - eps=1e-5, - bias=True, - norm_type="layer_norm", - ): - super().__init__() - self.silu = nn.SiLU() - self.linear = nn.Linear(conditioning_embedding_dim, embedding_dim * 2, bias=bias) - if norm_type == "layer_norm": - self.norm = LayerNorm(embedding_dim, eps, elementwise_affine, bias) - elif norm_type == "rms_norm": - self.norm = RMSNorm(embedding_dim, eps, elementwise_affine) - else: - raise ValueError(f"unknown norm_type {norm_type}") - - def forward(self, x: torch.Tensor, conditioning_embedding: torch.Tensor) -> torch.Tensor: - # convert back to the original dtype in case `conditioning_embedding`` is upcasted to float32 (needed for hunyuanDiT) - emb = self.linear(self.silu(conditioning_embedding).to(x.dtype)) - scale, shift = torch.chunk(emb, 2, dim=1) - x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :] - return x - - -class LuminaLayerNormContinuous(nn.Module): - def __init__( - self, - embedding_dim: int, - conditioning_embedding_dim: int, - # NOTE: It is a bit weird that the norm layer can be configured to have scale and shift parameters - # because the output is immediately scaled and shifted by the projected conditioning embeddings. - # Note that AdaLayerNorm does not let the norm layer have scale and shift parameters. - # However, this is how it was implemented in the original code, and it's rather likely you should - # set `elementwise_affine` to False. - elementwise_affine=True, - eps=1e-5, - bias=True, - norm_type="layer_norm", - out_dim: int | None = None, - ): - super().__init__() - - # AdaLN - self.silu = nn.SiLU() - self.linear_1 = nn.Linear(conditioning_embedding_dim, embedding_dim, bias=bias) - - if norm_type == "layer_norm": - self.norm = LayerNorm(embedding_dim, eps, elementwise_affine, bias) - elif norm_type == "rms_norm": - self.norm = RMSNorm(embedding_dim, eps=eps, elementwise_affine=elementwise_affine) - else: - raise ValueError(f"unknown norm_type {norm_type}") - - self.linear_2 = None - if out_dim is not None: - self.linear_2 = nn.Linear(embedding_dim, out_dim, bias=bias) - - def forward( - self, - x: torch.Tensor, - conditioning_embedding: torch.Tensor, - ) -> torch.Tensor: - # convert back to the original dtype in case `conditioning_embedding`` is upcasted to float32 (needed for hunyuanDiT) - emb = self.linear_1(self.silu(conditioning_embedding).to(x.dtype)) - scale = emb - x = self.norm(x) * (1 + scale)[:, None, :] - - if self.linear_2 is not None: - x = self.linear_2(x) - - return x - - -class CogView3PlusAdaLayerNormZeroTextImage(nn.Module): - r""" - Norm layer adaptive layer norm zero (adaLN-Zero). - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - num_embeddings (`int`): The size of the embeddings dictionary. - """ - - def __init__(self, embedding_dim: int, dim: int): - super().__init__() - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, 12 * dim, bias=True) - self.norm_x = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - self.norm_c = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - - def forward( - self, - x: torch.Tensor, - context: torch.Tensor, - emb: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - emb = self.linear(self.silu(emb)) - ( - shift_msa, - scale_msa, - gate_msa, - shift_mlp, - scale_mlp, - gate_mlp, - c_shift_msa, - c_scale_msa, - c_gate_msa, - c_shift_mlp, - c_scale_mlp, - c_gate_mlp, - ) = emb.chunk(12, dim=1) - normed_x = self.norm_x(x) - normed_context = self.norm_c(context) - x = normed_x * (1 + scale_msa[:, None]) + shift_msa[:, None] - context = normed_context * (1 + c_scale_msa[:, None]) + c_shift_msa[:, None] - return x, gate_msa, shift_mlp, scale_mlp, gate_mlp, context, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp - - -class CogVideoXLayerNormZero(nn.Module): - def __init__( - self, - conditioning_dim: int, - embedding_dim: int, - elementwise_affine: bool = True, - eps: float = 1e-5, - bias: bool = True, - ) -> None: - super().__init__() - - self.silu = nn.SiLU() - self.linear = nn.Linear(conditioning_dim, 6 * embedding_dim, bias=bias) - self.norm = nn.LayerNorm(embedding_dim, eps=eps, elementwise_affine=elementwise_affine) - - def forward( - self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor]: - shift, scale, gate, enc_shift, enc_scale, enc_gate = self.linear(self.silu(temb)).chunk(6, dim=1) - hidden_states = self.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :] - encoder_hidden_states = self.norm(encoder_hidden_states) * (1 + enc_scale)[:, None, :] + enc_shift[:, None, :] - return hidden_states, encoder_hidden_states, gate[:, None, :], enc_gate[:, None, :] - - -if is_torch_version(">=", "2.1.0"): - LayerNorm = nn.LayerNorm -else: - # Has optional bias parameter compared to torch layer norm - # TODO: replace with torch layernorm once min required torch version >= 2.1 - class LayerNorm(nn.Module): - r""" - LayerNorm with the bias parameter. - - Args: - dim (`int`): Dimensionality to use for the parameters. - eps (`float`, defaults to 1e-5): Epsilon factor. - elementwise_affine (`bool`, defaults to `True`): - Boolean flag to denote if affine transformation should be applied. - bias (`bias`, defaults to `True`): Boolean flag to denote if bias should be use. - """ - - def __init__(self, dim, eps: float = 1e-5, elementwise_affine: bool = True, bias: bool = True): - super().__init__() - - self.eps = eps - - if isinstance(dim, numbers.Integral): - dim = (dim,) - - self.dim = torch.Size(dim) - - if elementwise_affine: - self.weight = nn.Parameter(torch.ones(dim)) - self.bias = nn.Parameter(torch.zeros(dim)) if bias else None - else: - self.weight = None - self.bias = None - - def forward(self, input): - return F.layer_norm(input, self.dim, self.weight, self.bias, self.eps) - - -class RMSNorm(nn.Module): - r""" - RMS Norm as introduced in https://huggingface.co/papers/1910.07467 by Zhang et al. - - Args: - dim (`int`): Number of dimensions to use for `weights`. Only effective when `elementwise_affine` is True. - eps (`float`): Small value to use when calculating the reciprocal of the square-root. - elementwise_affine (`bool`, defaults to `True`): - Boolean flag to denote if affine transformation should be applied. - bias (`bool`, defaults to False): If also training the `bias` param. - """ - - def __init__(self, dim, eps: float, elementwise_affine: bool = True, bias: bool = False): - super().__init__() - - self.eps = eps - self.elementwise_affine = elementwise_affine - - if isinstance(dim, numbers.Integral): - dim = (dim,) - - self.dim = torch.Size(dim) - - self.weight = None - self.bias = None - - if elementwise_affine: - self.weight = nn.Parameter(torch.ones(dim)) - if bias: - self.bias = nn.Parameter(torch.zeros(dim)) - - def forward(self, hidden_states): - if is_torch_npu_available(): - import torch_npu - - if self.weight is not None: - # convert into half-precision if necessary - if self.weight.dtype in [torch.float16, torch.bfloat16]: - hidden_states = hidden_states.to(self.weight.dtype) - hidden_states = torch_npu.npu_rms_norm(hidden_states, self.weight, epsilon=self.eps)[0] - if self.bias is not None: - hidden_states = hidden_states + self.bias - else: - input_dtype = hidden_states.dtype - variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True) - hidden_states = hidden_states * torch.rsqrt(variance + self.eps) - - if self.weight is not None: - # convert into half-precision if necessary - if self.weight.dtype in [torch.float16, torch.bfloat16]: - hidden_states = hidden_states.to(self.weight.dtype) - hidden_states = hidden_states * self.weight - if self.bias is not None: - hidden_states = hidden_states + self.bias - else: - hidden_states = hidden_states.to(input_dtype) - - return hidden_states - - -# TODO: (Dhruv) This can be replaced with regular RMSNorm in Mochi once `_keep_in_fp32_modules` is supported -# for sharded checkpoints, see: https://github.com/huggingface/diffusers/issues/10013 -class MochiRMSNorm(nn.Module): - def __init__(self, dim, eps: float, elementwise_affine: bool = True): - super().__init__() - - self.eps = eps - - if isinstance(dim, numbers.Integral): - dim = (dim,) - - self.dim = torch.Size(dim) - - if elementwise_affine: - self.weight = nn.Parameter(torch.ones(dim)) - else: - self.weight = None - - def forward(self, hidden_states): - input_dtype = hidden_states.dtype - variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True) - hidden_states = hidden_states * torch.rsqrt(variance + self.eps) - - if self.weight is not None: - hidden_states = hidden_states * self.weight - hidden_states = hidden_states.to(input_dtype) - - return hidden_states - - -class GlobalResponseNorm(nn.Module): - r""" - Global response normalization as introduced in ConvNeXt-v2 (https://huggingface.co/papers/2301.00808). - - Args: - dim (`int`): Number of dimensions to use for the `gamma` and `beta`. - """ - - # Taken from https://github.com/facebookresearch/ConvNeXt-V2/blob/3608f67cc1dae164790c5d0aead7bf2d73d9719b/models/utils.py#L105 - def __init__(self, dim): - super().__init__() - self.gamma = nn.Parameter(torch.zeros(1, 1, 1, dim)) - self.beta = nn.Parameter(torch.zeros(1, 1, 1, dim)) - - def forward(self, x): - gx = torch.norm(x, p=2, dim=(1, 2), keepdim=True) - nx = gx / (gx.mean(dim=-1, keepdim=True) + 1e-6) - return self.gamma * (x * nx) + self.beta + x - - -class LpNorm(nn.Module): - def __init__(self, p: int = 2, dim: int = -1, eps: float = 1e-12): - super().__init__() - - self.p = p - self.dim = dim - self.eps = eps - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return F.normalize(hidden_states, p=self.p, dim=self.dim, eps=self.eps) - - -def get_normalization( - norm_type: str = "batch_norm", - num_features: int | None = None, - eps: float = 1e-5, - elementwise_affine: bool = True, - bias: bool = True, -) -> nn.Module: - if norm_type == "rms_norm": - norm = RMSNorm(num_features, eps=eps, elementwise_affine=elementwise_affine, bias=bias) - elif norm_type == "layer_norm": - norm = nn.LayerNorm(num_features, eps=eps, elementwise_affine=elementwise_affine, bias=bias) - elif norm_type == "batch_norm": - norm = nn.BatchNorm2d(num_features, eps=eps, affine=elementwise_affine) - else: - raise ValueError(f"{norm_type=} is not supported.") - return norm diff --git a/diffusers/models/resnet.py b/diffusers/models/resnet.py deleted file mode 100644 index d63e4fd0017be25518628e0e83f7f725f159017f..0000000000000000000000000000000000000000 --- a/diffusers/models/resnet.py +++ /dev/null @@ -1,801 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# `TemporalConvLayer` Copyright 2025 Alibaba DAMO-VILAB, The ModelScope Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from functools import partial - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ..utils import deprecate -from .activations import get_activation -from .attention_processor import SpatialNorm -from .downsampling import ( # noqa - Downsample1D, - Downsample2D, - FirDownsample2D, - KDownsample2D, - downsample_2d, -) -from .normalization import AdaGroupNorm -from .upsampling import ( # noqa - FirUpsample2D, - KUpsample2D, - Upsample1D, - Upsample2D, - upfirdn2d_native, - upsample_2d, -) - - -class ResnetBlockCondNorm2D(nn.Module): - r""" - A Resnet block that use normalization layer that incorporate conditioning information. - - Parameters: - in_channels (`int`): The number of channels in the input. - out_channels (`int`, *optional*, default to be `None`): - The number of output channels for the first conv2d layer. If None, same as `in_channels`. - dropout (`float`, *optional*, defaults to `0.0`): The dropout probability to use. - temb_channels (`int`, *optional*, default to `512`): the number of channels in timestep embedding. - groups (`int`, *optional*, default to `32`): The number of groups to use for the first normalization layer. - groups_out (`int`, *optional*, default to None): - The number of groups to use for the second normalization layer. if set to None, same as `groups`. - eps (`float`, *optional*, defaults to `1e-6`): The epsilon to use for the normalization. - non_linearity (`str`, *optional*, default to `"swish"`): the activation function to use. - time_embedding_norm (`str`, *optional*, default to `"ada_group"` ): - The normalization layer for time embedding `temb`. Currently only support "ada_group" or "spatial". - kernel (`torch.Tensor`, optional, default to None): FIR filter, see - [`~models.resnet.FirUpsample2D`] and [`~models.resnet.FirDownsample2D`]. - output_scale_factor (`float`, *optional*, default to be `1.0`): the scale factor to use for the output. - use_in_shortcut (`bool`, *optional*, default to `True`): - If `True`, add a 1x1 nn.conv2d layer for skip-connection. - up (`bool`, *optional*, default to `False`): If `True`, add an upsample layer. - down (`bool`, *optional*, default to `False`): If `True`, add a downsample layer. - conv_shortcut_bias (`bool`, *optional*, default to `True`): If `True`, adds a learnable bias to the - `conv_shortcut` output. - conv_2d_out_channels (`int`, *optional*, default to `None`): the number of channels in the output. - If None, same as `out_channels`. - """ - - def __init__( - self, - *, - in_channels: int, - out_channels: int | None = None, - conv_shortcut: bool = False, - dropout: float = 0.0, - temb_channels: int = 512, - groups: int = 32, - groups_out: int | None = None, - eps: float = 1e-6, - non_linearity: str = "swish", - time_embedding_norm: str = "ada_group", # ada_group, spatial - output_scale_factor: float = 1.0, - use_in_shortcut: bool | None = None, - up: bool = False, - down: bool = False, - conv_shortcut_bias: bool = True, - conv_2d_out_channels: int | None = None, - ): - super().__init__() - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - self.use_conv_shortcut = conv_shortcut - self.up = up - self.down = down - self.output_scale_factor = output_scale_factor - self.time_embedding_norm = time_embedding_norm - - if groups_out is None: - groups_out = groups - - if self.time_embedding_norm == "ada_group": # ada_group - self.norm1 = AdaGroupNorm(temb_channels, in_channels, groups, eps=eps) - elif self.time_embedding_norm == "spatial": - self.norm1 = SpatialNorm(in_channels, temb_channels) - else: - raise ValueError(f" unsupported time_embedding_norm: {self.time_embedding_norm}") - - self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) - - if self.time_embedding_norm == "ada_group": # ada_group - self.norm2 = AdaGroupNorm(temb_channels, out_channels, groups_out, eps=eps) - elif self.time_embedding_norm == "spatial": # spatial - self.norm2 = SpatialNorm(out_channels, temb_channels) - else: - raise ValueError(f" unsupported time_embedding_norm: {self.time_embedding_norm}") - - self.dropout = torch.nn.Dropout(dropout) - - conv_2d_out_channels = conv_2d_out_channels or out_channels - self.conv2 = nn.Conv2d(out_channels, conv_2d_out_channels, kernel_size=3, stride=1, padding=1) - - self.nonlinearity = get_activation(non_linearity) - - self.upsample = self.downsample = None - if self.up: - self.upsample = Upsample2D(in_channels, use_conv=False) - elif self.down: - self.downsample = Downsample2D(in_channels, use_conv=False, padding=1, name="op") - - self.use_in_shortcut = self.in_channels != conv_2d_out_channels if use_in_shortcut is None else use_in_shortcut - - self.conv_shortcut = None - if self.use_in_shortcut: - self.conv_shortcut = nn.Conv2d( - in_channels, - conv_2d_out_channels, - kernel_size=1, - stride=1, - padding=0, - bias=conv_shortcut_bias, - ) - - def forward(self, input_tensor: torch.Tensor, temb: torch.Tensor, *args, **kwargs) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - hidden_states = input_tensor - - hidden_states = self.norm1(hidden_states, temb) - - hidden_states = self.nonlinearity(hidden_states) - - if self.upsample is not None: - # upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984 - if hidden_states.shape[0] >= 64: - input_tensor = input_tensor.contiguous() - hidden_states = hidden_states.contiguous() - input_tensor = self.upsample(input_tensor) - hidden_states = self.upsample(hidden_states) - - elif self.downsample is not None: - input_tensor = self.downsample(input_tensor) - hidden_states = self.downsample(hidden_states) - - hidden_states = self.conv1(hidden_states) - - hidden_states = self.norm2(hidden_states, temb) - - hidden_states = self.nonlinearity(hidden_states) - - hidden_states = self.dropout(hidden_states) - hidden_states = self.conv2(hidden_states) - - if self.conv_shortcut is not None: - input_tensor = self.conv_shortcut(input_tensor) - - output_tensor = (input_tensor + hidden_states) / self.output_scale_factor - - return output_tensor - - -class ResnetBlock2D(nn.Module): - r""" - A Resnet block. - - Parameters: - in_channels (`int`): The number of channels in the input. - out_channels (`int`, *optional*, default to be `None`): - The number of output channels for the first conv2d layer. If None, same as `in_channels`. - dropout (`float`, *optional*, defaults to `0.0`): The dropout probability to use. - temb_channels (`int`, *optional*, default to `512`): the number of channels in timestep embedding. - groups (`int`, *optional*, default to `32`): The number of groups to use for the first normalization layer. - groups_out (`int`, *optional*, default to None): - The number of groups to use for the second normalization layer. if set to None, same as `groups`. - eps (`float`, *optional*, defaults to `1e-6`): The epsilon to use for the normalization. - non_linearity (`str`, *optional*, default to `"swish"`): the activation function to use. - time_embedding_norm (`str`, *optional*, default to `"default"` ): Time scale shift config. - By default, apply timestep embedding conditioning with a simple shift mechanism. Choose "scale_shift" for a - stronger conditioning with scale and shift. - kernel (`torch.Tensor`, optional, default to None): FIR filter, see - [`~models.resnet.FirUpsample2D`] and [`~models.resnet.FirDownsample2D`]. - output_scale_factor (`float`, *optional*, default to be `1.0`): the scale factor to use for the output. - use_in_shortcut (`bool`, *optional*, default to `True`): - If `True`, add a 1x1 nn.conv2d layer for skip-connection. - up (`bool`, *optional*, default to `False`): If `True`, add an upsample layer. - down (`bool`, *optional*, default to `False`): If `True`, add a downsample layer. - conv_shortcut_bias (`bool`, *optional*, default to `True`): If `True`, adds a learnable bias to the - `conv_shortcut` output. - conv_2d_out_channels (`int`, *optional*, default to `None`): the number of channels in the output. - If None, same as `out_channels`. - """ - - def __init__( - self, - *, - in_channels: int, - out_channels: int | None = None, - conv_shortcut: bool = False, - dropout: float = 0.0, - temb_channels: int = 512, - groups: int = 32, - groups_out: int | None = None, - pre_norm: bool = True, - eps: float = 1e-6, - non_linearity: str = "swish", - skip_time_act: bool = False, - time_embedding_norm: str = "default", # default, scale_shift, - kernel: torch.Tensor | None = None, - output_scale_factor: float = 1.0, - use_in_shortcut: bool | None = None, - up: bool = False, - down: bool = False, - conv_shortcut_bias: bool = True, - conv_2d_out_channels: int | None = None, - ): - super().__init__() - if time_embedding_norm == "ada_group": - raise ValueError( - "This class cannot be used with `time_embedding_norm==ada_group`, please use `ResnetBlockCondNorm2D` instead", - ) - if time_embedding_norm == "spatial": - raise ValueError( - "This class cannot be used with `time_embedding_norm==spatial`, please use `ResnetBlockCondNorm2D` instead", - ) - - self.pre_norm = True - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - self.use_conv_shortcut = conv_shortcut - self.up = up - self.down = down - self.output_scale_factor = output_scale_factor - self.time_embedding_norm = time_embedding_norm - self.skip_time_act = skip_time_act - - if groups_out is None: - groups_out = groups - - self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True) - - self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) - - if temb_channels is not None: - if self.time_embedding_norm == "default": - self.time_emb_proj = nn.Linear(temb_channels, out_channels) - elif self.time_embedding_norm == "scale_shift": - self.time_emb_proj = nn.Linear(temb_channels, 2 * out_channels) - else: - raise ValueError(f"unknown time_embedding_norm : {self.time_embedding_norm} ") - else: - self.time_emb_proj = None - - self.norm2 = torch.nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True) - - self.dropout = torch.nn.Dropout(dropout) - conv_2d_out_channels = conv_2d_out_channels or out_channels - self.conv2 = nn.Conv2d(out_channels, conv_2d_out_channels, kernel_size=3, stride=1, padding=1) - - self.nonlinearity = get_activation(non_linearity) - - self.upsample = self.downsample = None - if self.up: - if kernel == "fir": - fir_kernel = (1, 3, 3, 1) - self.upsample = lambda x: upsample_2d(x, kernel=fir_kernel) - elif kernel == "sde_vp": - self.upsample = partial(F.interpolate, scale_factor=2.0, mode="nearest") - else: - self.upsample = Upsample2D(in_channels, use_conv=False) - elif self.down: - if kernel == "fir": - fir_kernel = (1, 3, 3, 1) - self.downsample = lambda x: downsample_2d(x, kernel=fir_kernel) - elif kernel == "sde_vp": - self.downsample = partial(F.avg_pool2d, kernel_size=2, stride=2) - else: - self.downsample = Downsample2D(in_channels, use_conv=False, padding=1, name="op") - - self.use_in_shortcut = self.in_channels != conv_2d_out_channels if use_in_shortcut is None else use_in_shortcut - - self.conv_shortcut = None - if self.use_in_shortcut: - self.conv_shortcut = nn.Conv2d( - in_channels, - conv_2d_out_channels, - kernel_size=1, - stride=1, - padding=0, - bias=conv_shortcut_bias, - ) - - def forward(self, input_tensor: torch.Tensor, temb: torch.Tensor, *args, **kwargs) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - hidden_states = input_tensor - - hidden_states = self.norm1(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - - if self.upsample is not None: - # upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984 - if hidden_states.shape[0] >= 64: - input_tensor = input_tensor.contiguous() - hidden_states = hidden_states.contiguous() - input_tensor = self.upsample(input_tensor) - hidden_states = self.upsample(hidden_states) - elif self.downsample is not None: - input_tensor = self.downsample(input_tensor) - hidden_states = self.downsample(hidden_states) - - hidden_states = self.conv1(hidden_states) - - if self.time_emb_proj is not None: - if not self.skip_time_act: - temb = self.nonlinearity(temb) - temb = self.time_emb_proj(temb)[:, :, None, None] - - if self.time_embedding_norm == "default": - if temb is not None: - hidden_states = hidden_states + temb - hidden_states = self.norm2(hidden_states) - elif self.time_embedding_norm == "scale_shift": - if temb is None: - raise ValueError( - f" `temb` should not be None when `time_embedding_norm` is {self.time_embedding_norm}" - ) - time_scale, time_shift = torch.chunk(temb, 2, dim=1) - hidden_states = self.norm2(hidden_states) - hidden_states = hidden_states * (1 + time_scale) + time_shift - else: - hidden_states = self.norm2(hidden_states) - - hidden_states = self.nonlinearity(hidden_states) - - hidden_states = self.dropout(hidden_states) - hidden_states = self.conv2(hidden_states) - - if self.conv_shortcut is not None: - # Only use contiguous() during training to avoid DDP gradient stride mismatch warning. - # In inference mode (eval or no_grad), skip contiguous() for better performance, especially on CPU. - # Issue: https://github.com/huggingface/diffusers/issues/12975 - if self.training: - input_tensor = input_tensor.contiguous() - input_tensor = self.conv_shortcut(input_tensor) - - output_tensor = (input_tensor + hidden_states) / self.output_scale_factor - - return output_tensor - - -# unet_rl.py -def rearrange_dims(tensor: torch.Tensor) -> torch.Tensor: - if len(tensor.shape) == 2: - return tensor[:, :, None] - if len(tensor.shape) == 3: - return tensor[:, :, None, :] - elif len(tensor.shape) == 4: - return tensor[:, :, 0, :] - else: - raise ValueError(f"`len(tensor)`: {len(tensor)} has to be 2, 3 or 4.") - - -class Conv1dBlock(nn.Module): - """ - Conv1d --> GroupNorm --> Mish - - Parameters: - inp_channels (`int`): Number of input channels. - out_channels (`int`): Number of output channels. - kernel_size (`int` or `tuple`): Size of the convolving kernel. - n_groups (`int`, default `8`): Number of groups to separate the channels into. - activation (`str`, defaults to `mish`): Name of the activation function. - """ - - def __init__( - self, - inp_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int], - n_groups: int = 8, - activation: str = "mish", - ): - super().__init__() - - self.conv1d = nn.Conv1d(inp_channels, out_channels, kernel_size, padding=kernel_size // 2) - self.group_norm = nn.GroupNorm(n_groups, out_channels) - self.mish = get_activation(activation) - - def forward(self, inputs: torch.Tensor) -> torch.Tensor: - intermediate_repr = self.conv1d(inputs) - intermediate_repr = rearrange_dims(intermediate_repr) - intermediate_repr = self.group_norm(intermediate_repr) - intermediate_repr = rearrange_dims(intermediate_repr) - output = self.mish(intermediate_repr) - return output - - -# unet_rl.py -class ResidualTemporalBlock1D(nn.Module): - """ - Residual 1D block with temporal convolutions. - - Parameters: - inp_channels (`int`): Number of input channels. - out_channels (`int`): Number of output channels. - embed_dim (`int`): Embedding dimension. - kernel_size (`int` or `tuple`): Size of the convolving kernel. - activation (`str`, defaults `mish`): It is possible to choose the right activation function. - """ - - def __init__( - self, - inp_channels: int, - out_channels: int, - embed_dim: int, - kernel_size: int | tuple[int, int] = 5, - activation: str = "mish", - ): - super().__init__() - self.conv_in = Conv1dBlock(inp_channels, out_channels, kernel_size) - self.conv_out = Conv1dBlock(out_channels, out_channels, kernel_size) - - self.time_emb_act = get_activation(activation) - self.time_emb = nn.Linear(embed_dim, out_channels) - - self.residual_conv = ( - nn.Conv1d(inp_channels, out_channels, 1) if inp_channels != out_channels else nn.Identity() - ) - - def forward(self, inputs: torch.Tensor, t: torch.Tensor) -> torch.Tensor: - """ - Args: - inputs : [ batch_size x inp_channels x horizon ] - t : [ batch_size x embed_dim ] - - returns: - out : [ batch_size x out_channels x horizon ] - """ - t = self.time_emb_act(t) - t = self.time_emb(t) - out = self.conv_in(inputs) + rearrange_dims(t) - out = self.conv_out(out) - return out + self.residual_conv(inputs) - - -class TemporalConvLayer(nn.Module): - """ - Temporal convolutional layer that can be used for video (sequence of images) input Code mostly copied from: - https://github.com/modelscope/modelscope/blob/1509fdb973e5871f37148a4b5e5964cafd43e64d/modelscope/models/multi_modal/video_synthesis/unet_sd.py#L1016 - - Parameters: - in_dim (`int`): Number of input channels. - out_dim (`int`): Number of output channels. - dropout (`float`, *optional*, defaults to `0.0`): The dropout probability to use. - """ - - def __init__( - self, - in_dim: int, - out_dim: int | None = None, - dropout: float = 0.0, - norm_num_groups: int = 32, - ): - super().__init__() - out_dim = out_dim or in_dim - self.in_dim = in_dim - self.out_dim = out_dim - - # conv layers - self.conv1 = nn.Sequential( - nn.GroupNorm(norm_num_groups, in_dim), - nn.SiLU(), - nn.Conv3d(in_dim, out_dim, (3, 1, 1), padding=(1, 0, 0)), - ) - self.conv2 = nn.Sequential( - nn.GroupNorm(norm_num_groups, out_dim), - nn.SiLU(), - nn.Dropout(dropout), - nn.Conv3d(out_dim, in_dim, (3, 1, 1), padding=(1, 0, 0)), - ) - self.conv3 = nn.Sequential( - nn.GroupNorm(norm_num_groups, out_dim), - nn.SiLU(), - nn.Dropout(dropout), - nn.Conv3d(out_dim, in_dim, (3, 1, 1), padding=(1, 0, 0)), - ) - self.conv4 = nn.Sequential( - nn.GroupNorm(norm_num_groups, out_dim), - nn.SiLU(), - nn.Dropout(dropout), - nn.Conv3d(out_dim, in_dim, (3, 1, 1), padding=(1, 0, 0)), - ) - - # zero out the last layer params,so the conv block is identity - nn.init.zeros_(self.conv4[-1].weight) - nn.init.zeros_(self.conv4[-1].bias) - - def forward(self, hidden_states: torch.Tensor, num_frames: int = 1) -> torch.Tensor: - hidden_states = ( - hidden_states[None, :].reshape((-1, num_frames) + hidden_states.shape[1:]).permute(0, 2, 1, 3, 4) - ) - - identity = hidden_states - hidden_states = self.conv1(hidden_states) - hidden_states = self.conv2(hidden_states) - hidden_states = self.conv3(hidden_states) - hidden_states = self.conv4(hidden_states) - - hidden_states = identity + hidden_states - - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).reshape( - (hidden_states.shape[0] * hidden_states.shape[2], -1) + hidden_states.shape[3:] - ) - return hidden_states - - -class TemporalResnetBlock(nn.Module): - r""" - A Resnet block. - - Parameters: - in_channels (`int`): The number of channels in the input. - out_channels (`int`, *optional*, default to be `None`): - The number of output channels for the first conv2d layer. If None, same as `in_channels`. - temb_channels (`int`, *optional*, default to `512`): the number of channels in timestep embedding. - eps (`float`, *optional*, defaults to `1e-6`): The epsilon to use for the normalization. - """ - - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - temb_channels: int = 512, - eps: float = 1e-6, - ): - super().__init__() - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - - kernel_size = (3, 1, 1) - padding = [k // 2 for k in kernel_size] - - self.norm1 = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=eps, affine=True) - self.conv1 = nn.Conv3d( - in_channels, - out_channels, - kernel_size=kernel_size, - stride=1, - padding=padding, - ) - - if temb_channels is not None: - self.time_emb_proj = nn.Linear(temb_channels, out_channels) - else: - self.time_emb_proj = None - - self.norm2 = torch.nn.GroupNorm(num_groups=32, num_channels=out_channels, eps=eps, affine=True) - - self.dropout = torch.nn.Dropout(0.0) - self.conv2 = nn.Conv3d( - out_channels, - out_channels, - kernel_size=kernel_size, - stride=1, - padding=padding, - ) - - self.nonlinearity = get_activation("silu") - - self.use_in_shortcut = self.in_channels != out_channels - - self.conv_shortcut = None - if self.use_in_shortcut: - self.conv_shortcut = nn.Conv3d( - in_channels, - out_channels, - kernel_size=1, - stride=1, - padding=0, - ) - - def forward(self, input_tensor: torch.Tensor, temb: torch.Tensor) -> torch.Tensor: - hidden_states = input_tensor - - hidden_states = self.norm1(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv1(hidden_states) - - if self.time_emb_proj is not None: - temb = self.nonlinearity(temb) - temb = self.time_emb_proj(temb)[:, :, :, None, None] - temb = temb.permute(0, 2, 1, 3, 4) - hidden_states = hidden_states + temb - - hidden_states = self.norm2(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.dropout(hidden_states) - hidden_states = self.conv2(hidden_states) - - if self.conv_shortcut is not None: - input_tensor = self.conv_shortcut(input_tensor) - - output_tensor = input_tensor + hidden_states - - return output_tensor - - -# VideoResBlock -class SpatioTemporalResBlock(nn.Module): - r""" - A SpatioTemporal Resnet block. - - Parameters: - in_channels (`int`): The number of channels in the input. - out_channels (`int`, *optional*, default to be `None`): - The number of output channels for the first conv2d layer. If None, same as `in_channels`. - temb_channels (`int`, *optional*, default to `512`): the number of channels in timestep embedding. - eps (`float`, *optional*, defaults to `1e-6`): The epsilon to use for the spatial resenet. - temporal_eps (`float`, *optional*, defaults to `eps`): The epsilon to use for the temporal resnet. - merge_factor (`float`, *optional*, defaults to `0.5`): The merge factor to use for the temporal mixing. - merge_strategy (`str`, *optional*, defaults to `learned_with_images`): - The merge strategy to use for the temporal mixing. - switch_spatial_to_temporal_mix (`bool`, *optional*, defaults to `False`): - If `True`, switch the spatial and temporal mixing. - """ - - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - temb_channels: int = 512, - eps: float = 1e-6, - temporal_eps: float | None = None, - merge_factor: float = 0.5, - merge_strategy="learned_with_images", - switch_spatial_to_temporal_mix: bool = False, - ): - super().__init__() - - self.spatial_res_block = ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=eps, - ) - - self.temporal_res_block = TemporalResnetBlock( - in_channels=out_channels if out_channels is not None else in_channels, - out_channels=out_channels if out_channels is not None else in_channels, - temb_channels=temb_channels, - eps=temporal_eps if temporal_eps is not None else eps, - ) - - self.time_mixer = AlphaBlender( - alpha=merge_factor, - merge_strategy=merge_strategy, - switch_spatial_to_temporal_mix=switch_spatial_to_temporal_mix, - ) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - image_only_indicator: torch.Tensor | None = None, - ): - num_frames = image_only_indicator.shape[-1] - hidden_states = self.spatial_res_block(hidden_states, temb) - - batch_frames, channels, height, width = hidden_states.shape - batch_size = batch_frames // num_frames - - hidden_states_mix = ( - hidden_states[None, :].reshape(batch_size, num_frames, channels, height, width).permute(0, 2, 1, 3, 4) - ) - hidden_states = ( - hidden_states[None, :].reshape(batch_size, num_frames, channels, height, width).permute(0, 2, 1, 3, 4) - ) - - if temb is not None: - temb = temb.reshape(batch_size, num_frames, -1) - - hidden_states = self.temporal_res_block(hidden_states, temb) - hidden_states = self.time_mixer( - x_spatial=hidden_states_mix, - x_temporal=hidden_states, - image_only_indicator=image_only_indicator, - ) - - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).reshape(batch_frames, channels, height, width) - return hidden_states - - -class AlphaBlender(nn.Module): - r""" - A module to blend spatial and temporal features. - - Parameters: - alpha (`float`): The initial value of the blending factor. - merge_strategy (`str`, *optional*, defaults to `learned_with_images`): - The merge strategy to use for the temporal mixing. - switch_spatial_to_temporal_mix (`bool`, *optional*, defaults to `False`): - If `True`, switch the spatial and temporal mixing. - """ - - strategies = ["learned", "fixed", "learned_with_images"] - - def __init__( - self, - alpha: float, - merge_strategy: str = "learned_with_images", - switch_spatial_to_temporal_mix: bool = False, - ): - super().__init__() - self.merge_strategy = merge_strategy - self.switch_spatial_to_temporal_mix = switch_spatial_to_temporal_mix # For TemporalVAE - - if merge_strategy not in self.strategies: - raise ValueError(f"merge_strategy needs to be in {self.strategies}") - - if self.merge_strategy == "fixed": - self.register_buffer("mix_factor", torch.Tensor([alpha])) - elif self.merge_strategy == "learned" or self.merge_strategy == "learned_with_images": - self.register_parameter("mix_factor", torch.nn.Parameter(torch.Tensor([alpha]))) - else: - raise ValueError(f"Unknown merge strategy {self.merge_strategy}") - - def get_alpha(self, image_only_indicator: torch.Tensor, ndims: int) -> torch.Tensor: - if self.merge_strategy == "fixed": - alpha = self.mix_factor - - elif self.merge_strategy == "learned": - alpha = torch.sigmoid(self.mix_factor) - - elif self.merge_strategy == "learned_with_images": - if image_only_indicator is None: - raise ValueError("Please provide image_only_indicator to use learned_with_images merge strategy") - - alpha = torch.where( - image_only_indicator.bool(), - torch.ones(1, 1, device=image_only_indicator.device), - torch.sigmoid(self.mix_factor)[..., None], - ) - - # (batch, channel, frames, height, width) - if ndims == 5: - alpha = alpha[:, None, :, None, None] - # (batch*frames, height*width, channels) - elif ndims == 3: - alpha = alpha.reshape(-1)[:, None, None] - else: - raise ValueError(f"Unexpected ndims {ndims}. Dimensions should be 3 or 5") - - else: - raise NotImplementedError - - return alpha - - def forward( - self, - x_spatial: torch.Tensor, - x_temporal: torch.Tensor, - image_only_indicator: torch.Tensor | None = None, - ) -> torch.Tensor: - alpha = self.get_alpha(image_only_indicator, x_spatial.ndim) - alpha = alpha.to(x_spatial.dtype) - - if self.switch_spatial_to_temporal_mix: - alpha = 1.0 - alpha - - x = alpha * x_spatial + (1.0 - alpha) * x_temporal - return x diff --git a/diffusers/models/transformers/__init__.py b/diffusers/models/transformers/__init__.py deleted file mode 100644 index 7a1213639e3de21b1742eb658389d9fc4e689df4..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/__init__.py +++ /dev/null @@ -1,68 +0,0 @@ -from ...utils import is_torch_available - - -if is_torch_available(): - from .ace_step_transformer import AceStepTransformer1DModel - from .auraflow_transformer_2d import AuraFlowTransformer2DModel - from .cogvideox_transformer_3d import CogVideoXTransformer3DModel - from .consisid_transformer_3d import ConsisIDTransformer3DModel - from .dit_transformer_2d import DiTTransformer2DModel - from .dual_transformer_2d import DualTransformer2DModel - from .hunyuan_transformer_2d import HunyuanDiT2DModel - from .latte_transformer_3d import LatteTransformer3DModel - from .lumina_nextdit2d import LuminaNextDiT2DModel - from .pixart_transformer_2d import PixArtTransformer2DModel - from .prior_transformer import PriorTransformer - from .sana_transformer import SanaTransformer2DModel - from .stable_audio_transformer import StableAudioDiTModel - from .t5_film_transformer import T5FilmDecoder - from .transformer_2d import Transformer2DModel - from .transformer_2d_dreamlite import DreamLiteTransformer2DModel - from .transformer_allegro import AllegroTransformer3DModel - from .transformer_anyflow import AnyFlowTransformer3DModel - from .transformer_anyflow_far import AnyFlowFARTransformer3DModel - from .transformer_bria import BriaTransformer2DModel - from .transformer_bria_fibo import BriaFiboTransformer2DModel - from .transformer_chroma import ChromaTransformer2DModel - from .transformer_chronoedit import ChronoEditTransformer3DModel - from .transformer_cogview3plus import CogView3PlusTransformer2DModel - from .transformer_cogview4 import CogView4Transformer2DModel - from .transformer_cosmos import CosmosTransformer3DModel - from .transformer_cosmos3 import Cosmos3OmniTransformer - from .transformer_easyanimate import EasyAnimateTransformer3DModel - from .transformer_ernie_image import ErnieImageTransformer2DModel - from .transformer_flux import FluxTransformer2DModel - from .transformer_flux2 import Flux2Transformer2DModel - from .transformer_glm_image import GlmImageTransformer2DModel - from .transformer_helios import HeliosTransformer3DModel - from .transformer_hidream_image import HiDreamImageTransformer2DModel - from .transformer_hunyuan_video import HunyuanVideoTransformer3DModel - from .transformer_hunyuan_video15 import HunyuanVideo15Transformer3DModel - from .transformer_hunyuan_video_framepack import HunyuanVideoFramepackTransformer3DModel - from .transformer_hunyuanimage import HunyuanImageTransformer2DModel - from .transformer_ideogram4 import Ideogram4Transformer2DModel - from .transformer_joyimage import JoyImageEditTransformer3DModel - from .transformer_joyimage_edit_plus import JoyImageEditPlusTransformer3DModel - from .transformer_kandinsky import Kandinsky5Transformer3DModel - from .transformer_krea2 import Krea2Transformer2DModel - from .transformer_longcat_audio_dit import LongCatAudioDiTTransformer - from .transformer_longcat_image import LongCatImageTransformer2DModel - from .transformer_ltx import LTXVideoTransformer3DModel - from .transformer_ltx2 import LTX2VideoTransformer3DModel - from .transformer_lumina2 import Lumina2Transformer2DModel - from .transformer_minimax_h3 import MiniMaxH3Transformer3DModel - from .transformer_mochi import MochiTransformer3DModel - from .transformer_motif_video import MotifVideoTransformer3DModel - from .transformer_nucleusmoe_image import NucleusMoEImageTransformer2DModel - from .transformer_omnigen import OmniGenTransformer2DModel - from .transformer_ovis_image import OvisImageTransformer2DModel - from .transformer_prx import PRXTransformer2DModel - from .transformer_qwenimage import QwenImageTransformer2DModel - from .transformer_sana_video import SanaVideoTransformer3DModel - from .transformer_sd3 import SD3Transformer2DModel - from .transformer_skyreels_v2 import SkyReelsV2Transformer3DModel - from .transformer_temporal import TransformerTemporalModel - from .transformer_wan import WanTransformer3DModel - from .transformer_wan_animate import WanAnimateTransformer3DModel - from .transformer_wan_vace import WanVACETransformer3DModel - from .transformer_z_image import ZImageTransformer2DModel diff --git a/diffusers/models/transformers/ace_step_transformer.py b/diffusers/models/transformers/ace_step_transformer.py deleted file mode 100644 index 821c7ad1491a7042e6be302c45b538d09287abe8..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/ace_step_transformer.py +++ /dev/null @@ -1,632 +0,0 @@ -# Copyright 2025 The ACE-Step Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -"""Diffusion Transformer (DiT) for ACE-Step 1.5 music generation.""" - -import inspect -from typing import List, Optional, Tuple, Union - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ..attention import AttentionMixin, AttentionModuleMixin -from ..attention_dispatch import ( - AttentionBackendName, - _AttentionBackendRegistry, - dispatch_attention_fn, -) -from ..cache_utils import CacheMixin -from ..embeddings import Timesteps, apply_rotary_emb, get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -_FLASH_ATTENTION_BACKENDS = { - AttentionBackendName.FLASH, - AttentionBackendName.FLASH_HUB, - AttentionBackendName.FLASH_VARLEN, - AttentionBackendName.FLASH_VARLEN_HUB, -} - -_FLASH_ATTENTION_VARLEN_BACKENDS = { - AttentionBackendName.FLASH_VARLEN, - AttentionBackendName.FLASH_VARLEN_HUB, -} - - -def _get_current_attention_backend(processor: Optional["AceStepAttnProcessor2_0"] = None) -> AttentionBackendName: - backend = getattr(processor, "_attention_backend", None) - if backend is None: - backend, _ = _AttentionBackendRegistry.get_active_backend() - return AttentionBackendName(backend) - - -def _is_flash_attention_backend(processor: Optional["AceStepAttnProcessor2_0"] = None) -> bool: - return _get_current_attention_backend(processor) in _FLASH_ATTENTION_BACKENDS - - -# --------------------------------------------------------------------------- # -# attention-mask # -# --------------------------------------------------------------------------- # - - -def _create_4d_mask( - seq_len: int, - dtype: torch.dtype, - device: torch.device, - attention_mask: Optional[torch.Tensor] = None, - sliding_window: Optional[int] = None, - is_sliding_window: bool = False, - is_causal: bool = True, -) -> torch.Tensor: - """Build a `[B, 1, seq_len, seq_len]` additive mask (0.0 kept, -inf masked). - - Mirrors the mask construction in ``acestep/models/turbo/modeling_acestep_v15_turbo.py::create_4d_mask`` so the DiT - sees identical attention coverage regardless of whether SDPA, eager or flash attention is selected downstream. - """ - indices = torch.arange(seq_len, device=device) - diff = indices.unsqueeze(1) - indices.unsqueeze(0) - valid_mask = torch.ones((seq_len, seq_len), device=device, dtype=torch.bool) - - if is_causal: - valid_mask = valid_mask & (diff >= 0) - - if is_sliding_window and sliding_window is not None: - if is_causal: - valid_mask = valid_mask & (diff <= sliding_window) - else: - valid_mask = valid_mask & (torch.abs(diff) <= sliding_window) - - valid_mask = valid_mask.unsqueeze(0).unsqueeze(0) - - if attention_mask is not None: - padding_mask_4d = attention_mask.view(attention_mask.shape[0], 1, 1, seq_len).to(torch.bool) - valid_mask = valid_mask & padding_mask_4d - - min_dtype = torch.finfo(dtype).min - mask_tensor = torch.full(valid_mask.shape, min_dtype, dtype=dtype, device=device) - mask_tensor.masked_fill_(valid_mask, 0.0) - return mask_tensor - - -# --------------------------------------------------------------------------- # -# RoPE helpers # -# --------------------------------------------------------------------------- # - - -def _ace_step_rotary_freqs( - seq_len: int, head_dim: int, theta: float, device: torch.device, dtype: torch.dtype -) -> Tuple[torch.Tensor, torch.Tensor]: - """Build (cos, sin) freqs for ACE-Step RoPE using ``get_1d_rotary_pos_embed``. - - The original ACE-Step DiT reuses Qwen3's rotary layout: ``freqs = cat([freq_half, freq_half], dim=-1)`` (not - interleaved), and the rotate-half convention splits the last dim in two halves rather than unbinding pairs. That - matches ``get_1d_rotary_pos_embed(..., use_real=True, repeat_interleave_real=False)`` + ``apply_rotary_emb(..., - use_real_unbind_dim=-2)``. - """ - positions = torch.arange(seq_len, device=device, dtype=torch.float32) - cos, sin = get_1d_rotary_pos_embed(head_dim, positions, theta=theta, use_real=True, repeat_interleave_real=False) - return cos.to(dtype=dtype), sin.to(dtype=dtype) - - -# --------------------------------------------------------------------------- # -# building blocks # -# --------------------------------------------------------------------------- # - - -class AceStepMLP(nn.Module): - """SwiGLU MLP used in ACE-Step transformer blocks.""" - - def __init__(self, hidden_size: int, intermediate_size: int): - super().__init__() - self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False) - self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False) - self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) - - -class AceStepTimestepEmbedding(nn.Module): - """Sinusoidal timestep embedding + 2-layer MLP + 6-way AdaLN scale/shift projection. - - Matches the original ACE-Step checkpoint layout exactly (``linear_1``, ``linear_2``, ``time_proj``) so the - converter maps keys 1:1. The sinusoid itself is the shared ``Timesteps`` module (``flip_sin_to_cos=True`` for - ACE-Step's ``cat([cos, sin])`` convention). - """ - - def __init__(self, in_channels: int = 256, time_embed_dim: int = 2048, scale: float = 1000.0): - super().__init__() - self.in_channels = in_channels - self.scale = scale - self.time_sinusoid = Timesteps(num_channels=in_channels, flip_sin_to_cos=True, downscale_freq_shift=0) - self.linear_1 = nn.Linear(in_channels, time_embed_dim, bias=True) - self.act1 = nn.SiLU() - self.linear_2 = nn.Linear(time_embed_dim, time_embed_dim, bias=True) - self.act2 = nn.SiLU() - self.time_proj = nn.Linear(time_embed_dim, time_embed_dim * 6) - - def forward(self, t: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: - t_freq = self.time_sinusoid(t * self.scale) - temb = self.linear_1(t_freq.to(t.dtype)) - temb = self.act1(temb) - temb = self.linear_2(temb) - timestep_proj = self.time_proj(self.act2(temb)).unflatten(1, (6, -1)) - return temb, timestep_proj - - -class AceStepAttnProcessor2_0: - """Attention processor for ACE-Step GQA attention. - - Dispatches the actual attention call through ``dispatch_attention_fn`` so users can pick flash / sage / native - backends via ``model.set_attention_backend(...)`` or the ``attention_backend`` context manager. Uses the ``(B, L, - H, D)`` tensor layout that the diffusers attention backends consume directly. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("AceStepAttnProcessor2_0 requires PyTorch 2.0. Please upgrade your pytorch version.") - - def __call__( - self, - attn: "AceStepAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, - ) -> torch.Tensor: - is_cross = attn.is_cross_attention and encoder_hidden_states is not None - kv_input = encoder_hidden_states if is_cross else hidden_states - - # Project to (B, L, H, D). Q uses ``heads``; K/V use ``kv_heads`` (GQA). - query = attn.to_q(hidden_states).unflatten(-1, (attn.heads, attn.head_dim)) - key = attn.to_k(kv_input).unflatten(-1, (attn.kv_heads, attn.head_dim)) - value = attn.to_v(kv_input).unflatten(-1, (attn.kv_heads, attn.head_dim)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - # RoPE on self-attention only. Matches Qwen3 layout: - # freqs = cat([freq_half, freq_half], dim=-1); rotate-half splits last dim. - if not is_cross and image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, use_real=True, use_real_unbind_dim=-2, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, use_real=True, use_real_unbind_dim=-2, sequence_dim=1) - - attention_kwargs = None - backend = _get_current_attention_backend(self) - dispatch_backend = self._attention_backend - sliding_window = getattr(attn, "sliding_window", None) - - if backend in _FLASH_ATTENTION_BACKENDS: - if attention_mask is not None: - if attention_mask.ndim == 2: - padding_mask = attention_mask.to(torch.bool) - elif attention_mask.ndim == 4: - keep_mask = attention_mask if attention_mask.dtype == torch.bool else attention_mask == 0 - padding_mask = keep_mask.any(dim=(1, 2)) - else: - raise ValueError( - f"Unsupported ACE-Step attention mask shape for flash attention: {attention_mask.shape}" - ) - - has_padding = not torch.all(padding_mask).item() - if has_padding: - attention_mask = padding_mask - if backend not in _FLASH_ATTENTION_VARLEN_BACKENDS: - raise ValueError( - "ACE-Step flash attention received a padded attention mask. Use `flash_varlen` or " - "`flash_varlen_hub` for batched prompts with padding, or use an unpadded batch with `flash`." - ) - else: - attention_mask = None - - if not is_cross and sliding_window is not None and key.shape[1] > sliding_window: - # ACE-Step's dense mask keeps `abs(i - j) <= sliding_window`; flash-attn uses the same inclusive - # left/right window convention, so pass the configured value through directly. - attention_kwargs = {"window_size": (sliding_window, sliding_window)} - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=attn.dropout if attn.training else 0.0, - scale=attn.scaling, - enable_gqa=attn.heads != attn.kv_heads, - attention_kwargs=attention_kwargs, - backend=dispatch_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3).to(query.dtype) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class AceStepAttention(torch.nn.Module, AttentionModuleMixin): - """GQA attention with RMSNorm on query/key for ACE-Step 1.5. - - Uses the diffusers ``Attention`` + ``AttnProcessor`` split: this module holds the projections and Q/K norm; the - processor runs the attention dispatch. Self-attention applies RoPE on query/key; cross-attention reads K/V from - ``encoder_hidden_states`` and does not apply RoPE. - - GQA means Q has ``heads * head_dim`` output while K/V have ``kv_heads * head_dim`` — QKV fusion is therefore - disabled (``_supports_qkv_fusion = False``). - """ - - _default_processor_cls = AceStepAttnProcessor2_0 - _available_processors = [AceStepAttnProcessor2_0] - _supports_qkv_fusion = False - - def __init__( - self, - hidden_size: int, - num_attention_heads: int, - num_key_value_heads: int, - head_dim: int, - bias: bool = False, - dropout: float = 0.0, - eps: float = 1e-6, - sliding_window: Optional[int] = None, - is_cross_attention: bool = False, - processor: Optional[AceStepAttnProcessor2_0] = None, - ): - super().__init__() - self.heads = num_attention_heads - self.kv_heads = num_key_value_heads - self.head_dim = head_dim - self.dropout = dropout - self.scaling = head_dim**-0.5 - self.sliding_window = sliding_window - self.is_cross_attention = is_cross_attention - - self.to_q = nn.Linear(hidden_size, num_attention_heads * head_dim, bias=bias) - self.to_k = nn.Linear(hidden_size, num_key_value_heads * head_dim, bias=bias) - self.to_v = nn.Linear(hidden_size, num_key_value_heads * head_dim, bias=bias) - self.to_out = nn.ModuleList( - [nn.Linear(num_attention_heads * head_dim, hidden_size, bias=bias), nn.Dropout(0.0)] - ) - self.norm_q = RMSNorm(head_dim, eps=eps) - self.norm_k = RMSNorm(head_dim, eps=eps) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - kwargs = {k: v for k, v in kwargs.items() if k in attn_parameters} - return self.processor( - self, - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - **kwargs, - ) - - -class AceStepTransformerBlock(nn.Module): - """ACE-Step DiT transformer block: self-attn (AdaLN) → cross-attn → MLP (AdaLN). - - AdaLN parameters come from the shared ``scale_shift_table + timestep_proj`` chunked into 6 (3 for self-attn + 3 for - MLP). - """ - - def __init__( - self, - hidden_size: int, - num_attention_heads: int, - num_key_value_heads: int, - head_dim: int, - intermediate_size: int, - attention_bias: bool = False, - attention_dropout: float = 0.0, - rms_norm_eps: float = 1e-6, - sliding_window: Optional[int] = None, - use_cross_attention: bool = True, - ): - super().__init__() - self.self_attn_norm = RMSNorm(hidden_size, eps=rms_norm_eps) - self.self_attn = AceStepAttention( - hidden_size=hidden_size, - num_attention_heads=num_attention_heads, - num_key_value_heads=num_key_value_heads, - head_dim=head_dim, - bias=attention_bias, - dropout=attention_dropout, - eps=rms_norm_eps, - sliding_window=sliding_window, - is_cross_attention=False, - ) - - self.use_cross_attention = use_cross_attention - if self.use_cross_attention: - self.cross_attn_norm = RMSNorm(hidden_size, eps=rms_norm_eps) - self.cross_attn = AceStepAttention( - hidden_size=hidden_size, - num_attention_heads=num_attention_heads, - num_key_value_heads=num_key_value_heads, - head_dim=head_dim, - bias=attention_bias, - dropout=attention_dropout, - eps=rms_norm_eps, - is_cross_attention=True, - ) - - self.mlp_norm = RMSNorm(hidden_size, eps=rms_norm_eps) - self.mlp = AceStepMLP(hidden_size, intermediate_size) - - self.scale_shift_table = nn.Parameter(torch.randn(1, 6, hidden_size) / hidden_size**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - position_embeddings: Tuple[torch.Tensor, torch.Tensor], - temb: torch.Tensor, - attention_mask: Optional[torch.Tensor] = None, - encoder_hidden_states: Optional[torch.Tensor] = None, - encoder_attention_mask: Optional[torch.Tensor] = None, - ) -> torch.Tensor: - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (self.scale_shift_table + temb).chunk( - 6, dim=1 - ) - - # Self-attention with AdaLN. - norm_hidden_states = (self.self_attn_norm(hidden_states) * (1 + scale_msa) + shift_msa).type_as(hidden_states) - attn_output = self.self_attn( - hidden_states=norm_hidden_states, - image_rotary_emb=position_embeddings, - attention_mask=attention_mask, - ) - hidden_states = (hidden_states + attn_output * gate_msa).type_as(hidden_states) - - if self.use_cross_attention and encoder_hidden_states is not None: - norm_hidden_states = self.cross_attn_norm(hidden_states).type_as(hidden_states) - attn_output = self.cross_attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=encoder_attention_mask, - ) - hidden_states = hidden_states + attn_output - - norm_hidden_states = (self.mlp_norm(hidden_states) * (1 + c_scale_msa) + c_shift_msa).type_as(hidden_states) - ff_output = self.mlp(norm_hidden_states) - hidden_states = (hidden_states + ff_output * c_gate_msa).type_as(hidden_states) - return hidden_states - - -# --------------------------------------------------------------------------- # -# main DiT model # -# --------------------------------------------------------------------------- # - - -class AceStepTransformer1DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, AttentionMixin, CacheMixin): - """Diffusion Transformer for ACE-Step 1.5 music generation. - - Generates audio latents conditioned on text, lyrics, and timbre. Uses 1D patch embedding (`Conv1d` with stride - `patch_size`) followed by a stack of `AceStepTransformerBlock`s with alternating sliding-window / full attention on - the self-attention branch. Cross-attention consumes the packed `encoder_hidden_states` produced by - `AceStepConditionEncoder`. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - hidden_size: int = 2048, - intermediate_size: int = 6144, - num_hidden_layers: int = 24, - num_attention_heads: int = 16, - num_key_value_heads: int = 8, - head_dim: int = 128, - in_channels: int = 192, - audio_acoustic_hidden_dim: int = 64, - patch_size: int = 2, - rope_theta: float = 1000000.0, - attention_bias: bool = False, - attention_dropout: float = 0.0, - rms_norm_eps: float = 1e-6, - sliding_window: int = 128, - layer_types: Optional[List[str]] = None, - # Dim of the condition encoder's output. Equal to `hidden_size` on the - # non-XL turbo / base models, but the XL turbo has a smaller condition - # encoder (`encoder_hidden_size=2048`) feeding a wider DiT - # (`hidden_size=2560`), so `condition_embedder` needs to project it up. - encoder_hidden_size: Optional[int] = None, - # Variant metadata. Turbo models have guidance distilled into the weights and - # should run without CFG; base/SFT models require CFG with the learned - # `AceStepConditionEncoder.null_condition_emb`. The pipeline reads these to - # pick default `guidance_scale`, `shift`, and `num_inference_steps`. - is_turbo: bool = False, - model_version: Optional[str] = None, - ): - super().__init__() - if encoder_hidden_size is None: - encoder_hidden_size = hidden_size - self.patch_size = patch_size - self.head_dim = head_dim - self.rope_theta = rope_theta - - if layer_types is None: - layer_types = [ - "sliding_attention" if bool((i + 1) % 2) else "full_attention" for i in range(num_hidden_layers) - ] - self.layer_types = list(layer_types) - - self.layers = nn.ModuleList( - [ - AceStepTransformerBlock( - hidden_size=hidden_size, - num_attention_heads=num_attention_heads, - num_key_value_heads=num_key_value_heads, - head_dim=head_dim, - intermediate_size=intermediate_size, - attention_bias=attention_bias, - attention_dropout=attention_dropout, - rms_norm_eps=rms_norm_eps, - sliding_window=sliding_window if layer_types[i] == "sliding_attention" else None, - use_cross_attention=True, - ) - for i in range(num_hidden_layers) - ] - ) - - # Patchify: concat(src_latents, chunk_mask) on the channel dim then Conv1d with - # stride=patch_size lifts (B, T, in_channels) -> (B, T/patch_size, hidden_size). - self.proj_in_conv = nn.Conv1d( - in_channels=in_channels, - out_channels=hidden_size, - kernel_size=patch_size, - stride=patch_size, - padding=0, - ) - - # Dual-timestep conditioning: one path for `t`, one for `(t - r)` (mean-flow). - self.time_embed = AceStepTimestepEmbedding(in_channels=256, time_embed_dim=hidden_size) - self.time_embed_r = AceStepTimestepEmbedding(in_channels=256, time_embed_dim=hidden_size) - - self.condition_embedder = nn.Linear(encoder_hidden_size, hidden_size, bias=True) - - self.norm_out = RMSNorm(hidden_size, eps=rms_norm_eps) - self.proj_out_conv = nn.ConvTranspose1d( - in_channels=hidden_size, - out_channels=audio_acoustic_hidden_dim, - kernel_size=patch_size, - stride=patch_size, - padding=0, - ) - self.scale_shift_table = nn.Parameter(torch.randn(1, 2, hidden_size) / hidden_size**0.5) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - timestep_r: torch.Tensor, - encoder_hidden_states: torch.Tensor, - context_latents: torch.Tensor, - attention_kwargs: Optional[dict] = None, - return_dict: bool = True, - ) -> Union[torch.Tensor, Transformer2DModelOutput]: - """The [`AceStepTransformer1DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, seq_len, channels)`): - Noisy latent input for the diffusion process. - timestep (`torch.Tensor` of shape `(batch_size,)`): - Current diffusion timestep `t`. - timestep_r (`torch.Tensor` of shape `(batch_size,)`): - Reference timestep `r` (set equal to `t` for standard inference). - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, encoder_seq_len, hidden_size)`): - Conditioning embeddings from the condition encoder (text + lyrics + timbre). - context_latents (`torch.Tensor` of shape `(batch_size, seq_len, context_dim)`): - Context latents (source latents concatenated with chunk masks) — fed to the patchify conv alongside - `hidden_states`. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary passed along to the `AttentionProcessor`. Used to pass the LoRA scale via - `{"scale": float}`. - return_dict (`bool`, defaults to `True`): - Whether to return a `Transformer2DModelOutput` or a plain tuple. - - Returns: - `Transformer2DModelOutput` or `tuple`: The predicted velocity field. - """ - # Dual timestep embedding: t and (t - r). Sum both paths' AdaLN projections. - temb_t, timestep_proj_t = self.time_embed(timestep) - temb_r, timestep_proj_r = self.time_embed_r(timestep - timestep_r) - temb = temb_t + temb_r - timestep_proj = timestep_proj_t + timestep_proj_r - - # Context concatenation + padding to patch_size boundary + patchify. - hidden_states = torch.cat([context_latents, hidden_states], dim=-1) - original_seq_len = hidden_states.shape[1] - if hidden_states.shape[1] % self.patch_size != 0: - pad_length = self.patch_size - (hidden_states.shape[1] % self.patch_size) - hidden_states = F.pad(hidden_states, (0, 0, 0, pad_length), mode="constant", value=0) - hidden_states = self.proj_in_conv(hidden_states.transpose(1, 2)).transpose(1, 2) - encoder_hidden_states = self.condition_embedder(encoder_hidden_states) - - seq_len = hidden_states.shape[1] - dtype = hidden_states.dtype - device = hidden_states.device - - cos, sin = _ace_step_rotary_freqs(seq_len, self.head_dim, self.rope_theta, device, dtype) - position_embeddings = (cos, sin) - - sliding_attn_mask = None - if not _is_flash_attention_backend(self.layers[0].self_attn.processor): - sliding_attn_mask = _create_4d_mask( - seq_len=seq_len, - dtype=dtype, - device=device, - sliding_window=self.config.sliding_window, - is_sliding_window=True, - is_causal=False, - ) - - for i, layer_module in enumerate(self.layers): - # Full-attention layers see no mask; only the sliding-attention layers - # need the banded mask. Cross-attention uses no padding mask. - layer_attn_mask = sliding_attn_mask if self.layer_types[i] == "sliding_attention" else None - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - layer_module, - hidden_states, - position_embeddings, - timestep_proj, - layer_attn_mask, - encoder_hidden_states, - None, - ) - else: - hidden_states = layer_module( - hidden_states=hidden_states, - position_embeddings=position_embeddings, - temb=timestep_proj, - attention_mask=layer_attn_mask, - encoder_hidden_states=encoder_hidden_states, - encoder_attention_mask=None, - ) - - # Adaptive output normalization + de-patchify. - shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1) - hidden_states = (self.norm_out(hidden_states) * (1 + scale) + shift).type_as(hidden_states) - hidden_states = self.proj_out_conv(hidden_states.transpose(1, 2)).transpose(1, 2) - hidden_states = hidden_states[:, :original_seq_len, :] - - if not return_dict: - return (hidden_states,) - return Transformer2DModelOutput(sample=hidden_states) diff --git a/diffusers/models/transformers/auraflow_transformer_2d.py b/diffusers/models/transformers/auraflow_transformer_2d.py deleted file mode 100644 index ff6c0c78a53b5262a37a8ab4d30268f69838a5b4..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/auraflow_transformer_2d.py +++ /dev/null @@ -1,500 +0,0 @@ -# Copyright 2025 AuraFlow Authors, The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import AttentionMixin -from ..attention_processor import ( - Attention, - AuraFlowAttnProcessor2_0, - FusedAuraFlowAttnProcessor2_0, -) -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormZero, FP32LayerNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# Taken from the original aura flow inference code. -def find_multiple(n: int, k: int) -> int: - if n % k == 0: - return n - return n + k - (n % k) - - -# Aura Flow patch embed doesn't use convs for projections. -# Additionally, it uses learned positional embeddings. -class AuraFlowPatchEmbed(nn.Module): - def __init__( - self, - height=224, - width=224, - patch_size=16, - in_channels=3, - embed_dim=768, - pos_embed_max_size=None, - ): - super().__init__() - - self.num_patches = (height // patch_size) * (width // patch_size) - self.pos_embed_max_size = pos_embed_max_size - - self.proj = nn.Linear(patch_size * patch_size * in_channels, embed_dim) - self.pos_embed = nn.Parameter(torch.randn(1, pos_embed_max_size, embed_dim) * 0.1) - - self.patch_size = patch_size - self.height, self.width = height // patch_size, width // patch_size - self.base_size = height // patch_size - - def pe_selection_index_based_on_dim(self, h, w): - # select subset of positional embedding based on H, W, where H, W is size of latent - # PE will be viewed as 2d-grid, and H/p x W/p of the PE will be selected - # because original input are in flattened format, we have to flatten this 2d grid as well. - h_p, w_p = h // self.patch_size, w // self.patch_size - h_max, w_max = int(self.pos_embed_max_size**0.5), int(self.pos_embed_max_size**0.5) - - # Calculate the top-left corner indices for the centered patch grid - starth = h_max // 2 - h_p // 2 - startw = w_max // 2 - w_p // 2 - - # Generate the row and column indices for the desired patch grid - rows = torch.arange(starth, starth + h_p, device=self.pos_embed.device) - cols = torch.arange(startw, startw + w_p, device=self.pos_embed.device) - - # Create a 2D grid of indices - row_indices, col_indices = torch.meshgrid(rows, cols, indexing="ij") - - # Convert the 2D grid indices to flattened 1D indices - selected_indices = (row_indices * w_max + col_indices).flatten() - - return selected_indices - - def forward(self, latent) -> torch.Tensor: - batch_size, num_channels, height, width = latent.size() - latent = latent.view( - batch_size, - num_channels, - height // self.patch_size, - self.patch_size, - width // self.patch_size, - self.patch_size, - ) - latent = latent.permute(0, 2, 4, 1, 3, 5).flatten(-3).flatten(1, 2) - latent = self.proj(latent) - pe_index = self.pe_selection_index_based_on_dim(height, width) - return latent + self.pos_embed[:, pe_index] - - -# Taken from the original Aura flow inference code. -# Our feedforward only has GELU but Aura uses SiLU. -class AuraFlowFeedForward(nn.Module): - def __init__(self, dim, hidden_dim=None) -> None: - super().__init__() - if hidden_dim is None: - hidden_dim = 4 * dim - - final_hidden_dim = int(2 * hidden_dim / 3) - final_hidden_dim = find_multiple(final_hidden_dim, 256) - - self.linear_1 = nn.Linear(dim, final_hidden_dim, bias=False) - self.linear_2 = nn.Linear(dim, final_hidden_dim, bias=False) - self.out_projection = nn.Linear(final_hidden_dim, dim, bias=False) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = F.silu(self.linear_1(x)) * self.linear_2(x) - x = self.out_projection(x) - return x - - -class AuraFlowPreFinalBlock(nn.Module): - def __init__(self, embedding_dim: int, conditioning_embedding_dim: int): - super().__init__() - - self.silu = nn.SiLU() - self.linear = nn.Linear(conditioning_embedding_dim, embedding_dim * 2, bias=False) - - def forward(self, x: torch.Tensor, conditioning_embedding: torch.Tensor) -> torch.Tensor: - emb = self.linear(self.silu(conditioning_embedding).to(x.dtype)) - scale, shift = torch.chunk(emb, 2, dim=1) - x = x * (1 + scale)[:, None, :] + shift[:, None, :] - return x - - -@maybe_allow_in_graph -class AuraFlowSingleTransformerBlock(nn.Module): - """Similar to `AuraFlowJointTransformerBlock` with a single DiT instead of an MMDiT.""" - - def __init__(self, dim, num_attention_heads, attention_head_dim): - super().__init__() - - self.norm1 = AdaLayerNormZero(dim, bias=False, norm_type="fp32_layer_norm") - - processor = AuraFlowAttnProcessor2_0() - self.attn = Attention( - query_dim=dim, - cross_attention_dim=None, - dim_head=attention_head_dim, - heads=num_attention_heads, - qk_norm="fp32_layer_norm", - out_dim=dim, - bias=False, - out_bias=False, - processor=processor, - ) - - self.norm2 = FP32LayerNorm(dim, elementwise_affine=False, bias=False) - self.ff = AuraFlowFeedForward(dim, dim * 4) - - def forward( - self, - hidden_states: torch.FloatTensor, - temb: torch.FloatTensor, - attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor: - residual = hidden_states - attention_kwargs = attention_kwargs or {} - - # Norm + Projection. - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - - # Attention. - attn_output = self.attn(hidden_states=norm_hidden_states, **attention_kwargs) - - # Process attention outputs for the `hidden_states`. - hidden_states = self.norm2(residual + gate_msa.unsqueeze(1) * attn_output) - hidden_states = hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - ff_output = self.ff(hidden_states) - hidden_states = gate_mlp.unsqueeze(1) * ff_output - hidden_states = residual + hidden_states - - return hidden_states - - -@maybe_allow_in_graph -class AuraFlowJointTransformerBlock(nn.Module): - r""" - Transformer block for Aura Flow. Similar to SD3 MMDiT. Differences (non-exhaustive): - - * QK Norm in the attention blocks - * No bias in the attention blocks - * Most LayerNorms are in FP32 - - Parameters: - dim (`int`): The number of channels in the input and output. - num_attention_heads (`int`): The number of heads to use for multi-head attention. - attention_head_dim (`int`): The number of channels in each head. - is_last (`bool`): Boolean to determine if this is the last block in the model. - """ - - def __init__(self, dim, num_attention_heads, attention_head_dim): - super().__init__() - - self.norm1 = AdaLayerNormZero(dim, bias=False, norm_type="fp32_layer_norm") - self.norm1_context = AdaLayerNormZero(dim, bias=False, norm_type="fp32_layer_norm") - - processor = AuraFlowAttnProcessor2_0() - self.attn = Attention( - query_dim=dim, - cross_attention_dim=None, - added_kv_proj_dim=dim, - added_proj_bias=False, - dim_head=attention_head_dim, - heads=num_attention_heads, - qk_norm="fp32_layer_norm", - out_dim=dim, - bias=False, - out_bias=False, - processor=processor, - context_pre_only=False, - ) - - self.norm2 = FP32LayerNorm(dim, elementwise_affine=False, bias=False) - self.ff = AuraFlowFeedForward(dim, dim * 4) - self.norm2_context = FP32LayerNorm(dim, elementwise_affine=False, bias=False) - self.ff_context = AuraFlowFeedForward(dim, dim * 4) - - def forward( - self, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor, - temb: torch.FloatTensor, - attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - residual = hidden_states - residual_context = encoder_hidden_states - attention_kwargs = attention_kwargs or {} - - # Norm + Projection. - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - - # Attention. - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - **attention_kwargs, - ) - - # Process attention outputs for the `hidden_states`. - hidden_states = self.norm2(residual + gate_msa.unsqueeze(1) * attn_output) - hidden_states = hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - hidden_states = gate_mlp.unsqueeze(1) * self.ff(hidden_states) - hidden_states = residual + hidden_states - - # Process attention outputs for the `encoder_hidden_states`. - encoder_hidden_states = self.norm2_context(residual_context + c_gate_msa.unsqueeze(1) * context_attn_output) - encoder_hidden_states = encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - encoder_hidden_states = c_gate_mlp.unsqueeze(1) * self.ff_context(encoder_hidden_states) - encoder_hidden_states = residual_context + encoder_hidden_states - - return encoder_hidden_states, hidden_states - - -class AuraFlowTransformer2DModel(ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): - r""" - A 2D Transformer model as introduced in AuraFlow (https://blog.fal.ai/auraflow/). - - Parameters: - sample_size (`int`): The width of the latent images. This is fixed during training since - it is used to learn a number of position embeddings. - patch_size (`int`): Patch size to turn the input data into small patches. - in_channels (`int`, *optional*, defaults to 4): The number of channels in the input. - num_mmdit_layers (`int`, *optional*, defaults to 4): The number of layers of MMDiT Transformer blocks to use. - num_single_dit_layers (`int`, *optional*, defaults to 32): - The number of layers of Transformer blocks to use. These blocks use concatenated image and text - representations. - attention_head_dim (`int`, *optional*, defaults to 256): The number of channels in each head. - num_attention_heads (`int`, *optional*, defaults to 12): The number of heads to use for multi-head attention. - joint_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use. - caption_projection_dim (`int`): Number of dimensions to use when projecting the `encoder_hidden_states`. - out_channels (`int`, defaults to 4): Number of output channels. - pos_embed_max_size (`int`, defaults to 1024): Maximum positions to embed from the image latents. - """ - - _no_split_modules = ["AuraFlowJointTransformerBlock", "AuraFlowSingleTransformerBlock", "AuraFlowPatchEmbed"] - _skip_layerwise_casting_patterns = ["pos_embed", "norm"] - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - sample_size: int = 64, - patch_size: int = 2, - in_channels: int = 4, - num_mmdit_layers: int = 4, - num_single_dit_layers: int = 32, - attention_head_dim: int = 256, - num_attention_heads: int = 12, - joint_attention_dim: int = 2048, - caption_projection_dim: int = 3072, - out_channels: int = 4, - pos_embed_max_size: int = 1024, - ): - super().__init__() - default_out_channels = in_channels - self.out_channels = out_channels if out_channels is not None else default_out_channels - self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim - - self.pos_embed = AuraFlowPatchEmbed( - height=self.config.sample_size, - width=self.config.sample_size, - patch_size=self.config.patch_size, - in_channels=self.config.in_channels, - embed_dim=self.inner_dim, - pos_embed_max_size=pos_embed_max_size, - ) - - self.context_embedder = nn.Linear( - self.config.joint_attention_dim, self.config.caption_projection_dim, bias=False - ) - self.time_step_embed = Timesteps(num_channels=256, downscale_freq_shift=0, scale=1000, flip_sin_to_cos=True) - self.time_step_proj = TimestepEmbedding(in_channels=256, time_embed_dim=self.inner_dim) - - self.joint_transformer_blocks = nn.ModuleList( - [ - AuraFlowJointTransformerBlock( - dim=self.inner_dim, - num_attention_heads=self.config.num_attention_heads, - attention_head_dim=self.config.attention_head_dim, - ) - for i in range(self.config.num_mmdit_layers) - ] - ) - self.single_transformer_blocks = nn.ModuleList( - [ - AuraFlowSingleTransformerBlock( - dim=self.inner_dim, - num_attention_heads=self.config.num_attention_heads, - attention_head_dim=self.config.attention_head_dim, - ) - for _ in range(self.config.num_single_dit_layers) - ] - ) - - self.norm_out = AuraFlowPreFinalBlock(self.inner_dim, self.inner_dim) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=False) - - # https://huggingface.co/papers/2309.16588 - # prevents artifacts in the attention maps - self.register_tokens = nn.Parameter(torch.randn(1, 8, self.inner_dim) * 0.02) - - self.gradient_checkpointing = False - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections with FusedAttnProcessor2_0->FusedAuraFlowAttnProcessor2_0 - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedAuraFlowAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor = None, - timestep: torch.LongTensor = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`AuraFlowTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.FloatTensor` of shape `(batch size, channel, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.FloatTensor` of shape `(batch size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - height, width = hidden_states.shape[-2:] - - # Apply patch embedding, timestep embedding, and project the caption embeddings. - hidden_states = self.pos_embed(hidden_states) # takes care of adding positional embeddings too. - temb = self.time_step_embed(timestep).to(dtype=next(self.parameters()).dtype) - temb = self.time_step_proj(temb) - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - encoder_hidden_states = torch.cat( - [self.register_tokens.repeat(encoder_hidden_states.size(0), 1, 1), encoder_hidden_states], dim=1 - ) - - # MMDiT blocks. - for index_block, block in enumerate(self.joint_transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - attention_kwargs=attention_kwargs, - ) - - # Single DiT blocks that combine the `hidden_states` (image) and `encoder_hidden_states` (text) - if len(self.single_transformer_blocks) > 0: - encoder_seq_len = encoder_hidden_states.size(1) - combined_hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - for index_block, block in enumerate(self.single_transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - combined_hidden_states = self._gradient_checkpointing_func( - block, - combined_hidden_states, - temb, - ) - - else: - combined_hidden_states = block( - hidden_states=combined_hidden_states, temb=temb, attention_kwargs=attention_kwargs - ) - - hidden_states = combined_hidden_states[:, encoder_seq_len:] - - hidden_states = self.norm_out(hidden_states, temb) - hidden_states = self.proj_out(hidden_states) - - # unpatchify - patch_size = self.config.patch_size - out_channels = self.config.out_channels - height = height // patch_size - width = width // patch_size - - hidden_states = hidden_states.reshape( - shape=(hidden_states.shape[0], height, width, patch_size, patch_size, out_channels) - ) - hidden_states = torch.einsum("nhwpqc->nchpwq", hidden_states) - output = hidden_states.reshape( - shape=(hidden_states.shape[0], out_channels, height * patch_size, width * patch_size) - ) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/cogvideox_transformer_3d.py b/diffusers/models/transformers/cogvideox_transformer_3d.py deleted file mode 100644 index 08299f05e1b80f22000c705a8bd538e534cf9561..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/cogvideox_transformer_3d.py +++ /dev/null @@ -1,474 +0,0 @@ -# Copyright 2025 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import Attention, AttentionMixin, FeedForward -from ..attention_processor import CogVideoXAttnProcessor2_0, FusedCogVideoXAttnProcessor2_0 -from ..cache_utils import CacheMixin -from ..embeddings import CogVideoXPatchEmbed, TimestepEmbedding, Timesteps -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNorm, CogVideoXLayerNormZero - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@maybe_allow_in_graph -class CogVideoXBlock(nn.Module): - r""" - Transformer block used in [CogVideoX](https://github.com/THUDM/CogVideo) model. - - Parameters: - dim (`int`): - The number of channels in the input and output. - num_attention_heads (`int`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`): - The number of channels in each head. - time_embed_dim (`int`): - The number of channels in timestep embedding. - dropout (`float`, defaults to `0.0`): - The dropout probability to use. - activation_fn (`str`, defaults to `"gelu-approximate"`): - Activation function to be used in feed-forward. - attention_bias (`bool`, defaults to `False`): - Whether or not to use bias in attention projection layers. - qk_norm (`bool`, defaults to `True`): - Whether or not to use normalization after query and key projections in Attention. - norm_elementwise_affine (`bool`, defaults to `True`): - Whether to use learnable elementwise affine parameters for normalization. - norm_eps (`float`, defaults to `1e-5`): - Epsilon value for normalization layers. - final_dropout (`bool` defaults to `False`): - Whether to apply a final dropout after the last feed-forward layer. - ff_inner_dim (`int`, *optional*, defaults to `None`): - Custom hidden dimension of Feed-forward layer. If not provided, `4 * dim` is used. - ff_bias (`bool`, defaults to `True`): - Whether or not to use bias in Feed-forward layer. - attention_out_bias (`bool`, defaults to `True`): - Whether or not to use bias in Attention output projection layer. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - time_embed_dim: int, - dropout: float = 0.0, - activation_fn: str = "gelu-approximate", - attention_bias: bool = False, - qk_norm: bool = True, - norm_elementwise_affine: bool = True, - norm_eps: float = 1e-5, - final_dropout: bool = True, - ff_inner_dim: int | None = None, - ff_bias: bool = True, - attention_out_bias: bool = True, - ): - super().__init__() - - # 1. Self Attention - self.norm1 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True) - - self.attn1 = Attention( - query_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - qk_norm="layer_norm" if qk_norm else None, - eps=1e-6, - bias=attention_bias, - out_bias=attention_out_bias, - processor=CogVideoXAttnProcessor2_0(), - ) - - # 2. Feed Forward - self.norm2 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True) - - self.ff = FeedForward( - dim, - dropout=dropout, - activation_fn=activation_fn, - final_dropout=final_dropout, - inner_dim=ff_inner_dim, - bias=ff_bias, - ) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - text_seq_length = encoder_hidden_states.size(1) - attention_kwargs = attention_kwargs or {} - - # norm & modulate - norm_hidden_states, norm_encoder_hidden_states, gate_msa, enc_gate_msa = self.norm1( - hidden_states, encoder_hidden_states, temb - ) - - # attention - attn_hidden_states, attn_encoder_hidden_states = self.attn1( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - **attention_kwargs, - ) - - hidden_states = hidden_states + gate_msa * attn_hidden_states - encoder_hidden_states = encoder_hidden_states + enc_gate_msa * attn_encoder_hidden_states - - # norm & modulate - norm_hidden_states, norm_encoder_hidden_states, gate_ff, enc_gate_ff = self.norm2( - hidden_states, encoder_hidden_states, temb - ) - - # feed-forward - norm_hidden_states = torch.cat([norm_encoder_hidden_states, norm_hidden_states], dim=1) - ff_output = self.ff(norm_hidden_states) - - hidden_states = hidden_states + gate_ff * ff_output[:, text_seq_length:] - encoder_hidden_states = encoder_hidden_states + enc_gate_ff * ff_output[:, :text_seq_length] - - return hidden_states, encoder_hidden_states - - -class CogVideoXTransformer3DModel(ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin, CacheMixin): - """ - A Transformer model for video-like data in [CogVideoX](https://github.com/THUDM/CogVideo). - - Parameters: - num_attention_heads (`int`, defaults to `30`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `64`): - The number of channels in each head. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `16`): - The number of channels in the output. - flip_sin_to_cos (`bool`, defaults to `True`): - Whether to flip the sin to cos in the time embedding. - time_embed_dim (`int`, defaults to `512`): - Output dimension of timestep embeddings. - ofs_embed_dim (`int`, defaults to `512`): - Output dimension of "ofs" embeddings used in CogVideoX-5b-I2B in version 1.5 - text_embed_dim (`int`, defaults to `4096`): - Input dimension of text embeddings from the text encoder. - num_layers (`int`, defaults to `30`): - The number of layers of Transformer blocks to use. - dropout (`float`, defaults to `0.0`): - The dropout probability to use. - attention_bias (`bool`, defaults to `True`): - Whether to use bias in the attention projection layers. - sample_width (`int`, defaults to `90`): - The width of the input latents. - sample_height (`int`, defaults to `60`): - The height of the input latents. - sample_frames (`int`, defaults to `49`): - The number of frames in the input latents. Note that this parameter was incorrectly initialized to 49 - instead of 13 because CogVideoX processed 13 latent frames at once in its default and recommended settings, - but cannot be changed to the correct value to ensure backwards compatibility. To create a transformer with - K latent frames, the correct value to pass here would be: ((K - 1) * temporal_compression_ratio + 1). - patch_size (`int`, defaults to `2`): - The size of the patches to use in the patch embedding layer. - temporal_compression_ratio (`int`, defaults to `4`): - The compression ratio across the temporal dimension. See documentation for `sample_frames`. - max_text_seq_length (`int`, defaults to `226`): - The maximum sequence length of the input text embeddings. - activation_fn (`str`, defaults to `"gelu-approximate"`): - Activation function to use in feed-forward. - timestep_activation_fn (`str`, defaults to `"silu"`): - Activation function to use when generating the timestep embeddings. - norm_elementwise_affine (`bool`, defaults to `True`): - Whether to use elementwise affine in normalization layers. - norm_eps (`float`, defaults to `1e-5`): - The epsilon value to use in normalization layers. - spatial_interpolation_scale (`float`, defaults to `1.875`): - Scaling factor to apply in 3D positional embeddings across spatial dimensions. - temporal_interpolation_scale (`float`, defaults to `1.0`): - Scaling factor to apply in 3D positional embeddings across temporal dimensions. - """ - - _skip_layerwise_casting_patterns = ["patch_embed", "norm"] - _supports_gradient_checkpointing = True - _no_split_modules = ["CogVideoXBlock", "CogVideoXPatchEmbed"] - - @register_to_config - def __init__( - self, - num_attention_heads: int = 30, - attention_head_dim: int = 64, - in_channels: int = 16, - out_channels: int | None = 16, - flip_sin_to_cos: bool = True, - freq_shift: int = 0, - time_embed_dim: int = 512, - ofs_embed_dim: int | None = None, - text_embed_dim: int = 4096, - num_layers: int = 30, - dropout: float = 0.0, - attention_bias: bool = True, - sample_width: int = 90, - sample_height: int = 60, - sample_frames: int = 49, - patch_size: int = 2, - patch_size_t: int | None = None, - temporal_compression_ratio: int = 4, - max_text_seq_length: int = 226, - activation_fn: str = "gelu-approximate", - timestep_activation_fn: str = "silu", - norm_elementwise_affine: bool = True, - norm_eps: float = 1e-5, - spatial_interpolation_scale: float = 1.875, - temporal_interpolation_scale: float = 1.0, - use_rotary_positional_embeddings: bool = False, - use_learned_positional_embeddings: bool = False, - patch_bias: bool = True, - ): - super().__init__() - inner_dim = num_attention_heads * attention_head_dim - - if not use_rotary_positional_embeddings and use_learned_positional_embeddings: - raise ValueError( - "There are no CogVideoX checkpoints available with disable rotary embeddings and learned positional " - "embeddings. If you're using a custom model and/or believe this should be supported, please open an " - "issue at https://github.com/huggingface/diffusers/issues." - ) - - # 1. Patch embedding - self.patch_embed = CogVideoXPatchEmbed( - patch_size=patch_size, - patch_size_t=patch_size_t, - in_channels=in_channels, - embed_dim=inner_dim, - text_embed_dim=text_embed_dim, - bias=patch_bias, - sample_width=sample_width, - sample_height=sample_height, - sample_frames=sample_frames, - temporal_compression_ratio=temporal_compression_ratio, - max_text_seq_length=max_text_seq_length, - spatial_interpolation_scale=spatial_interpolation_scale, - temporal_interpolation_scale=temporal_interpolation_scale, - use_positional_embeddings=not use_rotary_positional_embeddings, - use_learned_positional_embeddings=use_learned_positional_embeddings, - ) - self.embedding_dropout = nn.Dropout(dropout) - - # 2. Time embeddings and ofs embedding(Only CogVideoX1.5-5B I2V have) - - self.time_proj = Timesteps(inner_dim, flip_sin_to_cos, freq_shift) - self.time_embedding = TimestepEmbedding(inner_dim, time_embed_dim, timestep_activation_fn) - - self.ofs_proj = None - self.ofs_embedding = None - if ofs_embed_dim: - self.ofs_proj = Timesteps(ofs_embed_dim, flip_sin_to_cos, freq_shift) - self.ofs_embedding = TimestepEmbedding( - ofs_embed_dim, ofs_embed_dim, timestep_activation_fn - ) # same as time embeddings, for ofs - - # 3. Define spatio-temporal transformers blocks - self.transformer_blocks = nn.ModuleList( - [ - CogVideoXBlock( - dim=inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - time_embed_dim=time_embed_dim, - dropout=dropout, - activation_fn=activation_fn, - attention_bias=attention_bias, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - ) - for _ in range(num_layers) - ] - ) - self.norm_final = nn.LayerNorm(inner_dim, norm_eps, norm_elementwise_affine) - - # 4. Output blocks - self.norm_out = AdaLayerNorm( - embedding_dim=time_embed_dim, - output_dim=2 * inner_dim, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - chunk_dim=1, - ) - - if patch_size_t is None: - # For CogVideox 1.0 - output_dim = patch_size * patch_size * out_channels - else: - # For CogVideoX 1.5 - output_dim = patch_size * patch_size * patch_size_t * out_channels - - self.proj_out = nn.Linear(inner_dim, output_dim) - - self.gradient_checkpointing = False - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections with FusedAttnProcessor2_0->FusedCogVideoXAttnProcessor2_0 - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedCogVideoXAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: int | float | torch.LongTensor, - timestep_cond: torch.Tensor | None = None, - ofs: int | float | torch.LongTensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`CogVideoXTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_frames, channels, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - timestep_cond (`torch.Tensor`, *optional*): - Conditional embeddings for timestep. If provided, the embeddings will be summed with the samples passed - through the `self.time_embedding` layer to obtain the final timestep embeddings. - ofs (`torch.Tensor`, *optional*): - Offset embeddings used in CogVideoX-5b-I2V. - image_rotary_emb (`tuple` of `torch.Tensor`, *optional*): - Pre-computed rotary positional embeddings. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_frames, channels, height, width = hidden_states.shape - - # 1. Time embedding - timesteps = timestep - t_emb = self.time_proj(timesteps) - - # timesteps does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=hidden_states.dtype) - emb = self.time_embedding(t_emb, timestep_cond) - - if self.ofs_embedding is not None: - ofs_emb = self.ofs_proj(ofs) - ofs_emb = ofs_emb.to(dtype=hidden_states.dtype) - ofs_emb = self.ofs_embedding(ofs_emb) - emb = emb + ofs_emb - - # 2. Patch embedding - hidden_states = self.patch_embed(encoder_hidden_states, hidden_states) - hidden_states = self.embedding_dropout(hidden_states) - - text_seq_length = encoder_hidden_states.shape[1] - encoder_hidden_states = hidden_states[:, :text_seq_length] - hidden_states = hidden_states[:, text_seq_length:] - - # 3. Transformer blocks - for i, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - emb, - image_rotary_emb, - attention_kwargs, - ) - else: - hidden_states, encoder_hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=emb, - image_rotary_emb=image_rotary_emb, - attention_kwargs=attention_kwargs, - ) - - hidden_states = self.norm_final(hidden_states) - - # 4. Final block - hidden_states = self.norm_out(hidden_states, temb=emb) - hidden_states = self.proj_out(hidden_states) - - # 5. Unpatchify - p = self.config.patch_size - p_t = self.config.patch_size_t - - if p_t is None: - output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p) - output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4) - else: - output = hidden_states.reshape( - batch_size, (num_frames + p_t - 1) // p_t, height // p, width // p, -1, p_t, p, p - ) - output = output.permute(0, 1, 5, 4, 2, 6, 3, 7).flatten(6, 7).flatten(4, 5).flatten(1, 2) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/consisid_transformer_3d.py b/diffusers/models/transformers/consisid_transformer_3d.py deleted file mode 100644 index e534f9479311b9f7c0b82bbeff74018adb14a08d..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/consisid_transformer_3d.py +++ /dev/null @@ -1,742 +0,0 @@ -# Copyright 2025 ConsisID Authors and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import Any - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import Attention, AttentionMixin, FeedForward -from ..attention_processor import CogVideoXAttnProcessor2_0 -from ..embeddings import CogVideoXPatchEmbed, TimestepEmbedding, Timesteps -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNorm, CogVideoXLayerNormZero - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class PerceiverAttention(nn.Module): - def __init__(self, dim: int, dim_head: int = 64, heads: int = 8, kv_dim: int | None = None): - super().__init__() - - self.scale = dim_head**-0.5 - self.dim_head = dim_head - self.heads = heads - inner_dim = dim_head * heads - - self.norm1 = nn.LayerNorm(dim if kv_dim is None else kv_dim) - self.norm2 = nn.LayerNorm(dim) - - self.to_q = nn.Linear(dim, inner_dim, bias=False) - self.to_kv = nn.Linear(dim if kv_dim is None else kv_dim, inner_dim * 2, bias=False) - self.to_out = nn.Linear(inner_dim, dim, bias=False) - - def forward(self, image_embeds: torch.Tensor, latents: torch.Tensor) -> torch.Tensor: - # Apply normalization - image_embeds = self.norm1(image_embeds) - latents = self.norm2(latents) - - batch_size, seq_len, _ = latents.shape # Get batch size and sequence length - - # Compute query, key, and value matrices - query = self.to_q(latents) - kv_input = torch.cat((image_embeds, latents), dim=-2) - key, value = self.to_kv(kv_input).chunk(2, dim=-1) - - # Reshape the tensors for multi-head attention - query = query.reshape(query.size(0), -1, self.heads, self.dim_head).transpose(1, 2) - key = key.reshape(key.size(0), -1, self.heads, self.dim_head).transpose(1, 2) - value = value.reshape(value.size(0), -1, self.heads, self.dim_head).transpose(1, 2) - - # attention - scale = 1 / math.sqrt(math.sqrt(self.dim_head)) - weight = (query * scale) @ (key * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards - weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) - output = weight @ value - - # Reshape and return the final output - output = output.permute(0, 2, 1, 3).reshape(batch_size, seq_len, -1) - - return self.to_out(output) - - -class LocalFacialExtractor(nn.Module): - def __init__( - self, - id_dim: int = 1280, - vit_dim: int = 1024, - depth: int = 10, - dim_head: int = 64, - heads: int = 16, - num_id_token: int = 5, - num_queries: int = 32, - output_dim: int = 2048, - ff_mult: int = 4, - num_scale: int = 5, - ): - super().__init__() - - # Storing identity token and query information - self.num_id_token = num_id_token - self.vit_dim = vit_dim - self.num_queries = num_queries - assert depth % num_scale == 0 - self.depth = depth // num_scale - self.num_scale = num_scale - scale = vit_dim**-0.5 - - # Learnable latent query embeddings - self.latents = nn.Parameter(torch.randn(1, num_queries, vit_dim) * scale) - # Projection layer to map the latent output to the desired dimension - self.proj_out = nn.Parameter(scale * torch.randn(vit_dim, output_dim)) - - # Attention and ConsisIDFeedForward layer stack - self.layers = nn.ModuleList([]) - for _ in range(depth): - self.layers.append( - nn.ModuleList( - [ - PerceiverAttention(dim=vit_dim, dim_head=dim_head, heads=heads), # Perceiver Attention layer - nn.Sequential( - nn.LayerNorm(vit_dim), - nn.Linear(vit_dim, vit_dim * ff_mult, bias=False), - nn.GELU(), - nn.Linear(vit_dim * ff_mult, vit_dim, bias=False), - ), # ConsisIDFeedForward layer - ] - ) - ) - - # Mappings for each of the 5 different ViT features - for i in range(num_scale): - setattr( - self, - f"mapping_{i}", - nn.Sequential( - nn.Linear(vit_dim, vit_dim), - nn.LayerNorm(vit_dim), - nn.LeakyReLU(), - nn.Linear(vit_dim, vit_dim), - nn.LayerNorm(vit_dim), - nn.LeakyReLU(), - nn.Linear(vit_dim, vit_dim), - ), - ) - - # Mapping for identity embedding vectors - self.id_embedding_mapping = nn.Sequential( - nn.Linear(id_dim, vit_dim), - nn.LayerNorm(vit_dim), - nn.LeakyReLU(), - nn.Linear(vit_dim, vit_dim), - nn.LayerNorm(vit_dim), - nn.LeakyReLU(), - nn.Linear(vit_dim, vit_dim * num_id_token), - ) - - def forward(self, id_embeds: torch.Tensor, vit_hidden_states: list[torch.Tensor]) -> torch.Tensor: - # Repeat latent queries for the batch size - latents = self.latents.repeat(id_embeds.size(0), 1, 1) - - # Map the identity embedding to tokens - id_embeds = self.id_embedding_mapping(id_embeds) - id_embeds = id_embeds.reshape(-1, self.num_id_token, self.vit_dim) - - # Concatenate identity tokens with the latent queries - latents = torch.cat((latents, id_embeds), dim=1) - - # Process each of the num_scale visual feature inputs - for i in range(self.num_scale): - vit_feature = getattr(self, f"mapping_{i}")(vit_hidden_states[i]) - ctx_feature = torch.cat((id_embeds, vit_feature), dim=1) - - # Pass through the PerceiverAttention and ConsisIDFeedForward layers - for attn, ff in self.layers[i * self.depth : (i + 1) * self.depth]: - latents = attn(ctx_feature, latents) + latents - latents = ff(latents) + latents - - # Retain only the query latents - latents = latents[:, : self.num_queries] - # Project the latents to the output dimension - latents = latents @ self.proj_out - return latents - - -class PerceiverCrossAttention(nn.Module): - def __init__(self, dim: int = 3072, dim_head: int = 128, heads: int = 16, kv_dim: int = 2048): - super().__init__() - - self.scale = dim_head**-0.5 - self.dim_head = dim_head - self.heads = heads - inner_dim = dim_head * heads - - # Layer normalization to stabilize training - self.norm1 = nn.LayerNorm(dim if kv_dim is None else kv_dim) - self.norm2 = nn.LayerNorm(dim) - - # Linear transformations to produce queries, keys, and values - self.to_q = nn.Linear(dim, inner_dim, bias=False) - self.to_kv = nn.Linear(dim if kv_dim is None else kv_dim, inner_dim * 2, bias=False) - self.to_out = nn.Linear(inner_dim, dim, bias=False) - - def forward(self, image_embeds: torch.Tensor, hidden_states: torch.Tensor) -> torch.Tensor: - # Apply layer normalization to the input image and latent features - image_embeds = self.norm1(image_embeds) - hidden_states = self.norm2(hidden_states) - - batch_size, seq_len, _ = hidden_states.shape - - # Compute queries, keys, and values - query = self.to_q(hidden_states) - key, value = self.to_kv(image_embeds).chunk(2, dim=-1) - - # Reshape tensors to split into attention heads - query = query.reshape(query.size(0), -1, self.heads, self.dim_head).transpose(1, 2) - key = key.reshape(key.size(0), -1, self.heads, self.dim_head).transpose(1, 2) - value = value.reshape(value.size(0), -1, self.heads, self.dim_head).transpose(1, 2) - - # Compute attention weights - scale = 1 / math.sqrt(math.sqrt(self.dim_head)) - weight = (query * scale) @ (key * scale).transpose(-2, -1) # More stable scaling than post-division - weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) - - # Compute the output via weighted combination of values - out = weight @ value - - # Reshape and permute to prepare for final linear transformation - out = out.permute(0, 2, 1, 3).reshape(batch_size, seq_len, -1) - - return self.to_out(out) - - -@maybe_allow_in_graph -class ConsisIDBlock(nn.Module): - r""" - Transformer block used in [ConsisID](https://github.com/PKU-YuanGroup/ConsisID) model. - - Parameters: - dim (`int`): - The number of channels in the input and output. - num_attention_heads (`int`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`): - The number of channels in each head. - time_embed_dim (`int`): - The number of channels in timestep embedding. - dropout (`float`, defaults to `0.0`): - The dropout probability to use. - activation_fn (`str`, defaults to `"gelu-approximate"`): - Activation function to be used in feed-forward. - attention_bias (`bool`, defaults to `False`): - Whether or not to use bias in attention projection layers. - qk_norm (`bool`, defaults to `True`): - Whether or not to use normalization after query and key projections in Attention. - norm_elementwise_affine (`bool`, defaults to `True`): - Whether to use learnable elementwise affine parameters for normalization. - norm_eps (`float`, defaults to `1e-5`): - Epsilon value for normalization layers. - final_dropout (`bool` defaults to `False`): - Whether to apply a final dropout after the last feed-forward layer. - ff_inner_dim (`int`, *optional*, defaults to `None`): - Custom hidden dimension of Feed-forward layer. If not provided, `4 * dim` is used. - ff_bias (`bool`, defaults to `True`): - Whether or not to use bias in Feed-forward layer. - attention_out_bias (`bool`, defaults to `True`): - Whether or not to use bias in Attention output projection layer. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - time_embed_dim: int, - dropout: float = 0.0, - activation_fn: str = "gelu-approximate", - attention_bias: bool = False, - qk_norm: bool = True, - norm_elementwise_affine: bool = True, - norm_eps: float = 1e-5, - final_dropout: bool = True, - ff_inner_dim: int | None = None, - ff_bias: bool = True, - attention_out_bias: bool = True, - ): - super().__init__() - - # 1. Self Attention - self.norm1 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True) - - self.attn1 = Attention( - query_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - qk_norm="layer_norm" if qk_norm else None, - eps=1e-6, - bias=attention_bias, - out_bias=attention_out_bias, - processor=CogVideoXAttnProcessor2_0(), - ) - - # 2. Feed Forward - self.norm2 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True) - - self.ff = FeedForward( - dim, - dropout=dropout, - activation_fn=activation_fn, - final_dropout=final_dropout, - inner_dim=ff_inner_dim, - bias=ff_bias, - ) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - text_seq_length = encoder_hidden_states.size(1) - - # norm & modulate - norm_hidden_states, norm_encoder_hidden_states, gate_msa, enc_gate_msa = self.norm1( - hidden_states, encoder_hidden_states, temb - ) - - # attention - attn_hidden_states, attn_encoder_hidden_states = self.attn1( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - ) - - hidden_states = hidden_states + gate_msa * attn_hidden_states - encoder_hidden_states = encoder_hidden_states + enc_gate_msa * attn_encoder_hidden_states - - # norm & modulate - norm_hidden_states, norm_encoder_hidden_states, gate_ff, enc_gate_ff = self.norm2( - hidden_states, encoder_hidden_states, temb - ) - - # feed-forward - norm_hidden_states = torch.cat([norm_encoder_hidden_states, norm_hidden_states], dim=1) - ff_output = self.ff(norm_hidden_states) - - hidden_states = hidden_states + gate_ff * ff_output[:, text_seq_length:] - encoder_hidden_states = encoder_hidden_states + enc_gate_ff * ff_output[:, :text_seq_length] - - return hidden_states, encoder_hidden_states - - -class ConsisIDTransformer3DModel(ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin): - """ - A Transformer model for video-like data in [ConsisID](https://github.com/PKU-YuanGroup/ConsisID). - - Parameters: - num_attention_heads (`int`, defaults to `30`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `64`): - The number of channels in each head. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `16`): - The number of channels in the output. - flip_sin_to_cos (`bool`, defaults to `True`): - Whether to flip the sin to cos in the time embedding. - time_embed_dim (`int`, defaults to `512`): - Output dimension of timestep embeddings. - text_embed_dim (`int`, defaults to `4096`): - Input dimension of text embeddings from the text encoder. - num_layers (`int`, defaults to `30`): - The number of layers of Transformer blocks to use. - dropout (`float`, defaults to `0.0`): - The dropout probability to use. - attention_bias (`bool`, defaults to `True`): - Whether to use bias in the attention projection layers. - sample_width (`int`, defaults to `90`): - The width of the input latents. - sample_height (`int`, defaults to `60`): - The height of the input latents. - sample_frames (`int`, defaults to `49`): - The number of frames in the input latents. Note that this parameter was incorrectly initialized to 49 - instead of 13 because ConsisID processed 13 latent frames at once in its default and recommended settings, - but cannot be changed to the correct value to ensure backwards compatibility. To create a transformer with - K latent frames, the correct value to pass here would be: ((K - 1) * temporal_compression_ratio + 1). - patch_size (`int`, defaults to `2`): - The size of the patches to use in the patch embedding layer. - temporal_compression_ratio (`int`, defaults to `4`): - The compression ratio across the temporal dimension. See documentation for `sample_frames`. - max_text_seq_length (`int`, defaults to `226`): - The maximum sequence length of the input text embeddings. - activation_fn (`str`, defaults to `"gelu-approximate"`): - Activation function to use in feed-forward. - timestep_activation_fn (`str`, defaults to `"silu"`): - Activation function to use when generating the timestep embeddings. - norm_elementwise_affine (`bool`, defaults to `True`): - Whether to use elementwise affine in normalization layers. - norm_eps (`float`, defaults to `1e-5`): - The epsilon value to use in normalization layers. - spatial_interpolation_scale (`float`, defaults to `1.875`): - Scaling factor to apply in 3D positional embeddings across spatial dimensions. - temporal_interpolation_scale (`float`, defaults to `1.0`): - Scaling factor to apply in 3D positional embeddings across temporal dimensions. - is_train_face (`bool`, defaults to `False`): - Whether to use enable the identity-preserving module during the training process. When set to `True`, the - model will focus on identity-preserving tasks. - is_kps (`bool`, defaults to `False`): - Whether to enable keypoint for global facial extractor. If `True`, keypoints will be in the model. - cross_attn_interval (`int`, defaults to `2`): - The interval between cross-attention layers in the Transformer architecture. A larger value may reduce the - frequency of cross-attention computations, which can help reduce computational overhead. - cross_attn_dim_head (`int`, optional, defaults to `128`): - The dimensionality of each attention head in the cross-attention layers of the Transformer architecture. A - larger value increases the capacity to attend to more complex patterns, but also increases memory and - computation costs. - cross_attn_num_heads (`int`, optional, defaults to `16`): - The number of attention heads in the cross-attention layers. More heads allow for more parallel attention - mechanisms, capturing diverse relationships between different components of the input, but can also - increase computational requirements. - LFE_id_dim (`int`, optional, defaults to `1280`): - The dimensionality of the identity vector used in the Local Facial Extractor (LFE). This vector represents - the identity features of a face, which are important for tasks like face recognition and identity - preservation across different frames. - LFE_vit_dim (`int`, optional, defaults to `1024`): - The dimension of the vision transformer (ViT) output used in the Local Facial Extractor (LFE). This value - dictates the size of the transformer-generated feature vectors that will be processed for facial feature - extraction. - LFE_depth (`int`, optional, defaults to `10`): - The number of layers in the Local Facial Extractor (LFE). Increasing the depth allows the model to capture - more complex representations of facial features, but also increases the computational load. - LFE_dim_head (`int`, optional, defaults to `64`): - The dimensionality of each attention head in the Local Facial Extractor (LFE). This parameter affects how - finely the model can process and focus on different parts of the facial features during the extraction - process. - LFE_num_heads (`int`, optional, defaults to `16`): - The number of attention heads in the Local Facial Extractor (LFE). More heads can improve the model's - ability to capture diverse facial features, but at the cost of increased computational complexity. - LFE_num_id_token (`int`, optional, defaults to `5`): - The number of identity tokens used in the Local Facial Extractor (LFE). This defines how many - identity-related tokens the model will process to ensure face identity preservation during feature - extraction. - LFE_num_querie (`int`, optional, defaults to `32`): - The number of query tokens used in the Local Facial Extractor (LFE). These tokens are used to capture - high-frequency face-related information that aids in accurate facial feature extraction. - LFE_output_dim (`int`, optional, defaults to `2048`): - The output dimension of the Local Facial Extractor (LFE). This dimension determines the size of the feature - vectors produced by the LFE module, which will be used for subsequent tasks such as face recognition or - tracking. - LFE_ff_mult (`int`, optional, defaults to `4`): - The multiplication factor applied to the feed-forward network's hidden layer size in the Local Facial - Extractor (LFE). A higher value increases the model's capacity to learn more complex facial feature - transformations, but also increases the computation and memory requirements. - LFE_num_scale (`int`, optional, defaults to `5`): - The number of different scales visual feature. A higher value increases the model's capacity to learn more - complex facial feature transformations, but also increases the computation and memory requirements. - local_face_scale (`float`, defaults to `1.0`): - A scaling factor used to adjust the importance of local facial features in the model. This can influence - how strongly the model focuses on high frequency face-related content. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - num_attention_heads: int = 30, - attention_head_dim: int = 64, - in_channels: int = 16, - out_channels: int | None = 16, - flip_sin_to_cos: bool = True, - freq_shift: int = 0, - time_embed_dim: int = 512, - text_embed_dim: int = 4096, - num_layers: int = 30, - dropout: float = 0.0, - attention_bias: bool = True, - sample_width: int = 90, - sample_height: int = 60, - sample_frames: int = 49, - patch_size: int = 2, - temporal_compression_ratio: int = 4, - max_text_seq_length: int = 226, - activation_fn: str = "gelu-approximate", - timestep_activation_fn: str = "silu", - norm_elementwise_affine: bool = True, - norm_eps: float = 1e-5, - spatial_interpolation_scale: float = 1.875, - temporal_interpolation_scale: float = 1.0, - use_rotary_positional_embeddings: bool = False, - use_learned_positional_embeddings: bool = False, - is_train_face: bool = False, - is_kps: bool = False, - cross_attn_interval: int = 2, - cross_attn_dim_head: int = 128, - cross_attn_num_heads: int = 16, - LFE_id_dim: int = 1280, - LFE_vit_dim: int = 1024, - LFE_depth: int = 10, - LFE_dim_head: int = 64, - LFE_num_heads: int = 16, - LFE_num_id_token: int = 5, - LFE_num_querie: int = 32, - LFE_output_dim: int = 2048, - LFE_ff_mult: int = 4, - LFE_num_scale: int = 5, - local_face_scale: float = 1.0, - ): - super().__init__() - inner_dim = num_attention_heads * attention_head_dim - - if not use_rotary_positional_embeddings and use_learned_positional_embeddings: - raise ValueError( - "There are no ConsisID checkpoints available with disable rotary embeddings and learned positional " - "embeddings. If you're using a custom model and/or believe this should be supported, please open an " - "issue at https://github.com/huggingface/diffusers/issues." - ) - - # 1. Patch embedding - self.patch_embed = CogVideoXPatchEmbed( - patch_size=patch_size, - in_channels=in_channels, - embed_dim=inner_dim, - text_embed_dim=text_embed_dim, - bias=True, - sample_width=sample_width, - sample_height=sample_height, - sample_frames=sample_frames, - temporal_compression_ratio=temporal_compression_ratio, - max_text_seq_length=max_text_seq_length, - spatial_interpolation_scale=spatial_interpolation_scale, - temporal_interpolation_scale=temporal_interpolation_scale, - use_positional_embeddings=not use_rotary_positional_embeddings, - use_learned_positional_embeddings=use_learned_positional_embeddings, - ) - self.embedding_dropout = nn.Dropout(dropout) - - # 2. Time embeddings - self.time_proj = Timesteps(inner_dim, flip_sin_to_cos, freq_shift) - self.time_embedding = TimestepEmbedding(inner_dim, time_embed_dim, timestep_activation_fn) - - # 3. Define spatio-temporal transformers blocks - self.transformer_blocks = nn.ModuleList( - [ - ConsisIDBlock( - dim=inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - time_embed_dim=time_embed_dim, - dropout=dropout, - activation_fn=activation_fn, - attention_bias=attention_bias, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - ) - for _ in range(num_layers) - ] - ) - self.norm_final = nn.LayerNorm(inner_dim, norm_eps, norm_elementwise_affine) - - # 4. Output blocks - self.norm_out = AdaLayerNorm( - embedding_dim=time_embed_dim, - output_dim=2 * inner_dim, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - chunk_dim=1, - ) - self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels) - - self.is_train_face = is_train_face - self.is_kps = is_kps - - # 5. Define identity-preserving config - if is_train_face: - # LFE configs - self.LFE_id_dim = LFE_id_dim - self.LFE_vit_dim = LFE_vit_dim - self.LFE_depth = LFE_depth - self.LFE_dim_head = LFE_dim_head - self.LFE_num_heads = LFE_num_heads - self.LFE_num_id_token = LFE_num_id_token - self.LFE_num_querie = LFE_num_querie - self.LFE_output_dim = LFE_output_dim - self.LFE_ff_mult = LFE_ff_mult - self.LFE_num_scale = LFE_num_scale - # cross configs - self.inner_dim = inner_dim - self.cross_attn_interval = cross_attn_interval - self.num_cross_attn = num_layers // cross_attn_interval - self.cross_attn_dim_head = cross_attn_dim_head - self.cross_attn_num_heads = cross_attn_num_heads - self.cross_attn_kv_dim = int(self.inner_dim / 3 * 2) - self.local_face_scale = local_face_scale - # face modules - self._init_face_inputs() - - self.gradient_checkpointing = False - - def _init_face_inputs(self): - self.local_facial_extractor = LocalFacialExtractor( - id_dim=self.LFE_id_dim, - vit_dim=self.LFE_vit_dim, - depth=self.LFE_depth, - dim_head=self.LFE_dim_head, - heads=self.LFE_num_heads, - num_id_token=self.LFE_num_id_token, - num_queries=self.LFE_num_querie, - output_dim=self.LFE_output_dim, - ff_mult=self.LFE_ff_mult, - num_scale=self.LFE_num_scale, - ) - self.perceiver_cross_attention = nn.ModuleList( - [ - PerceiverCrossAttention( - dim=self.inner_dim, - dim_head=self.cross_attn_dim_head, - heads=self.cross_attn_num_heads, - kv_dim=self.cross_attn_kv_dim, - ) - for _ in range(self.num_cross_attn) - ] - ) - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: int | float | torch.LongTensor, - timestep_cond: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_kwargs: dict[str, Any] | None = None, - id_cond: torch.Tensor | None = None, - id_vit_hidden: torch.Tensor | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`ConsisIDTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_frames, channels, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - timestep_cond (`torch.Tensor`, *optional*): - Conditional embeddings for timestep. If provided, the embeddings will be summed with the samples passed - through the `self.time_embedding` layer to obtain the final timestep embeddings. - image_rotary_emb (`tuple` of `torch.Tensor`, *optional*): - Pre-computed rotary positional embeddings. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - id_cond (`torch.Tensor`, *optional*): - The face embedding extracted by the local facial extractor used for identity conditioning. - id_vit_hidden (`torch.Tensor`, *optional*): - The ViT hidden states extracted from face images used for identity conditioning. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - # fuse clip and insightface - valid_face_emb = None - if self.is_train_face: - id_cond = id_cond.to(device=hidden_states.device, dtype=hidden_states.dtype) - id_vit_hidden = [ - tensor.to(device=hidden_states.device, dtype=hidden_states.dtype) for tensor in id_vit_hidden - ] - valid_face_emb = self.local_facial_extractor( - id_cond, id_vit_hidden - ) # torch.Size([1, 1280]), list[5](torch.Size([1, 577, 1024])) -> torch.Size([1, 32, 2048]) - - batch_size, num_frames, channels, height, width = hidden_states.shape - - # 1. Time embedding - timesteps = timestep - t_emb = self.time_proj(timesteps) - - # timesteps does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=hidden_states.dtype) - emb = self.time_embedding(t_emb, timestep_cond) - - # 2. Patch embedding - # torch.Size([1, 226, 4096]) torch.Size([1, 13, 32, 60, 90]) - hidden_states = self.patch_embed(encoder_hidden_states, hidden_states) # torch.Size([1, 17776, 3072]) - hidden_states = self.embedding_dropout(hidden_states) # torch.Size([1, 17776, 3072]) - - text_seq_length = encoder_hidden_states.shape[1] - encoder_hidden_states = hidden_states[:, :text_seq_length] # torch.Size([1, 226, 3072]) - hidden_states = hidden_states[:, text_seq_length:] # torch.Size([1, 17550, 3072]) - - # 3. Transformer blocks - ca_idx = 0 - for i, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - emb, - image_rotary_emb, - ) - else: - hidden_states, encoder_hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=emb, - image_rotary_emb=image_rotary_emb, - ) - - if self.is_train_face: - if i % self.cross_attn_interval == 0 and valid_face_emb is not None: - hidden_states = hidden_states + self.local_face_scale * self.perceiver_cross_attention[ca_idx]( - valid_face_emb, hidden_states - ) # torch.Size([2, 32, 2048]) torch.Size([2, 17550, 3072]) - ca_idx += 1 - - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - hidden_states = self.norm_final(hidden_states) - hidden_states = hidden_states[:, text_seq_length:] - - # 4. Final block - hidden_states = self.norm_out(hidden_states, temb=emb) - hidden_states = self.proj_out(hidden_states) - - # 5. Unpatchify - # Note: we use `-1` instead of `channels`: - # - It is okay to `channels` use for ConsisID (number of input channels is equal to output channels) - p = self.config.patch_size - output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p) - output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/dit_transformer_2d.py b/diffusers/models/transformers/dit_transformer_2d.py deleted file mode 100644 index 0457acf771087540b14781afc04b1cd9da713be1..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/dit_transformer_2d.py +++ /dev/null @@ -1,226 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from typing import Any - -import torch -import torch.nn.functional as F -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ..attention import BasicTransformerBlock -from ..embeddings import PatchEmbed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class DiTTransformer2DModel(ModelMixin, ConfigMixin): - r""" - A 2D Transformer model as introduced in DiT (https://huggingface.co/papers/2212.09748). - - Parameters: - num_attention_heads (int, optional, defaults to 16): The number of heads to use for multi-head attention. - attention_head_dim (int, optional, defaults to 72): The number of channels in each head. - in_channels (int, defaults to 4): The number of channels in the input. - out_channels (int, optional): - The number of channels in the output. Specify this parameter if the output channel number differs from the - input. - num_layers (int, optional, defaults to 28): The number of layers of Transformer blocks to use. - dropout (float, optional, defaults to 0.0): The dropout probability to use within the Transformer blocks. - norm_num_groups (int, optional, defaults to 32): - Number of groups for group normalization within Transformer blocks. - attention_bias (bool, optional, defaults to True): - Configure if the Transformer blocks' attention should contain a bias parameter. - sample_size (int, defaults to 32): - The width of the latent images. This parameter is fixed during training. - patch_size (int, defaults to 2): - Size of the patches the model processes, relevant for architectures working on non-sequential data. - activation_fn (str, optional, defaults to "gelu-approximate"): - Activation function to use in feed-forward networks within Transformer blocks. - num_embeds_ada_norm (int, optional, defaults to 1000): - Number of embeddings for AdaLayerNorm, fixed during training and affects the maximum denoising steps during - inference. - upcast_attention (bool, optional, defaults to False): - If true, upcasts the attention mechanism dimensions for potentially improved performance. - norm_type (str, optional, defaults to "ada_norm_zero"): - Specifies the type of normalization used, can be 'ada_norm_zero'. - norm_elementwise_affine (bool, optional, defaults to False): - If true, enables element-wise affine parameters in the normalization layers. - norm_eps (float, optional, defaults to 1e-5): - A small constant added to the denominator in normalization layers to prevent division by zero. - """ - - _skip_layerwise_casting_patterns = ["pos_embed", "norm"] - _supports_gradient_checkpointing = True - _supports_group_offloading = False - - @register_to_config - def __init__( - self, - num_attention_heads: int = 16, - attention_head_dim: int = 72, - in_channels: int = 4, - out_channels: int | None = None, - num_layers: int = 28, - dropout: float = 0.0, - norm_num_groups: int = 32, - attention_bias: bool = True, - sample_size: int = 32, - patch_size: int = 2, - activation_fn: str = "gelu-approximate", - num_embeds_ada_norm: int | None = 1000, - upcast_attention: bool = False, - norm_type: str = "ada_norm_zero", - norm_elementwise_affine: bool = False, - norm_eps: float = 1e-5, - ): - super().__init__() - - # Validate inputs. - if norm_type != "ada_norm_zero": - raise NotImplementedError( - f"Forward pass is not implemented when `patch_size` is not None and `norm_type` is '{norm_type}'." - ) - elif norm_type == "ada_norm_zero" and num_embeds_ada_norm is None: - raise ValueError( - f"When using a `patch_size` and this `norm_type` ({norm_type}), `num_embeds_ada_norm` cannot be None." - ) - - # Set some common variables used across the board. - self.attention_head_dim = attention_head_dim - self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim - self.out_channels = in_channels if out_channels is None else out_channels - self.gradient_checkpointing = False - - # 2. Initialize the position embedding and transformer blocks. - self.height = self.config.sample_size - self.width = self.config.sample_size - - self.patch_size = self.config.patch_size - self.pos_embed = PatchEmbed( - height=self.config.sample_size, - width=self.config.sample_size, - patch_size=self.config.patch_size, - in_channels=self.config.in_channels, - embed_dim=self.inner_dim, - ) - - self.transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlock( - self.inner_dim, - self.config.num_attention_heads, - self.config.attention_head_dim, - dropout=self.config.dropout, - activation_fn=self.config.activation_fn, - num_embeds_ada_norm=self.config.num_embeds_ada_norm, - attention_bias=self.config.attention_bias, - upcast_attention=self.config.upcast_attention, - norm_type=norm_type, - norm_elementwise_affine=self.config.norm_elementwise_affine, - norm_eps=self.config.norm_eps, - ) - for _ in range(self.config.num_layers) - ] - ) - - # 3. Output blocks. - self.norm_out = nn.LayerNorm(self.inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out_1 = nn.Linear(self.inner_dim, 2 * self.inner_dim) - self.proj_out_2 = nn.Linear( - self.inner_dim, self.config.patch_size * self.config.patch_size * self.out_channels - ) - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor | None = None, - class_labels: torch.LongTensor | None = None, - cross_attention_kwargs: dict[str, Any] = None, - return_dict: bool = True, - ): - """ - The [`DiTTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.LongTensor` of shape `(batch size, num latent pixels)` if discrete, `torch.FloatTensor` of shape `(batch size, channel, height, width)` if continuous): - Input `hidden_states`. - timestep ( `torch.LongTensor`, *optional*): - Used to indicate denoising step. Optional timestep to be applied as an embedding in `AdaLayerNorm`. - class_labels ( `torch.LongTensor` of shape `(batch size, num classes)`, *optional*): - Used to indicate class labels conditioning. Optional class labels to be applied as an embedding in - `AdaLayerZeroNorm`. - cross_attention_kwargs ( `dict[str, Any]`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unets.unet_2d_condition.UNet2DConditionOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - # 1. Input - height, width = hidden_states.shape[-2] // self.patch_size, hidden_states.shape[-1] // self.patch_size - hidden_states = self.pos_embed(hidden_states) - - # 2. Blocks - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - None, - None, - None, - timestep, - cross_attention_kwargs, - class_labels, - ) - else: - hidden_states = block( - hidden_states, - attention_mask=None, - encoder_hidden_states=None, - encoder_attention_mask=None, - timestep=timestep, - cross_attention_kwargs=cross_attention_kwargs, - class_labels=class_labels, - ) - - # 3. Output - conditioning = self.transformer_blocks[0].norm1.emb(timestep, class_labels, hidden_dtype=hidden_states.dtype) - shift, scale = self.proj_out_1(F.silu(conditioning)).chunk(2, dim=1) - hidden_states = self.norm_out(hidden_states) * (1 + scale[:, None]) + shift[:, None] - hidden_states = self.proj_out_2(hidden_states) - - # unpatchify - height = width = int(hidden_states.shape[1] ** 0.5) - hidden_states = hidden_states.reshape( - shape=(-1, height, width, self.patch_size, self.patch_size, self.out_channels) - ) - hidden_states = torch.einsum("nhwpqc->nchpwq", hidden_states) - output = hidden_states.reshape( - shape=(-1, self.out_channels, height * self.patch_size, width * self.patch_size) - ) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/dual_transformer_2d.py b/diffusers/models/transformers/dual_transformer_2d.py deleted file mode 100644 index 778d5128ee23a699d629fb55a88b9607c4aca5df..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/dual_transformer_2d.py +++ /dev/null @@ -1,154 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from torch import nn - -from ..modeling_outputs import Transformer2DModelOutput -from .transformer_2d import Transformer2DModel - - -class DualTransformer2DModel(nn.Module): - """ - Dual transformer wrapper that combines two `Transformer2DModel`s for mixed inference. - - Parameters: - num_attention_heads (`int`, *optional*, defaults to 16): The number of heads to use for multi-head attention. - attention_head_dim (`int`, *optional*, defaults to 88): The number of channels in each head. - in_channels (`int`, *optional*): - Pass if the input is continuous. The number of channels in the input and output. - num_layers (`int`, *optional*, defaults to 1): The number of layers of Transformer blocks to use. - dropout (`float`, *optional*, defaults to 0.1): The dropout probability to use. - cross_attention_dim (`int`, *optional*): The number of encoder_hidden_states dimensions to use. - sample_size (`int`, *optional*): Pass if the input is discrete. The width of the latent images. - Note that this is fixed at training time as it is used for learning a number of position embeddings. See - `ImagePositionalEmbeddings`. - num_vector_embeds (`int`, *optional*): - Pass if the input is discrete. The number of classes of the vector embeddings of the latent pixels. - Includes the class for the masked latent pixel. - activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward. - num_embeds_ada_norm ( `int`, *optional*): Pass if at least one of the norm_layers is `AdaLayerNorm`. - The number of diffusion steps used during training. Note that this is fixed at training time as it is used - to learn a number of embeddings that are added to the hidden states. During inference, you can denoise for - up to but not more than steps than `num_embeds_ada_norm`. - attention_bias (`bool`, *optional*): - Configure if the TransformerBlocks' attention should contain a bias parameter. - """ - - def __init__( - self, - num_attention_heads: int = 16, - attention_head_dim: int = 88, - in_channels: int | None = None, - num_layers: int = 1, - dropout: float = 0.0, - norm_num_groups: int = 32, - cross_attention_dim: int | None = None, - attention_bias: bool = False, - sample_size: int | None = None, - num_vector_embeds: int | None = None, - activation_fn: str = "geglu", - num_embeds_ada_norm: int | None = None, - ): - super().__init__() - self.transformers = nn.ModuleList( - [ - Transformer2DModel( - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - in_channels=in_channels, - num_layers=num_layers, - dropout=dropout, - norm_num_groups=norm_num_groups, - cross_attention_dim=cross_attention_dim, - attention_bias=attention_bias, - sample_size=sample_size, - num_vector_embeds=num_vector_embeds, - activation_fn=activation_fn, - num_embeds_ada_norm=num_embeds_ada_norm, - ) - for _ in range(2) - ] - ) - - # Variables that can be set by a pipeline: - - # The ratio of transformer1 to transformer2's output states to be combined during inference - self.mix_ratio = 0.5 - - # The shape of `encoder_hidden_states` is expected to be - # `(batch_size, condition_lengths[0]+condition_lengths[1], num_features)` - self.condition_lengths = [77, 257] - - # Which transformer to use to encode which condition. - # E.g. `(1, 0)` means that we'll use `transformers[1](conditions[0])` and `transformers[0](conditions[1])` - self.transformer_index_for_condition = [1, 0] - - def forward( - self, - hidden_states, - encoder_hidden_states, - timestep=None, - attention_mask=None, - cross_attention_kwargs=None, - return_dict: bool = True, - ): - """ - Args: - hidden_states ( When discrete, `torch.LongTensor` of shape `(batch size, num latent pixels)`. - When continuous, `torch.Tensor` of shape `(batch size, channel, height, width)`): Input hidden_states. - encoder_hidden_states ( `torch.LongTensor` of shape `(batch size, encoder_hidden_states dim)`, *optional*): - Conditional embeddings for cross attention layer. If not given, cross-attention defaults to - self-attention. - timestep ( `torch.long`, *optional*): - Optional timestep to be applied as an embedding in AdaLayerNorm's. Used to indicate denoising step. - attention_mask (`torch.Tensor`, *optional*): - Optional attention mask to be applied in Attention. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`models.unets.unet_2d_condition.UNet2DConditionOutput`] instead of a plain - tuple. - - Returns: - [`~models.transformers.transformer_2d.Transformer2DModelOutput`] or `tuple`: - [`~models.transformers.transformer_2d.Transformer2DModelOutput`] if `return_dict` is True, otherwise a - `tuple`. When returning a tuple, the first element is the sample tensor. - """ - input_states = hidden_states - - encoded_states = [] - tokens_start = 0 - # attention_mask is not used yet - for i in range(2): - # for each of the two transformers, pass the corresponding condition tokens - condition_state = encoder_hidden_states[:, tokens_start : tokens_start + self.condition_lengths[i]] - transformer_index = self.transformer_index_for_condition[i] - encoded_state = self.transformers[transformer_index]( - input_states, - encoder_hidden_states=condition_state, - timestep=timestep, - cross_attention_kwargs=cross_attention_kwargs, - return_dict=False, - )[0] - encoded_states.append(encoded_state - input_states) - tokens_start += self.condition_lengths[i] - - output_states = encoded_states[0] * self.mix_ratio + encoded_states[1] * (1 - self.mix_ratio) - output_states = output_states + input_states - - if not return_dict: - return (output_states,) - - return Transformer2DModelOutput(sample=output_states) diff --git a/diffusers/models/transformers/hunyuan_transformer_2d.py b/diffusers/models/transformers/hunyuan_transformer_2d.py deleted file mode 100644 index 83b3797c4fc3998af963dd05c92f93e575282ca0..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/hunyuan_transformer_2d.py +++ /dev/null @@ -1,511 +0,0 @@ -# Copyright 2025 HunyuanDiT Authors, Qixun Wang and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import AttentionMixin, FeedForward -from ..attention_processor import Attention, FusedHunyuanAttnProcessor2_0, HunyuanAttnProcessor2_0 -from ..embeddings import ( - HunyuanCombinedTimestepTextSizeStyleEmbedding, - PatchEmbed, - PixArtAlphaTextProjection, -) -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, FP32LayerNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class AdaLayerNormShift(nn.Module): - r""" - Norm layer modified to incorporate timestep embeddings. - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - num_embeddings (`int`): The size of the embeddings dictionary. - """ - - def __init__(self, embedding_dim: int, elementwise_affine=True, eps=1e-6): - super().__init__() - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, embedding_dim) - self.norm = FP32LayerNorm(embedding_dim, elementwise_affine=elementwise_affine, eps=eps) - - def forward(self, x: torch.Tensor, emb: torch.Tensor) -> torch.Tensor: - shift = self.linear(self.silu(emb.to(torch.float32)).to(emb.dtype)) - x = self.norm(x) + shift.unsqueeze(dim=1) - return x - - -@maybe_allow_in_graph -class HunyuanDiTBlock(nn.Module): - r""" - Transformer block used in Hunyuan-DiT model (https://github.com/Tencent/HunyuanDiT). Allow skip connection and - QKNorm - - Parameters: - dim (`int`): - The number of channels in the input and output. - num_attention_heads (`int`): - The number of headsto use for multi-head attention. - cross_attention_dim (`int`,*optional*): - The size of the encoder_hidden_states vector for cross attention. - dropout(`float`, *optional*, defaults to 0.0): - The dropout probability to use. - activation_fn (`str`,*optional*, defaults to `"geglu"`): - Activation function to be used in feed-forward. . - norm_elementwise_affine (`bool`, *optional*, defaults to `True`): - Whether to use learnable elementwise affine parameters for normalization. - norm_eps (`float`, *optional*, defaults to 1e-6): - A small constant added to the denominator in normalization layers to prevent division by zero. - final_dropout (`bool` *optional*, defaults to False): - Whether to apply a final dropout after the last feed-forward layer. - ff_inner_dim (`int`, *optional*): - The size of the hidden layer in the feed-forward block. Defaults to `None`. - ff_bias (`bool`, *optional*, defaults to `True`): - Whether to use bias in the feed-forward block. - skip (`bool`, *optional*, defaults to `False`): - Whether to use skip connection. Defaults to `False` for down-blocks and mid-blocks. - qk_norm (`bool`, *optional*, defaults to `True`): - Whether to use normalization in QK calculation. Defaults to `True`. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - cross_attention_dim: int = 1024, - dropout=0.0, - activation_fn: str = "geglu", - norm_elementwise_affine: bool = True, - norm_eps: float = 1e-6, - final_dropout: bool = False, - ff_inner_dim: int | None = None, - ff_bias: bool = True, - skip: bool = False, - qk_norm: bool = True, - ): - super().__init__() - - # Define 3 blocks. Each block has its own normalization layer. - # NOTE: when new version comes, check norm2 and norm 3 - # 1. Self-Attn - self.norm1 = AdaLayerNormShift(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps) - - self.attn1 = Attention( - query_dim=dim, - cross_attention_dim=None, - dim_head=dim // num_attention_heads, - heads=num_attention_heads, - qk_norm="layer_norm" if qk_norm else None, - eps=1e-6, - bias=True, - processor=HunyuanAttnProcessor2_0(), - ) - - # 2. Cross-Attn - self.norm2 = FP32LayerNorm(dim, norm_eps, norm_elementwise_affine) - - self.attn2 = Attention( - query_dim=dim, - cross_attention_dim=cross_attention_dim, - dim_head=dim // num_attention_heads, - heads=num_attention_heads, - qk_norm="layer_norm" if qk_norm else None, - eps=1e-6, - bias=True, - processor=HunyuanAttnProcessor2_0(), - ) - # 3. Feed-forward - self.norm3 = FP32LayerNorm(dim, norm_eps, norm_elementwise_affine) - - self.ff = FeedForward( - dim, - dropout=dropout, ### 0.0 - activation_fn=activation_fn, ### approx GeLU - final_dropout=final_dropout, ### 0.0 - inner_dim=ff_inner_dim, ### int(dim * mlp_ratio) - bias=ff_bias, - ) - - # 4. Skip Connection - if skip: - self.skip_norm = FP32LayerNorm(2 * dim, norm_eps, elementwise_affine=True) - self.skip_linear = nn.Linear(2 * dim, dim) - else: - self.skip_linear = None - - # let chunk size default to None - self._chunk_size = None - self._chunk_dim = 0 - - # Copied from diffusers.models.attention.BasicTransformerBlock.set_chunk_feed_forward - def set_chunk_feed_forward(self, chunk_size: int | None, dim: int = 0): - # Sets chunk feed-forward - self._chunk_size = chunk_size - self._chunk_dim = dim - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - image_rotary_emb=None, - skip=None, - ) -> torch.Tensor: - # Notice that normalization is always applied before the real computation in the following blocks. - # 0. Long Skip Connection - if self.skip_linear is not None: - cat = torch.cat([hidden_states, skip], dim=-1) - cat = self.skip_norm(cat) - hidden_states = self.skip_linear(cat) - - # 1. Self-Attention - norm_hidden_states = self.norm1(hidden_states, temb) ### checked: self.norm1 is correct - attn_output = self.attn1( - norm_hidden_states, - image_rotary_emb=image_rotary_emb, - ) - hidden_states = hidden_states + attn_output - - # 2. Cross-Attention - hidden_states = hidden_states + self.attn2( - self.norm2(hidden_states), - encoder_hidden_states=encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - ) - - # FFN Layer ### TODO: switch norm2 and norm3 in the state dict - mlp_inputs = self.norm3(hidden_states) - hidden_states = hidden_states + self.ff(mlp_inputs) - - return hidden_states - - -class HunyuanDiT2DModel(ModelMixin, AttentionMixin, ConfigMixin): - """ - HunYuanDiT: Diffusion model with a Transformer backbone. - - Inherit ModelMixin and ConfigMixin to be compatible with the sampler StableDiffusionPipeline of diffusers. - - Parameters: - num_attention_heads (`int`, *optional*, defaults to 16): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, *optional*, defaults to 88): - The number of channels in each head. - in_channels (`int`, *optional*): - The number of channels in the input and output (specify if the input is **continuous**). - patch_size (`int`, *optional*): - The size of the patch to use for the input. - activation_fn (`str`, *optional*, defaults to `"geglu"`): - Activation function to use in feed-forward. - sample_size (`int`, *optional*): - The width of the latent images. This is fixed during training since it is used to learn a number of - position embeddings. - dropout (`float`, *optional*, defaults to 0.0): - The dropout probability to use. - cross_attention_dim (`int`, *optional*): - The number of dimension in the clip text embedding. - hidden_size (`int`, *optional*): - The size of hidden layer in the conditioning embedding layers. - num_layers (`int`, *optional*, defaults to 1): - The number of layers of Transformer blocks to use. - mlp_ratio (`float`, *optional*, defaults to 4.0): - The ratio of the hidden layer size to the input size. - learn_sigma (`bool`, *optional*, defaults to `True`): - Whether to predict variance. - cross_attention_dim_t5 (`int`, *optional*): - The number dimensions in t5 text embedding. - pooled_projection_dim (`int`, *optional*): - The size of the pooled projection. - text_len (`int`, *optional*): - The length of the clip text embedding. - text_len_t5 (`int`, *optional*): - The length of the T5 text embedding. - use_style_cond_and_image_meta_size (`bool`, *optional*): - Whether or not to use style condition and image meta size. True for version <=1.1, False for version >= 1.2 - """ - - _skip_layerwise_casting_patterns = ["pos_embed", "norm", "pooler"] - _supports_group_offloading = False - - @register_to_config - def __init__( - self, - num_attention_heads: int = 16, - attention_head_dim: int = 88, - in_channels: int | None = None, - patch_size: int | None = None, - activation_fn: str = "gelu-approximate", - sample_size=32, - hidden_size=1152, - num_layers: int = 28, - mlp_ratio: float = 4.0, - learn_sigma: bool = True, - cross_attention_dim: int = 1024, - norm_type: str = "layer_norm", - cross_attention_dim_t5: int = 2048, - pooled_projection_dim: int = 1024, - text_len: int = 77, - text_len_t5: int = 256, - use_style_cond_and_image_meta_size: bool = True, - ): - super().__init__() - self.out_channels = in_channels * 2 if learn_sigma else in_channels - self.num_heads = num_attention_heads - self.inner_dim = num_attention_heads * attention_head_dim - - self.text_embedder = PixArtAlphaTextProjection( - in_features=cross_attention_dim_t5, - hidden_size=cross_attention_dim_t5 * 4, - out_features=cross_attention_dim, - act_fn="silu_fp32", - ) - - self.text_embedding_padding = nn.Parameter(torch.randn(text_len + text_len_t5, cross_attention_dim)) - - self.pos_embed = PatchEmbed( - height=sample_size, - width=sample_size, - in_channels=in_channels, - embed_dim=hidden_size, - patch_size=patch_size, - pos_embed_type=None, - ) - - self.time_extra_emb = HunyuanCombinedTimestepTextSizeStyleEmbedding( - hidden_size, - pooled_projection_dim=pooled_projection_dim, - seq_len=text_len_t5, - cross_attention_dim=cross_attention_dim_t5, - use_style_cond_and_image_meta_size=use_style_cond_and_image_meta_size, - ) - - # HunyuanDiT Blocks - self.blocks = nn.ModuleList( - [ - HunyuanDiTBlock( - dim=self.inner_dim, - num_attention_heads=self.config.num_attention_heads, - activation_fn=activation_fn, - ff_inner_dim=int(self.inner_dim * mlp_ratio), - cross_attention_dim=cross_attention_dim, - qk_norm=True, # See https://huggingface.co/papers/2302.05442 for details. - skip=layer > num_layers // 2, - ) - for layer in range(num_layers) - ] - ) - - self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections with FusedAttnProcessor2_0->FusedHunyuanAttnProcessor2_0 - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedHunyuanAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) - - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - self.set_attn_processor(HunyuanAttnProcessor2_0()) - - def forward( - self, - hidden_states, - timestep, - encoder_hidden_states=None, - text_embedding_mask=None, - encoder_hidden_states_t5=None, - text_embedding_mask_t5=None, - image_meta_size=None, - style=None, - image_rotary_emb=None, - controlnet_block_samples=None, - return_dict=True, - ): - """ - The [`HunyuanDiT2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch size, dim, height, width)`): - The input tensor. - timestep ( `torch.LongTensor`, *optional*): - Used to indicate denoising step. - encoder_hidden_states ( `torch.Tensor` of shape `(batch size, sequence len, embed dims)`, *optional*): - Conditional embeddings for cross attention layer. This is the output of `BertModel`. - text_embedding_mask: torch.Tensor - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. This is the output - of `BertModel`. - encoder_hidden_states_t5 ( `torch.Tensor` of shape `(batch size, sequence len, embed dims)`, *optional*): - Conditional embeddings for cross attention layer. This is the output of T5 Text Encoder. - text_embedding_mask_t5: torch.Tensor - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. This is the output - of T5 Text Encoder. - image_meta_size (torch.Tensor): - Conditional embedding indicate the image sizes - style: torch.Tensor: - Conditional embedding indicate the style - image_rotary_emb (`torch.Tensor`): - The image rotary embeddings to apply on query and key tensors during attention calculation. - controlnet_block_samples (`list` of `torch.Tensor`, *optional*): - A list of tensors that if specified are added to the residuals of transformer blocks. - return_dict: bool - Whether to return a dictionary. - """ - - height, width = hidden_states.shape[-2:] - - hidden_states = self.pos_embed(hidden_states) - - temb = self.time_extra_emb( - timestep, encoder_hidden_states_t5, image_meta_size, style, hidden_dtype=timestep.dtype - ) # [B, D] - - # text projection - batch_size, sequence_length, _ = encoder_hidden_states_t5.shape - encoder_hidden_states_t5 = self.text_embedder( - encoder_hidden_states_t5.view(-1, encoder_hidden_states_t5.shape[-1]) - ) - encoder_hidden_states_t5 = encoder_hidden_states_t5.view(batch_size, sequence_length, -1) - - encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states_t5], dim=1) - text_embedding_mask = torch.cat([text_embedding_mask, text_embedding_mask_t5], dim=-1) - text_embedding_mask = text_embedding_mask.unsqueeze(2).bool() - - encoder_hidden_states = torch.where(text_embedding_mask, encoder_hidden_states, self.text_embedding_padding) - - skips = [] - for layer, block in enumerate(self.blocks): - if layer > self.config.num_layers // 2: - if controlnet_block_samples is not None: - skip = skips.pop() + controlnet_block_samples.pop() - else: - skip = skips.pop() - hidden_states = block( - hidden_states, - temb=temb, - encoder_hidden_states=encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - skip=skip, - ) # (N, L, D) - else: - hidden_states = block( - hidden_states, - temb=temb, - encoder_hidden_states=encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - ) # (N, L, D) - - if layer < (self.config.num_layers // 2 - 1): - skips.append(hidden_states) - - if controlnet_block_samples is not None and len(controlnet_block_samples) != 0: - raise ValueError("The number of controls is not equal to the number of skip connections.") - - # final layer - hidden_states = self.norm_out(hidden_states, temb.to(torch.float32)) - hidden_states = self.proj_out(hidden_states) - # (N, L, patch_size ** 2 * out_channels) - - # unpatchify: (N, out_channels, H, W) - patch_size = self.pos_embed.patch_size - height = height // patch_size - width = width // patch_size - - hidden_states = hidden_states.reshape( - shape=(hidden_states.shape[0], height, width, patch_size, patch_size, self.out_channels) - ) - hidden_states = torch.einsum("nhwpqc->nchpwq", hidden_states) - output = hidden_states.reshape( - shape=(hidden_states.shape[0], self.out_channels, height * patch_size, width * patch_size) - ) - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) - - # Copied from diffusers.models.unets.unet_3d_condition.UNet3DConditionModel.enable_forward_chunking - def enable_forward_chunking(self, chunk_size: int | None = None, dim: int = 0) -> None: - """ - Sets the attention processor to use [feed forward - chunking](https://huggingface.co/blog/reformer#2-chunked-feed-forward-layers). - - Parameters: - chunk_size (`int`, *optional*): - The chunk size of the feed-forward layers. If not specified, will run feed-forward layer individually - over each tensor of dim=`dim`. - dim (`int`, *optional*, defaults to `0`): - The dimension over which the feed-forward computation should be chunked. Choose between dim=0 (batch) - or dim=1 (sequence length). - """ - if dim not in [0, 1]: - raise ValueError(f"Make sure to set `dim` to either 0 or 1, not {dim}") - - # By default chunk size is 1 - chunk_size = chunk_size or 1 - - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, chunk_size, dim) - - # Copied from diffusers.models.unets.unet_3d_condition.UNet3DConditionModel.disable_forward_chunking - def disable_forward_chunking(self): - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, None, 0) diff --git a/diffusers/models/transformers/latte_transformer_3d.py b/diffusers/models/transformers/latte_transformer_3d.py deleted file mode 100644 index 01a1e608a927f968848d98549532bb184e08ebef..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/latte_transformer_3d.py +++ /dev/null @@ -1,329 +0,0 @@ -# Copyright 2025 the Latte Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ..attention import BasicTransformerBlock -from ..cache_utils import CacheMixin -from ..embeddings import PatchEmbed, PixArtAlphaTextProjection, get_1d_sincos_pos_embed_from_grid -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormSingle - - -class LatteTransformer3DModel(ModelMixin, ConfigMixin, CacheMixin): - _supports_gradient_checkpointing = True - - """ - A 3D Transformer model for video-like data, paper: https://huggingface.co/papers/2401.03048, official code: - https://github.com/Vchitect/Latte - - Parameters: - num_attention_heads (`int`, *optional*, defaults to 16): The number of heads to use for multi-head attention. - attention_head_dim (`int`, *optional*, defaults to 88): The number of channels in each head. - in_channels (`int`, *optional*): - The number of channels in the input. - out_channels (`int`, *optional*): - The number of channels in the output. - num_layers (`int`, *optional*, defaults to 1): The number of layers of Transformer blocks to use. - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - cross_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use. - attention_bias (`bool`, *optional*): - Configure if the `TransformerBlocks` attention should contain a bias parameter. - sample_size (`int`, *optional*): The width of the latent images (specify if the input is **discrete**). - This is fixed during training since it is used to learn a number of position embeddings. - patch_size (`int`, *optional*): - The size of the patches to use in the patch embedding layer. - activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to use in feed-forward. - num_embeds_ada_norm ( `int`, *optional*): - The number of diffusion steps used during training. Pass if at least one of the norm_layers is - `AdaLayerNorm`. This is fixed during training since it is used to learn a number of embeddings that are - added to the hidden states. During inference, you can denoise for up to but not more steps than - `num_embeds_ada_norm`. - norm_type (`str`, *optional*, defaults to `"layer_norm"`): - The type of normalization to use. Options are `"layer_norm"` or `"ada_layer_norm"`. - norm_elementwise_affine (`bool`, *optional*, defaults to `True`): - Whether or not to use elementwise affine in normalization layers. - norm_eps (`float`, *optional*, defaults to 1e-5): The epsilon value to use in normalization layers. - caption_channels (`int`, *optional*): - The number of channels in the caption embeddings. - video_length (`int`, *optional*): - The number of frames in the video-like data. - """ - - _skip_layerwise_casting_patterns = ["pos_embed", "norm"] - - @register_to_config - def __init__( - self, - num_attention_heads: int = 16, - attention_head_dim: int = 88, - in_channels: int | None = None, - out_channels: int | None = None, - num_layers: int = 1, - dropout: float = 0.0, - cross_attention_dim: int | None = None, - attention_bias: bool = False, - sample_size: int = 64, - patch_size: int | None = None, - activation_fn: str = "geglu", - num_embeds_ada_norm: int | None = None, - norm_type: str = "layer_norm", - norm_elementwise_affine: bool = True, - norm_eps: float = 1e-5, - caption_channels: int = None, - video_length: int = 16, - ): - super().__init__() - inner_dim = num_attention_heads * attention_head_dim - - # 1. Define input layers - self.height = sample_size - self.width = sample_size - - interpolation_scale = self.config.sample_size // 64 - interpolation_scale = max(interpolation_scale, 1) - self.pos_embed = PatchEmbed( - height=sample_size, - width=sample_size, - patch_size=patch_size, - in_channels=in_channels, - embed_dim=inner_dim, - interpolation_scale=interpolation_scale, - ) - - # 2. Define spatial transformers blocks - self.transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlock( - inner_dim, - num_attention_heads, - attention_head_dim, - dropout=dropout, - cross_attention_dim=cross_attention_dim, - activation_fn=activation_fn, - num_embeds_ada_norm=num_embeds_ada_norm, - attention_bias=attention_bias, - norm_type=norm_type, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - ) - for d in range(num_layers) - ] - ) - - # 3. Define temporal transformers blocks - self.temporal_transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlock( - inner_dim, - num_attention_heads, - attention_head_dim, - dropout=dropout, - cross_attention_dim=None, - activation_fn=activation_fn, - num_embeds_ada_norm=num_embeds_ada_norm, - attention_bias=attention_bias, - norm_type=norm_type, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - ) - for d in range(num_layers) - ] - ) - - # 4. Define output layers - self.out_channels = in_channels if out_channels is None else out_channels - self.norm_out = nn.LayerNorm(inner_dim, elementwise_affine=False, eps=1e-6) - self.scale_shift_table = nn.Parameter(torch.randn(2, inner_dim) / inner_dim**0.5) - self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * self.out_channels) - - # 5. Latte other blocks. - self.adaln_single = AdaLayerNormSingle(inner_dim, use_additional_conditions=False) - self.caption_projection = PixArtAlphaTextProjection(in_features=caption_channels, hidden_size=inner_dim) - - # define temporal positional embedding - temp_pos_embed = get_1d_sincos_pos_embed_from_grid( - inner_dim, torch.arange(0, video_length).unsqueeze(1), output_type="pt" - ) # 1152 hidden size - self.register_buffer("temp_pos_embed", temp_pos_embed.float().unsqueeze(0), persistent=False) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - enable_temporal_attentions: bool = True, - return_dict: bool = True, - ): - """ - The [`LatteTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch size, channel, num_frame, height, width)`): - Input `hidden_states`. - timestep ( `torch.LongTensor`, *optional*): - Used to indicate denoising step. Optional timestep to be applied as an embedding in `AdaLayerNorm`. - encoder_hidden_states ( `torch.FloatTensor` of shape `(batch size, sequence len, embed dims)`, *optional*): - Conditional embeddings for cross attention layer. If not given, cross-attention defaults to - self-attention. - encoder_attention_mask ( `torch.Tensor`, *optional*): - Cross-attention mask applied to `encoder_hidden_states`. Two formats supported: - - * Mask `(batcheight, sequence_length)` True = keep, False = discard. - * Bias `(batcheight, 1, sequence_length)` 0 = keep, -10000 = discard. - - If `ndim == 2`: will be interpreted as a mask, then converted into a bias consistent with the format - above. This bias will be added to the cross-attention scores. - enable_temporal_attentions: - (`bool`, *optional*, defaults to `True`): Whether to enable temporal attentions. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unet_2d_condition.UNet2DConditionOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - - # Reshape hidden states - batch_size, channels, num_frame, height, width = hidden_states.shape - # batch_size channels num_frame height width -> (batch_size * num_frame) channels height width - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).reshape(-1, channels, height, width) - - # Input - height, width = ( - hidden_states.shape[-2] // self.config.patch_size, - hidden_states.shape[-1] // self.config.patch_size, - ) - num_patches = height * width - - hidden_states = self.pos_embed(hidden_states) # already add positional embeddings - - added_cond_kwargs = {"resolution": None, "aspect_ratio": None} - timestep, embedded_timestep = self.adaln_single( - timestep, added_cond_kwargs=added_cond_kwargs, batch_size=batch_size, hidden_dtype=hidden_states.dtype - ) - - # Prepare text embeddings for spatial block - # batch_size num_tokens hidden_size -> (batch_size * num_frame) num_tokens hidden_size - encoder_hidden_states = self.caption_projection(encoder_hidden_states) # 3 120 1152 - encoder_hidden_states_spatial = encoder_hidden_states.repeat_interleave( - num_frame, dim=0, output_size=encoder_hidden_states.shape[0] * num_frame - ).view(-1, encoder_hidden_states.shape[-2], encoder_hidden_states.shape[-1]) - - # Prepare timesteps for spatial and temporal block - timestep_spatial = timestep.repeat_interleave( - num_frame, dim=0, output_size=timestep.shape[0] * num_frame - ).view(-1, timestep.shape[-1]) - timestep_temp = timestep.repeat_interleave( - num_patches, dim=0, output_size=timestep.shape[0] * num_patches - ).view(-1, timestep.shape[-1]) - - # Spatial and temporal transformer blocks - for i, (spatial_block, temp_block) in enumerate( - zip(self.transformer_blocks, self.temporal_transformer_blocks) - ): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - spatial_block, - hidden_states, - None, # attention_mask - encoder_hidden_states_spatial, - encoder_attention_mask, - timestep_spatial, - None, # cross_attention_kwargs - None, # class_labels - ) - else: - hidden_states = spatial_block( - hidden_states, - None, # attention_mask - encoder_hidden_states_spatial, - encoder_attention_mask, - timestep_spatial, - None, # cross_attention_kwargs - None, # class_labels - ) - - if enable_temporal_attentions: - # (batch_size * num_frame) num_tokens hidden_size -> (batch_size * num_tokens) num_frame hidden_size - hidden_states = hidden_states.reshape( - batch_size, -1, hidden_states.shape[-2], hidden_states.shape[-1] - ).permute(0, 2, 1, 3) - hidden_states = hidden_states.reshape(-1, hidden_states.shape[-2], hidden_states.shape[-1]) - - if i == 0 and num_frame > 1: - hidden_states = hidden_states + self.temp_pos_embed.to(hidden_states.dtype) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - temp_block, - hidden_states, - None, # attention_mask - None, # encoder_hidden_states - None, # encoder_attention_mask - timestep_temp, - None, # cross_attention_kwargs - None, # class_labels - ) - else: - hidden_states = temp_block( - hidden_states, - None, # attention_mask - None, # encoder_hidden_states - None, # encoder_attention_mask - timestep_temp, - None, # cross_attention_kwargs - None, # class_labels - ) - - # (batch_size * num_tokens) num_frame hidden_size -> (batch_size * num_frame) num_tokens hidden_size - hidden_states = hidden_states.reshape( - batch_size, -1, hidden_states.shape[-2], hidden_states.shape[-1] - ).permute(0, 2, 1, 3) - hidden_states = hidden_states.reshape(-1, hidden_states.shape[-2], hidden_states.shape[-1]) - - embedded_timestep = embedded_timestep.repeat_interleave( - num_frame, dim=0, output_size=embedded_timestep.shape[0] * num_frame - ).view(-1, embedded_timestep.shape[-1]) - shift, scale = (self.scale_shift_table[None] + embedded_timestep[:, None]).chunk(2, dim=1) - hidden_states = self.norm_out(hidden_states) - # Modulation - hidden_states = hidden_states * (1 + scale) + shift - hidden_states = self.proj_out(hidden_states) - - # unpatchify - if self.adaln_single is None: - height = width = int(hidden_states.shape[1] ** 0.5) - hidden_states = hidden_states.reshape( - shape=(-1, height, width, self.config.patch_size, self.config.patch_size, self.out_channels) - ) - hidden_states = torch.einsum("nhwpqc->nchpwq", hidden_states) - output = hidden_states.reshape( - shape=(-1, self.out_channels, height * self.config.patch_size, width * self.config.patch_size) - ) - output = output.reshape(batch_size, -1, output.shape[-3], output.shape[-2], output.shape[-1]).permute( - 0, 2, 1, 3, 4 - ) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/lumina_nextdit2d.py b/diffusers/models/transformers/lumina_nextdit2d.py deleted file mode 100644 index 73468b5d853fb67fc13db48caa71e1e8d8235daf..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/lumina_nextdit2d.py +++ /dev/null @@ -1,356 +0,0 @@ -# Copyright 2025 Alpha-VLLM Authors and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ..attention import LuminaFeedForward -from ..attention_processor import Attention, LuminaAttnProcessor2_0 -from ..embeddings import ( - LuminaCombinedTimestepCaptionEmbedding, - LuminaPatchEmbed, -) -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import LuminaLayerNormContinuous, LuminaRMSNormZero, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class LuminaNextDiTBlock(nn.Module): - """ - A LuminaNextDiTBlock for LuminaNextDiT2DModel. - - Parameters: - dim (`int`): Embedding dimension of the input features. - num_attention_heads (`int`): Number of attention heads. - num_kv_heads (`int`): - Number of attention heads in key and value features (if using GQA), or set to None for the same as query. - multiple_of (`int`): The number of multiple of ffn layer. - ffn_dim_multiplier (`float`): The multiplier factor of ffn layer dimension. - norm_eps (`float`): The eps for norm layer. - qk_norm (`bool`): normalization for query and key. - cross_attention_dim (`int`): Cross attention embedding dimension of the input text prompt hidden_states. - norm_elementwise_affine (`bool`, *optional*, defaults to True), - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - num_kv_heads: int, - multiple_of: int, - ffn_dim_multiplier: float, - norm_eps: float, - qk_norm: bool, - cross_attention_dim: int, - norm_elementwise_affine: bool = True, - ) -> None: - super().__init__() - self.head_dim = dim // num_attention_heads - - self.gate = nn.Parameter(torch.zeros([num_attention_heads])) - - # Self-attention - self.attn1 = Attention( - query_dim=dim, - cross_attention_dim=None, - dim_head=dim // num_attention_heads, - qk_norm="layer_norm_across_heads" if qk_norm else None, - heads=num_attention_heads, - kv_heads=num_kv_heads, - eps=1e-5, - bias=False, - out_bias=False, - processor=LuminaAttnProcessor2_0(), - ) - self.attn1.to_out = nn.Identity() - - # Cross-attention - self.attn2 = Attention( - query_dim=dim, - cross_attention_dim=cross_attention_dim, - dim_head=dim // num_attention_heads, - qk_norm="layer_norm_across_heads" if qk_norm else None, - heads=num_attention_heads, - kv_heads=num_kv_heads, - eps=1e-5, - bias=False, - out_bias=False, - processor=LuminaAttnProcessor2_0(), - ) - - self.feed_forward = LuminaFeedForward( - dim=dim, - inner_dim=int(4 * 2 * dim / 3), - multiple_of=multiple_of, - ffn_dim_multiplier=ffn_dim_multiplier, - ) - - self.norm1 = LuminaRMSNormZero( - embedding_dim=dim, - norm_eps=norm_eps, - norm_elementwise_affine=norm_elementwise_affine, - ) - self.ffn_norm1 = RMSNorm(dim, eps=norm_eps, elementwise_affine=norm_elementwise_affine) - - self.norm2 = RMSNorm(dim, eps=norm_eps, elementwise_affine=norm_elementwise_affine) - self.ffn_norm2 = RMSNorm(dim, eps=norm_eps, elementwise_affine=norm_elementwise_affine) - - self.norm1_context = RMSNorm(cross_attention_dim, eps=norm_eps, elementwise_affine=norm_elementwise_affine) - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor, - image_rotary_emb: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_mask: torch.Tensor, - temb: torch.Tensor, - cross_attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor: - """ - Perform a forward pass through the LuminaNextDiTBlock. - - Parameters: - hidden_states (`torch.Tensor`): The input of hidden_states for LuminaNextDiTBlock. - attention_mask (`torch.Tensor): The input of hidden_states corresponse attention mask. - image_rotary_emb (`torch.Tensor`): Precomputed cosine and sine frequencies. - encoder_hidden_states: (`torch.Tensor`): The hidden_states of text prompt are processed by Gemma encoder. - encoder_mask (`torch.Tensor`): The hidden_states of text prompt attention mask. - temb (`torch.Tensor`): Timestep embedding with text prompt embedding. - cross_attention_kwargs (`dict[str, Any]`): kwargs for cross attention. - """ - residual = hidden_states - - # Self-attention - norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(hidden_states, temb) - self_attn_output = self.attn1( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_hidden_states, - attention_mask=attention_mask, - query_rotary_emb=image_rotary_emb, - key_rotary_emb=image_rotary_emb, - **cross_attention_kwargs, - ) - - # Cross-attention - norm_encoder_hidden_states = self.norm1_context(encoder_hidden_states) - cross_attn_output = self.attn2( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=encoder_mask, - query_rotary_emb=image_rotary_emb, - key_rotary_emb=None, - **cross_attention_kwargs, - ) - cross_attn_output = cross_attn_output * self.gate.tanh().view(1, 1, -1, 1) - mixed_attn_output = self_attn_output + cross_attn_output - mixed_attn_output = mixed_attn_output.flatten(-2) - # linear proj - hidden_states = self.attn2.to_out[0](mixed_attn_output) - - hidden_states = residual + gate_msa.unsqueeze(1).tanh() * self.norm2(hidden_states) - - mlp_output = self.feed_forward(self.ffn_norm1(hidden_states) * (1 + scale_mlp.unsqueeze(1))) - - hidden_states = hidden_states + gate_mlp.unsqueeze(1).tanh() * self.ffn_norm2(mlp_output) - - return hidden_states - - -class LuminaNextDiT2DModel(ModelMixin, ConfigMixin): - """ - LuminaNextDiT: Diffusion model with a Transformer backbone. - - Inherit ModelMixin and ConfigMixin to be compatible with the sampler StableDiffusionPipeline of diffusers. - - Parameters: - sample_size (`int`): The width of the latent images. This is fixed during training since - it is used to learn a number of position embeddings. - patch_size (`int`, *optional*, (`int`, *optional*, defaults to 2): - The size of each patch in the image. This parameter defines the resolution of patches fed into the model. - in_channels (`int`, *optional*, defaults to 4): - The number of input channels for the model. Typically, this matches the number of channels in the input - images. - hidden_size (`int`, *optional*, defaults to 4096): - The dimensionality of the hidden layers in the model. This parameter determines the width of the model's - hidden representations. - num_layers (`int`, *optional*, default to 32): - The number of layers in the model. This defines the depth of the neural network. - num_attention_heads (`int`, *optional*, defaults to 32): - The number of attention heads in each attention layer. This parameter specifies how many separate attention - mechanisms are used. - num_kv_heads (`int`, *optional*, defaults to 8): - The number of key-value heads in the attention mechanism, if different from the number of attention heads. - If None, it defaults to num_attention_heads. - multiple_of (`int`, *optional*, defaults to 256): - A factor that the hidden size should be a multiple of. This can help optimize certain hardware - configurations. - ffn_dim_multiplier (`float`, *optional*): - A multiplier for the dimensionality of the feed-forward network. If None, it uses a default value based on - the model configuration. - norm_eps (`float`, *optional*, defaults to 1e-5): - A small value added to the denominator for numerical stability in normalization layers. - learn_sigma (`bool`, *optional*, defaults to True): - Whether the model should learn the sigma parameter, which might be related to uncertainty or variance in - predictions. - qk_norm (`bool`, *optional*, defaults to True): - Indicates if the queries and keys in the attention mechanism should be normalized. - cross_attention_dim (`int`, *optional*, defaults to 2048): - The dimensionality of the text embeddings. This parameter defines the size of the text representations used - in the model. - scaling_factor (`float`, *optional*, defaults to 1.0): - A scaling factor applied to certain parameters or layers in the model. This can be used for adjusting the - overall scale of the model's operations. - """ - - _skip_layerwise_casting_patterns = ["patch_embedder", "norm", "ffn_norm"] - - @register_to_config - def __init__( - self, - sample_size: int = 128, - patch_size: int | None = 2, - in_channels: int | None = 4, - hidden_size: int | None = 2304, - num_layers: int | None = 32, - num_attention_heads: int | None = 32, - num_kv_heads: int | None = None, - multiple_of: int | None = 256, - ffn_dim_multiplier: float | None = None, - norm_eps: float | None = 1e-5, - learn_sigma: bool | None = True, - qk_norm: bool | None = True, - cross_attention_dim: int | None = 2048, - scaling_factor: float | None = 1.0, - ) -> None: - super().__init__() - self.sample_size = sample_size - self.patch_size = patch_size - self.in_channels = in_channels - self.out_channels = in_channels * 2 if learn_sigma else in_channels - self.hidden_size = hidden_size - self.num_attention_heads = num_attention_heads - self.head_dim = hidden_size // num_attention_heads - self.scaling_factor = scaling_factor - - self.patch_embedder = LuminaPatchEmbed( - patch_size=patch_size, in_channels=in_channels, embed_dim=hidden_size, bias=True - ) - - self.pad_token = nn.Parameter(torch.empty(hidden_size)) - - self.time_caption_embed = LuminaCombinedTimestepCaptionEmbedding( - hidden_size=min(hidden_size, 1024), cross_attention_dim=cross_attention_dim - ) - - self.layers = nn.ModuleList( - [ - LuminaNextDiTBlock( - hidden_size, - num_attention_heads, - num_kv_heads, - multiple_of, - ffn_dim_multiplier, - norm_eps, - qk_norm, - cross_attention_dim, - ) - for _ in range(num_layers) - ] - ) - self.norm_out = LuminaLayerNormContinuous( - embedding_dim=hidden_size, - conditioning_embedding_dim=min(hidden_size, 1024), - elementwise_affine=False, - eps=1e-6, - bias=True, - out_dim=patch_size * patch_size * self.out_channels, - ) - # self.final_layer = LuminaFinalLayer(hidden_size, patch_size, self.out_channels) - - assert (hidden_size // num_attention_heads) % 4 == 0, "2d rope needs head dim to be divisible by 4" - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_mask: torch.Tensor, - image_rotary_emb: torch.Tensor, - cross_attention_kwargs: dict[str, Any] = None, - return_dict=True, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - Forward pass of LuminaNextDiT. - - Parameters: - hidden_states (torch.Tensor): Input tensor of shape (N, C, H, W). - timestep (torch.Tensor): Tensor of diffusion timesteps of shape (N,). - encoder_hidden_states (torch.Tensor): Tensor of caption features of shape (N, D). - encoder_mask (torch.Tensor): Tensor of caption masks of shape (N, L). - image_rotary_emb (`torch.Tensor`): - Pre-computed rotary positional embeddings. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - [`~models.transformer_2d.Transformer2DModelOutput`] or `tuple`: - If `return_dict` is True, a [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise - a plain `tuple` is returned. - """ - hidden_states, mask, img_size, image_rotary_emb = self.patch_embedder(hidden_states, image_rotary_emb) - image_rotary_emb = image_rotary_emb.to(hidden_states.device) - - temb = self.time_caption_embed(timestep, encoder_hidden_states, encoder_mask) - - encoder_mask = encoder_mask.bool() - for layer in self.layers: - hidden_states = layer( - hidden_states, - mask, - image_rotary_emb, - encoder_hidden_states, - encoder_mask, - temb=temb, - cross_attention_kwargs=cross_attention_kwargs, - ) - - hidden_states = self.norm_out(hidden_states, temb) - - # unpatchify - height_tokens = width_tokens = self.patch_size - height, width = img_size[0] - batch_size = hidden_states.size(0) - sequence_length = (height // height_tokens) * (width // width_tokens) - hidden_states = hidden_states[:, :sequence_length].view( - batch_size, height // height_tokens, width // width_tokens, height_tokens, width_tokens, self.out_channels - ) - output = hidden_states.permute(0, 5, 1, 3, 2, 4).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/pixart_transformer_2d.py b/diffusers/models/transformers/pixart_transformer_2d.py deleted file mode 100644 index e5e6178eaf4a7885fc99587db7a399f1113124b3..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/pixart_transformer_2d.py +++ /dev/null @@ -1,362 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from typing import Any - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ..attention import AttentionMixin, BasicTransformerBlock -from ..attention_processor import Attention, AttnProcessor, FusedAttnProcessor2_0 -from ..embeddings import PatchEmbed, PixArtAlphaTextProjection -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormSingle - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class PixArtTransformer2DModel(ModelMixin, AttentionMixin, ConfigMixin): - r""" - A 2D Transformer model as introduced in PixArt family of models (https://huggingface.co/papers/2310.00426, - https://huggingface.co/papers/2403.04692). - - Parameters: - num_attention_heads (int, optional, defaults to 16): The number of heads to use for multi-head attention. - attention_head_dim (int, optional, defaults to 72): The number of channels in each head. - in_channels (int, defaults to 4): The number of channels in the input. - out_channels (int, optional): - The number of channels in the output. Specify this parameter if the output channel number differs from the - input. - num_layers (int, optional, defaults to 28): The number of layers of Transformer blocks to use. - dropout (float, optional, defaults to 0.0): The dropout probability to use within the Transformer blocks. - norm_num_groups (int, optional, defaults to 32): - Number of groups for group normalization within Transformer blocks. - cross_attention_dim (int, optional): - The dimensionality for cross-attention layers, typically matching the encoder's hidden dimension. - attention_bias (bool, optional, defaults to True): - Configure if the Transformer blocks' attention should contain a bias parameter. - sample_size (int, defaults to 128): - The width of the latent images. This parameter is fixed during training. - patch_size (int, defaults to 2): - Size of the patches the model processes, relevant for architectures working on non-sequential data. - activation_fn (str, optional, defaults to "gelu-approximate"): - Activation function to use in feed-forward networks within Transformer blocks. - num_embeds_ada_norm (int, optional, defaults to 1000): - Number of embeddings for AdaLayerNorm, fixed during training and affects the maximum denoising steps during - inference. - upcast_attention (bool, optional, defaults to False): - If true, upcasts the attention mechanism dimensions for potentially improved performance. - norm_type (str, optional, defaults to "ada_norm_zero"): - Specifies the type of normalization used, can be 'ada_norm_zero'. - norm_elementwise_affine (bool, optional, defaults to False): - If true, enables element-wise affine parameters in the normalization layers. - norm_eps (float, optional, defaults to 1e-6): - A small constant added to the denominator in normalization layers to prevent division by zero. - interpolation_scale (int, optional): Scale factor to use during interpolating the position embeddings. - use_additional_conditions (bool, optional): If we're using additional conditions as inputs. - attention_type (str, optional, defaults to "default"): Kind of attention mechanism to be used. - caption_channels (int, optional, defaults to None): - Number of channels to use for projecting the caption embeddings. - use_linear_projection (bool, optional, defaults to False): - Deprecated argument. Will be removed in a future version. - num_vector_embeds (bool, optional, defaults to False): - Deprecated argument. Will be removed in a future version. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["BasicTransformerBlock", "PatchEmbed"] - _skip_layerwise_casting_patterns = ["pos_embed", "norm", "adaln_single"] - - @register_to_config - def __init__( - self, - num_attention_heads: int = 16, - attention_head_dim: int = 72, - in_channels: int = 4, - out_channels: int | None = 8, - num_layers: int = 28, - dropout: float = 0.0, - norm_num_groups: int = 32, - cross_attention_dim: int | None = 1152, - attention_bias: bool = True, - sample_size: int = 128, - patch_size: int = 2, - activation_fn: str = "gelu-approximate", - num_embeds_ada_norm: int | None = 1000, - upcast_attention: bool = False, - norm_type: str = "ada_norm_single", - norm_elementwise_affine: bool = False, - norm_eps: float = 1e-6, - interpolation_scale: int | None = None, - use_additional_conditions: bool | None = None, - caption_channels: int | None = None, - attention_type: str | None = "default", - ): - super().__init__() - - # Validate inputs. - if norm_type != "ada_norm_single": - raise NotImplementedError( - f"Forward pass is not implemented when `patch_size` is not None and `norm_type` is '{norm_type}'." - ) - elif norm_type == "ada_norm_single" and num_embeds_ada_norm is None: - raise ValueError( - f"When using a `patch_size` and this `norm_type` ({norm_type}), `num_embeds_ada_norm` cannot be None." - ) - - # Set some common variables used across the board. - self.attention_head_dim = attention_head_dim - self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim - self.out_channels = in_channels if out_channels is None else out_channels - if use_additional_conditions is None: - if sample_size == 128: - use_additional_conditions = True - else: - use_additional_conditions = False - self.use_additional_conditions = use_additional_conditions - - self.gradient_checkpointing = False - - # 2. Initialize the position embedding and transformer blocks. - self.height = self.config.sample_size - self.width = self.config.sample_size - - interpolation_scale = ( - self.config.interpolation_scale - if self.config.interpolation_scale is not None - else max(self.config.sample_size // 64, 1) - ) - self.pos_embed = PatchEmbed( - height=self.config.sample_size, - width=self.config.sample_size, - patch_size=self.config.patch_size, - in_channels=self.config.in_channels, - embed_dim=self.inner_dim, - interpolation_scale=interpolation_scale, - ) - - self.transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlock( - self.inner_dim, - self.config.num_attention_heads, - self.config.attention_head_dim, - dropout=self.config.dropout, - cross_attention_dim=self.config.cross_attention_dim, - activation_fn=self.config.activation_fn, - num_embeds_ada_norm=self.config.num_embeds_ada_norm, - attention_bias=self.config.attention_bias, - upcast_attention=self.config.upcast_attention, - norm_type=norm_type, - norm_elementwise_affine=self.config.norm_elementwise_affine, - norm_eps=self.config.norm_eps, - attention_type=self.config.attention_type, - ) - for _ in range(self.config.num_layers) - ] - ) - - # 3. Output blocks. - self.norm_out = nn.LayerNorm(self.inner_dim, elementwise_affine=False, eps=1e-6) - self.scale_shift_table = nn.Parameter(torch.randn(2, self.inner_dim) / self.inner_dim**0.5) - self.proj_out = nn.Linear(self.inner_dim, self.config.patch_size * self.config.patch_size * self.out_channels) - - self.adaln_single = AdaLayerNormSingle( - self.inner_dim, use_additional_conditions=self.use_additional_conditions - ) - self.caption_projection = None - if self.config.caption_channels is not None: - self.caption_projection = PixArtAlphaTextProjection( - in_features=self.config.caption_channels, hidden_size=self.inner_dim - ) - - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - - Safe to just use `AttnProcessor()` as PixArt doesn't have any exotic attention processors in default model. - """ - self.set_attn_processor(AttnProcessor()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - timestep: torch.LongTensor | None = None, - added_cond_kwargs: dict[str, torch.Tensor] = None, - cross_attention_kwargs: dict[str, Any] = None, - attention_mask: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - return_dict: bool = True, - ): - """ - The [`PixArtTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.FloatTensor` of shape `(batch size, channel, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.FloatTensor` of shape `(batch size, sequence len, embed dims)`, *optional*): - Conditional embeddings for cross attention layer. If not given, cross-attention defaults to - self-attention. - timestep (`torch.LongTensor`, *optional*): - Used to indicate denoising step. Optional timestep to be applied as an embedding in `AdaLayerNorm`. - added_cond_kwargs: (`dict[str, Any]`, *optional*): Additional conditions to be used as inputs. - cross_attention_kwargs ( `dict[str, Any]`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - attention_mask ( `torch.Tensor`, *optional*): - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask - is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large - negative values to the attention scores corresponding to "discard" tokens. - encoder_attention_mask ( `torch.Tensor`, *optional*): - Cross-attention mask applied to `encoder_hidden_states`. Two formats supported: - - * Mask `(batch, sequence_length)` True = keep, False = discard. - * Bias `(batch, 1, sequence_length)` 0 = keep, -10000 = discard. - - If `ndim == 2`: will be interpreted as a mask, then converted into a bias consistent with the format - above. This bias will be added to the cross-attention scores. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unets.unet_2d_condition.UNet2DConditionOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - if self.use_additional_conditions and added_cond_kwargs is None: - raise ValueError("`added_cond_kwargs` cannot be None when using additional conditions for `adaln_single`.") - - # ensure attention_mask is a bias, and give it a singleton query_tokens dimension. - # we may have done this conversion already, e.g. if we came here via UNet2DConditionModel#forward. - # we can tell by counting dims; if ndim == 2: it's a mask rather than a bias. - # expects mask of shape: - # [batch, key_tokens] - # adds singleton query_tokens dimension: - # [batch, 1, key_tokens] - # this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes: - # [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn) - # [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn) - if attention_mask is not None and attention_mask.ndim == 2: - # assume that mask is expressed as: - # (1 = keep, 0 = discard) - # convert mask into a bias that can be added to attention scores: - # (keep = +0, discard = -10000.0) - attention_mask = (1 - attention_mask.to(hidden_states.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # convert encoder_attention_mask to a bias the same way we do for attention_mask - if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2: - encoder_attention_mask = (1 - encoder_attention_mask.to(hidden_states.dtype)) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - # 1. Input - batch_size = hidden_states.shape[0] - height, width = ( - hidden_states.shape[-2] // self.config.patch_size, - hidden_states.shape[-1] // self.config.patch_size, - ) - hidden_states = self.pos_embed(hidden_states) - - timestep, embedded_timestep = self.adaln_single( - timestep, added_cond_kwargs, batch_size=batch_size, hidden_dtype=hidden_states.dtype - ) - - if self.caption_projection is not None: - encoder_hidden_states = self.caption_projection(encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states.view(batch_size, -1, hidden_states.shape[-1]) - - # 2. Blocks - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - attention_mask, - encoder_hidden_states, - encoder_attention_mask, - timestep, - cross_attention_kwargs, - None, - ) - else: - hidden_states = block( - hidden_states, - attention_mask=attention_mask, - encoder_hidden_states=encoder_hidden_states, - encoder_attention_mask=encoder_attention_mask, - timestep=timestep, - cross_attention_kwargs=cross_attention_kwargs, - class_labels=None, - ) - - # 3. Output - shift, scale = ( - self.scale_shift_table[None] + embedded_timestep[:, None].to(self.scale_shift_table.device) - ).chunk(2, dim=1) - hidden_states = self.norm_out(hidden_states) - # Modulation - hidden_states = hidden_states * (1 + scale.to(hidden_states.device)) + shift.to(hidden_states.device) - hidden_states = self.proj_out(hidden_states) - hidden_states = hidden_states.squeeze(1) - - # unpatchify - hidden_states = hidden_states.reshape( - shape=(-1, height, width, self.config.patch_size, self.config.patch_size, self.out_channels) - ) - hidden_states = torch.einsum("nhwpqc->nchpwq", hidden_states) - output = hidden_states.reshape( - shape=(-1, self.out_channels, height * self.config.patch_size, width * self.config.patch_size) - ) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/prior_transformer.py b/diffusers/models/transformers/prior_transformer.py deleted file mode 100644 index ace2b529c4f2109e1069a5c635feb898e8752245..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/prior_transformer.py +++ /dev/null @@ -1,322 +0,0 @@ -from dataclasses import dataclass - -import torch -import torch.nn.functional as F -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin, UNet2DConditionLoadersMixin -from ...utils import BaseOutput -from ..attention import AttentionMixin, BasicTransformerBlock -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - AttnAddedKVProcessor, - AttnProcessor, -) -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin - - -@dataclass -class PriorTransformerOutput(BaseOutput): - """ - The output of [`PriorTransformer`]. - - Args: - predicted_image_embedding (`torch.Tensor` of shape `(batch_size, embedding_dim)`): - The predicted CLIP image embedding conditioned on the CLIP text embedding input. - """ - - predicted_image_embedding: torch.Tensor - - -class PriorTransformer(ModelMixin, AttentionMixin, ConfigMixin, UNet2DConditionLoadersMixin, PeftAdapterMixin): - """ - A Prior Transformer model. - - Parameters: - num_attention_heads (`int`, *optional*, defaults to 32): The number of heads to use for multi-head attention. - attention_head_dim (`int`, *optional*, defaults to 64): The number of channels in each head. - num_layers (`int`, *optional*, defaults to 20): The number of layers of Transformer blocks to use. - embedding_dim (`int`, *optional*, defaults to 768): The dimension of the model input `hidden_states` - num_embeddings (`int`, *optional*, defaults to 77): - The number of embeddings of the model input `hidden_states` - additional_embeddings (`int`, *optional*, defaults to 4): The number of additional tokens appended to the - projected `hidden_states`. The actual length of the used `hidden_states` is `num_embeddings + - additional_embeddings`. - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - time_embed_act_fn (`str`, *optional*, defaults to 'silu'): - The activation function to use to create timestep embeddings. - norm_in_type (`str`, *optional*, defaults to None): The normalization layer to apply on hidden states before - passing to Transformer blocks. Set it to `None` if normalization is not needed. - embedding_proj_norm_type (`str`, *optional*, defaults to None): - The normalization layer to apply on the input `proj_embedding`. Set it to `None` if normalization is not - needed. - encoder_hid_proj_type (`str`, *optional*, defaults to `linear`): - The projection layer to apply on the input `encoder_hidden_states`. Set it to `None` if - `encoder_hidden_states` is `None`. - added_emb_type (`str`, *optional*, defaults to `prd`): Additional embeddings to condition the model. - Choose from `prd` or `None`. if choose `prd`, it will prepend a token indicating the (quantized) dot - product between the text embedding and image embedding as proposed in the unclip paper - https://huggingface.co/papers/2204.06125 If it is `None`, no additional embeddings will be prepended. - time_embed_dim (`int, *optional*, defaults to None): The dimension of timestep embeddings. - If None, will be set to `num_attention_heads * attention_head_dim` - embedding_proj_dim (`int`, *optional*, default to None): - The dimension of `proj_embedding`. If None, will be set to `embedding_dim`. - clip_embed_dim (`int`, *optional*, default to None): - The dimension of the output. If None, will be set to `embedding_dim`. - """ - - @register_to_config - def __init__( - self, - num_attention_heads: int = 32, - attention_head_dim: int = 64, - num_layers: int = 20, - embedding_dim: int = 768, - num_embeddings=77, - additional_embeddings=4, - dropout: float = 0.0, - time_embed_act_fn: str = "silu", - norm_in_type: str | None = None, # layer - embedding_proj_norm_type: str | None = None, # layer - encoder_hid_proj_type: str | None = "linear", # linear - added_emb_type: str | None = "prd", # prd - time_embed_dim: int | None = None, - embedding_proj_dim: int | None = None, - clip_embed_dim: int | None = None, - ): - super().__init__() - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - inner_dim = num_attention_heads * attention_head_dim - self.additional_embeddings = additional_embeddings - - time_embed_dim = time_embed_dim or inner_dim - embedding_proj_dim = embedding_proj_dim or embedding_dim - clip_embed_dim = clip_embed_dim or embedding_dim - - self.time_proj = Timesteps(inner_dim, True, 0) - self.time_embedding = TimestepEmbedding(inner_dim, time_embed_dim, out_dim=inner_dim, act_fn=time_embed_act_fn) - - self.proj_in = nn.Linear(embedding_dim, inner_dim) - - if embedding_proj_norm_type is None: - self.embedding_proj_norm = None - elif embedding_proj_norm_type == "layer": - self.embedding_proj_norm = nn.LayerNorm(embedding_proj_dim) - else: - raise ValueError(f"unsupported embedding_proj_norm_type: {embedding_proj_norm_type}") - - self.embedding_proj = nn.Linear(embedding_proj_dim, inner_dim) - - if encoder_hid_proj_type is None: - self.encoder_hidden_states_proj = None - elif encoder_hid_proj_type == "linear": - self.encoder_hidden_states_proj = nn.Linear(embedding_dim, inner_dim) - else: - raise ValueError(f"unsupported encoder_hid_proj_type: {encoder_hid_proj_type}") - - self.positional_embedding = nn.Parameter(torch.zeros(1, num_embeddings + additional_embeddings, inner_dim)) - - if added_emb_type == "prd": - self.prd_embedding = nn.Parameter(torch.zeros(1, 1, inner_dim)) - elif added_emb_type is None: - self.prd_embedding = None - else: - raise ValueError( - f"`added_emb_type`: {added_emb_type} is not supported. Make sure to choose one of `'prd'` or `None`." - ) - - self.transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlock( - inner_dim, - num_attention_heads, - attention_head_dim, - dropout=dropout, - activation_fn="gelu", - attention_bias=True, - ) - for d in range(num_layers) - ] - ) - - if norm_in_type == "layer": - self.norm_in = nn.LayerNorm(inner_dim) - elif norm_in_type is None: - self.norm_in = None - else: - raise ValueError(f"Unsupported norm_in_type: {norm_in_type}.") - - self.norm_out = nn.LayerNorm(inner_dim) - - self.proj_to_clip_embeddings = nn.Linear(inner_dim, clip_embed_dim) - - causal_attention_mask = torch.full( - [num_embeddings + additional_embeddings, num_embeddings + additional_embeddings], -10000.0 - ) - causal_attention_mask.triu_(1) - causal_attention_mask = causal_attention_mask[None, ...] - self.register_buffer("causal_attention_mask", causal_attention_mask, persistent=False) - - self.clip_mean = nn.Parameter(torch.zeros(1, clip_embed_dim)) - self.clip_std = nn.Parameter(torch.zeros(1, clip_embed_dim)) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - def forward( - self, - hidden_states, - timestep: torch.Tensor | float | int, - proj_embedding: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.BoolTensor | None = None, - return_dict: bool = True, - ): - """ - The [`PriorTransformer`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, embedding_dim)`): - The currently predicted image embeddings. - timestep (`torch.LongTensor`): - Current denoising step. - proj_embedding (`torch.Tensor` of shape `(batch_size, embedding_dim)`): - Projected embedding vector the denoising process is conditioned on. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, num_embeddings, embedding_dim)`): - Hidden states of the text embeddings the denoising process is conditioned on. - attention_mask (`torch.BoolTensor` of shape `(batch_size, num_embeddings)`): - Text mask for the text embeddings. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformers.prior_transformer.PriorTransformerOutput`] instead of - a plain tuple. - - Returns: - [`~models.transformers.prior_transformer.PriorTransformerOutput`] or `tuple`: - If return_dict is True, a [`~models.transformers.prior_transformer.PriorTransformerOutput`] is - returned, otherwise a tuple is returned where the first element is the sample tensor. - """ - batch_size = hidden_states.shape[0] - - timesteps = timestep - if not torch.is_tensor(timesteps): - timesteps = torch.tensor([timesteps], dtype=torch.long, device=hidden_states.device) - elif torch.is_tensor(timesteps) and len(timesteps.shape) == 0: - timesteps = timesteps[None].to(hidden_states.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - timesteps = timesteps * torch.ones(batch_size, dtype=timesteps.dtype, device=timesteps.device) - - timesteps_projected = self.time_proj(timesteps) - - # timesteps does not contain any weights and will always return f32 tensors - # but time_embedding might be fp16, so we need to cast here. - timesteps_projected = timesteps_projected.to(dtype=self.dtype) - time_embeddings = self.time_embedding(timesteps_projected) - - if self.embedding_proj_norm is not None: - proj_embedding = self.embedding_proj_norm(proj_embedding) - - proj_embeddings = self.embedding_proj(proj_embedding) - if self.encoder_hidden_states_proj is not None and encoder_hidden_states is not None: - encoder_hidden_states = self.encoder_hidden_states_proj(encoder_hidden_states) - elif self.encoder_hidden_states_proj is not None and encoder_hidden_states is None: - raise ValueError("`encoder_hidden_states_proj` requires `encoder_hidden_states` to be set") - - hidden_states = self.proj_in(hidden_states) - - positional_embeddings = self.positional_embedding.to(hidden_states.dtype) - - additional_embeds = [] - additional_embeddings_len = 0 - - if encoder_hidden_states is not None: - additional_embeds.append(encoder_hidden_states) - additional_embeddings_len += encoder_hidden_states.shape[1] - - if len(proj_embeddings.shape) == 2: - proj_embeddings = proj_embeddings[:, None, :] - - if len(hidden_states.shape) == 2: - hidden_states = hidden_states[:, None, :] - - additional_embeds = additional_embeds + [ - proj_embeddings, - time_embeddings[:, None, :], - hidden_states, - ] - - if self.prd_embedding is not None: - prd_embedding = self.prd_embedding.to(hidden_states.dtype).expand(batch_size, -1, -1) - additional_embeds.append(prd_embedding) - - hidden_states = torch.cat( - additional_embeds, - dim=1, - ) - - # Allow positional_embedding to not include the `addtional_embeddings` and instead pad it with zeros for these additional tokens - additional_embeddings_len = additional_embeddings_len + proj_embeddings.shape[1] + 1 - if positional_embeddings.shape[1] < hidden_states.shape[1]: - positional_embeddings = F.pad( - positional_embeddings, - ( - 0, - 0, - additional_embeddings_len, - self.prd_embedding.shape[1] if self.prd_embedding is not None else 0, - ), - value=0.0, - ) - - hidden_states = hidden_states + positional_embeddings - - if attention_mask is not None: - attention_mask = (1 - attention_mask.to(hidden_states.dtype)) * -10000.0 - attention_mask = F.pad(attention_mask, (0, self.additional_embeddings), value=0.0) - attention_mask = (attention_mask[:, None, :] + self.causal_attention_mask).to(hidden_states.dtype) - attention_mask = attention_mask.repeat_interleave( - self.config.num_attention_heads, - dim=0, - output_size=attention_mask.shape[0] * self.config.num_attention_heads, - ) - - if self.norm_in is not None: - hidden_states = self.norm_in(hidden_states) - - for block in self.transformer_blocks: - hidden_states = block(hidden_states, attention_mask=attention_mask) - - hidden_states = self.norm_out(hidden_states) - - if self.prd_embedding is not None: - hidden_states = hidden_states[:, -1] - else: - hidden_states = hidden_states[:, additional_embeddings_len:] - - predicted_image_embedding = self.proj_to_clip_embeddings(hidden_states) - - if not return_dict: - return (predicted_image_embedding,) - - return PriorTransformerOutput(predicted_image_embedding=predicted_image_embedding) - - def post_process_latents(self, prior_latents): - prior_latents = (prior_latents * self.clip_std) + self.clip_mean - return prior_latents diff --git a/diffusers/models/transformers/sana_transformer.py b/diffusers/models/transformers/sana_transformer.py deleted file mode 100644 index 1451750d50ef9c38da66345c6a5b46589103d7cb..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/sana_transformer.py +++ /dev/null @@ -1,549 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch -import torch.nn.functional as F -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ..attention import AttentionMixin -from ..attention_processor import ( - Attention, - SanaLinearAttnProcessor2_0, -) -from ..embeddings import PatchEmbed, PixArtAlphaTextProjection, TimestepEmbedding, Timesteps -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormSingle, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class GLUMBConv(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - expand_ratio: float = 4, - norm_type: str | None = None, - residual_connection: bool = True, - ) -> None: - super().__init__() - - hidden_channels = int(expand_ratio * in_channels) - self.norm_type = norm_type - self.residual_connection = residual_connection - - self.nonlinearity = nn.SiLU() - self.conv_inverted = nn.Conv2d(in_channels, hidden_channels * 2, 1, 1, 0) - self.conv_depth = nn.Conv2d(hidden_channels * 2, hidden_channels * 2, 3, 1, 1, groups=hidden_channels * 2) - self.conv_point = nn.Conv2d(hidden_channels, out_channels, 1, 1, 0, bias=False) - - self.norm = None - if norm_type == "rms_norm": - self.norm = RMSNorm(out_channels, eps=1e-5, elementwise_affine=True, bias=True) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.residual_connection: - residual = hidden_states - - hidden_states = self.conv_inverted(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - - hidden_states = self.conv_depth(hidden_states) - hidden_states, gate = torch.chunk(hidden_states, 2, dim=1) - hidden_states = hidden_states * self.nonlinearity(gate) - - hidden_states = self.conv_point(hidden_states) - - if self.norm_type == "rms_norm": - # move channel to the last dimension so we apply RMSnorm across channel dimension - hidden_states = self.norm(hidden_states.movedim(1, -1)).movedim(-1, 1) - - if self.residual_connection: - hidden_states = hidden_states + residual - - return hidden_states - - -class SanaModulatedNorm(nn.Module): - def __init__(self, dim: int, elementwise_affine: bool = False, eps: float = 1e-6): - super().__init__() - self.norm = nn.LayerNorm(dim, elementwise_affine=elementwise_affine, eps=eps) - - def forward( - self, hidden_states: torch.Tensor, temb: torch.Tensor, scale_shift_table: torch.Tensor - ) -> torch.Tensor: - hidden_states = self.norm(hidden_states) - shift, scale = (scale_shift_table[None] + temb[:, None].to(scale_shift_table.device)).chunk(2, dim=1) - hidden_states = hidden_states * (1 + scale) + shift - return hidden_states - - -class SanaCombinedTimestepGuidanceEmbeddings(nn.Module): - def __init__(self, embedding_dim): - super().__init__() - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - self.guidance_condition_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.guidance_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=True) - - def forward(self, timestep: torch.Tensor, guidance: torch.Tensor = None, hidden_dtype: torch.dtype = None): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) - - guidance_proj = self.guidance_condition_proj(guidance) - guidance_emb = self.guidance_embedder(guidance_proj.to(dtype=hidden_dtype)) - conditioning = timesteps_emb + guidance_emb - - return self.linear(self.silu(conditioning)), conditioning - - -class SanaAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("SanaAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class SanaTransformerBlock(nn.Module): - r""" - Transformer block introduced in [Sana](https://huggingface.co/papers/2410.10629). - """ - - def __init__( - self, - dim: int = 2240, - num_attention_heads: int = 70, - attention_head_dim: int = 32, - dropout: float = 0.0, - num_cross_attention_heads: int | None = 20, - cross_attention_head_dim: int | None = 112, - cross_attention_dim: int | None = 2240, - attention_bias: bool = True, - norm_elementwise_affine: bool = False, - norm_eps: float = 1e-6, - attention_out_bias: bool = True, - mlp_ratio: float = 2.5, - qk_norm: str | None = None, - ) -> None: - super().__init__() - - # 1. Self Attention - self.norm1 = nn.LayerNorm(dim, elementwise_affine=False, eps=norm_eps) - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - kv_heads=num_attention_heads if qk_norm is not None else None, - qk_norm=qk_norm, - dropout=dropout, - bias=attention_bias, - cross_attention_dim=None, - processor=SanaLinearAttnProcessor2_0(), - ) - - # 2. Cross Attention - if cross_attention_dim is not None: - self.norm2 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps) - self.attn2 = Attention( - query_dim=dim, - qk_norm=qk_norm, - kv_heads=num_cross_attention_heads if qk_norm is not None else None, - cross_attention_dim=cross_attention_dim, - heads=num_cross_attention_heads, - dim_head=cross_attention_head_dim, - dropout=dropout, - bias=True, - out_bias=attention_out_bias, - processor=SanaAttnProcessor2_0(), - ) - - # 3. Feed-forward - self.ff = GLUMBConv(dim, dim, mlp_ratio, norm_type=None, residual_connection=False) - - self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - timestep: torch.LongTensor | None = None, - height: int = None, - width: int = None, - ) -> torch.Tensor: - batch_size = hidden_states.shape[0] - - # 1. Modulation - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( - self.scale_shift_table[None] + timestep.reshape(batch_size, 6, -1) - ).chunk(6, dim=1) - - # 2. Self Attention - norm_hidden_states = self.norm1(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_msa) + shift_msa - norm_hidden_states = norm_hidden_states.to(hidden_states.dtype) - - attn_output = self.attn1(norm_hidden_states) - hidden_states = hidden_states + gate_msa * attn_output - - # 3. Cross Attention - if self.attn2 is not None: - attn_output = self.attn2( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=encoder_attention_mask, - ) - hidden_states = attn_output + hidden_states - - # 4. Feed-forward - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp - - norm_hidden_states = norm_hidden_states.unflatten(1, (height, width)).permute(0, 3, 1, 2) - ff_output = self.ff(norm_hidden_states) - ff_output = ff_output.flatten(2, 3).permute(0, 2, 1) - hidden_states = hidden_states + gate_mlp * ff_output - - return hidden_states - - -class SanaTransformer2DModel(ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): - r""" - A 2D Transformer model introduced in [Sana](https://huggingface.co/papers/2410.10629) family of models. - - Args: - in_channels (`int`, defaults to `32`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `32`): - The number of channels in the output. - num_attention_heads (`int`, defaults to `70`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `32`): - The number of channels in each head. - num_layers (`int`, defaults to `20`): - The number of layers of Transformer blocks to use. - num_cross_attention_heads (`int`, *optional*, defaults to `20`): - The number of heads to use for cross-attention. - cross_attention_head_dim (`int`, *optional*, defaults to `112`): - The number of channels in each head for cross-attention. - cross_attention_dim (`int`, *optional*, defaults to `2240`): - The number of channels in the cross-attention output. - caption_channels (`int`, defaults to `2304`): - The number of channels in the caption embeddings. - mlp_ratio (`float`, defaults to `2.5`): - The expansion ratio to use in the GLUMBConv layer. - dropout (`float`, defaults to `0.0`): - The dropout probability. - attention_bias (`bool`, defaults to `False`): - Whether to use bias in the attention layer. - sample_size (`int`, defaults to `32`): - The base size of the input latent. - patch_size (`int`, defaults to `1`): - The size of the patches to use in the patch embedding layer. - norm_elementwise_affine (`bool`, defaults to `False`): - Whether to use elementwise affinity in the normalization layer. - norm_eps (`float`, defaults to `1e-6`): - The epsilon value for the normalization layer. - qk_norm (`str`, *optional*, defaults to `None`): - The normalization to use for the query and key. - timestep_scale (`float`, defaults to `1.0`): - The scale to use for the timesteps. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["SanaTransformerBlock", "PatchEmbed", "SanaModulatedNorm"] - _skip_layerwise_casting_patterns = ["patch_embed", "norm"] - - @register_to_config - def __init__( - self, - in_channels: int = 32, - out_channels: int | None = 32, - num_attention_heads: int = 70, - attention_head_dim: int = 32, - num_layers: int = 20, - num_cross_attention_heads: int | None = 20, - cross_attention_head_dim: int | None = 112, - cross_attention_dim: int | None = 2240, - caption_channels: int = 2304, - mlp_ratio: float = 2.5, - dropout: float = 0.0, - attention_bias: bool = False, - sample_size: int = 32, - patch_size: int = 1, - norm_elementwise_affine: bool = False, - norm_eps: float = 1e-6, - interpolation_scale: int | None = None, - guidance_embeds: bool = False, - guidance_embeds_scale: float = 0.1, - qk_norm: str | None = None, - timestep_scale: float = 1.0, - ) -> None: - super().__init__() - - out_channels = out_channels or in_channels - inner_dim = num_attention_heads * attention_head_dim - - # 1. Patch Embedding - self.patch_embed = PatchEmbed( - height=sample_size, - width=sample_size, - patch_size=patch_size, - in_channels=in_channels, - embed_dim=inner_dim, - interpolation_scale=interpolation_scale, - pos_embed_type="sincos" if interpolation_scale is not None else None, - ) - - # 2. Additional condition embeddings - if guidance_embeds: - self.time_embed = SanaCombinedTimestepGuidanceEmbeddings(inner_dim) - else: - self.time_embed = AdaLayerNormSingle(inner_dim) - - self.caption_projection = PixArtAlphaTextProjection(in_features=caption_channels, hidden_size=inner_dim) - self.caption_norm = RMSNorm(inner_dim, eps=1e-5, elementwise_affine=True) - - # 3. Transformer blocks - self.transformer_blocks = nn.ModuleList( - [ - SanaTransformerBlock( - inner_dim, - num_attention_heads, - attention_head_dim, - dropout=dropout, - num_cross_attention_heads=num_cross_attention_heads, - cross_attention_head_dim=cross_attention_head_dim, - cross_attention_dim=cross_attention_dim, - attention_bias=attention_bias, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - mlp_ratio=mlp_ratio, - qk_norm=qk_norm, - ) - for _ in range(num_layers) - ] - ) - - # 4. Output blocks - self.scale_shift_table = nn.Parameter(torch.randn(2, inner_dim) / inner_dim**0.5) - self.norm_out = SanaModulatedNorm(inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: torch.Tensor, - guidance: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - attention_kwargs: dict[str, Any] | None = None, - controlnet_block_samples: tuple[torch.Tensor] | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor, ...] | Transformer2DModelOutput: - """ - The [`SanaTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, in_channels, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - guidance (`torch.Tensor`, *optional*): - Guidance scale embedding. - encoder_attention_mask (`torch.Tensor`, *optional*): - Cross-attention mask applied to `encoder_hidden_states`. - attention_mask (`torch.Tensor`, *optional*): - Self-attention mask applied to `hidden_states`. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - controlnet_block_samples (`tuple` of `torch.Tensor`, *optional*): - A list of tensors that if specified are added to the residuals of transformer blocks. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - # ensure attention_mask is a bias, and give it a singleton query_tokens dimension. - # we may have done this conversion already, e.g. if we came here via UNet2DConditionModel#forward. - # we can tell by counting dims; if ndim == 2: it's a mask rather than a bias. - # expects mask of shape: - # [batch, key_tokens] - # adds singleton query_tokens dimension: - # [batch, 1, key_tokens] - # this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes: - # [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn) - # [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn) - if attention_mask is not None and attention_mask.ndim == 2: - # assume that mask is expressed as: - # (1 = keep, 0 = discard) - # convert mask into a bias that can be added to attention scores: - # (keep = +0, discard = -10000.0) - attention_mask = (1 - attention_mask.to(hidden_states.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # convert encoder_attention_mask to a bias the same way we do for attention_mask - if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2: - encoder_attention_mask = (1 - encoder_attention_mask.to(hidden_states.dtype)) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - # 1. Input - batch_size, num_channels, height, width = hidden_states.shape - p = self.config.patch_size - post_patch_height, post_patch_width = height // p, width // p - - hidden_states = self.patch_embed(hidden_states) - - if guidance is not None: - timestep, embedded_timestep = self.time_embed( - timestep, guidance=guidance, hidden_dtype=hidden_states.dtype - ) - else: - timestep, embedded_timestep = self.time_embed( - timestep, batch_size=batch_size, hidden_dtype=hidden_states.dtype - ) - - encoder_hidden_states = self.caption_projection(encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states.view(batch_size, -1, hidden_states.shape[-1]) - - encoder_hidden_states = self.caption_norm(encoder_hidden_states) - - # 2. Transformer blocks - if torch.is_grad_enabled() and self.gradient_checkpointing: - for index_block, block in enumerate(self.transformer_blocks): - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - attention_mask, - encoder_hidden_states, - encoder_attention_mask, - timestep, - post_patch_height, - post_patch_width, - ) - if controlnet_block_samples is not None and 0 < index_block <= len(controlnet_block_samples): - hidden_states = hidden_states + controlnet_block_samples[index_block - 1] - - else: - for index_block, block in enumerate(self.transformer_blocks): - hidden_states = block( - hidden_states, - attention_mask, - encoder_hidden_states, - encoder_attention_mask, - timestep, - post_patch_height, - post_patch_width, - ) - if controlnet_block_samples is not None and 0 < index_block <= len(controlnet_block_samples): - hidden_states = hidden_states + controlnet_block_samples[index_block - 1] - - # 3. Normalization - hidden_states = self.norm_out(hidden_states, embedded_timestep, self.scale_shift_table) - - hidden_states = self.proj_out(hidden_states) - - # 5. Unpatchify - hidden_states = hidden_states.reshape( - batch_size, post_patch_height, post_patch_width, self.config.patch_size, self.config.patch_size, -1 - ) - hidden_states = hidden_states.permute(0, 5, 1, 3, 2, 4) - output = hidden_states.reshape(batch_size, -1, post_patch_height * p, post_patch_width * p) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/stable_audio_transformer.py b/diffusers/models/transformers/stable_audio_transformer.py deleted file mode 100644 index f4974926ec7279a1691e35f41c86d427d9123dea..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/stable_audio_transformer.py +++ /dev/null @@ -1,376 +0,0 @@ -# Copyright 2025 Stability AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import numpy as np -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import AttentionMixin, FeedForward -from ..attention_processor import Attention, StableAudioAttnProcessor2_0 -from ..modeling_utils import ModelMixin -from ..transformers.transformer_2d import Transformer2DModelOutput - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class StableAudioGaussianFourierProjection(nn.Module): - """Gaussian Fourier embeddings for noise levels.""" - - # Copied from diffusers.models.embeddings.GaussianFourierProjection.__init__ - def __init__( - self, embedding_size: int = 256, scale: float = 1.0, set_W_to_weight=True, log=True, flip_sin_to_cos=False - ): - super().__init__() - self.weight = nn.Parameter(torch.randn(embedding_size) * scale, requires_grad=False) - self.log = log - self.flip_sin_to_cos = flip_sin_to_cos - - if set_W_to_weight: - # to delete later - del self.weight - self.W = nn.Parameter(torch.randn(embedding_size) * scale, requires_grad=False) - self.weight = self.W - del self.W - - def forward(self, x): - if self.log: - x = torch.log(x) - - x_proj = 2 * np.pi * x[:, None] @ self.weight[None, :] - - if self.flip_sin_to_cos: - out = torch.cat([torch.cos(x_proj), torch.sin(x_proj)], dim=-1) - else: - out = torch.cat([torch.sin(x_proj), torch.cos(x_proj)], dim=-1) - return out - - -@maybe_allow_in_graph -class StableAudioDiTBlock(nn.Module): - r""" - Transformer block used in Stable Audio model (https://github.com/Stability-AI/stable-audio-tools). Allow skip - connection and QKNorm - - Parameters: - dim (`int`): The number of channels in the input and output. - num_attention_heads (`int`): The number of heads to use for the query states. - num_key_value_attention_heads (`int`): The number of heads to use for the key and value states. - attention_head_dim (`int`): The number of channels in each head. - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention. - upcast_attention (`bool`, *optional*): - Whether to upcast the attention computation to float32. This is useful for mixed precision training. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - num_key_value_attention_heads: int, - attention_head_dim: int, - dropout=0.0, - cross_attention_dim: int | None = None, - upcast_attention: bool = False, - norm_eps: float = 1e-5, - ff_inner_dim: int | None = None, - ): - super().__init__() - # Define 3 blocks. Each block has its own normalization layer. - # 1. Self-Attn - self.norm1 = nn.LayerNorm(dim, elementwise_affine=True, eps=norm_eps) - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=False, - upcast_attention=upcast_attention, - out_bias=False, - processor=StableAudioAttnProcessor2_0(), - ) - - # 2. Cross-Attn - self.norm2 = nn.LayerNorm(dim, norm_eps, True) - - self.attn2 = Attention( - query_dim=dim, - cross_attention_dim=cross_attention_dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - kv_heads=num_key_value_attention_heads, - dropout=dropout, - bias=False, - upcast_attention=upcast_attention, - out_bias=False, - processor=StableAudioAttnProcessor2_0(), - ) # is self-attn if encoder_hidden_states is none - - # 3. Feed-forward - self.norm3 = nn.LayerNorm(dim, norm_eps, True) - self.ff = FeedForward( - dim, - dropout=dropout, - activation_fn="swiglu", - final_dropout=False, - inner_dim=ff_inner_dim, - bias=True, - ) - - # let chunk size default to None - self._chunk_size = None - self._chunk_dim = 0 - - def set_chunk_feed_forward(self, chunk_size: int | None, dim: int = 0): - # Sets chunk feed-forward - self._chunk_size = chunk_size - self._chunk_dim = dim - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - rotary_embedding: torch.FloatTensor | None = None, - ) -> torch.Tensor: - # Notice that normalization is always applied before the real computation in the following blocks. - # 0. Self-Attention - norm_hidden_states = self.norm1(hidden_states) - - attn_output = self.attn1( - norm_hidden_states, - attention_mask=attention_mask, - rotary_emb=rotary_embedding, - ) - - hidden_states = attn_output + hidden_states - - # 2. Cross-Attention - norm_hidden_states = self.norm2(hidden_states) - - attn_output = self.attn2( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=encoder_attention_mask, - ) - hidden_states = attn_output + hidden_states - - # 3. Feed-forward - norm_hidden_states = self.norm3(hidden_states) - ff_output = self.ff(norm_hidden_states) - - hidden_states = ff_output + hidden_states - - return hidden_states - - -class StableAudioDiTModel(ModelMixin, AttentionMixin, ConfigMixin): - """ - The Diffusion Transformer model introduced in Stable Audio. - - Reference: https://github.com/Stability-AI/stable-audio-tools - - Parameters: - sample_size ( `int`, *optional*, defaults to 1024): The size of the input sample. - in_channels (`int`, *optional*, defaults to 64): The number of channels in the input. - num_layers (`int`, *optional*, defaults to 24): The number of layers of Transformer blocks to use. - attention_head_dim (`int`, *optional*, defaults to 64): The number of channels in each head. - num_attention_heads (`int`, *optional*, defaults to 24): The number of heads to use for the query states. - num_key_value_attention_heads (`int`, *optional*, defaults to 12): - The number of heads to use for the key and value states. - out_channels (`int`, defaults to 64): Number of output channels. - cross_attention_dim ( `int`, *optional*, defaults to 768): Dimension of the cross-attention projection. - time_proj_dim ( `int`, *optional*, defaults to 256): Dimension of the timestep inner projection. - global_states_input_dim ( `int`, *optional*, defaults to 1536): - Input dimension of the global hidden states projection. - cross_attention_input_dim ( `int`, *optional*, defaults to 768): - Input dimension of the cross-attention projection - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["preprocess_conv", "postprocess_conv", "^proj_in$", "^proj_out$", "norm"] - - @register_to_config - def __init__( - self, - sample_size: int = 1024, - in_channels: int = 64, - num_layers: int = 24, - attention_head_dim: int = 64, - num_attention_heads: int = 24, - num_key_value_attention_heads: int = 12, - out_channels: int = 64, - cross_attention_dim: int = 768, - time_proj_dim: int = 256, - global_states_input_dim: int = 1536, - cross_attention_input_dim: int = 768, - ): - super().__init__() - self.sample_size = sample_size - self.out_channels = out_channels - self.inner_dim = num_attention_heads * attention_head_dim - - self.time_proj = StableAudioGaussianFourierProjection( - embedding_size=time_proj_dim // 2, - flip_sin_to_cos=True, - log=False, - set_W_to_weight=False, - ) - - self.timestep_proj = nn.Sequential( - nn.Linear(time_proj_dim, self.inner_dim, bias=True), - nn.SiLU(), - nn.Linear(self.inner_dim, self.inner_dim, bias=True), - ) - - self.global_proj = nn.Sequential( - nn.Linear(global_states_input_dim, self.inner_dim, bias=False), - nn.SiLU(), - nn.Linear(self.inner_dim, self.inner_dim, bias=False), - ) - - self.cross_attention_proj = nn.Sequential( - nn.Linear(cross_attention_input_dim, cross_attention_dim, bias=False), - nn.SiLU(), - nn.Linear(cross_attention_dim, cross_attention_dim, bias=False), - ) - - self.preprocess_conv = nn.Conv1d(in_channels, in_channels, 1, bias=False) - self.proj_in = nn.Linear(in_channels, self.inner_dim, bias=False) - - self.transformer_blocks = nn.ModuleList( - [ - StableAudioDiTBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - num_key_value_attention_heads=num_key_value_attention_heads, - attention_head_dim=attention_head_dim, - cross_attention_dim=cross_attention_dim, - ) - for i in range(num_layers) - ] - ) - - self.proj_out = nn.Linear(self.inner_dim, self.out_channels, bias=False) - self.postprocess_conv = nn.Conv1d(self.out_channels, self.out_channels, 1, bias=False) - - self.gradient_checkpointing = False - - # Copied from diffusers.models.transformers.hunyuan_transformer_2d.HunyuanDiT2DModel.set_default_attn_processor with Hunyuan->StableAudio - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - self.set_attn_processor(StableAudioAttnProcessor2_0()) - - def forward( - self, - hidden_states: torch.FloatTensor, - timestep: torch.LongTensor = None, - encoder_hidden_states: torch.FloatTensor = None, - global_hidden_states: torch.FloatTensor = None, - rotary_embedding: torch.FloatTensor = None, - return_dict: bool = True, - attention_mask: torch.LongTensor | None = None, - encoder_attention_mask: torch.LongTensor | None = None, - ) -> torch.FloatTensor | Transformer2DModelOutput: - """ - The [`StableAudioDiTModel`] forward method. - - Args: - hidden_states (`torch.FloatTensor` of shape `(batch size, in_channels, sequence_len)`): - Input `hidden_states`. - timestep ( `torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.FloatTensor` of shape `(batch size, encoder_sequence_len, cross_attention_input_dim)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - global_hidden_states (`torch.FloatTensor` of shape `(batch size, global_sequence_len, global_states_input_dim)`): - Global embeddings that will be prepended to the hidden states. - rotary_embedding (`torch.Tensor`): - The rotary embeddings to apply on query and key tensors during attention calculation. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - attention_mask (`torch.Tensor` of shape `(batch_size, sequence_len)`, *optional*): - Mask to avoid performing attention on padding token indices, formed by concatenating the attention - masks - for the two text encoders together. Mask values selected in `[0, 1]`: - - - 1 for tokens that are **not masked**, - - 0 for tokens that are **masked**. - encoder_attention_mask (`torch.Tensor` of shape `(batch_size, sequence_len)`, *optional*): - Mask to avoid performing attention on padding token cross-attention indices, formed by concatenating - the attention masks - for the two text encoders together. Mask values selected in `[0, 1]`: - - - 1 for tokens that are **not masked**, - - 0 for tokens that are **masked**. - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - cross_attention_hidden_states = self.cross_attention_proj(encoder_hidden_states) - global_hidden_states = self.global_proj(global_hidden_states) - time_hidden_states = self.timestep_proj(self.time_proj(timestep.to(self.dtype))) - - global_hidden_states = global_hidden_states + time_hidden_states.unsqueeze(1) - - hidden_states = self.preprocess_conv(hidden_states) + hidden_states - # (batch_size, dim, sequence_length) -> (batch_size, sequence_length, dim) - hidden_states = hidden_states.transpose(1, 2) - - hidden_states = self.proj_in(hidden_states) - - # prepend global states to hidden states - hidden_states = torch.cat([global_hidden_states, hidden_states], dim=-2) - if attention_mask is not None: - prepend_mask = torch.ones((hidden_states.shape[0], 1), device=hidden_states.device, dtype=torch.bool) - attention_mask = torch.cat([prepend_mask, attention_mask], dim=-1) - - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - attention_mask, - cross_attention_hidden_states, - encoder_attention_mask, - rotary_embedding, - ) - - else: - hidden_states = block( - hidden_states=hidden_states, - attention_mask=attention_mask, - encoder_hidden_states=cross_attention_hidden_states, - encoder_attention_mask=encoder_attention_mask, - rotary_embedding=rotary_embedding, - ) - - hidden_states = self.proj_out(hidden_states) - - # (batch_size, sequence_length, dim) -> (batch_size, dim, sequence_length) - # remove prepend length that has been added by global hidden states - hidden_states = hidden_states.transpose(1, 2)[:, :, 1:] - hidden_states = self.postprocess_conv(hidden_states) + hidden_states - - if not return_dict: - return (hidden_states,) - - return Transformer2DModelOutput(sample=hidden_states) diff --git a/diffusers/models/transformers/t5_film_transformer.py b/diffusers/models/transformers/t5_film_transformer.py deleted file mode 100644 index 547e720899908ebb424f7b8feb0469e7436eb46c..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/t5_film_transformer.py +++ /dev/null @@ -1,447 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import math - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ..attention_processor import Attention -from ..embeddings import get_timestep_embedding -from ..modeling_utils import ModelMixin - - -class T5FilmDecoder(ModelMixin, ConfigMixin): - r""" - T5 style decoder with FiLM conditioning. - - Args: - input_dims (`int`, *optional*, defaults to `128`): - The number of input dimensions. - targets_length (`int`, *optional*, defaults to `256`): - The length of the targets. - d_model (`int`, *optional*, defaults to `768`): - Size of the input hidden states. - num_layers (`int`, *optional*, defaults to `12`): - The number of `DecoderLayer`'s to use. - num_heads (`int`, *optional*, defaults to `12`): - The number of attention heads to use. - d_kv (`int`, *optional*, defaults to `64`): - Size of the key-value projection vectors. - d_ff (`int`, *optional*, defaults to `2048`): - The number of dimensions in the intermediate feed-forward layer of `DecoderLayer`'s. - dropout_rate (`float`, *optional*, defaults to `0.1`): - Dropout probability. - """ - - @register_to_config - def __init__( - self, - input_dims: int = 128, - targets_length: int = 256, - max_decoder_noise_time: float = 2000.0, - d_model: int = 768, - num_layers: int = 12, - num_heads: int = 12, - d_kv: int = 64, - d_ff: int = 2048, - dropout_rate: float = 0.1, - ): - super().__init__() - - self.conditioning_emb = nn.Sequential( - nn.Linear(d_model, d_model * 4, bias=False), - nn.SiLU(), - nn.Linear(d_model * 4, d_model * 4, bias=False), - nn.SiLU(), - ) - - self.position_encoding = nn.Embedding(targets_length, d_model) - self.position_encoding.weight.requires_grad = False - - self.continuous_inputs_projection = nn.Linear(input_dims, d_model, bias=False) - - self.dropout = nn.Dropout(p=dropout_rate) - - self.decoders = nn.ModuleList() - for lyr_num in range(num_layers): - # FiLM conditional T5 decoder - lyr = DecoderLayer(d_model=d_model, d_kv=d_kv, num_heads=num_heads, d_ff=d_ff, dropout_rate=dropout_rate) - self.decoders.append(lyr) - - self.decoder_norm = T5LayerNorm(d_model) - - self.post_dropout = nn.Dropout(p=dropout_rate) - self.spec_out = nn.Linear(d_model, input_dims, bias=False) - - def encoder_decoder_mask(self, query_input: torch.Tensor, key_input: torch.Tensor) -> torch.Tensor: - mask = torch.mul(query_input.unsqueeze(-1), key_input.unsqueeze(-2)) - return mask.unsqueeze(-3) - - def forward(self, encodings_and_masks, decoder_input_tokens, decoder_noise_time): - """ - The [`T5FilmDecoder`] forward method. - - Args: - encodings_and_masks (`list` of `tuple` of `torch.Tensor`): - A list of `(encoding, mask)` tuples produced by upstream encoders. The encodings are concatenated and - cross-attended to by the decoder. - decoder_input_tokens (`torch.Tensor` of shape `(batch_size, seq_length, input_dims)`): - Input tokens for the decoder. - decoder_noise_time (`torch.Tensor` of shape `(batch_size,)`): - Diffusion timesteps in `[0, 1)` used to condition the decoder. - """ - batch, _, _ = decoder_input_tokens.shape - assert decoder_noise_time.shape == (batch,) - - # decoder_noise_time is in [0, 1), so rescale to expected timing range. - time_steps = get_timestep_embedding( - decoder_noise_time * self.config.max_decoder_noise_time, - embedding_dim=self.config.d_model, - max_period=self.config.max_decoder_noise_time, - ).to(dtype=self.dtype) - - conditioning_emb = self.conditioning_emb(time_steps).unsqueeze(1) - - assert conditioning_emb.shape == (batch, 1, self.config.d_model * 4) - - seq_length = decoder_input_tokens.shape[1] - - # If we want to use relative positions for audio context, we can just offset - # this sequence by the length of encodings_and_masks. - decoder_positions = torch.broadcast_to( - torch.arange(seq_length, device=decoder_input_tokens.device), - (batch, seq_length), - ) - - position_encodings = self.position_encoding(decoder_positions) - - inputs = self.continuous_inputs_projection(decoder_input_tokens) - inputs += position_encodings - y = self.dropout(inputs) - - # decoder: No padding present. - decoder_mask = torch.ones( - decoder_input_tokens.shape[:2], device=decoder_input_tokens.device, dtype=inputs.dtype - ) - - # Translate encoding masks to encoder-decoder masks. - encodings_and_encdec_masks = [(x, self.encoder_decoder_mask(decoder_mask, y)) for x, y in encodings_and_masks] - - # cross attend style: concat encodings - encoded = torch.cat([x[0] for x in encodings_and_encdec_masks], dim=1) - encoder_decoder_mask = torch.cat([x[1] for x in encodings_and_encdec_masks], dim=-1) - - for lyr in self.decoders: - y = lyr( - y, - conditioning_emb=conditioning_emb, - encoder_hidden_states=encoded, - encoder_attention_mask=encoder_decoder_mask, - )[0] - - y = self.decoder_norm(y) - y = self.post_dropout(y) - - spec_out = self.spec_out(y) - return spec_out - - -class DecoderLayer(nn.Module): - r""" - T5 decoder layer. - - Args: - d_model (`int`): - Size of the input hidden states. - d_kv (`int`): - Size of the key-value projection vectors. - num_heads (`int`): - Number of attention heads. - d_ff (`int`): - Size of the intermediate feed-forward layer. - dropout_rate (`float`): - Dropout probability. - layer_norm_epsilon (`float`, *optional*, defaults to `1e-6`): - A small value used for numerical stability to avoid dividing by zero. - """ - - def __init__( - self, d_model: int, d_kv: int, num_heads: int, d_ff: int, dropout_rate: float, layer_norm_epsilon: float = 1e-6 - ): - super().__init__() - self.layer = nn.ModuleList() - - # cond self attention: layer 0 - self.layer.append( - T5LayerSelfAttentionCond(d_model=d_model, d_kv=d_kv, num_heads=num_heads, dropout_rate=dropout_rate) - ) - - # cross attention: layer 1 - self.layer.append( - T5LayerCrossAttention( - d_model=d_model, - d_kv=d_kv, - num_heads=num_heads, - dropout_rate=dropout_rate, - layer_norm_epsilon=layer_norm_epsilon, - ) - ) - - # Film Cond MLP + dropout: last layer - self.layer.append( - T5LayerFFCond(d_model=d_model, d_ff=d_ff, dropout_rate=dropout_rate, layer_norm_epsilon=layer_norm_epsilon) - ) - - def forward( - self, - hidden_states: torch.Tensor, - conditioning_emb: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - encoder_decoder_position_bias=None, - ) -> tuple[torch.Tensor]: - hidden_states = self.layer[0]( - hidden_states, - conditioning_emb=conditioning_emb, - attention_mask=attention_mask, - ) - - if encoder_hidden_states is not None: - encoder_extended_attention_mask = torch.where(encoder_attention_mask > 0, 0, -1e10).to( - encoder_hidden_states.dtype - ) - - hidden_states = self.layer[1]( - hidden_states, - key_value_states=encoder_hidden_states, - attention_mask=encoder_extended_attention_mask, - ) - - # Apply Film Conditional Feed Forward layer - hidden_states = self.layer[-1](hidden_states, conditioning_emb) - - return (hidden_states,) - - -class T5LayerSelfAttentionCond(nn.Module): - r""" - T5 style self-attention layer with conditioning. - - Args: - d_model (`int`): - Size of the input hidden states. - d_kv (`int`): - Size of the key-value projection vectors. - num_heads (`int`): - Number of attention heads. - dropout_rate (`float`): - Dropout probability. - """ - - def __init__(self, d_model: int, d_kv: int, num_heads: int, dropout_rate: float): - super().__init__() - self.layer_norm = T5LayerNorm(d_model) - self.FiLMLayer = T5FiLMLayer(in_features=d_model * 4, out_features=d_model) - self.attention = Attention(query_dim=d_model, heads=num_heads, dim_head=d_kv, out_bias=False, scale_qk=False) - self.dropout = nn.Dropout(dropout_rate) - - def forward( - self, - hidden_states: torch.Tensor, - conditioning_emb: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - # pre_self_attention_layer_norm - normed_hidden_states = self.layer_norm(hidden_states) - - if conditioning_emb is not None: - normed_hidden_states = self.FiLMLayer(normed_hidden_states, conditioning_emb) - - # Self-attention block - attention_output = self.attention(normed_hidden_states) - - hidden_states = hidden_states + self.dropout(attention_output) - - return hidden_states - - -class T5LayerCrossAttention(nn.Module): - r""" - T5 style cross-attention layer. - - Args: - d_model (`int`): - Size of the input hidden states. - d_kv (`int`): - Size of the key-value projection vectors. - num_heads (`int`): - Number of attention heads. - dropout_rate (`float`): - Dropout probability. - layer_norm_epsilon (`float`): - A small value used for numerical stability to avoid dividing by zero. - """ - - def __init__(self, d_model: int, d_kv: int, num_heads: int, dropout_rate: float, layer_norm_epsilon: float): - super().__init__() - self.attention = Attention(query_dim=d_model, heads=num_heads, dim_head=d_kv, out_bias=False, scale_qk=False) - self.layer_norm = T5LayerNorm(d_model, eps=layer_norm_epsilon) - self.dropout = nn.Dropout(dropout_rate) - - def forward( - self, - hidden_states: torch.Tensor, - key_value_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - normed_hidden_states = self.layer_norm(hidden_states) - attention_output = self.attention( - normed_hidden_states, - encoder_hidden_states=key_value_states, - attention_mask=attention_mask.squeeze(1), - ) - layer_output = hidden_states + self.dropout(attention_output) - return layer_output - - -class T5LayerFFCond(nn.Module): - r""" - T5 style feed-forward conditional layer. - - Args: - d_model (`int`): - Size of the input hidden states. - d_ff (`int`): - Size of the intermediate feed-forward layer. - dropout_rate (`float`): - Dropout probability. - layer_norm_epsilon (`float`): - A small value used for numerical stability to avoid dividing by zero. - """ - - def __init__(self, d_model: int, d_ff: int, dropout_rate: float, layer_norm_epsilon: float): - super().__init__() - self.DenseReluDense = T5DenseGatedActDense(d_model=d_model, d_ff=d_ff, dropout_rate=dropout_rate) - self.film = T5FiLMLayer(in_features=d_model * 4, out_features=d_model) - self.layer_norm = T5LayerNorm(d_model, eps=layer_norm_epsilon) - self.dropout = nn.Dropout(dropout_rate) - - def forward(self, hidden_states: torch.Tensor, conditioning_emb: torch.Tensor | None = None) -> torch.Tensor: - forwarded_states = self.layer_norm(hidden_states) - if conditioning_emb is not None: - forwarded_states = self.film(forwarded_states, conditioning_emb) - - forwarded_states = self.DenseReluDense(forwarded_states) - hidden_states = hidden_states + self.dropout(forwarded_states) - return hidden_states - - -class T5DenseGatedActDense(nn.Module): - r""" - T5 style feed-forward layer with gated activations and dropout. - - Args: - d_model (`int`): - Size of the input hidden states. - d_ff (`int`): - Size of the intermediate feed-forward layer. - dropout_rate (`float`): - Dropout probability. - """ - - def __init__(self, d_model: int, d_ff: int, dropout_rate: float): - super().__init__() - self.wi_0 = nn.Linear(d_model, d_ff, bias=False) - self.wi_1 = nn.Linear(d_model, d_ff, bias=False) - self.wo = nn.Linear(d_ff, d_model, bias=False) - self.dropout = nn.Dropout(dropout_rate) - self.act = NewGELUActivation() - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_gelu = self.act(self.wi_0(hidden_states)) - hidden_linear = self.wi_1(hidden_states) - hidden_states = hidden_gelu * hidden_linear - hidden_states = self.dropout(hidden_states) - - hidden_states = self.wo(hidden_states) - return hidden_states - - -class T5LayerNorm(nn.Module): - r""" - T5 style layer normalization module. - - Args: - hidden_size (`int`): - Size of the input hidden states. - eps (`float`, `optional`, defaults to `1e-6`): - A small value used for numerical stability to avoid dividing by zero. - """ - - def __init__(self, hidden_size: int, eps: float = 1e-6): - """ - Construct a layernorm module in the T5 style. No bias and no subtraction of mean. - """ - super().__init__() - self.weight = nn.Parameter(torch.ones(hidden_size)) - self.variance_epsilon = eps - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - # T5 uses a layer_norm which only scales and doesn't shift, which is also known as Root Mean - # Square Layer Normalization https://huggingface.co/papers/1910.07467 thus variance is calculated - # w/o mean and there is no bias. Additionally we want to make sure that the accumulation for - # half-precision inputs is done in fp32 - - variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True) - hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) - - # convert into half-precision if necessary - if self.weight.dtype in [torch.float16, torch.bfloat16]: - hidden_states = hidden_states.to(self.weight.dtype) - - return self.weight * hidden_states - - -class NewGELUActivation(nn.Module): - """ - Implementation of the GELU activation function currently in Google BERT repo (identical to OpenAI GPT). Also see - the Gaussian Error Linear Units paper: https://huggingface.co/papers/1606.08415 - """ - - def forward(self, input: torch.Tensor) -> torch.Tensor: - return 0.5 * input * (1.0 + torch.tanh(math.sqrt(2.0 / math.pi) * (input + 0.044715 * torch.pow(input, 3.0)))) - - -class T5FiLMLayer(nn.Module): - """ - T5 style FiLM Layer. - - Args: - in_features (`int`): - Number of input features. - out_features (`int`): - Number of output features. - """ - - def __init__(self, in_features: int, out_features: int): - super().__init__() - self.scale_bias = nn.Linear(in_features, out_features * 2, bias=False) - - def forward(self, x: torch.Tensor, conditioning_emb: torch.Tensor) -> torch.Tensor: - emb = self.scale_bias(conditioning_emb) - scale, shift = torch.chunk(emb, 2, -1) - x = x * (1 + scale) + shift - return x diff --git a/diffusers/models/transformers/transformer_2d.py b/diffusers/models/transformers/transformer_2d.py deleted file mode 100644 index 6714383b77abac4013ceafe1ddd24989618d090c..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_2d.py +++ /dev/null @@ -1,551 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from typing import Any - -import torch -import torch.nn.functional as F -from torch import nn - -from ...configuration_utils import LegacyConfigMixin, register_to_config -from ...utils import deprecate, logging -from ..attention import BasicTransformerBlock -from ..embeddings import ImagePositionalEmbeddings, PatchEmbed, PixArtAlphaTextProjection -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import LegacyModelMixin -from ..normalization import AdaLayerNormSingle - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class Transformer2DModelOutput(Transformer2DModelOutput): - def __init__(self, *args, **kwargs): - deprecation_message = "Importing `Transformer2DModelOutput` from `diffusers.models.transformer_2d` is deprecated and this will be removed in a future version. Please use `from diffusers.models.modeling_outputs import Transformer2DModelOutput`, instead." - deprecate("Transformer2DModelOutput", "1.0.0", deprecation_message) - super().__init__(*args, **kwargs) - - -class Transformer2DModel(LegacyModelMixin, LegacyConfigMixin): - """ - A 2D Transformer model for image-like data. - - Parameters: - num_attention_heads (`int`, *optional*, defaults to 16): The number of heads to use for multi-head attention. - attention_head_dim (`int`, *optional*, defaults to 88): The number of channels in each head. - in_channels (`int`, *optional*): - The number of channels in the input and output (specify if the input is **continuous**). - num_layers (`int`, *optional*, defaults to 1): The number of layers of Transformer blocks to use. - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - cross_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use. - sample_size (`int`, *optional*): The width of the latent images (specify if the input is **discrete**). - This is fixed during training since it is used to learn a number of position embeddings. - num_vector_embeds (`int`, *optional*): - The number of classes of the vector embeddings of the latent pixels (specify if the input is **discrete**). - Includes the class for the masked latent pixel. - activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to use in feed-forward. - num_embeds_ada_norm ( `int`, *optional*): - The number of diffusion steps used during training. Pass if at least one of the norm_layers is - `AdaLayerNorm`. This is fixed during training since it is used to learn a number of embeddings that are - added to the hidden states. - - During inference, you can denoise for up to but not more steps than `num_embeds_ada_norm`. - attention_bias (`bool`, *optional*): - Configure if the `TransformerBlocks` attention should contain a bias parameter. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["BasicTransformerBlock"] - _skip_layerwise_casting_patterns = ["latent_image_embedding", "norm"] - - @register_to_config - def __init__( - self, - num_attention_heads: int = 16, - attention_head_dim: int = 88, - in_channels: int | None = None, - out_channels: int | None = None, - num_layers: int = 1, - dropout: float = 0.0, - norm_num_groups: int = 32, - cross_attention_dim: int | None = None, - attention_bias: bool = False, - sample_size: int | None = None, - num_vector_embeds: int | None = None, - patch_size: int | None = None, - activation_fn: str = "geglu", - num_embeds_ada_norm: int | None = None, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - double_self_attention: bool = False, - upcast_attention: bool = False, - norm_type: str = "layer_norm", # 'layer_norm', 'ada_norm', 'ada_norm_zero', 'ada_norm_single', 'ada_norm_continuous', 'layer_norm_i2vgen' - norm_elementwise_affine: bool = True, - norm_eps: float = 1e-5, - attention_type: str = "default", - caption_channels: int = None, - interpolation_scale: float = None, - use_additional_conditions: bool | None = None, - ): - super().__init__() - - # Validate inputs. - if patch_size is not None: - if norm_type not in ["ada_norm", "ada_norm_zero", "ada_norm_single"]: - raise NotImplementedError( - f"Forward pass is not implemented when `patch_size` is not None and `norm_type` is '{norm_type}'." - ) - elif norm_type in ["ada_norm", "ada_norm_zero"] and num_embeds_ada_norm is None: - raise ValueError( - f"When using a `patch_size` and this `norm_type` ({norm_type}), `num_embeds_ada_norm` cannot be None." - ) - - # 1. Transformer2DModel can process both standard continuous images of shape `(batch_size, num_channels, width, height)` as well as quantized image embeddings of shape `(batch_size, num_image_vectors)` - # Define whether input is continuous or discrete depending on configuration - self.is_input_continuous = (in_channels is not None) and (patch_size is None) - self.is_input_vectorized = num_vector_embeds is not None - self.is_input_patches = in_channels is not None and patch_size is not None - - if self.is_input_continuous and self.is_input_vectorized: - raise ValueError( - f"Cannot define both `in_channels`: {in_channels} and `num_vector_embeds`: {num_vector_embeds}. Make" - " sure that either `in_channels` or `num_vector_embeds` is None." - ) - elif self.is_input_vectorized and self.is_input_patches: - raise ValueError( - f"Cannot define both `num_vector_embeds`: {num_vector_embeds} and `patch_size`: {patch_size}. Make" - " sure that either `num_vector_embeds` or `num_patches` is None." - ) - elif not self.is_input_continuous and not self.is_input_vectorized and not self.is_input_patches: - raise ValueError( - f"Has to define `in_channels`: {in_channels}, `num_vector_embeds`: {num_vector_embeds}, or patch_size:" - f" {patch_size}. Make sure that `in_channels`, `num_vector_embeds` or `num_patches` is not None." - ) - - if norm_type == "layer_norm" and num_embeds_ada_norm is not None: - deprecation_message = ( - f"The configuration file of this model: {self.__class__} is outdated. `norm_type` is either not set or" - " incorrectly set to `'layer_norm'`. Make sure to set `norm_type` to `'ada_norm'` in the config." - " Please make sure to update the config accordingly as leaving `norm_type` might led to incorrect" - " results in future versions. If you have downloaded this checkpoint from the Hugging Face Hub, it" - " would be very nice if you could open a Pull request for the `transformer/config.json` file" - ) - deprecate("norm_type!=num_embeds_ada_norm", "1.0.0", deprecation_message, standard_warn=False) - norm_type = "ada_norm" - - # Set some common variables used across the board. - self.use_linear_projection = use_linear_projection - self.interpolation_scale = interpolation_scale - self.caption_channels = caption_channels - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim - self.in_channels = in_channels - self.out_channels = in_channels if out_channels is None else out_channels - self.gradient_checkpointing = False - - if use_additional_conditions is None: - if norm_type == "ada_norm_single" and sample_size == 128: - use_additional_conditions = True - else: - use_additional_conditions = False - self.use_additional_conditions = use_additional_conditions - - # 2. Initialize the right blocks. - # These functions follow a common structure: - # a. Initialize the input blocks. b. Initialize the transformer blocks. - # c. Initialize the output blocks and other projection blocks when necessary. - if self.is_input_continuous: - self._init_continuous_input(norm_type=norm_type) - elif self.is_input_vectorized: - self._init_vectorized_inputs(norm_type=norm_type) - elif self.is_input_patches: - self._init_patched_inputs(norm_type=norm_type) - - def _init_continuous_input(self, norm_type): - self.norm = torch.nn.GroupNorm( - num_groups=self.config.norm_num_groups, num_channels=self.in_channels, eps=1e-6, affine=True - ) - if self.use_linear_projection: - self.proj_in = torch.nn.Linear(self.in_channels, self.inner_dim) - else: - self.proj_in = torch.nn.Conv2d(self.in_channels, self.inner_dim, kernel_size=1, stride=1, padding=0) - - self.transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlock( - self.inner_dim, - self.config.num_attention_heads, - self.config.attention_head_dim, - dropout=self.config.dropout, - cross_attention_dim=self.config.cross_attention_dim, - activation_fn=self.config.activation_fn, - num_embeds_ada_norm=self.config.num_embeds_ada_norm, - attention_bias=self.config.attention_bias, - only_cross_attention=self.config.only_cross_attention, - double_self_attention=self.config.double_self_attention, - upcast_attention=self.config.upcast_attention, - norm_type=norm_type, - norm_elementwise_affine=self.config.norm_elementwise_affine, - norm_eps=self.config.norm_eps, - attention_type=self.config.attention_type, - ) - for _ in range(self.config.num_layers) - ] - ) - - if self.use_linear_projection: - self.proj_out = torch.nn.Linear(self.inner_dim, self.out_channels) - else: - self.proj_out = torch.nn.Conv2d(self.inner_dim, self.out_channels, kernel_size=1, stride=1, padding=0) - - def _init_vectorized_inputs(self, norm_type): - assert self.config.sample_size is not None, "Transformer2DModel over discrete input must provide sample_size" - assert self.config.num_vector_embeds is not None, ( - "Transformer2DModel over discrete input must provide num_embed" - ) - - self.height = self.config.sample_size - self.width = self.config.sample_size - self.num_latent_pixels = self.height * self.width - - self.latent_image_embedding = ImagePositionalEmbeddings( - num_embed=self.config.num_vector_embeds, embed_dim=self.inner_dim, height=self.height, width=self.width - ) - - self.transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlock( - self.inner_dim, - self.config.num_attention_heads, - self.config.attention_head_dim, - dropout=self.config.dropout, - cross_attention_dim=self.config.cross_attention_dim, - activation_fn=self.config.activation_fn, - num_embeds_ada_norm=self.config.num_embeds_ada_norm, - attention_bias=self.config.attention_bias, - only_cross_attention=self.config.only_cross_attention, - double_self_attention=self.config.double_self_attention, - upcast_attention=self.config.upcast_attention, - norm_type=norm_type, - norm_elementwise_affine=self.config.norm_elementwise_affine, - norm_eps=self.config.norm_eps, - attention_type=self.config.attention_type, - ) - for _ in range(self.config.num_layers) - ] - ) - - self.norm_out = nn.LayerNorm(self.inner_dim) - self.out = nn.Linear(self.inner_dim, self.config.num_vector_embeds - 1) - - def _init_patched_inputs(self, norm_type): - assert self.config.sample_size is not None, "Transformer2DModel over patched input must provide sample_size" - - self.height = self.config.sample_size - self.width = self.config.sample_size - - self.patch_size = self.config.patch_size - interpolation_scale = ( - self.config.interpolation_scale - if self.config.interpolation_scale is not None - else max(self.config.sample_size // 64, 1) - ) - self.pos_embed = PatchEmbed( - height=self.config.sample_size, - width=self.config.sample_size, - patch_size=self.config.patch_size, - in_channels=self.in_channels, - embed_dim=self.inner_dim, - interpolation_scale=interpolation_scale, - ) - - self.transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlock( - self.inner_dim, - self.config.num_attention_heads, - self.config.attention_head_dim, - dropout=self.config.dropout, - cross_attention_dim=self.config.cross_attention_dim, - activation_fn=self.config.activation_fn, - num_embeds_ada_norm=self.config.num_embeds_ada_norm, - attention_bias=self.config.attention_bias, - only_cross_attention=self.config.only_cross_attention, - double_self_attention=self.config.double_self_attention, - upcast_attention=self.config.upcast_attention, - norm_type=norm_type, - norm_elementwise_affine=self.config.norm_elementwise_affine, - norm_eps=self.config.norm_eps, - attention_type=self.config.attention_type, - ) - for _ in range(self.config.num_layers) - ] - ) - - if self.config.norm_type != "ada_norm_single": - self.norm_out = nn.LayerNorm(self.inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out_1 = nn.Linear(self.inner_dim, 2 * self.inner_dim) - self.proj_out_2 = nn.Linear( - self.inner_dim, self.config.patch_size * self.config.patch_size * self.out_channels - ) - elif self.config.norm_type == "ada_norm_single": - self.norm_out = nn.LayerNorm(self.inner_dim, elementwise_affine=False, eps=1e-6) - self.scale_shift_table = nn.Parameter(torch.randn(2, self.inner_dim) / self.inner_dim**0.5) - self.proj_out = nn.Linear( - self.inner_dim, self.config.patch_size * self.config.patch_size * self.out_channels - ) - - # PixArt-Alpha blocks. - self.adaln_single = None - if self.config.norm_type == "ada_norm_single": - # TODO(Sayak, PVP) clean this, for now we use sample size to determine whether to use - # additional conditions until we find better name - self.adaln_single = AdaLayerNormSingle( - self.inner_dim, use_additional_conditions=self.use_additional_conditions - ) - - self.caption_projection = None - if self.caption_channels is not None: - self.caption_projection = PixArtAlphaTextProjection( - in_features=self.caption_channels, hidden_size=self.inner_dim - ) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - timestep: torch.LongTensor | None = None, - added_cond_kwargs: dict[str, torch.Tensor] = None, - class_labels: torch.LongTensor | None = None, - cross_attention_kwargs: dict[str, Any] = None, - attention_mask: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - return_dict: bool = True, - ): - """ - The [`Transformer2DModel`] forward method. - - Args: - hidden_states (`torch.LongTensor` of shape `(batch size, num latent pixels)` if discrete, `torch.Tensor` of shape `(batch size, channel, height, width)` if continuous): - Input `hidden_states`. - encoder_hidden_states ( `torch.Tensor` of shape `(batch size, sequence len, embed dims)`, *optional*): - Conditional embeddings for cross attention layer. If not given, cross-attention defaults to - self-attention. - timestep ( `torch.LongTensor`, *optional*): - Used to indicate denoising step. Optional timestep to be applied as an embedding in `AdaLayerNorm`. - class_labels ( `torch.LongTensor` of shape `(batch size, num classes)`, *optional*): - Used to indicate class labels conditioning. Optional class labels to be applied as an embedding in - `AdaLayerZeroNorm`. - cross_attention_kwargs ( `dict[str, Any]`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - attention_mask ( `torch.Tensor`, *optional*): - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask - is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large - negative values to the attention scores corresponding to "discard" tokens. - encoder_attention_mask ( `torch.Tensor`, *optional*): - Cross-attention mask applied to `encoder_hidden_states`. Two formats supported: - - * Mask `(batch, sequence_length)` True = keep, False = discard. - * Bias `(batch, 1, sequence_length)` 0 = keep, -10000 = discard. - - If `ndim == 2`: will be interpreted as a mask, then converted into a bias consistent with the format - above. This bias will be added to the cross-attention scores. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unets.unet_2d_condition.UNet2DConditionOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformers.transformer_2d.Transformer2DModelOutput`] is returned, - otherwise a `tuple` where the first element is the sample tensor. - """ - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - # ensure attention_mask is a bias, and give it a singleton query_tokens dimension. - # we may have done this conversion already, e.g. if we came here via UNet2DConditionModel#forward. - # we can tell by counting dims; if ndim == 2: it's a mask rather than a bias. - # expects mask of shape: - # [batch, key_tokens] - # adds singleton query_tokens dimension: - # [batch, 1, key_tokens] - # this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes: - # [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn) - # [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn) - if attention_mask is not None and attention_mask.ndim == 2: - # assume that mask is expressed as: - # (1 = keep, 0 = discard) - # convert mask into a bias that can be added to attention scores: - # (keep = +0, discard = -10000.0) - attention_mask = (1 - attention_mask.to(hidden_states.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # convert encoder_attention_mask to a bias the same way we do for attention_mask - if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2: - encoder_attention_mask = (1 - encoder_attention_mask.to(hidden_states.dtype)) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - # 1. Input - if self.is_input_continuous: - batch_size, _, height, width = hidden_states.shape - residual = hidden_states - hidden_states, inner_dim = self._operate_on_continuous_inputs(hidden_states) - elif self.is_input_vectorized: - hidden_states = self.latent_image_embedding(hidden_states) - elif self.is_input_patches: - height, width = hidden_states.shape[-2] // self.patch_size, hidden_states.shape[-1] // self.patch_size - hidden_states, encoder_hidden_states, timestep, embedded_timestep = self._operate_on_patched_inputs( - hidden_states, encoder_hidden_states, timestep, added_cond_kwargs - ) - - # 2. Blocks - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - attention_mask, - encoder_hidden_states, - encoder_attention_mask, - timestep, - cross_attention_kwargs, - class_labels, - ) - else: - hidden_states = block( - hidden_states, - attention_mask=attention_mask, - encoder_hidden_states=encoder_hidden_states, - encoder_attention_mask=encoder_attention_mask, - timestep=timestep, - cross_attention_kwargs=cross_attention_kwargs, - class_labels=class_labels, - ) - - # 3. Output - if self.is_input_continuous: - output = self._get_output_for_continuous_inputs( - hidden_states=hidden_states, - residual=residual, - batch_size=batch_size, - height=height, - width=width, - inner_dim=inner_dim, - ) - elif self.is_input_vectorized: - output = self._get_output_for_vectorized_inputs(hidden_states) - elif self.is_input_patches: - output = self._get_output_for_patched_inputs( - hidden_states=hidden_states, - timestep=timestep, - class_labels=class_labels, - embedded_timestep=embedded_timestep, - height=height, - width=width, - ) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) - - def _operate_on_continuous_inputs(self, hidden_states): - batch, _, height, width = hidden_states.shape - hidden_states = self.norm(hidden_states) - - if not self.use_linear_projection: - hidden_states = self.proj_in(hidden_states) - inner_dim = hidden_states.shape[1] - hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * width, inner_dim) - else: - inner_dim = hidden_states.shape[1] - hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * width, inner_dim) - hidden_states = self.proj_in(hidden_states) - - return hidden_states, inner_dim - - def _operate_on_patched_inputs(self, hidden_states, encoder_hidden_states, timestep, added_cond_kwargs): - batch_size = hidden_states.shape[0] - hidden_states = self.pos_embed(hidden_states) - embedded_timestep = None - - if self.adaln_single is not None: - if self.use_additional_conditions and added_cond_kwargs is None: - raise ValueError( - "`added_cond_kwargs` cannot be None when using additional conditions for `adaln_single`." - ) - timestep, embedded_timestep = self.adaln_single( - timestep, added_cond_kwargs, batch_size=batch_size, hidden_dtype=hidden_states.dtype - ) - - if self.caption_projection is not None: - encoder_hidden_states = self.caption_projection(encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states.view(batch_size, -1, hidden_states.shape[-1]) - - return hidden_states, encoder_hidden_states, timestep, embedded_timestep - - def _get_output_for_continuous_inputs(self, hidden_states, residual, batch_size, height, width, inner_dim): - if not self.use_linear_projection: - hidden_states = ( - hidden_states.reshape(batch_size, height, width, inner_dim).permute(0, 3, 1, 2).contiguous() - ) - hidden_states = self.proj_out(hidden_states) - else: - hidden_states = self.proj_out(hidden_states) - hidden_states = ( - hidden_states.reshape(batch_size, height, width, inner_dim).permute(0, 3, 1, 2).contiguous() - ) - - output = hidden_states + residual - return output - - def _get_output_for_vectorized_inputs(self, hidden_states): - hidden_states = self.norm_out(hidden_states) - logits = self.out(hidden_states) - # (batch, self.num_vector_embeds - 1, self.num_latent_pixels) - logits = logits.permute(0, 2, 1) - # log(p(x_0)) - output = F.log_softmax(logits.double(), dim=1).float() - return output - - def _get_output_for_patched_inputs( - self, hidden_states, timestep, class_labels, embedded_timestep, height=None, width=None - ): - if self.config.norm_type != "ada_norm_single": - conditioning = self.transformer_blocks[0].norm1.emb( - timestep, class_labels, hidden_dtype=hidden_states.dtype - ) - shift, scale = self.proj_out_1(F.silu(conditioning)).chunk(2, dim=1) - hidden_states = self.norm_out(hidden_states) * (1 + scale[:, None]) + shift[:, None] - hidden_states = self.proj_out_2(hidden_states) - elif self.config.norm_type == "ada_norm_single": - shift, scale = (self.scale_shift_table[None] + embedded_timestep[:, None]).chunk(2, dim=1) - hidden_states = self.norm_out(hidden_states) - # Modulation - hidden_states = hidden_states * (1 + scale) + shift - hidden_states = self.proj_out(hidden_states) - hidden_states = hidden_states.squeeze(1) - - # unpatchify - if self.adaln_single is None: - height = width = int(hidden_states.shape[1] ** 0.5) - hidden_states = hidden_states.reshape( - shape=(-1, height, width, self.patch_size, self.patch_size, self.out_channels) - ) - hidden_states = torch.einsum("nhwpqc->nchpwq", hidden_states) - output = hidden_states.reshape( - shape=(-1, self.out_channels, height * self.patch_size, width * self.patch_size) - ) - return output diff --git a/diffusers/models/transformers/transformer_2d_dreamlite.py b/diffusers/models/transformers/transformer_2d_dreamlite.py deleted file mode 100644 index 9d66eeafbd002ac79793cc29ea0997a52f823226..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_2d_dreamlite.py +++ /dev/null @@ -1,598 +0,0 @@ -# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -"""DreamLite 2D transformer. - -This module is intentionally self-contained: it defines - -* ``BasicTransformerBlockDreamLite`` — a DreamLite-flavoured variant of - :class:`~diffusers.models.attention.BasicTransformerBlock` with four additional knobs (``use_self_attention``, - ``qk_norm``, ``num_kv_heads``, ``ff_mult``); and -* ``DreamLiteTransformer2DModel`` — a continuous-input-only counterpart of - :class:`~diffusers.models.transformers.transformer_2d.Transformer2DModel` that wires those knobs all the way down to - each block. - -Keeping everything here means the DreamLite integration never touches the upstream ``attention.py`` / -``transformer_2d.py``, which is the convention followed by other ported pipelines (SD3, Flux, Chroma, …). - -The numerical behaviour mirrors the original DreamLite reference implementation at ``dreamlite/models/{attention.py, -transformers/transformer_2d.py}`` — specifically, when ``use_self_attention=False`` the block keeps ``norm1``'s output -as the post-self-attn hidden state instead of running ``attn1``, matching the "Remove self-attention" path used by -DreamLite's ``DreamLiteCrossAttnNoSelfAttnDownBlock2D`` and ``DreamLiteCrossAttnNoSelfAttnUpBlock2D``. -""" - -from typing import Any - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ..attention import FeedForward, GatedSelfAttentionDense, _chunked_feed_forward -from ..attention_processor import Attention -from ..embeddings import SinusoidalPositionalEmbedding -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNorm, AdaLayerNormContinuous, AdaLayerNormZero -from .transformer_2d import Transformer2DModelOutput - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class BasicTransformerBlockDreamLite(nn.Module): - r"""DreamLite variant of :class:`BasicTransformerBlock`. - - Adds four constructor knobs on top of the upstream block: - - * ``use_self_attention`` — when ``False``, ``attn1`` is *not* instantiated and the self-attention residual branch - in ``forward`` is replaced by ``norm1``'s output (no add-residual). This implements DreamLite's "Remove - self-attention" trick used inside ``DreamLiteCrossAttnNoSelfAttnDownBlock2D`` / - ``DreamLiteCrossAttnNoSelfAttnUpBlock2D``. - * ``qk_norm`` — propagated to both attention layers' ``qk_norm``. - * ``num_kv_heads`` — propagated to both attention layers' ``kv_heads`` (enables Grouped-Query Attention). - * ``ff_mult`` — propagated to :class:`FeedForward.mult` (DreamLite uses a non-default expansion factor). - - Only the ``norm_type`` values actually exercised by DreamLite are supported in detail (``layer_norm`` and - ``ada_norm``); the other branches are preserved verbatim from the upstream block so that callers writing new - variants do not have to re-port them. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - dropout: float = 0.0, - cross_attention_dim: int | None = None, - activation_fn: str = "geglu", - num_embeds_ada_norm: int | None = None, - attention_bias: bool = False, - only_cross_attention: bool = False, - double_self_attention: bool = False, - upcast_attention: bool = False, - norm_elementwise_affine: bool = True, - norm_type: str = "layer_norm", - norm_eps: float = 1e-5, - final_dropout: bool = False, - attention_type: str = "default", - positional_embeddings: str | None = None, - num_positional_embeddings: int | None = None, - ada_norm_continous_conditioning_embedding_dim: int | None = None, - ada_norm_bias: int | None = None, - ff_inner_dim: int | None = None, - ff_bias: bool = True, - attention_out_bias: bool = True, - use_self_attention: bool = True, - qk_norm: str | None = None, - num_kv_heads: int | None = None, - ff_mult: int = 4, - ): - super().__init__() - self.dim = dim - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - self.dropout = dropout - self.cross_attention_dim = cross_attention_dim - self.activation_fn = activation_fn - self.attention_bias = attention_bias - self.double_self_attention = double_self_attention - self.norm_elementwise_affine = norm_elementwise_affine - self.positional_embeddings = positional_embeddings - self.num_positional_embeddings = num_positional_embeddings - self.only_cross_attention = only_cross_attention - self.use_self_attention = use_self_attention - - if not use_self_attention and norm_type in ("ada_norm_zero", "ada_norm_single"): - raise ValueError( - f"`use_self_attention=False` is incompatible with `norm_type={norm_type}` because " - "the gate/shift/scale modulation tuple is derived from `norm1`. " - "Use `norm_type='layer_norm'` or `'ada_norm'` instead." - ) - - # Backward-compatible boolean flags (kept for parity with BasicTransformerBlock). - self.use_ada_layer_norm_zero = (num_embeds_ada_norm is not None) and norm_type == "ada_norm_zero" - self.use_ada_layer_norm = (num_embeds_ada_norm is not None) and norm_type == "ada_norm" - self.use_ada_layer_norm_single = norm_type == "ada_norm_single" - self.use_layer_norm = norm_type == "layer_norm" - self.use_ada_layer_norm_continuous = norm_type == "ada_norm_continuous" - - if norm_type in ("ada_norm", "ada_norm_zero") and num_embeds_ada_norm is None: - raise ValueError( - f"`norm_type` is set to {norm_type}, but `num_embeds_ada_norm` is not defined. " - f"Please make sure to define `num_embeds_ada_norm` if setting `norm_type` to {norm_type}." - ) - - self.norm_type = norm_type - self.num_embeds_ada_norm = num_embeds_ada_norm - - if positional_embeddings and (num_positional_embeddings is None): - raise ValueError( - "If `positional_embedding` type is defined, `num_positition_embeddings` must also be defined." - ) - - if positional_embeddings == "sinusoidal": - self.pos_embed = SinusoidalPositionalEmbedding(dim, max_seq_length=num_positional_embeddings) - else: - self.pos_embed = None - - # 1. Self-Attn (or its replacement) - if norm_type == "ada_norm": - self.norm1 = AdaLayerNorm(dim, num_embeds_ada_norm) - elif norm_type == "ada_norm_zero": - self.norm1 = AdaLayerNormZero(dim, num_embeds_ada_norm) - elif norm_type == "ada_norm_continuous": - self.norm1 = AdaLayerNormContinuous( - dim, - ada_norm_continous_conditioning_embedding_dim, - norm_elementwise_affine, - norm_eps, - ada_norm_bias, - "rms_norm", - ) - else: - self.norm1 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps) - - if use_self_attention: - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - cross_attention_dim=cross_attention_dim if only_cross_attention else None, - upcast_attention=upcast_attention, - out_bias=attention_out_bias, - qk_norm=qk_norm, - kv_heads=num_kv_heads, - ) - else: - self.attn1 = None - - # 2. Cross-Attn - if cross_attention_dim is not None or double_self_attention: - if norm_type == "ada_norm": - self.norm2 = AdaLayerNorm(dim, num_embeds_ada_norm) - elif norm_type == "ada_norm_continuous": - self.norm2 = AdaLayerNormContinuous( - dim, - ada_norm_continous_conditioning_embedding_dim, - norm_elementwise_affine, - norm_eps, - ada_norm_bias, - "rms_norm", - ) - else: - self.norm2 = nn.LayerNorm(dim, norm_eps, norm_elementwise_affine) - - self.attn2 = Attention( - query_dim=dim, - cross_attention_dim=cross_attention_dim if not double_self_attention else None, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - upcast_attention=upcast_attention, - out_bias=attention_out_bias, - qk_norm=qk_norm, - kv_heads=num_kv_heads, - ) - else: - if norm_type == "ada_norm_single": - self.norm2 = nn.LayerNorm(dim, norm_eps, norm_elementwise_affine) - else: - self.norm2 = None - self.attn2 = None - - # 3. Feed-forward - if norm_type == "ada_norm_continuous": - self.norm3 = AdaLayerNormContinuous( - dim, - ada_norm_continous_conditioning_embedding_dim, - norm_elementwise_affine, - norm_eps, - ada_norm_bias, - "layer_norm", - ) - elif norm_type in ["ada_norm_zero", "ada_norm", "layer_norm"]: - self.norm3 = nn.LayerNorm(dim, norm_eps, norm_elementwise_affine) - elif norm_type == "layer_norm_i2vgen": - self.norm3 = None - - self.ff = FeedForward( - dim, - dropout=dropout, - activation_fn=activation_fn, - final_dropout=final_dropout, - inner_dim=ff_inner_dim, - bias=ff_bias, - mult=ff_mult, - ) - - # 4. Fuser - if attention_type == "gated" or attention_type == "gated-text-image": - self.fuser = GatedSelfAttentionDense(dim, cross_attention_dim, num_attention_heads, attention_head_dim) - - # 5. Scale-shift for PixArt-Alpha (kept for completeness; DreamLite does not use it). - if norm_type == "ada_norm_single": - self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5) - - # let chunk size default to None - self._chunk_size = None - self._chunk_dim = 0 - - def set_chunk_feed_forward(self, chunk_size: int | None, dim: int = 0): - self._chunk_size = chunk_size - self._chunk_dim = dim - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - timestep: torch.LongTensor | None = None, - cross_attention_kwargs: dict[str, Any] = None, - class_labels: torch.LongTensor | None = None, - added_cond_kwargs: dict[str, torch.Tensor] | None = None, - ) -> torch.Tensor: - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - # 0. Self-Attention norm - batch_size = hidden_states.shape[0] - - if self.norm_type == "ada_norm": - norm_hidden_states = self.norm1(hidden_states, timestep) - elif self.norm_type == "ada_norm_zero": - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1( - hidden_states, timestep, class_labels, hidden_dtype=hidden_states.dtype - ) - elif self.norm_type in ["layer_norm", "layer_norm_i2vgen"]: - norm_hidden_states = self.norm1(hidden_states) - elif self.norm_type == "ada_norm_continuous": - norm_hidden_states = self.norm1(hidden_states, added_cond_kwargs["pooled_text_emb"]) - elif self.norm_type == "ada_norm_single": - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( - self.scale_shift_table[None] + timestep.reshape(batch_size, 6, -1) - ).chunk(6, dim=1) - norm_hidden_states = self.norm1(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_msa) + shift_msa - else: - raise ValueError("Incorrect norm used") - - if self.pos_embed is not None: - norm_hidden_states = self.pos_embed(norm_hidden_states) - - # 1. GLIGEN kwargs split - cross_attention_kwargs = cross_attention_kwargs.copy() if cross_attention_kwargs is not None else {} - gligen_kwargs = cross_attention_kwargs.pop("gligen", None) - - if self.use_self_attention: - attn_output = self.attn1( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states if self.only_cross_attention else None, - attention_mask=attention_mask, - **cross_attention_kwargs, - ) - - if self.norm_type == "ada_norm_zero": - attn_output = gate_msa.unsqueeze(1) * attn_output - elif self.norm_type == "ada_norm_single": - attn_output = gate_msa * attn_output - - hidden_states = attn_output + hidden_states - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - else: - # DreamLite "Remove self-attention" path: drop attn1 entirely and let - # the normalized state propagate as-is to cross-attn / FF. Matches - # upstream DreamLite `BasicTransformerBlock.forward` when - # `use_self_attention=False`. - hidden_states = norm_hidden_states - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - - # 1.2 GLIGEN control - if gligen_kwargs is not None: - hidden_states = self.fuser(hidden_states, gligen_kwargs["objs"]) - - # 3. Cross-Attention - if self.attn2 is not None: - if self.norm_type == "ada_norm": - norm_hidden_states = self.norm2(hidden_states, timestep) - elif self.norm_type in ["ada_norm_zero", "layer_norm", "layer_norm_i2vgen"]: - norm_hidden_states = self.norm2(hidden_states) - elif self.norm_type == "ada_norm_single": - norm_hidden_states = hidden_states - elif self.norm_type == "ada_norm_continuous": - norm_hidden_states = self.norm2(hidden_states, added_cond_kwargs["pooled_text_emb"]) - else: - raise ValueError("Incorrect norm") - - if self.pos_embed is not None and self.norm_type != "ada_norm_single": - norm_hidden_states = self.pos_embed(norm_hidden_states) - - attn_output = self.attn2( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=encoder_attention_mask, - **cross_attention_kwargs, - ) - hidden_states = attn_output + hidden_states - - # 4. Feed-forward - if self.norm_type == "ada_norm_continuous": - norm_hidden_states = self.norm3(hidden_states, added_cond_kwargs["pooled_text_emb"]) - elif not self.norm_type == "ada_norm_single": - norm_hidden_states = self.norm3(hidden_states) - - if self.norm_type == "ada_norm_zero": - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - - if self.norm_type == "ada_norm_single": - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp - - if self._chunk_size is not None: - ff_output = _chunked_feed_forward(self.ff, norm_hidden_states, self._chunk_dim, self._chunk_size) - else: - ff_output = self.ff(norm_hidden_states) - - if self.norm_type == "ada_norm_zero": - ff_output = gate_mlp.unsqueeze(1) * ff_output - elif self.norm_type == "ada_norm_single": - ff_output = gate_mlp * ff_output - - hidden_states = ff_output + hidden_states - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - - return hidden_states - - -class DreamLiteTransformer2DModel(ModelMixin, ConfigMixin): - r"""Continuous-input 2D transformer used by the DreamLite U-Net. - - Equivalent to :class:`Transformer2DModel` restricted to the ``is_input_continuous`` branch (``in_channels`` set, - ``patch_size`` and ``num_vector_embeds`` both ``None``), with four extra knobs that are propagated into every - :class:`BasicTransformerBlockDreamLite`: - - * ``use_self_attention`` — set ``False`` from ``CrossAttn*RemoveSelfAttnBlock2D*DreamLite`` to enable DreamLite's - "Remove self-attention" path. - * ``qk_norm`` — RMS/LayerNorm applied to Q and K projections. - * ``num_kv_heads`` — enables Grouped-Query Attention when fewer than ``num_attention_heads``. - * ``ff_mult`` — feed-forward expansion factor (DreamLite uses a non-default value). - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["BasicTransformerBlockDreamLite"] - _skip_layerwise_casting_patterns = ["norm"] - - @register_to_config - def __init__( - self, - num_attention_heads: int = 16, - attention_head_dim: int = 88, - in_channels: int | None = None, - out_channels: int | None = None, - num_layers: int = 1, - dropout: float = 0.0, - norm_num_groups: int = 32, - cross_attention_dim: int | None = None, - attention_bias: bool = False, - activation_fn: str = "geglu", - num_embeds_ada_norm: int | None = None, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - double_self_attention: bool = False, - upcast_attention: bool = False, - norm_type: str = "layer_norm", - norm_elementwise_affine: bool = True, - norm_eps: float = 1e-5, - attention_type: str = "default", - use_self_attention: bool = True, - qk_norm: str | None = None, - num_kv_heads: int | None = None, - ff_mult: int = 4, - ): - super().__init__() - - if in_channels is None: - raise ValueError( - "`DreamLiteTransformer2DModel` only supports continuous inputs; `in_channels` must be provided." - ) - - self.use_linear_projection = use_linear_projection - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim - self.in_channels = in_channels - self.out_channels = in_channels if out_channels is None else out_channels - self.gradient_checkpointing = False - - self.norm = torch.nn.GroupNorm( - num_groups=self.config.norm_num_groups, num_channels=self.in_channels, eps=1e-6, affine=True - ) - if self.use_linear_projection: - self.proj_in = torch.nn.Linear(self.in_channels, self.inner_dim) - else: - self.proj_in = torch.nn.Conv2d(self.in_channels, self.inner_dim, kernel_size=1, stride=1, padding=0) - - self.transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlockDreamLite( - self.inner_dim, - self.config.num_attention_heads, - self.config.attention_head_dim, - dropout=self.config.dropout, - cross_attention_dim=self.config.cross_attention_dim, - activation_fn=self.config.activation_fn, - num_embeds_ada_norm=self.config.num_embeds_ada_norm, - attention_bias=self.config.attention_bias, - only_cross_attention=self.config.only_cross_attention, - double_self_attention=self.config.double_self_attention, - upcast_attention=self.config.upcast_attention, - norm_type=norm_type, - norm_elementwise_affine=self.config.norm_elementwise_affine, - norm_eps=self.config.norm_eps, - attention_type=self.config.attention_type, - use_self_attention=self.config.use_self_attention, - qk_norm=self.config.qk_norm, - num_kv_heads=self.config.num_kv_heads, - ff_mult=self.config.ff_mult, - ) - for _ in range(self.config.num_layers) - ] - ) - - if self.use_linear_projection: - self.proj_out = torch.nn.Linear(self.inner_dim, self.out_channels) - else: - self.proj_out = torch.nn.Conv2d(self.inner_dim, self.out_channels, kernel_size=1, stride=1, padding=0) - - def _operate_on_continuous_inputs(self, hidden_states): - batch, _, height, width = hidden_states.shape - hidden_states = self.norm(hidden_states) - - if not self.use_linear_projection: - hidden_states = self.proj_in(hidden_states) - inner_dim = hidden_states.shape[1] - hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * width, inner_dim) - else: - inner_dim = hidden_states.shape[1] - hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * width, inner_dim) - hidden_states = self.proj_in(hidden_states) - - return hidden_states, inner_dim - - def _get_output_for_continuous_inputs(self, hidden_states, residual, batch_size, height, width, inner_dim): - if not self.use_linear_projection: - hidden_states = ( - hidden_states.reshape(batch_size, height, width, inner_dim).permute(0, 3, 1, 2).contiguous() - ) - hidden_states = self.proj_out(hidden_states) - else: - hidden_states = self.proj_out(hidden_states) - hidden_states = ( - hidden_states.reshape(batch_size, height, width, inner_dim).permute(0, 3, 1, 2).contiguous() - ) - - output = hidden_states + residual - return output - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - timestep: torch.LongTensor | None = None, - added_cond_kwargs: dict[str, torch.Tensor] = None, - class_labels: torch.LongTensor | None = None, - cross_attention_kwargs: dict[str, Any] = None, - attention_mask: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - return_dict: bool = True, - ): - """Forward pass of :class:`DreamLiteTransformer2DModel`. - - Args: - hidden_states: Input latent tensor of shape ``(batch, channels, height, width)``. - encoder_hidden_states: Cross-attention conditioning embeddings. - timestep: Diffusion timestep(s); broadcast to batch if scalar. - added_cond_kwargs: Optional extra conditioning (e.g. ``text_embeds``, ``time_ids``). - class_labels: Optional class labels for class-conditional generation. - cross_attention_kwargs: Optional kwargs forwarded to the cross-attention processor. - Note: passing ``scale`` is deprecated and will be ignored. - attention_mask: Optional self-attention mask; 2D masks are converted to additive biases. - encoder_attention_mask: Optional cross-attention mask; 2D masks are converted to additive biases. - return_dict: If ``True``, returns a :class:`Transformer2DModelOutput`; otherwise a 1-tuple ``(sample,)``. - - Returns: - :class:`~diffusers.models.transformers.transformer_2d.Transformer2DModelOutput` (or a 1-tuple of the - sample) — kept output-compatible with the upstream class so callers don't have to special-case DreamLite. - """ - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - # Keep masks as bool tensors — dispatch_attention_fn handles per-backend conversion - # internally. Dense additive float masks would hard-raise on flash / sage backends. - if attention_mask is not None and attention_mask.ndim == 2: - attention_mask = attention_mask.bool() - - if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2: - encoder_attention_mask = encoder_attention_mask.bool() - - # 1. Input - batch_size, _, height, width = hidden_states.shape - residual = hidden_states - hidden_states, inner_dim = self._operate_on_continuous_inputs(hidden_states) - - # 2. Blocks - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - attention_mask, - encoder_hidden_states, - encoder_attention_mask, - timestep, - cross_attention_kwargs, - class_labels, - ) - else: - hidden_states = block( - hidden_states, - attention_mask=attention_mask, - encoder_hidden_states=encoder_hidden_states, - encoder_attention_mask=encoder_attention_mask, - timestep=timestep, - cross_attention_kwargs=cross_attention_kwargs, - class_labels=class_labels, - ) - - # 3. Output - output = self._get_output_for_continuous_inputs( - hidden_states=hidden_states, - residual=residual, - batch_size=batch_size, - height=height, - width=width, - inner_dim=inner_dim, - ) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_allegro.py b/diffusers/models/transformers/transformer_allegro.py deleted file mode 100644 index abe82ab578debdb47c51cbbe3d59088dcc067a45..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_allegro.py +++ /dev/null @@ -1,436 +0,0 @@ -# Copyright 2025 The RhymesAI and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import FeedForward -from ..attention_processor import AllegroAttnProcessor2_0, Attention -from ..cache_utils import CacheMixin -from ..embeddings import PatchEmbed, PixArtAlphaTextProjection -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormSingle - - -logger = logging.get_logger(__name__) - - -@maybe_allow_in_graph -class AllegroTransformerBlock(nn.Module): - r""" - Transformer block used in [Allegro](https://github.com/rhymes-ai/Allegro) model. - - Args: - dim (`int`): - The number of channels in the input and output. - num_attention_heads (`int`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`): - The number of channels in each head. - dropout (`float`, defaults to `0.0`): - The dropout probability to use. - cross_attention_dim (`int`, defaults to `2304`): - The dimension of the cross attention features. - activation_fn (`str`, defaults to `"gelu-approximate"`): - Activation function to be used in feed-forward. - attention_bias (`bool`, defaults to `False`): - Whether or not to use bias in attention projection layers. - only_cross_attention (`bool`, defaults to `False`): - norm_elementwise_affine (`bool`, defaults to `True`): - Whether to use learnable elementwise affine parameters for normalization. - norm_eps (`float`, defaults to `1e-5`): - Epsilon value for normalization layers. - final_dropout (`bool` defaults to `False`): - Whether to apply a final dropout after the last feed-forward layer. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - dropout=0.0, - cross_attention_dim: int | None = None, - activation_fn: str = "geglu", - attention_bias: bool = False, - norm_elementwise_affine: bool = True, - norm_eps: float = 1e-5, - ): - super().__init__() - - # 1. Self Attention - self.norm1 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps) - - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - cross_attention_dim=None, - processor=AllegroAttnProcessor2_0(), - ) - - # 2. Cross Attention - self.norm2 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps) - self.attn2 = Attention( - query_dim=dim, - cross_attention_dim=cross_attention_dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - processor=AllegroAttnProcessor2_0(), - ) - - # 3. Feed Forward - self.norm3 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps) - - self.ff = FeedForward( - dim, - dropout=dropout, - activation_fn=activation_fn, - ) - - # 4. Scale-shift - self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - temb: torch.LongTensor | None = None, - attention_mask: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - image_rotary_emb=None, - ) -> torch.Tensor: - # 0. Self-Attention - batch_size = hidden_states.shape[0] - - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( - self.scale_shift_table[None] + temb.reshape(batch_size, 6, -1) - ).chunk(6, dim=1) - norm_hidden_states = self.norm1(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_msa) + shift_msa - norm_hidden_states = norm_hidden_states.squeeze(1) - - attn_output = self.attn1( - norm_hidden_states, - encoder_hidden_states=None, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - attn_output = gate_msa * attn_output - - hidden_states = attn_output + hidden_states - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - - # 1. Cross-Attention - if self.attn2 is not None: - norm_hidden_states = hidden_states - - attn_output = self.attn2( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=encoder_attention_mask, - image_rotary_emb=None, - ) - hidden_states = attn_output + hidden_states - - # 2. Feed-forward - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp - - ff_output = self.ff(norm_hidden_states) - ff_output = gate_mlp * ff_output - - hidden_states = ff_output + hidden_states - - # TODO(aryan): maybe following line is not required - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - - return hidden_states - - -class AllegroTransformer3DModel(ModelMixin, ConfigMixin, CacheMixin): - _supports_gradient_checkpointing = True - - """ - A 3D Transformer model for video-like data. - - Args: - patch_size (`int`, defaults to `2`): - The size of spatial patches to use in the patch embedding layer. - patch_size_t (`int`, defaults to `1`): - The size of temporal patches to use in the patch embedding layer. - num_attention_heads (`int`, defaults to `24`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `96`): - The number of channels in each head. - in_channels (`int`, defaults to `4`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `4`): - The number of channels in the output. - num_layers (`int`, defaults to `32`): - The number of layers of Transformer blocks to use. - dropout (`float`, defaults to `0.0`): - The dropout probability to use. - cross_attention_dim (`int`, defaults to `2304`): - The dimension of the cross attention features. - attention_bias (`bool`, defaults to `True`): - Whether or not to use bias in the attention projection layers. - sample_height (`int`, defaults to `90`): - The height of the input latents. - sample_width (`int`, defaults to `160`): - The width of the input latents. - sample_frames (`int`, defaults to `22`): - The number of frames in the input latents. - activation_fn (`str`, defaults to `"gelu-approximate"`): - Activation function to use in feed-forward. - norm_elementwise_affine (`bool`, defaults to `False`): - Whether or not to use elementwise affine in normalization layers. - norm_eps (`float`, defaults to `1e-6`): - The epsilon value to use in normalization layers. - caption_channels (`int`, defaults to `4096`): - Number of channels to use for projecting the caption embeddings. - interpolation_scale_h (`float`, defaults to `2.0`): - Scaling factor to apply in 3D positional embeddings across height dimension. - interpolation_scale_w (`float`, defaults to `2.0`): - Scaling factor to apply in 3D positional embeddings across width dimension. - interpolation_scale_t (`float`, defaults to `2.2`): - Scaling factor to apply in 3D positional embeddings across time dimension. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["pos_embed", "norm", "adaln_single"] - - @register_to_config - def __init__( - self, - patch_size: int = 2, - patch_size_t: int = 1, - num_attention_heads: int = 24, - attention_head_dim: int = 96, - in_channels: int = 4, - out_channels: int = 4, - num_layers: int = 32, - dropout: float = 0.0, - cross_attention_dim: int = 2304, - attention_bias: bool = True, - sample_height: int = 90, - sample_width: int = 160, - sample_frames: int = 22, - activation_fn: str = "gelu-approximate", - norm_elementwise_affine: bool = False, - norm_eps: float = 1e-6, - caption_channels: int = 4096, - interpolation_scale_h: float = 2.0, - interpolation_scale_w: float = 2.0, - interpolation_scale_t: float = 2.2, - ): - super().__init__() - - self.inner_dim = num_attention_heads * attention_head_dim - - interpolation_scale_t = ( - interpolation_scale_t - if interpolation_scale_t is not None - else ((sample_frames - 1) // 16 + 1) - if sample_frames % 2 == 1 - else sample_frames // 16 - ) - interpolation_scale_h = interpolation_scale_h if interpolation_scale_h is not None else sample_height / 30 - interpolation_scale_w = interpolation_scale_w if interpolation_scale_w is not None else sample_width / 40 - - # 1. Patch embedding - self.pos_embed = PatchEmbed( - height=sample_height, - width=sample_width, - patch_size=patch_size, - in_channels=in_channels, - embed_dim=self.inner_dim, - pos_embed_type=None, - ) - - # 2. Transformer blocks - self.transformer_blocks = nn.ModuleList( - [ - AllegroTransformerBlock( - self.inner_dim, - num_attention_heads, - attention_head_dim, - dropout=dropout, - cross_attention_dim=cross_attention_dim, - activation_fn=activation_fn, - attention_bias=attention_bias, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - ) - for _ in range(num_layers) - ] - ) - - # 3. Output projection & norm - self.norm_out = nn.LayerNorm(self.inner_dim, elementwise_affine=False, eps=1e-6) - self.scale_shift_table = nn.Parameter(torch.randn(2, self.inner_dim) / self.inner_dim**0.5) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * out_channels) - - # 4. Timestep embeddings - self.adaln_single = AdaLayerNormSingle(self.inner_dim, use_additional_conditions=False) - - # 5. Caption projection - self.caption_projection = PixArtAlphaTextProjection(in_features=caption_channels, hidden_size=self.inner_dim) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: torch.LongTensor, - attention_mask: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - return_dict: bool = True, - ): - """ - The [`AllegroTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - attention_mask (`torch.Tensor`, *optional*): - Self-attention mask applied to `hidden_states`. - encoder_attention_mask (`torch.Tensor`, *optional*): - Cross-attention mask applied to `encoder_hidden_states`. - image_rotary_emb (`tuple` of `torch.Tensor`, *optional*): - Pre-computed rotary positional embeddings. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t = self.config.patch_size_t - p = self.config.patch_size - - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p - post_patch_width = width // p - - # ensure attention_mask is a bias, and give it a singleton query_tokens dimension. - # we may have done this conversion already, e.g. if we came here via UNet2DConditionModel#forward. - # we can tell by counting dims; if ndim == 2: it's a mask rather than a bias. - # expects mask of shape: - # [batch, key_tokens] - # adds singleton query_tokens dimension: - # [batch, 1, key_tokens] - # this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes: - # [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn) - # [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn) attention_mask_vid, attention_mask_img = None, None - if attention_mask is not None and attention_mask.ndim == 4: - # assume that mask is expressed as: - # (1 = keep, 0 = discard) - # convert mask into a bias that can be added to attention scores: - # (keep = +0, discard = -10000.0) - # b, frame+use_image_num, h, w -> a video with images - # b, 1, h, w -> only images - attention_mask = attention_mask.to(hidden_states.dtype) - attention_mask = attention_mask[:, :num_frames] # [batch_size, num_frames, height, width] - - if attention_mask.numel() > 0: - attention_mask = attention_mask.unsqueeze(1) # [batch_size, 1, num_frames, height, width] - attention_mask = F.max_pool3d(attention_mask, kernel_size=(p_t, p, p), stride=(p_t, p, p)) - attention_mask = attention_mask.flatten(1).view(batch_size, 1, -1) - - attention_mask = ( - (1 - attention_mask.bool().to(hidden_states.dtype)) * -10000.0 if attention_mask.numel() > 0 else None - ) - - # convert encoder_attention_mask to a bias the same way we do for attention_mask - if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2: - encoder_attention_mask = (1 - encoder_attention_mask.to(self.dtype)) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - # 1. Timestep embeddings - timestep, embedded_timestep = self.adaln_single( - timestep, batch_size=batch_size, hidden_dtype=hidden_states.dtype - ) - - # 2. Patch embeddings - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) - hidden_states = self.pos_embed(hidden_states) - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(1, 2) - - encoder_hidden_states = self.caption_projection(encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states.view(batch_size, -1, encoder_hidden_states.shape[-1]) - - # 3. Transformer blocks - for i, block in enumerate(self.transformer_blocks): - # TODO(aryan): Implement gradient checkpointing - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - timestep, - attention_mask, - encoder_attention_mask, - image_rotary_emb, - ) - else: - hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=timestep, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - image_rotary_emb=image_rotary_emb, - ) - - # 4. Output normalization & projection - shift, scale = (self.scale_shift_table[None] + embedded_timestep[:, None]).chunk(2, dim=1) - hidden_states = self.norm_out(hidden_states) - - # Modulation - hidden_states = hidden_states * (1 + scale) + shift - hidden_states = self.proj_out(hidden_states) - hidden_states = hidden_states.squeeze(1) - - # 5. Unpatchify - hidden_states = hidden_states.reshape( - batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p, p, -1 - ) - hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) - output = hidden_states.reshape(batch_size, -1, num_frames, height, width) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_anyflow.py b/diffusers/models/transformers/transformer_anyflow.py deleted file mode 100644 index 6b0872ffdb01e44201ff3f40f33c4e208bdfc475..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_anyflow.py +++ /dev/null @@ -1,726 +0,0 @@ -# Copyright 2026 The AnyFlow Team, NVIDIA Corp., and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -# -# This file derives from the FAR architecture (arXiv:2503.19325) and adds the -# AnyFlow dual-timestep flow-map embedding (AnyFlowDualTimestepTextImageEmbedding) introduced in -# AnyFlow (arXiv:2605.13724). The base 3D DiT structure is adapted from the -# v0.35.1 Wan2.1 transformer (transformer_wan.py); upstream Wan has since been refactored, so -# this file is intentionally self-contained rather than annotated with `# Copied from`. - -import math -from typing import Any, Dict, Optional, Tuple, Union - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..attention import AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps, get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import FP32LayerNorm, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def apply_rotary_emb(hidden_states: torch.Tensor, freqs: torch.Tensor): - # MPS / NPU backends do not support complex128 / float64; fall back to float32 on those devices. - rotary_dtype = maybe_adjust_dtype_for_device(torch.float64, hidden_states.device) - x_rotated = torch.view_as_complex(hidden_states.to(rotary_dtype).unflatten(3, (-1, 2))) - x_out = torch.view_as_real(x_rotated * freqs).flatten(3, 4) - return x_out.type_as(hidden_states) - - -class AnyFlowAttnProcessor: - """ - Bidirectional self-attention processor for AnyFlow. Routes through - :func:`~diffusers.models.attention_dispatch.dispatch_attention_fn` so any SDPA-compatible backend is supported - (SDPA, flash-attn, xformers, flex, …). FAR causal generation lives in - :class:`~diffusers.models.transformers.transformer_anyflow_far.AnyFlowCausalAttnProcessor`. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "AnyFlowAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0 or higher." - ) - - def __call__( - self, - attn: "AnyFlowAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[Any] = None, - rotary_emb: Optional[Dict[str, torch.Tensor]] = None, - ) -> torch.Tensor: - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Layout (B, H, L, D) for rotary application; transposed to (B, L, H, D) before dispatch. - query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2) - key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2) - value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2) - - if rotary_emb is not None: - query = apply_rotary_emb(query, rotary_emb["query"]) - key = apply_rotary_emb(key, rotary_emb["key"]) - - hidden_states = dispatch_attention_fn( - query.transpose(1, 2), - key.transpose(1, 2), - value.transpose(1, 2), - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class AnyFlowCrossAttnProcessor: - """ - Cross-attention processor for AnyFlow. Always uses the dispatched SDPA-compatible backend; no rotary embedding or - KV cache is applied to the text→video cross-attention path. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "AnyFlowCrossAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0 or higher." - ) - - def __call__( - self, - attn: "AnyFlowAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - ) -> torch.Tensor: - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # (B, L, H, D) layout for dispatch_attention_fn. - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class AnyFlowAttention(torch.nn.Module, AttentionModuleMixin): - """ - Attention module used by :class:`AnyFlowTransformerBlock`. Layout matches the legacy - :class:`~diffusers.models.attention_processor.Attention` so existing AnyFlow checkpoints load bit-exactly into this - class. - """ - - _default_processor_cls = AnyFlowAttnProcessor - _available_processors = [AnyFlowAttnProcessor, AnyFlowCrossAttnProcessor] - - def __init__( - self, - dim: int, - heads: int, - dim_head: int, - eps: float = 1e-6, - processor: Optional[Any] = None, - ): - super().__init__() - self.heads = heads - self.inner_dim = heads * dim_head - - self.to_q = torch.nn.Linear(dim, self.inner_dim, bias=True) - self.to_k = torch.nn.Linear(dim, self.inner_dim, bias=True) - self.to_v = torch.nn.Linear(dim, self.inner_dim, bias=True) - self.to_out = torch.nn.ModuleList( - [ - torch.nn.Linear(self.inner_dim, dim, bias=True), - torch.nn.Dropout(0.0), - ] - ) - # ``rms_norm_across_heads`` per-axis: normalize Q and K across the entire ``heads * dim_head`` - # channel axis. We use diffusers' RMSNorm (rather than ``torch.nn.RMSNorm``) so the numerics - # match the legacy Attention class that produced the released checkpoints. - self.norm_q = RMSNorm(self.inner_dim, eps=eps) - self.norm_k = RMSNorm(self.inner_dim, eps=eps) - - self.set_processor(processor if processor is not None else self._default_processor_cls()) - - def forward(self, hidden_states: torch.Tensor, **kwargs) -> torch.Tensor: - return self.processor(self, hidden_states, **kwargs) - - -class AnyFlowImageEmbedding(torch.nn.Module): - def __init__(self, in_features: int, out_features: int): - super().__init__() - - self.norm1 = FP32LayerNorm(in_features) - self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu") - self.norm2 = FP32LayerNorm(out_features) - - def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor: - hidden_states = self.norm1(encoder_hidden_states_image) - hidden_states = self.ff(hidden_states) - hidden_states = self.norm2(hidden_states) - return hidden_states - - -class AnyFlowDualTimestepTextImageEmbedding(nn.Module): - def __init__( - self, - dim: int, - gate_value: float, - deltatime_type: str, - time_freq_dim: int, - time_proj_dim: int, - text_embed_dim: int, - image_embed_dim: Optional[int] = None, - ): - super().__init__() - - self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim) - self.delta_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim) - self.act_fn = nn.SiLU() - self.time_proj = nn.Linear(dim, time_proj_dim) - self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh") - - self.image_embedder = None - if image_embed_dim is not None: - self.image_embedder = AnyFlowImageEmbedding(image_embed_dim, dim) - - self.register_buffer("delta_emb_gate", torch.tensor([gate_value], dtype=torch.float32), persistent=False) - self.deltatime_type = deltatime_type - - def forward_timestep( - self, timestep: torch.Tensor, delta_timestep: torch.Tensor, encoder_hidden_states, token_per_frame - ): - batch_size, num_frames = timestep.shape - timestep = timestep.reshape(-1) - delta_timestep = delta_timestep.reshape(-1) - - timestep = self.timesteps_proj(timestep) - - time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype - if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8: - timestep = timestep.to(time_embedder_dtype) - temb = self.time_embedder(timestep).type_as(encoder_hidden_states) - - delta_timestep = self.timesteps_proj(delta_timestep) - - delta_embedder_dtype = next(iter(self.delta_embedder.parameters())).dtype - if delta_timestep.dtype != delta_embedder_dtype and delta_embedder_dtype != torch.int8: - delta_timestep = delta_timestep.to(delta_embedder_dtype) - delta_emb = self.delta_embedder(delta_timestep).type_as(encoder_hidden_states) - - gate = self.delta_emb_gate.to(delta_embedder_dtype) - - rt_emb = (1 - gate) * temb + gate * delta_emb - timestep_proj = self.time_proj(self.act_fn(rt_emb)) - - rt_emb = rt_emb.unflatten(0, (batch_size, num_frames)).repeat_interleave(token_per_frame, dim=1) - timestep_proj = timestep_proj.unflatten(0, (batch_size, num_frames)).repeat_interleave(token_per_frame, dim=1) - - return rt_emb, timestep_proj - - def forward( - self, - timestep: torch.Tensor, - r_timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: Optional[torch.Tensor] = None, - layout_cfg=None, - ): - if self.deltatime_type == "r": - delta_timestep = r_timestep - elif self.deltatime_type == "t-r": - delta_timestep = timestep - r_timestep - else: - raise NotImplementedError - - timestep, timestep_proj = self.forward_timestep( - timestep, delta_timestep, encoder_hidden_states, layout_cfg["full_token_per_frame"] - ) - - encoder_hidden_states = self.text_embedder(encoder_hidden_states) - if encoder_hidden_states_image is not None: - encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image) - - return timestep, timestep_proj, encoder_hidden_states, encoder_hidden_states_image - - -class AnyFlowRotaryPosEmbed(nn.Module): - """Rotary positional embedding for the bidirectional AnyFlow transformer. - - The FAR causal variant lives in :mod:`~diffusers.models.transformers.transformer_anyflow_far` and additionally - handles compressed-frame chunks; this bidi class produces frequencies for the single full-resolution token grid - only. - """ - - def __init__( - self, - attention_head_dim: int, - patch_size: Tuple[int, int, int], - max_seq_len: int, - theta: float = 10000.0, - ): - super().__init__() - - self.attention_head_dim = attention_head_dim - self.patch_size = patch_size - self.max_seq_len = max_seq_len - self.theta = theta - - # Frequency table is lazily built per-device in ``_build_freqs``: MPS / NPU don't support - # complex128, so we downcast to complex64 there. - self._freqs_cache: Optional[Tuple[Any, torch.Tensor]] = None - - def _build_freqs(self, device: torch.device) -> torch.Tensor: - # Skip the cache read/write inside torch.compile: mutating ``self._freqs_cache`` between calls - # becomes a Dynamo guard and forces recompilation on the second invocation. - is_compiling = torch.compiler.is_compiling() - cache_key = (device.type, str(device)) - if not is_compiling and self._freqs_cache is not None and self._freqs_cache[0] == cache_key: - return self._freqs_cache[1] - - freqs_dtype = maybe_adjust_dtype_for_device(torch.float64, device) - - h_dim = w_dim = 2 * (self.attention_head_dim // 6) - t_dim = self.attention_head_dim - h_dim - w_dim - - freqs_list = [] - for dim in (t_dim, h_dim, w_dim): - f = get_1d_rotary_pos_embed( - dim, - self.max_seq_len, - self.theta, - use_real=False, - repeat_interleave_real=False, - freqs_dtype=freqs_dtype, - ) - freqs_list.append(f.to(device)) - freqs = torch.cat(freqs_list, dim=1) - if not is_compiling: - self._freqs_cache = (cache_key, freqs) - return freqs - - def _forward_full_frame(self, num_frames, height, width, device) -> torch.Tensor: - ppf, pph, ppw = num_frames, height, width - - freqs_full = self._build_freqs(device) - if min(ppf, pph, ppw) <= 0: - freq_channels = self.attention_head_dim // 2 - return torch.empty((ppf, pph, ppw, freq_channels), dtype=freqs_full.dtype, device=device) - - freqs = freqs_full.split_with_sizes( - [ - self.attention_head_dim // 2 - 2 * (self.attention_head_dim // 6), - self.attention_head_dim // 6, - self.attention_head_dim // 6, - ], - dim=1, - ) - - freqs_f = freqs[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_h = freqs[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1) - return freqs - - def forward(self, layout_cfg, device): - freqs = self._forward_full_frame( - num_frames=layout_cfg["total_frames"], - height=layout_cfg["full_frame_shape"][0], - width=layout_cfg["full_frame_shape"][1], - device=device, - ) - freqs = freqs.flatten(start_dim=0, end_dim=2) - freqs = freqs[None, None, ...] - return {"query": freqs, "key": freqs} - - -class AnyFlowTransformerBlock(nn.Module): - """AnyFlow transformer block. - - The self-attention processor is chosen at construction by ``is_causal``: the bidirectional transformer passes - ``is_causal=False`` (the default), the FAR causal transformer passes ``is_causal=True``. The forward pass is - identical in both modes — only the processor differs, so all causal-specific machinery (BlockMask, KV cache) lives - inside the processor. - """ - - def __init__( - self, - dim: int, - ffn_dim: int, - num_heads: int, - cross_attn_norm: bool = False, - eps: float = 1e-6, - is_causal: bool = False, - ): - super().__init__() - - self.is_causal = is_causal - - # 1. Self-attention. The causal processor lives in the FAR sibling module; lazy-import to - # avoid a circular import at module load time. - if is_causal: - from .transformer_anyflow_far import AnyFlowCausalAttnProcessor - - self_attn_processor = AnyFlowCausalAttnProcessor() - else: - self_attn_processor = AnyFlowAttnProcessor() - - self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False) - self.attn1 = AnyFlowAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - processor=self_attn_processor, - ) - - # 2. Cross-attention - self.attn2 = AnyFlowAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - processor=AnyFlowCrossAttnProcessor(), - ) - self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity() - - # 3. Feed-forward - self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate") - self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False) - - self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - rotary_emb: torch.Tensor, - attention_mask: torch.Tensor, - kv_cache=None, - kv_cache_flag=None, - ) -> torch.Tensor: - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table + temb.float() - ).chunk(6, dim=2) - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - shift_msa.squeeze(2), - scale_msa.squeeze(2), - gate_msa.squeeze(2), - c_shift_msa.squeeze(2), - c_scale_msa.squeeze(2), - c_gate_msa.squeeze(2), - ) # noqa: E501 - - # 1. Self-attention - norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states) - attn1_kwargs = { - "hidden_states": norm_hidden_states, - "rotary_emb": rotary_emb, - "attention_mask": attention_mask, - } - # KV cache kwargs are only consumed by the FAR causal processor; the bidi processor - # doesn't accept them, so we forward them only when they're actually populated. - if kv_cache is not None: - attn1_kwargs["kv_cache"] = kv_cache - attn1_kwargs["kv_cache_flag"] = kv_cache_flag - attn_output = self.attn1(**attn1_kwargs) - hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states) - - # 2. Cross-attention - norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states) - attn_output = self.attn2(hidden_states=norm_hidden_states, encoder_hidden_states=encoder_hidden_states) - hidden_states = hidden_states + attn_output - - # 3. Feed-forward - norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as( - hidden_states - ) - ff_output = self.ffn(norm_hidden_states) - hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states) - - return hidden_states - - -class AnyFlowTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): - r""" - Bidirectional 3D Transformer for AnyFlow flow-map sampling. - - The architecture is the v0.35.1 Wan2.1 3D DiT backbone with one structural change: the timestep embedder is - replaced by ``AnyFlowDualTimestepTextImageEmbedding`` so that every forward call conditions on both the source - timestep ``t`` and the target timestep ``r``. This is the embedding required to learn the flow map - :math:`\Phi_{r\leftarrow t}` introduced in [AnyFlow](https://huggingface.co/papers/2605.13724). - - For chunk-wise autoregressive (FAR causal) generation, use ``AnyFlowFARTransformer3DModel`` instead; that variant - adds the FAR causal block-mask and a compressed-frame patch embedding on top of the same backbone. - - Args: - patch_size (`Tuple[int]`, defaults to `(1, 2, 2)`): - 3D patch dimensions for video embedding (t_patch, h_patch, w_patch). - num_attention_heads (`int`, defaults to `40`): - Number of attention heads. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - in_channels (`int`, defaults to `16`): - The number of channels in the input latent. - out_channels (`int`, defaults to `16`): - The number of channels in the output latent. - text_dim (`int`, defaults to `4096`): - Input dimension for text embeddings (UMT5). - freq_dim (`int`, defaults to `256`): - Dimension for sinusoidal time embeddings. - ffn_dim (`int`, defaults to `13824`): - Intermediate dimension in feed-forward network. - num_layers (`int`, defaults to `40`): - Number of transformer blocks. - cross_attn_norm (`bool`, defaults to `True`): - Enable cross-attention normalization. - eps (`float`, defaults to `1e-6`): - Epsilon for normalization layers. - image_dim (`Optional[int]`, *optional*, defaults to `None`): - Image embedding dimension for I2V conditioning (`1280` for the original Wan2.1-I2V model). - rope_max_seq_len (`int`, defaults to `1024`): - Maximum sequence length used to precompute rotary position frequencies. - gate_value (`float`, defaults to `0.25`): - Mixing gate between source-timestep and delta-timestep embeddings (the AnyFlow paper's :math:`g` parameter, - fixed at 0.25 in stage-1 distillation). - deltatime_type (`str`, defaults to `'r'`): - Either ``"r"`` (delta is the target timestep) or ``"t-r"`` (delta is the absolute interval). - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"] - _no_split_modules = ["AnyFlowTransformerBlock"] - _keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"] - _repeated_blocks = ["AnyFlowTransformerBlock"] - - @register_to_config - def __init__( - self, - patch_size: Tuple[int] = (1, 2, 2), - num_attention_heads: int = 40, - attention_head_dim: int = 128, - in_channels: int = 16, - out_channels: int = 16, - text_dim: int = 4096, - freq_dim: int = 256, - ffn_dim: int = 13824, - num_layers: int = 40, - cross_attn_norm: bool = True, - eps: float = 1e-6, - image_dim: Optional[int] = None, - rope_max_seq_len: int = 1024, - gate_value: float = 0.25, - deltatime_type: str = "r", - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - # 1. Patch & position embedding (full-frame only). - self.rope = AnyFlowRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len) - self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - - # 2. Condition embedding (always dual-timestep for AnyFlow distilled checkpoints). - self.condition_embedder = AnyFlowDualTimestepTextImageEmbedding( - dim=inner_dim, - gate_value=gate_value, - deltatime_type=deltatime_type, - time_freq_dim=freq_dim, - time_proj_dim=inner_dim * 6, - text_embed_dim=text_dim, - image_embed_dim=image_dim, - ) - - # 3. Transformer blocks - self.blocks = nn.ModuleList( - [ - AnyFlowTransformerBlock(inner_dim, ffn_dim, num_attention_heads, cross_attn_norm, eps) - for _ in range(num_layers) - ] - ) - - # 4. Output norm & projection - self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size)) - self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5) - - self.gradient_checkpointing = False - - def _unpack_latent_sequence(self, latents, num_frames, height, width, patch_size): - batch_size, num_patches, channels = latents.shape - height, width = height // patch_size, width // patch_size - - latents = latents.view( - batch_size * num_frames, height, width, patch_size, patch_size, channels // (patch_size * patch_size) - ) - latents = latents.permute(0, 5, 1, 3, 2, 4) - latents = latents.reshape( - batch_size, num_frames, channels // (patch_size * patch_size), height * patch_size, width * patch_size - ) - return latents - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - r_timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: Optional[torch.Tensor] = None, - attention_kwargs: Optional[Dict[str, Any]] = None, - return_dict: bool = True, - ) -> Union[Transformer2DModelOutput, Tuple]: - """ - Bidirectional flow-map forward pass. ``hidden_states`` is laid out as ``(B, F, C, H, W)`` (per-frame latents). - The input is patchified with the standard ``patch_embedding`` (kernel = stride = ``patch_size``) and denoised - with global bidirectional self-attention over the resulting flat token sequence. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_frames, num_channels, height, width)`): - Input video latents. - timestep (`torch.Tensor`): - Source (noisier) flow-map timestep `t`. - r_timestep (`torch.Tensor`): - Target (cleaner) flow-map timestep `r`; defines the destination of the flow-map step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Text-conditioning embeddings. - encoder_hidden_states_image (`torch.Tensor`, *optional*): - Image-conditioning embeddings; concatenated before the text tokens when provided. - attention_kwargs (`dict`, *optional*): - Kwargs forwarded to the `AttentionProcessor` as defined under `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain tuple. - - Returns: - [`~models.transformer_2d.Transformer2DModelOutput`] if `return_dict` is True, otherwise a `tuple` whose - first element is the predicted velocity tensor. - """ - hidden_states = hidden_states.permute(0, 2, 1, 3, 4) - batch_size, num_channels, num_frames, height, width = hidden_states.shape - - full_token_per_frame = (height * width) // (self.config.patch_size[1] * self.config.patch_size[2]) - - layout_cfg = { - "total_frames": num_frames, - "full_frame_shape": (height // self.config.patch_size[1], width // self.config.patch_size[2]), - "full_token_per_frame": full_token_per_frame, - } - - rotary_emb = self.rope(layout_cfg=layout_cfg, device=hidden_states.device) - - hidden_states = self.patch_embedding(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) - - temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder( - timestep, - r_timestep, - encoder_hidden_states, - encoder_hidden_states_image, - layout_cfg=layout_cfg, - ) - timestep_proj = timestep_proj.unflatten(2, (6, -1)) - - attention_mask = None - - if encoder_hidden_states_image is not None: - encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.blocks: - hidden_states = self._gradient_checkpointing_func( - block, hidden_states, encoder_hidden_states, timestep_proj, rotary_emb, attention_mask - ) - else: - for block in self.blocks: - hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb, attention_mask) - - # Output norm, projection & unpatchify. - # `temb` is always 3D from `condition_embedder.forward()` (broadcast over total tokens). - shift, scale = (self.scale_shift_table.unsqueeze(0) + temb.unsqueeze(2)).chunk(2, dim=2) - shift = shift.squeeze(2) - scale = scale.squeeze(2) - - # Move shift/scale to hidden_states' device for multi-GPU accelerate inference. - shift = shift.to(hidden_states.device) - scale = scale.to(hidden_states.device) - - hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states) - hidden_states = self.proj_out(hidden_states) - - output = self._unpack_latent_sequence( - hidden_states, - num_frames=layout_cfg["total_frames"], - height=height, - width=width, - patch_size=self.config.patch_size[1], - ) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_anyflow_far.py b/diffusers/models/transformers/transformer_anyflow_far.py deleted file mode 100644 index 9ecc16bd04e08c3d0d43458a93f4babe3c984fd8..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_anyflow_far.py +++ /dev/null @@ -1,1622 +0,0 @@ -# Copyright 2026 The AnyFlow Team, NVIDIA Corp., and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -# -# This file is the FAR causal sibling of `transformer_anyflow.py`. Shared submodules are duplicated -# via `# Copied from` so `make fix-copies` keeps both files in sync; this keeps each transformer -# variant readable in isolation. The FAR architecture comes from FAR -# (arXiv:2503.19325); the dual-timestep flow-map embedding is AnyFlow's contribution -# (arXiv:2605.13724). - -import math -from dataclasses import dataclass -from typing import Any, Dict, List, Optional, Tuple, Union - -import torch -import torch.nn as nn -import torch.nn.functional as F -from torch.nn.attention.flex_attention import BlockMask, create_block_mask - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import BaseOutput, apply_lora_scale, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..attention import AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps, get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import FP32LayerNorm, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# Copied from diffusers.models.transformers.transformer_anyflow.apply_rotary_emb -def apply_rotary_emb(hidden_states: torch.Tensor, freqs: torch.Tensor): - # MPS / NPU backends do not support complex128 / float64; fall back to float32 on those devices. - rotary_dtype = maybe_adjust_dtype_for_device(torch.float64, hidden_states.device) - x_rotated = torch.view_as_complex(hidden_states.to(rotary_dtype).unflatten(3, (-1, 2))) - x_out = torch.view_as_real(x_rotated * freqs).flatten(3, 4) - return x_out.type_as(hidden_states) - - -@dataclass -class AnyFlowFARTransformerOutput(BaseOutput): - """ - Output dataclass for ``AnyFlowFARTransformer3DModel``'s causal forward paths. - - Args: - sample (`torch.Tensor` or `None`): - Predicted denoising target for the autoregressive chunk. ``None`` for the cache-prefill path, which only - writes the KV cache and produces no usable sample. - kv_cache (`list[dict[str, torch.Tensor]]`, *optional*): - Per-block KV cache state used by subsequent autoregressive steps. - """ - - sample: Optional[torch.Tensor] = None - kv_cache: Optional[List[Dict[str, torch.Tensor]]] = None - - -class AnyFlowCausalAttnProcessor: - """ - Causal self-attention processor for AnyFlow FAR. Routes through - :func:`~diffusers.models.attention_dispatch.dispatch_attention_fn` with the ``flex`` backend and a precomputed - :class:`~torch.nn.attention.flex_attention.BlockMask`. Supports KV-cache prefill (cache-write step) and - autoregressive read (cache-read step). - - Requires the ``flex`` attention backend — the ``BlockMask`` produced by - :meth:`AnyFlowFARTransformer3DModel.build_attention_mask` is consumed only by the flex backend. A clear - :class:`ValueError` is raised if a non-flex backend is configured via ``_attention_backend``. - """ - - _attention_backend = "flex" - _parallel_config = None - - _SUPPORTED_BACKENDS = ("flex", "_native_flex") - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "AnyFlowCausalAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0 or higher." - ) - - def __call__( - self, - attn, - hidden_states: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[Any] = None, - rotary_emb: Optional[Dict[str, torch.Tensor]] = None, - kv_cache: Optional[Dict[str, torch.Tensor]] = None, - kv_cache_flag: Optional[Dict[str, Any]] = None, - ) -> torch.Tensor: - if self._attention_backend not in self._SUPPORTED_BACKENDS: - raise ValueError( - f"AnyFlowCausalAttnProcessor requires the 'flex' attention backend " - f"(got {self._attention_backend!r}). FAR causal generation builds a " - f"flex_attention.BlockMask which is only consumed by the flex backend in " - f"`dispatch_attention_fn`." - ) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - target_dtype = hidden_states.dtype # Effective compute dtype - - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # norm_q and norm_k upcast query and key to FP32 due to the use of RMSNorm, so cast them back to the effective - # compute dtype. - query = query.to(target_dtype) - key = key.to(target_dtype) - - # Layout (B, H, L, D) is required by KV-cache slicing and rotary application. - query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2) - key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2) - value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2) - - if kv_cache is not None: - if kv_cache_flag["is_cache_step"]: - kv_cache["compressed_cache"][0, :, :, : kv_cache_flag["num_compressed_tokens"], :] = key[ - :, :, : kv_cache_flag["num_compressed_tokens"] - ] - kv_cache["compressed_cache"][1, :, :, : kv_cache_flag["num_compressed_tokens"], :] = value[ - :, :, : kv_cache_flag["num_compressed_tokens"] - ] - kv_cache["full_cache"][0, :, :, : kv_cache_flag["num_full_tokens"], :] = key[ - :, :, kv_cache_flag["num_compressed_tokens"] : - ] - kv_cache["full_cache"][1, :, :, : kv_cache_flag["num_full_tokens"], :] = value[ - :, :, kv_cache_flag["num_compressed_tokens"] : - ] - else: - key = torch.cat( - [ - kv_cache["compressed_cache"][0, :, :, : kv_cache_flag["num_cached_compressed_tokens"], :], - kv_cache["full_cache"][0, :, :, : kv_cache_flag["num_cached_full_tokens"], :], - key, - ], - dim=2, - ) - value = torch.cat( - [ - kv_cache["compressed_cache"][1, :, :, : kv_cache_flag["num_cached_compressed_tokens"], :], - kv_cache["full_cache"][1, :, :, : kv_cache_flag["num_cached_full_tokens"], :], - value, - ], - dim=2, - ) - - if rotary_emb is not None: - query = apply_rotary_emb(query, rotary_emb["query"]) - key = apply_rotary_emb(key, rotary_emb["key"]) - - # BlockMask block-size is 128 — pad seq_len to a multiple of 128. Tiny dummy components may - # have head_dim < 16; flex_attention requires head_dim >= 16, so right-pad q/k/v on the head - # dim with zeros and override `scale` so the result matches the original head_dim. - seq_len = query.shape[2] - head_dim = query.shape[3] - padded_length = int(math.ceil(seq_len / 128.0) * 128.0 - seq_len) - if padded_length > 0: - pad_shape = [query.shape[0], query.shape[1], padded_length, head_dim] - query = torch.cat([query, torch.zeros(pad_shape, device=query.device, dtype=query.dtype)], dim=2) - key = torch.cat([key, torch.zeros(pad_shape, device=key.device, dtype=key.dtype)], dim=2) - value = torch.cat([value, torch.zeros(pad_shape, device=value.device, dtype=value.dtype)], dim=2) - - head_pad = max(0, 16 - head_dim) - scale = 1.0 / (head_dim**0.5) if head_pad > 0 else None - if head_pad > 0: - query = F.pad(query, (0, head_pad)) - key = F.pad(key, (0, head_pad)) - value = F.pad(value, (0, head_pad)) - - # `dispatch_attention_fn` expects (B, L, H, D); the flex backend permutes back to - # (B, H, L, D) internally before calling flex_attention — same kernel call as the bare - # flex_attention path, same numerics. Verified against - # `attention_dispatch._native_flex_attention`. - hidden_states = dispatch_attention_fn( - query.transpose(1, 2), - key.transpose(1, 2), - value.transpose(1, 2), - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - scale=scale, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - # `dispatch_attention_fn` returns (B, L, H, D). Trim head pad on the last axis, then trim - # seq pad on dim=1, then fold heads back into the channel dim. - if head_pad > 0: - hidden_states = hidden_states[..., :head_dim] - if padded_length > 0: - hidden_states = hidden_states[:, :seq_len, :, :] - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -# Copied from diffusers.models.transformers.transformer_anyflow.AnyFlowAttnProcessor -class AnyFlowAttnProcessor: - """ - Bidirectional self-attention processor for AnyFlow. Routes through - :func:`~diffusers.models.attention_dispatch.dispatch_attention_fn` so any SDPA-compatible backend is supported - (SDPA, flash-attn, xformers, flex, …). FAR causal generation lives in - :class:`~diffusers.models.transformers.transformer_anyflow_far.AnyFlowCausalAttnProcessor`. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "AnyFlowAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0 or higher." - ) - - def __call__( - self, - attn: "AnyFlowAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[Any] = None, - rotary_emb: Optional[Dict[str, torch.Tensor]] = None, - ) -> torch.Tensor: - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Layout (B, H, L, D) for rotary application; transposed to (B, L, H, D) before dispatch. - query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2) - key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2) - value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2) - - if rotary_emb is not None: - query = apply_rotary_emb(query, rotary_emb["query"]) - key = apply_rotary_emb(key, rotary_emb["key"]) - - hidden_states = dispatch_attention_fn( - query.transpose(1, 2), - key.transpose(1, 2), - value.transpose(1, 2), - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -# Copied from diffusers.models.transformers.transformer_anyflow.AnyFlowCrossAttnProcessor -class AnyFlowCrossAttnProcessor: - """ - Cross-attention processor for AnyFlow. Always uses the dispatched SDPA-compatible backend; no rotary embedding or - KV cache is applied to the text→video cross-attention path. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "AnyFlowCrossAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0 or higher." - ) - - def __call__( - self, - attn: "AnyFlowAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - ) -> torch.Tensor: - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # (B, L, H, D) layout for dispatch_attention_fn. - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -# Copied from diffusers.models.transformers.transformer_anyflow.AnyFlowAttention with AnyFlowAttnProcessor->AnyFlowCausalAttnProcessor -class AnyFlowAttention(torch.nn.Module, AttentionModuleMixin): - """ - Attention module used by :class:`AnyFlowTransformerBlock`. Layout matches the legacy - :class:`~diffusers.models.attention_processor.Attention` so existing AnyFlow checkpoints load bit-exactly into this - class. - """ - - _default_processor_cls = AnyFlowCausalAttnProcessor - _available_processors = [AnyFlowCausalAttnProcessor, AnyFlowCrossAttnProcessor] - - def __init__( - self, - dim: int, - heads: int, - dim_head: int, - eps: float = 1e-6, - processor: Optional[Any] = None, - ): - super().__init__() - self.heads = heads - self.inner_dim = heads * dim_head - - self.to_q = torch.nn.Linear(dim, self.inner_dim, bias=True) - self.to_k = torch.nn.Linear(dim, self.inner_dim, bias=True) - self.to_v = torch.nn.Linear(dim, self.inner_dim, bias=True) - self.to_out = torch.nn.ModuleList( - [ - torch.nn.Linear(self.inner_dim, dim, bias=True), - torch.nn.Dropout(0.0), - ] - ) - # ``rms_norm_across_heads`` per-axis: normalize Q and K across the entire ``heads * dim_head`` - # channel axis. We use diffusers' RMSNorm (rather than ``torch.nn.RMSNorm``) so the numerics - # match the legacy Attention class that produced the released checkpoints. - self.norm_q = RMSNorm(self.inner_dim, eps=eps) - self.norm_k = RMSNorm(self.inner_dim, eps=eps) - - self.set_processor(processor if processor is not None else self._default_processor_cls()) - - def forward(self, hidden_states: torch.Tensor, **kwargs) -> torch.Tensor: - return self.processor(self, hidden_states, **kwargs) - - -# Copied from diffusers.models.transformers.transformer_anyflow.AnyFlowImageEmbedding -class AnyFlowImageEmbedding(torch.nn.Module): - def __init__(self, in_features: int, out_features: int): - super().__init__() - - self.norm1 = FP32LayerNorm(in_features) - self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu") - self.norm2 = FP32LayerNorm(out_features) - - def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor: - hidden_states = self.norm1(encoder_hidden_states_image) - hidden_states = self.ff(hidden_states) - hidden_states = self.norm2(hidden_states) - return hidden_states - - -class AnyFlowDualTimestepTextImageEmbeddingCausal(nn.Module): - """Causal variant of :class:`AnyFlowDualTimestepTextImageEmbedding`. - - Splits the per-frame timestep stream into a full-resolution suffix (length ``far_cfg["num_full_frames"]``) and a - FAR-compressed prefix, expanding each segment by its own ``token_per_frame`` factor so the assembled time embedding - aligns with the chunk-mixed token sequence. Optionally concatenates a ``clean_timestep`` embedding for the training - rollout. - """ - - def __init__( - self, - dim: int, - gate_value: float, - deltatime_type: str, - time_freq_dim: int, - time_proj_dim: int, - text_embed_dim: int, - image_embed_dim: Optional[int] = None, - ): - super().__init__() - - self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim) - self.delta_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim) - self.act_fn = nn.SiLU() - self.time_proj = nn.Linear(dim, time_proj_dim) - self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh") - - self.image_embedder = None - if image_embed_dim is not None: - self.image_embedder = AnyFlowImageEmbedding(image_embed_dim, dim) - - self.register_buffer("delta_emb_gate", torch.tensor([gate_value], dtype=torch.float32), persistent=False) - self.deltatime_type = deltatime_type - - # Copied from diffusers.models.transformers.transformer_anyflow.AnyFlowDualTimestepTextImageEmbedding.forward_timestep - def forward_timestep( - self, timestep: torch.Tensor, delta_timestep: torch.Tensor, encoder_hidden_states, token_per_frame - ): - batch_size, num_frames = timestep.shape - timestep = timestep.reshape(-1) - delta_timestep = delta_timestep.reshape(-1) - - timestep = self.timesteps_proj(timestep) - - time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype - if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8: - timestep = timestep.to(time_embedder_dtype) - temb = self.time_embedder(timestep).type_as(encoder_hidden_states) - - delta_timestep = self.timesteps_proj(delta_timestep) - - delta_embedder_dtype = next(iter(self.delta_embedder.parameters())).dtype - if delta_timestep.dtype != delta_embedder_dtype and delta_embedder_dtype != torch.int8: - delta_timestep = delta_timestep.to(delta_embedder_dtype) - delta_emb = self.delta_embedder(delta_timestep).type_as(encoder_hidden_states) - - gate = self.delta_emb_gate.to(delta_embedder_dtype) - - rt_emb = (1 - gate) * temb + gate * delta_emb - timestep_proj = self.time_proj(self.act_fn(rt_emb)) - - rt_emb = rt_emb.unflatten(0, (batch_size, num_frames)).repeat_interleave(token_per_frame, dim=1) - timestep_proj = timestep_proj.unflatten(0, (batch_size, num_frames)).repeat_interleave(token_per_frame, dim=1) - - return rt_emb, timestep_proj - - def forward( - self, - timestep: torch.Tensor, - r_timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: Optional[torch.Tensor] = None, - far_cfg=None, - clean_timestep=None, - ): - if self.deltatime_type == "r": - delta_timestep = r_timestep - elif self.deltatime_type == "t-r": - delta_timestep = timestep - r_timestep - else: - raise NotImplementedError - - full_frame_timestep, full_frame_timestep_proj = self.forward_timestep( - timestep[:, -far_cfg["num_full_frames"] :], - delta_timestep[:, -far_cfg["num_full_frames"] :], - encoder_hidden_states, - far_cfg["full_token_per_frame"], - ) - compressed_frame_timestep, compressed_frame_timestep_proj = self.forward_timestep( - timestep[:, : -far_cfg["num_full_frames"]], - delta_timestep[:, : -far_cfg["num_full_frames"]], - encoder_hidden_states, - far_cfg["compressed_token_per_frame"], - ) - - if clean_timestep is not None: - clean_timestep, clean_timestep_proj = self.forward_timestep( - clean_timestep, clean_timestep, encoder_hidden_states, far_cfg["full_token_per_frame"] - ) - timestep = torch.cat([compressed_frame_timestep, full_frame_timestep, clean_timestep], dim=1) - timestep_proj = torch.cat( - [compressed_frame_timestep_proj, full_frame_timestep_proj, clean_timestep_proj], dim=1 - ) - else: - timestep = torch.cat([compressed_frame_timestep, full_frame_timestep], dim=1) - timestep_proj = torch.cat([compressed_frame_timestep_proj, full_frame_timestep_proj], dim=1) - - encoder_hidden_states = self.text_embedder(encoder_hidden_states) - if encoder_hidden_states_image is not None: - encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image) - - return timestep, timestep_proj, encoder_hidden_states, encoder_hidden_states_image - - -# Copied from diffusers.models.transformers.transformer_anyflow.AnyFlowTransformerBlock -class AnyFlowTransformerBlock(nn.Module): - """AnyFlow transformer block. - - The self-attention processor is chosen at construction by ``is_causal``: the bidirectional transformer passes - ``is_causal=False`` (the default), the FAR causal transformer passes ``is_causal=True``. The forward pass is - identical in both modes — only the processor differs, so all causal-specific machinery (BlockMask, KV cache) lives - inside the processor. - """ - - def __init__( - self, - dim: int, - ffn_dim: int, - num_heads: int, - cross_attn_norm: bool = False, - eps: float = 1e-6, - is_causal: bool = False, - ): - super().__init__() - - self.is_causal = is_causal - - # 1. Self-attention. The causal processor lives in the FAR sibling module; lazy-import to - # avoid a circular import at module load time. - if is_causal: - from .transformer_anyflow_far import AnyFlowCausalAttnProcessor - - self_attn_processor = AnyFlowCausalAttnProcessor() - else: - self_attn_processor = AnyFlowAttnProcessor() - - self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False) - self.attn1 = AnyFlowAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - processor=self_attn_processor, - ) - - # 2. Cross-attention - self.attn2 = AnyFlowAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - processor=AnyFlowCrossAttnProcessor(), - ) - self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity() - - # 3. Feed-forward - self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate") - self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False) - - self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - rotary_emb: torch.Tensor, - attention_mask: torch.Tensor, - kv_cache=None, - kv_cache_flag=None, - ) -> torch.Tensor: - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table + temb.float() - ).chunk(6, dim=2) - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - shift_msa.squeeze(2), - scale_msa.squeeze(2), - gate_msa.squeeze(2), - c_shift_msa.squeeze(2), - c_scale_msa.squeeze(2), - c_gate_msa.squeeze(2), - ) # noqa: E501 - - # 1. Self-attention - norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states) - attn1_kwargs = { - "hidden_states": norm_hidden_states, - "rotary_emb": rotary_emb, - "attention_mask": attention_mask, - } - # KV cache kwargs are only consumed by the FAR causal processor; the bidi processor - # doesn't accept them, so we forward them only when they're actually populated. - if kv_cache is not None: - attn1_kwargs["kv_cache"] = kv_cache - attn1_kwargs["kv_cache_flag"] = kv_cache_flag - attn_output = self.attn1(**attn1_kwargs) - hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states) - - # 2. Cross-attention - norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states) - attn_output = self.attn2(hidden_states=norm_hidden_states, encoder_hidden_states=encoder_hidden_states) - hidden_states = hidden_states + attn_output - - # 3. Feed-forward - norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as( - hidden_states - ) - ff_output = self.ffn(norm_hidden_states) - hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states) - - return hidden_states - - -class AnyFlowCausalRotaryPosEmbed(nn.Module): - """ - Rotary positional embedding for the FAR causal transformer. - - Produces position frequencies for both the full-resolution noisy chunk(s) and the FAR-compressed context chunk(s); - the compressed branch downscales the per-axis frequency table via complex average pooling so the compressed grid - stays aligned with the full grid. - """ - - def __init__( - self, - attention_head_dim: int, - patch_size: Tuple[int, int, int], - compressed_patch_size: Tuple[int, int, int], - max_seq_len: int, - theta: float = 10000.0, - ): - super().__init__() - - self.attention_head_dim = attention_head_dim - self.patch_size = patch_size - self.compressed_patch_size = compressed_patch_size - self.max_seq_len = max_seq_len - self.theta = theta - - # Frequency table is lazily built per-device in ``_build_freqs``: MPS / NPU don't support - # complex128, so we downcast to complex64 there. - self._freqs_cache: Optional[Tuple[Any, torch.Tensor]] = None - - # Copied from diffusers.models.transformers.transformer_anyflow.AnyFlowRotaryPosEmbed._build_freqs - def _build_freqs(self, device: torch.device) -> torch.Tensor: - # Skip the cache read/write inside torch.compile: mutating ``self._freqs_cache`` between calls - # becomes a Dynamo guard and forces recompilation on the second invocation. - is_compiling = torch.compiler.is_compiling() - cache_key = (device.type, str(device)) - if not is_compiling and self._freqs_cache is not None and self._freqs_cache[0] == cache_key: - return self._freqs_cache[1] - - freqs_dtype = maybe_adjust_dtype_for_device(torch.float64, device) - - h_dim = w_dim = 2 * (self.attention_head_dim // 6) - t_dim = self.attention_head_dim - h_dim - w_dim - - freqs_list = [] - for dim in (t_dim, h_dim, w_dim): - f = get_1d_rotary_pos_embed( - dim, - self.max_seq_len, - self.theta, - use_real=False, - repeat_interleave_real=False, - freqs_dtype=freqs_dtype, - ) - freqs_list.append(f.to(device)) - freqs = torch.cat(freqs_list, dim=1) - if not is_compiling: - self._freqs_cache = (cache_key, freqs) - return freqs - - def avg_pool_complex(self, freq: torch.Tensor, kernel_size: int, stride: int): - real = freq.real # [B, C, L], float - real = real.transpose(0, 1).unsqueeze(0) - imag = freq.imag # [B, C, L], float - imag = imag.transpose(0, 1).unsqueeze(0) - - pr = F.avg_pool1d(real, kernel_size, stride) - pi = F.avg_pool1d(imag, kernel_size, stride) - - pr = pr.squeeze(0).transpose(0, 1) - pi = pi.squeeze(0).transpose(0, 1) - - norm = torch.sqrt(pr**2 + pi**2) - pr_unit = pr / norm - pi_unit = pi / norm - - return torch.complex(pr_unit, pi_unit) - - def _forward_compressed_frame(self, num_frames, height, width, device): - ppf, pph, ppw = num_frames, height, width - # Tiny dummy components (e.g. height=16/width=16 with compressed_patch_size=(1,4,4) and - # an upstream VAE stride of 8) can produce 0-element grids; the .view(0, k, 1, -1) reshape - # below would be ambiguous. Real ckpts use 60x104 latents and never hit this path. - freqs_full = self._build_freqs(device) - if min(ppf, pph, ppw) <= 0: - freq_channels = self.attention_head_dim // 2 - return torch.empty((ppf, pph, ppw, freq_channels), dtype=freqs_full.dtype, device=device) - downscale = [self.compressed_patch_size[i] // self.patch_size[i] for i in range(len(self.patch_size))] - - freqs = freqs_full.split_with_sizes( - [ - self.attention_head_dim // 2 - 2 * (self.attention_head_dim // 6), - self.attention_head_dim // 6, - self.attention_head_dim // 6, - ], - dim=1, - ) - - freqs_f = self.avg_pool_complex(freqs[0], kernel_size=downscale[0], stride=downscale[0]) - freqs_h = self.avg_pool_complex(freqs[1], kernel_size=downscale[1], stride=downscale[1]) - freqs_w = self.avg_pool_complex(freqs[2], kernel_size=downscale[2], stride=downscale[2]) - - freqs_f = freqs_f[:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_h = freqs_h[:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_w = freqs_w[:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - - freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1) - return freqs - - # Copied from diffusers.models.transformers.transformer_anyflow.AnyFlowRotaryPosEmbed._forward_full_frame - def _forward_full_frame(self, num_frames, height, width, device) -> torch.Tensor: - ppf, pph, ppw = num_frames, height, width - - freqs_full = self._build_freqs(device) - if min(ppf, pph, ppw) <= 0: - freq_channels = self.attention_head_dim // 2 - return torch.empty((ppf, pph, ppw, freq_channels), dtype=freqs_full.dtype, device=device) - - freqs = freqs_full.split_with_sizes( - [ - self.attention_head_dim // 2 - 2 * (self.attention_head_dim // 6), - self.attention_head_dim // 6, - self.attention_head_dim // 6, - ], - dim=1, - ) - - freqs_f = freqs[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_h = freqs[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1) - return freqs - - def forward(self, far_cfg, device, clean_hidden_states=None): - full_frame_freqs = self._forward_full_frame( - num_frames=far_cfg["total_frames"], - height=far_cfg["full_frame_shape"][0], - width=far_cfg["full_frame_shape"][1], - device=device, - ) - compressed_frame_freqs = self._forward_compressed_frame( - num_frames=far_cfg["total_frames"], - height=far_cfg["compressed_frame_shape"][0], - width=far_cfg["compressed_frame_shape"][1], - device=device, - ) - - compressed_frame_freqs, full_frame_freqs = ( - compressed_frame_freqs[: far_cfg["num_compressed_frames"]], - full_frame_freqs[far_cfg["num_compressed_frames"] :], - ) - - compressed_frame_freqs = compressed_frame_freqs.flatten(start_dim=0, end_dim=2) - full_frame_freqs = full_frame_freqs.flatten(start_dim=0, end_dim=2) - - if clean_hidden_states is not None: - freqs = torch.cat([compressed_frame_freqs, full_frame_freqs, full_frame_freqs], dim=0) - else: - freqs = torch.cat([compressed_frame_freqs, full_frame_freqs], dim=0) - - freqs = freqs[None, None, ...] - - return {"query": freqs, "key": freqs} - - -def _build_anyflow_far_causal_block_mask( - chunk_partition: List[int], - height: int, - width: int, - patch_size: Tuple[int, int, int], - compressed_patch_size: Tuple[int, int, int], - full_chunk_limit: int, - *, - mode: str = "train", - has_clean_context: bool = False, - device: Optional[torch.device] = None, -) -> BlockMask: - r"""Build the causal :class:`~torch.nn.attention.flex_attention.BlockMask` for the FAR transformer. - - Provided as a standalone function so callers can construct the mask *outside* the transformer's compiled region, - which is required to wrap the forward in ``torch.compile(fullgraph=True)`` (``flex_attention.create_block_mask`` - itself uses ``_compile=False`` internally and breaks the graph when invoked inside the compiled scope). - - Two modes are exposed, mirroring the FAR forward paths that actually consume a mask. The autoregressive - ``_forward_inference`` path attends through the KV cache and does not use a full BlockMask, so it has no - corresponding mode here. - - Args: - chunk_partition: per-chunk frame counts; must sum to the number of latent frames. - height, width: latent spatial dimensions. - patch_size, compressed_patch_size, full_chunk_limit: must match the transformer config. - mode: ``"train"`` (strict ``>`` comparison against ``full_chunk_limit``, matches - :meth:`AnyFlowFARTransformer3DModel._forward_train`) or ``"cache"`` (``>=`` comparison via the - ``full_chunk_limit - 1`` offset used by :meth:`AnyFlowFARTransformer3DModel._forward_cache`). - has_clean_context: ``True`` when ``clean_hidden_states`` is being threaded through the - transformer (training V2V/I2V). - device: device for the resulting BlockMask. Defaults to CPU. - """ - if mode not in {"train", "cache"}: - raise ValueError(f"Unknown mode {mode!r}; expected 'train' or 'cache'.") - full_token_per_frame = (height // patch_size[1]) * (width // patch_size[2]) - compressed_token_per_frame = (height // compressed_patch_size[1]) * (width // compressed_patch_size[2]) - - # `cache` uses `full_chunk_limit - 1` (an effective `>= full_chunk_limit` comparison); `train` uses a strict `>`. - total_chunks = len(chunk_partition) - threshold = full_chunk_limit - 1 if mode == "cache" else full_chunk_limit - if total_chunks > threshold: - num_full_chunk = threshold - num_compressed_chunk = total_chunks - threshold - else: - num_full_chunk, num_compressed_chunk = total_chunks, 0 - - far_cfg = { - "num_full_chunk": num_full_chunk, - "num_compressed_chunk": num_compressed_chunk, - "num_full_frames": sum(chunk_partition[num_compressed_chunk:]), - "num_compressed_frames": sum(chunk_partition[:num_compressed_chunk]), - "full_token_per_frame": full_token_per_frame, - "compressed_token_per_frame": compressed_token_per_frame, - "chunk_partition": chunk_partition, - } - return _build_far_block_mask_from_far_cfg(far_cfg, has_clean=has_clean_context, device=device) - - -def _build_far_block_mask_from_far_cfg(far_cfg, has_clean, device): - """Internal: build a BlockMask given an already-computed ``far_cfg`` dict. - - Factored out of :class:`AnyFlowFARTransformer3DModel` so it can be shared between - :func:`_build_anyflow_far_causal_block_mask` (the user-facing entry point) and the in-forward fallback path used - when no pre-built ``attention_mask`` is passed. - """ - chunk_partition = far_cfg["chunk_partition"] - - noise_seq_len = clean_seq_len = far_cfg["num_full_frames"] * far_cfg["full_token_per_frame"] - context_seq_len = far_cfg["num_compressed_frames"] * far_cfg["compressed_token_per_frame"] - - noise_start = context_seq_len - noise_end = noise_start + noise_seq_len - - clean_start = context_seq_len + noise_seq_len - clean_end = clean_start + clean_seq_len - - if has_clean: - real_seq_len = context_seq_len + noise_seq_len + clean_seq_len - else: - real_seq_len = context_seq_len + noise_seq_len - - padded_seq_len = int(math.ceil(real_seq_len / 128.0) * 128.0) - - context_chunk_partition, noise_chunk_partition = ( - chunk_partition[: far_cfg["num_compressed_chunk"]], - chunk_partition[far_cfg["num_compressed_chunk"] :], - ) - - if len(context_chunk_partition) != 0: - context_frame_idx = torch.cat( - [ - torch.ones(chunk_len * far_cfg["compressed_token_per_frame"], device=device) * chunk_idx - for chunk_idx, chunk_len in enumerate(context_chunk_partition) - ] - ) - else: - context_frame_idx = None - - if has_clean: - noise_frame_idx = clean_frame_idx = torch.cat( - [ - torch.ones(chunk_len * far_cfg["full_token_per_frame"], device=device) - * (chunk_idx + len(context_chunk_partition)) - for chunk_idx, chunk_len in enumerate(noise_chunk_partition) - ] - ) - pad_frame_idx = torch.zeros(padded_seq_len - real_seq_len, device=device) - - if len(context_chunk_partition) != 0: - frame_idx = torch.cat([context_frame_idx, noise_frame_idx, clean_frame_idx, pad_frame_idx], dim=0) - else: - frame_idx = torch.cat([noise_frame_idx, clean_frame_idx, pad_frame_idx], dim=0) - - def mask_mod(b, h, q_idx, kv_idx): - # 1) is padding - is_padding = (q_idx >= real_seq_len) | (kv_idx >= real_seq_len) - - # 2) chunk causal - base = frame_idx[q_idx] >= frame_idx[kv_idx] - - # 3) interval mask - q_is_noise = (q_idx >= noise_start) & (q_idx < noise_end) - q_is_clean = (q_idx >= clean_start) & (q_idx < clean_end) - - k_is_noise = (kv_idx >= noise_start) & (kv_idx < noise_end) - k_is_clean = (kv_idx >= clean_start) & (kv_idx < clean_end) - - # 4) clean -> noise: disallowed - is_clean_to_noise = q_is_clean & k_is_noise - - # 5) noise -> noise: only same frame - same_frame_idx = frame_idx[q_idx] == frame_idx[kv_idx] - - noise_to_noise = q_is_noise & k_is_noise - noise_to_clean = q_is_noise & k_is_clean - - noise_to_noise_allow = noise_to_noise & same_frame_idx - noise_to_noise_mask = (~noise_to_noise) | noise_to_noise_allow - - noise_to_clean_same = noise_to_clean & same_frame_idx - noise_to_clean_disallow = noise_to_clean_same - - allowed = base & ~is_padding & ~is_clean_to_noise & noise_to_noise_mask & ~noise_to_clean_disallow - return allowed - - else: - noise_frame_idx = torch.cat( - [ - torch.ones(chunk_len * far_cfg["full_token_per_frame"], device=device) - * (chunk_idx + len(context_chunk_partition)) - for chunk_idx, chunk_len in enumerate(noise_chunk_partition) - ] - ) - pad_frame_idx = torch.zeros(padded_seq_len - real_seq_len, device=device) - - if len(context_chunk_partition) != 0: - frame_idx = torch.cat([context_frame_idx, noise_frame_idx, pad_frame_idx], dim=0) - else: - frame_idx = torch.cat([noise_frame_idx, pad_frame_idx], dim=0) - - def mask_mod(b, h, q_idx, kv_idx): - is_padding = (q_idx >= real_seq_len) | (kv_idx >= real_seq_len) - base = frame_idx[q_idx] >= frame_idx[kv_idx] - return base & ~is_padding - - return create_block_mask( - mask_mod, - B=None, - H=None, - Q_LEN=padded_seq_len, - KV_LEN=padded_seq_len, - device=device, - _compile=False, - ) - - -class AnyFlowFARTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): - r""" - Causal (FAR) 3D Transformer for AnyFlow flow-map sampling with chunk-wise autoregressive generation. - - Extends the v0.35.1 Wan2.1 backbone with: - - * **FAR causal block-mask** via :func:`torch.nn.attention.flex_attention`, supporting chunk-wise autoregressive - generation ([FAR](https://huggingface.co/papers/2503.19325)). - * **Compressed-frame patch embedding** ``far_patch_embedding`` for context (already-generated) frames, initialized - from ``patch_embedding`` via trilinear interpolation so a freshly constructed model is already at a reasonable - starting point even before LoRA fine-tuning. - * **Dual-timestep flow-map embedding** for any-step sampling (same as ``AnyFlowTransformer3DModel``). - - Use ``AnyFlowTransformer3DModel`` instead for plain bidirectional T2V — that variant skips the FAR causal masking - and ``far_patch_embedding`` and is ~5–10% smaller. - - Args: - patch_size (`Tuple[int]`, defaults to `(1, 2, 2)`): - 3D patch dimensions for full-resolution chunks. - compressed_patch_size (`Tuple[int]`, defaults to `(1, 4, 4)`): - Larger patch dimensions for the FAR-compressed (context) chunks. - full_chunk_limit (`int`, defaults to `3`): - Maximum number of full-resolution chunks before earlier chunks are demoted to compressed FAR context. The - released checkpoints use ``3``. - num_attention_heads (`int`, defaults to `40`): - Number of attention heads. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - in_channels (`int`, defaults to `16`): - The number of channels in the input latent. - out_channels (`int`, defaults to `16`): - The number of channels in the output latent. - text_dim (`int`, defaults to `4096`): - Input dimension for text embeddings (UMT5). - freq_dim (`int`, defaults to `256`): - Dimension for sinusoidal time embeddings. - ffn_dim (`int`, defaults to `13824`): - Intermediate dimension in feed-forward network. - num_layers (`int`, defaults to `40`): - Number of transformer blocks. - cross_attn_norm (`bool`, defaults to `True`): - Enable cross-attention normalization. - eps (`float`, defaults to `1e-6`): - Epsilon for normalization layers. - image_dim (`Optional[int]`, *optional*, defaults to `None`): - Image embedding dimension for I2V conditioning. - rope_max_seq_len (`int`, defaults to `1024`): - Maximum sequence length used to precompute rotary position frequencies. - gate_value (`float`, defaults to `0.25`): - Mixing gate between source-timestep and delta-timestep embeddings. - deltatime_type (`str`, defaults to `'r'`): - Either ``"r"`` (delta is the target timestep) or ``"t-r"`` (delta is the absolute interval). - chunk_partition (`Tuple[int, ...]`, defaults to `(1, 3, 3, 3, 3, 3, 3, 2)`): - Default per-chunk frame counts used by the pipeline. The released NVIDIA AnyFlow-FAR checkpoints target - ``num_frames=81`` (21 latent frames at VAE temporal stride 4) split as ``1 + 3*6 + 2``. A different - ``num_frames`` requires a matching ``chunk_partition`` override passed to - :meth:`AnyFlowFARPipeline.__call__` (and likewise to :meth:`forward`). - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["patch_embedding", "far_patch_embedding", "condition_embedder", "norm"] - _no_split_modules = ["AnyFlowTransformerBlock"] - _keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"] - _repeated_blocks = ["AnyFlowTransformerBlock"] - - @register_to_config - def __init__( - self, - patch_size: Tuple[int] = (1, 2, 2), - compressed_patch_size: Tuple[int] = (1, 4, 4), - full_chunk_limit: int = 3, - num_attention_heads: int = 40, - attention_head_dim: int = 128, - in_channels: int = 16, - out_channels: int = 16, - text_dim: int = 4096, - freq_dim: int = 256, - ffn_dim: int = 13824, - num_layers: int = 40, - cross_attn_norm: bool = True, - eps: float = 1e-6, - image_dim: Optional[int] = None, - rope_max_seq_len: int = 1024, - gate_value: float = 0.25, - deltatime_type: str = "r", - chunk_partition: Tuple[int, ...] = (1, 3, 3, 3, 3, 3, 3, 2), - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - # 1. Patch & position embedding (full + FAR-compressed branches). - self.rope = AnyFlowCausalRotaryPosEmbed( - attention_head_dim, patch_size, compressed_patch_size, rope_max_seq_len - ) - self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - - self.far_patch_embedding = nn.Conv3d( - in_channels, inner_dim, kernel_size=compressed_patch_size, stride=compressed_patch_size - ) - # Warm-start the compressed branch from the full-resolution branch by trilinear interpolation. This - # matches FAR-Dev's `setup_far_model()` initialization. State-dict loading will overwrite these - # weights for trained checkpoints; the warm-start only matters when constructing a fresh model. - original_weight = self.patch_embedding.weight.data.view(-1, 1, *patch_size) - new_weight = F.interpolate(original_weight, size=compressed_patch_size, mode="trilinear", align_corners=False) - new_weight = new_weight.view(inner_dim, in_channels, *compressed_patch_size) - with torch.no_grad(): - self.far_patch_embedding.weight.copy_(new_weight) - self.far_patch_embedding.bias.copy_(self.patch_embedding.bias) - - # 2. Condition embedding (always dual-timestep for AnyFlow distilled checkpoints). - self.condition_embedder = AnyFlowDualTimestepTextImageEmbeddingCausal( - dim=inner_dim, - gate_value=gate_value, - deltatime_type=deltatime_type, - time_freq_dim=freq_dim, - time_proj_dim=inner_dim * 6, - text_embed_dim=text_dim, - image_embed_dim=image_dim, - ) - - # 3. Transformer blocks (causal self-attn processor) - self.blocks = nn.ModuleList( - [ - AnyFlowTransformerBlock(inner_dim, ffn_dim, num_attention_heads, cross_attn_norm, eps, is_causal=True) - for _ in range(num_layers) - ] - ) - - # 4. Output norm & projection - self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size)) - self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - r_timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - chunk_partition: List[int], - encoder_hidden_states_image: Optional[torch.Tensor] = None, - clean_hidden_states: Optional[torch.Tensor] = None, - clean_timestep: Optional[torch.Tensor] = None, - kv_cache: Optional[List[Dict[str, torch.Tensor]]] = None, - kv_cache_flag: Optional[Dict[str, Any]] = None, - attention_mask: Optional[BlockMask] = None, - attention_kwargs: Optional[Dict[str, Any]] = None, - return_dict: bool = True, - ) -> Union[Transformer2DModelOutput, AnyFlowFARTransformerOutput, Tuple]: - """ - FAR causal forward pass. Dispatches to one of three internal paths: - - * ``kv_cache is None`` → causal training rollout (returns :class:`Transformer2DModelOutput`). - * ``kv_cache is not None`` and ``kv_cache_flag["is_cache_step"]`` → cache-prefill (returns - :class:`AnyFlowFARTransformerOutput` with ``sample=None``). - * Otherwise → autoregressive inference step (returns :class:`AnyFlowFARTransformerOutput`). - - Args: - hidden_states (`torch.Tensor`): - Latent input of shape ``(B, F, C, H, W)``. - timestep (`torch.Tensor`): - Source (noisier) flow-map timestep `t`. - r_timestep (`torch.Tensor`): - Target (cleaner) flow-map timestep `r`. - encoder_hidden_states (`torch.Tensor`): - UMT5 text embeddings. - chunk_partition (`List[int]`): - Per-chunk frame counts; total must match the number of latent frames in ``hidden_states``. - encoder_hidden_states_image (`torch.Tensor`, *optional*): - I2V image embedding; concatenated before text tokens when provided. - clean_hidden_states (`torch.Tensor`, *optional*): - Clean (noise-free) conditioning frames used by the training rollout. - clean_timestep (`torch.Tensor`, *optional*): - Timesteps for the clean conditioning frames in the training rollout. - kv_cache (`List[Dict[str, torch.Tensor]]`, *optional*): - Per-block KV cache for autoregressive inference. `None` selects the training path. - kv_cache_flag (`Dict[str, Any]`, *optional*): - KV-cache metadata (e.g. ``is_cache_step`` flag and token counts). - attention_mask (`BlockMask`, *optional*): - Pre-built causal mask, typically constructed via :meth:`build_attention_mask`. Consumed by the train - and KV-cache prefill paths; the autoregressive inference path attends through the KV cache and does not - use a full mask. When ``None``, the train / cache paths build the mask internally; that fallback is not - compile-safe (the underlying ``flex_attention.create_block_mask`` breaks the graph under - ``fullgraph=True``), so pass a pre-built mask whenever wrapping ``forward`` in ``torch.compile``. - attention_kwargs (`dict`, *optional*): - Forwarded to the attention processors. - return_dict (`bool`, *optional*, defaults to `True`): - If `False`, returns positional tuples instead of an output dataclass. - - Returns: - [`~models.transformer_2d.Transformer2DModelOutput`], [`AnyFlowFARTransformerOutput`] or `tuple`: - When `return_dict` is `False`, a plain `tuple` is returned. Otherwise, the causal training rollout - (`kv_cache is None`) returns a [`~models.transformer_2d.Transformer2DModelOutput`], while the - cache-prefill and autoregressive inference paths return an [`AnyFlowFARTransformerOutput`]. - """ - # `attention_kwargs` is consumed by the @apply_lora_scale decorator on this method; - # it does not need to thread through to the inner _forward_* paths. - common = { - "hidden_states": hidden_states, - "chunk_partition": chunk_partition, - "timestep": timestep, - "r_timestep": r_timestep, - "encoder_hidden_states": encoder_hidden_states, - "encoder_hidden_states_image": encoder_hidden_states_image, - "return_dict": return_dict, - } - if kv_cache is not None: - common["kv_cache"] = kv_cache - common["kv_cache_flag"] = kv_cache_flag - if kv_cache_flag is not None and kv_cache_flag.get("is_cache_step"): - return self._forward_cache( - clean_hidden_states=clean_hidden_states, - clean_timestep=clean_timestep, - attention_mask=attention_mask, - **common, - ) - return self._forward_inference(**common) - return self._forward_train( - clean_hidden_states=clean_hidden_states, - clean_timestep=clean_timestep, - attention_mask=attention_mask, - **common, - ) - - def _unpack_latent_sequence(self, latents, num_frames, height, width, patch_size): - batch_size, num_patches, channels = latents.shape - height, width = height // patch_size, width // patch_size - - latents = latents.view( - batch_size * num_frames, height, width, patch_size, patch_size, channels // (patch_size * patch_size) - ) - - latents = latents.permute(0, 5, 1, 3, 2, 4) - latents = latents.reshape( - batch_size, num_frames, channels // (patch_size * patch_size), height * patch_size, width * patch_size - ) - return latents - - def _forward_far_patchify(self, hidden_states, far_cfg, clean_hidden_states=None): - full_hidden_states, compressed_hidden_states = ( - hidden_states[:, :, far_cfg["num_compressed_frames"] :], - hidden_states[:, :, : far_cfg["num_compressed_frames"]], - ) # noqa: E501 - - patchified_full_hidden_states = ( - self.patch_embedding(full_hidden_states).flatten(start_dim=2, end_dim=4).transpose(1, 2) - ) - if clean_hidden_states is not None: - clean_hidden_states = ( - self.patch_embedding(clean_hidden_states).flatten(start_dim=2, end_dim=4).transpose(1, 2) - ) - patchified_full_hidden_states = torch.cat([patchified_full_hidden_states, clean_hidden_states], dim=1) - - if far_cfg["num_compressed_frames"] > 0: - patchified_compressed_hidden_states = ( - self.far_patch_embedding(compressed_hidden_states).flatten(start_dim=2, end_dim=4).transpose(1, 2) - ) - hidden_states = torch.cat([patchified_compressed_hidden_states, patchified_full_hidden_states], dim=1) - else: - hidden_states = patchified_full_hidden_states - return hidden_states - - def _forward_far_patchify_inference(self, hidden_states): - hidden_states = self.patch_embedding(hidden_states).flatten(start_dim=2, end_dim=4).transpose(1, 2) - return hidden_states - - def build_attention_mask( - self, - *, - chunk_partition: List[int], - height: int, - width: int, - has_clean_context: bool = False, - device: Optional[torch.device] = None, - mode: str = "train", - ) -> BlockMask: - r"""Pre-build the causal :class:`~torch.nn.attention.flex_attention.BlockMask` outside ``forward``. - - Pass the result via :meth:`forward`'s ``attention_mask`` kwarg to make the whole transformer compatible with - ``torch.compile(fullgraph=True)``. Without a pre-built mask, ``forward`` falls back to constructing it - internally — that path uses ``flex_attention.create_block_mask(_compile=False)`` and breaks the compile graph. - - Args: - chunk_partition: per-chunk frame counts (must sum to the number of latent frames). - height, width: latent spatial dimensions. - has_clean_context: ``True`` when ``clean_hidden_states`` will be threaded through :meth:`forward` - (training V2V/I2V); only this presence flag affects the mask layout. - device: device for the resulting :class:`BlockMask`. The mask is not auto-moved by - ``device_map="auto"``; build it on the same device the transformer's inputs will live on. - mode: ``"train"`` (matches :meth:`_forward_train`) or ``"cache"`` (matches :meth:`_forward_cache`). - The autoregressive ``_forward_inference`` path attends through the KV cache and has no mode here. - - Returns: - :class:`~torch.nn.attention.flex_attention.BlockMask`: causal mask spanning the FAR layout, padded to a - multiple of 128 along the sequence dimension (the BlockMask block-size requirement). - - Raises: - ValueError: if ``mode`` is neither ``"train"`` nor ``"cache"``. - """ - return _build_anyflow_far_causal_block_mask( - chunk_partition=chunk_partition, - height=height, - width=width, - patch_size=self.config.patch_size, - compressed_patch_size=self.config.compressed_patch_size, - full_chunk_limit=self.config.full_chunk_limit, - mode=mode, - has_clean_context=has_clean_context, - device=device, - ) - - def _forward_inference( - self, - hidden_states: torch.Tensor, - chunk_partition, - timestep: torch.LongTensor, - r_timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: Optional[torch.Tensor] = None, - return_dict: bool = True, - kv_cache=None, - kv_cache_flag=None, - ) -> Union[torch.Tensor, Dict[str, torch.Tensor]]: - hidden_states = hidden_states.permute(0, 2, 1, 3, 4) - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - - full_token_per_frame = (height // self.config.patch_size[1]) * (width // self.config.patch_size[2]) - compressed_token_per_frame = (height // self.config.compressed_patch_size[1]) * ( - width // self.config.compressed_patch_size[2] - ) - - total_chunks = 1 + kv_cache_flag["num_cached_chunks"] - - if total_chunks >= self.config.full_chunk_limit: - num_full_chunk, num_compressed_chunk = ( - self.config.full_chunk_limit, - total_chunks - self.config.full_chunk_limit, - ) - else: - num_full_chunk, num_compressed_chunk = total_chunks, 0 - - kv_cache_flag["num_cached_full_tokens"] = ( - sum(chunk_partition[num_compressed_chunk : num_compressed_chunk + (num_full_chunk - 1)]) - * full_token_per_frame - ) # noqa: E501 - kv_cache_flag["num_cached_compressed_tokens"] = ( - sum(chunk_partition[:num_compressed_chunk]) * compressed_token_per_frame - ) - - far_cfg = { - "total_frames": sum(chunk_partition), - "num_full_frames": sum(chunk_partition[num_compressed_chunk:]), - "num_compressed_frames": sum(chunk_partition[:num_compressed_chunk]), - "full_frame_shape": (height // self.config.patch_size[1], width // self.config.patch_size[2]), - "compressed_frame_shape": ( - height // self.config.compressed_patch_size[1], - width // self.config.compressed_patch_size[2], - ), - "full_token_per_frame": full_token_per_frame, - "compressed_token_per_frame": compressed_token_per_frame, - } - - attention_mask = None - hidden_states = self._forward_far_patchify_inference(hidden_states) - - rotary_emb = self.rope(far_cfg=far_cfg, device=hidden_states.device) - rotary_emb["query"] = rotary_emb["query"][:, :, -hidden_states.shape[1] :] - - temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder( - timestep, - r_timestep, - encoder_hidden_states, - encoder_hidden_states_image, - far_cfg=far_cfg, # noqa: E501 - ) - timestep_proj = timestep_proj.unflatten(2, (6, -1)) - - if encoder_hidden_states_image is not None: - encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) - - # 4. Transformer blocks - for index_block, block in enumerate(self.blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - timestep_proj, - rotary_emb, - attention_mask, - kv_cache[index_block], - kv_cache_flag, - ) - else: - hidden_states = block( - hidden_states, - encoder_hidden_states, - timestep_proj, - rotary_emb, - attention_mask, - kv_cache[index_block], - kv_cache_flag, - ) - - # 5. Output norm, projection & unpatchify - shift, scale = (self.scale_shift_table + temb.unsqueeze(2)).chunk(2, dim=2) - shift, scale = shift.squeeze(2), scale.squeeze(2) - - # Move the shift and scale tensors to the same device as hidden_states. - # When using multi-GPU inference via accelerate these will be on the - # first device rather than the last device, which hidden_states ends up - # on. - shift = shift.to(hidden_states.device) - scale = scale.to(hidden_states.device) - - hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states) - - output = self.proj_out(hidden_states) - output = self._unpack_latent_sequence( - output, num_frames=chunk_partition[-1], height=height, width=width, patch_size=self.config.patch_size[1] - ) - - if not return_dict: - return output, kv_cache - - return AnyFlowFARTransformerOutput(sample=output, kv_cache=kv_cache) - - def _forward_cache( - self, - hidden_states: torch.Tensor, - chunk_partition, - timestep: torch.LongTensor, - r_timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: Optional[torch.Tensor] = None, - attention_mask: Optional[BlockMask] = None, - return_dict: bool = True, - clean_hidden_states=None, - clean_timestep=None, - kv_cache=None, - kv_cache_flag=None, - ) -> Union[torch.Tensor, Dict[str, torch.Tensor]]: - hidden_states = hidden_states.permute(0, 2, 1, 3, 4) - if clean_hidden_states is not None: - clean_hidden_states = clean_hidden_states.permute(0, 2, 1, 3, 4) - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - - full_token_per_frame = (height // self.config.patch_size[1]) * (width // self.config.patch_size[2]) - compressed_token_per_frame = (height // self.config.compressed_patch_size[1]) * ( - width // self.config.compressed_patch_size[2] - ) - total_chunks = len(chunk_partition) - - full_chunk_limit = self.config.full_chunk_limit - 1 - - if total_chunks > full_chunk_limit: - num_full_chunk, num_compressed_chunk = full_chunk_limit, total_chunks - full_chunk_limit - else: - num_full_chunk, num_compressed_chunk = total_chunks, 0 - - far_cfg = { - "total_frames": sum(chunk_partition), - "num_full_chunk": num_full_chunk, - "num_full_frames": sum(chunk_partition[num_compressed_chunk:]), - "num_compressed_chunk": num_compressed_chunk, - "num_compressed_frames": sum(chunk_partition[:num_compressed_chunk]), - "full_frame_shape": (height // self.config.patch_size[1], width // self.config.patch_size[2]), - "compressed_frame_shape": ( - height // self.config.compressed_patch_size[1], - width // self.config.compressed_patch_size[2], - ), - "full_token_per_frame": full_token_per_frame, - "compressed_token_per_frame": compressed_token_per_frame, - "chunk_partition": chunk_partition, - } - - kv_cache_flag["num_full_tokens"] = far_cfg["num_full_frames"] * far_cfg["full_token_per_frame"] - kv_cache_flag["num_compressed_tokens"] = ( - far_cfg["num_compressed_frames"] * far_cfg["compressed_token_per_frame"] - ) - - if attention_mask is None: - attention_mask = _build_far_block_mask_from_far_cfg( - far_cfg, has_clean=clean_hidden_states is not None, device=hidden_states.device - ) - - rotary_emb = self.rope(far_cfg=far_cfg, clean_hidden_states=clean_hidden_states, device=hidden_states.device) - hidden_states = self._forward_far_patchify( - hidden_states, far_cfg=far_cfg, clean_hidden_states=clean_hidden_states - ) - - temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder( - timestep, - r_timestep, - encoder_hidden_states, - encoder_hidden_states_image, - far_cfg=far_cfg, - clean_timestep=clean_timestep, - ) - timestep_proj = timestep_proj.unflatten(2, (6, -1)) - - if encoder_hidden_states_image is not None: - encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) - - # 4. Transformer blocks - for index_block, block in enumerate(self.blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - timestep_proj, - rotary_emb, - attention_mask, - kv_cache[index_block], - kv_cache_flag, - ) - else: - hidden_states = block( - hidden_states, - encoder_hidden_states, - timestep_proj, - rotary_emb, - attention_mask, - kv_cache[index_block], - kv_cache_flag, - ) - - if not return_dict: - return None, kv_cache - - return AnyFlowFARTransformerOutput(sample=None, kv_cache=kv_cache) - - def _forward_train( - self, - hidden_states: torch.Tensor, - chunk_partition, - timestep: torch.LongTensor, - r_timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: Optional[torch.Tensor] = None, - attention_mask: Optional[BlockMask] = None, - return_dict: bool = True, - clean_hidden_states=None, - clean_timestep=None, - ) -> Union[torch.Tensor, Dict[str, torch.Tensor]]: - hidden_states = hidden_states.permute(0, 2, 1, 3, 4) - if clean_hidden_states is not None: - clean_hidden_states = clean_hidden_states.permute(0, 2, 1, 3, 4) - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - - full_token_per_frame = (height // self.config.patch_size[1]) * (width // self.config.patch_size[2]) - compressed_token_per_frame = (height // self.config.compressed_patch_size[1]) * ( - width // self.config.compressed_patch_size[2] - ) - total_chunks = len(chunk_partition) - - if total_chunks > self.config.full_chunk_limit: - num_full_chunk, num_compressed_chunk = ( - self.config.full_chunk_limit, - total_chunks - self.config.full_chunk_limit, - ) - else: - num_full_chunk, num_compressed_chunk = total_chunks, 0 - - far_cfg = { - "total_frames": sum(chunk_partition), - "num_full_chunk": num_full_chunk, - "num_full_frames": sum(chunk_partition[num_compressed_chunk:]), - "num_compressed_chunk": num_compressed_chunk, - "num_compressed_frames": sum(chunk_partition[:num_compressed_chunk]), - "full_frame_shape": (height // self.config.patch_size[1], width // self.config.patch_size[2]), - "compressed_frame_shape": ( - height // self.config.compressed_patch_size[1], - width // self.config.compressed_patch_size[2], - ), - "full_token_per_frame": full_token_per_frame, - "compressed_token_per_frame": compressed_token_per_frame, - "chunk_partition": chunk_partition, - } - - if attention_mask is None: - # Fallback for callers that don't pre-build an attention mask (e.g. training scripts). This will introduce - # a graph break, which will cause an error if `torch.compile(fullgraph=True)` is used. In this case, - # pre-build the mask using `build_attention_mask` and pass it via the `attention_mask` argument. - attention_mask = _build_far_block_mask_from_far_cfg( - far_cfg, has_clean=clean_hidden_states is not None, device=hidden_states.device - ) - - rotary_emb = self.rope(far_cfg=far_cfg, clean_hidden_states=clean_hidden_states, device=hidden_states.device) - - hidden_states = self._forward_far_patchify( - hidden_states, far_cfg=far_cfg, clean_hidden_states=clean_hidden_states - ) - - temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder( - timestep, - r_timestep, - encoder_hidden_states, - encoder_hidden_states_image, - far_cfg=far_cfg, - clean_timestep=clean_timestep, - ) - timestep_proj = timestep_proj.unflatten(2, (6, -1)) - - if encoder_hidden_states_image is not None: - encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) - - # 4. Transformer blocks - for index_block, block in enumerate(self.blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - timestep_proj, - rotary_emb, - attention_mask, - ) - else: - hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb, attention_mask) - - # 5. Output norm, projection & unpatchify - shift, scale = (self.scale_shift_table + temb.unsqueeze(2)).chunk(2, dim=2) - shift, scale = shift.squeeze(2), scale.squeeze(2) - - # Move the shift and scale tensors to the same device as hidden_states. - # When using multi-GPU inference via accelerate these will be on the - # first device rather than the last device, which hidden_states ends up - # on. - shift = shift.to(hidden_states.device) - scale = scale.to(hidden_states.device) - - hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states) - - if clean_hidden_states is not None: - hidden_states = hidden_states[ - :, : -(far_cfg["num_full_frames"] * far_cfg["full_token_per_frame"]) - ] # remove clean copy - output = self.proj_out( - hidden_states[:, far_cfg["num_compressed_frames"] * far_cfg["compressed_token_per_frame"] :] - ) # remove far context - output = self._unpack_latent_sequence( - output, - num_frames=far_cfg["num_full_frames"], - height=height, - width=width, - patch_size=self.config.patch_size[1], - ) # noqa: E501 - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_bria.py b/diffusers/models/transformers/transformer_bria.py deleted file mode 100644 index ff4261343ab28c46bb6bd59c08b7695e5c9b4b7c..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_bria.py +++ /dev/null @@ -1,714 +0,0 @@ -import inspect -from typing import Any - -import numpy as np -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device, maybe_allow_in_graph -from ..attention import AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import TimestepEmbedding, apply_rotary_emb, get_timestep_embedding -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _get_projections(attn: "BriaAttention", hidden_states, encoder_hidden_states=None): - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - encoder_query = encoder_key = encoder_value = None - if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None: - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_fused_projections(attn: "BriaAttention", hidden_states, encoder_hidden_states=None): - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - - encoder_query = encoder_key = encoder_value = (None,) - if encoder_hidden_states is not None and hasattr(attn, "to_added_qkv"): - encoder_query, encoder_key, encoder_value = attn.to_added_qkv(encoder_hidden_states).chunk(3, dim=-1) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_qkv_projections(attn: "BriaAttention", hidden_states, encoder_hidden_states=None): - if attn.fused_projections: - return _get_fused_projections(attn, hidden_states, encoder_hidden_states) - return _get_projections(attn, hidden_states, encoder_hidden_states) - - -def get_1d_rotary_pos_embed( - dim: int, - pos: np.ndarray | int, - theta: float = 10000.0, - use_real=False, - linear_factor=1.0, - ntk_factor=1.0, - repeat_interleave_real=True, - freqs_dtype=torch.float32, # torch.float32, torch.float64 (flux) -): - """ - Precompute the frequency tensor for complex exponentials (cis) with given dimensions. - - This function calculates a frequency tensor with complex exponentials using the given dimension 'dim' and the end - index 'end'. The 'theta' parameter scales the frequencies. The returned tensor contains complex values in complex64 - data type. - - Args: - dim (`int`): Dimension of the frequency tensor. - pos (`np.ndarray` or `int`): Position indices for the frequency tensor. [S] or scalar - theta (`float`, *optional*, defaults to 10000.0): - Scaling factor for frequency computation. Defaults to 10000.0. - use_real (`bool`, *optional*): - If True, return real part and imaginary part separately. Otherwise, return complex numbers. - linear_factor (`float`, *optional*, defaults to 1.0): - Scaling factor for the context extrapolation. Defaults to 1.0. - ntk_factor (`float`, *optional*, defaults to 1.0): - Scaling factor for the NTK-Aware RoPE. Defaults to 1.0. - repeat_interleave_real (`bool`, *optional*, defaults to `True`): - If `True` and `use_real`, real part and imaginary part are each interleaved with themselves to reach `dim`. - Otherwise, they are concateanted with themselves. - freqs_dtype (`torch.float32` or `torch.float64`, *optional*, defaults to `torch.float32`): - the dtype of the frequency tensor. - Returns: - `torch.Tensor`: Precomputed frequency tensor with complex exponentials. [S, D/2] - """ - assert dim % 2 == 0 - - if isinstance(pos, int): - pos = torch.arange(pos) - if isinstance(pos, np.ndarray): - pos = torch.from_numpy(pos) # type: ignore # [S] - - theta = theta * ntk_factor - freqs = ( - 1.0 - / (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim)) - / linear_factor - ) # [D/2] - freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2] - if use_real and repeat_interleave_real: - # bria - freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D] - freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D] - return freqs_cos, freqs_sin - elif use_real: - # stable audio, allegro - freqs_cos = torch.cat([freqs.cos(), freqs.cos()], dim=-1).float() # [S, D] - freqs_sin = torch.cat([freqs.sin(), freqs.sin()], dim=-1).float() # [S, D] - return freqs_cos, freqs_sin - else: - # lumina - freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2] - return freqs_cis - - -class BriaAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError(f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version.") - - def __call__( - self, - attn: "BriaAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - query, key, value, encoder_query, encoder_key, encoder_value = _get_qkv_projections( - attn, hidden_states, encoder_hidden_states - ) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if attn.added_kv_proj_dim is not None: - encoder_query = encoder_query.unflatten(-1, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(-1, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(-1, (attn.heads, -1)) - - encoder_query = attn.norm_added_q(encoder_query) - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - if encoder_hidden_states is not None: - encoder_hidden_states, hidden_states = hidden_states.split_with_sizes( - [encoder_hidden_states.shape[1], hidden_states.shape[1] - encoder_hidden_states.shape[1]], dim=1 - ) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -class BriaAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = BriaAttnProcessor - _available_processors = [ - BriaAttnProcessor, - ] - - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - added_kv_proj_dim: int | None = None, - added_proj_bias: bool | None = True, - out_bias: bool = True, - eps: float = 1e-5, - out_dim: int = None, - context_pre_only: bool | None = None, - pre_only: bool = False, - elementwise_affine: bool = True, - processor=None, - ): - super().__init__() - - self.head_dim = dim_head - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.query_dim = query_dim - self.use_bias = bias - self.dropout = dropout - self.out_dim = out_dim if out_dim is not None else query_dim - self.context_pre_only = context_pre_only - self.pre_only = pre_only - self.heads = out_dim // dim_head if out_dim is not None else heads - self.added_kv_proj_dim = added_kv_proj_dim - self.added_proj_bias = added_proj_bias - - self.norm_q = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_v = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - - if not self.pre_only: - self.to_out = torch.nn.ModuleList([]) - self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(torch.nn.Dropout(dropout)) - - if added_kv_proj_dim is not None: - self.norm_added_q = torch.nn.RMSNorm(dim_head, eps=eps) - self.norm_added_k = torch.nn.RMSNorm(dim_head, eps=eps) - self.add_q_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.to_add_out = torch.nn.Linear(self.inner_dim, query_dim, bias=out_bias) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - quiet_attn_parameters = {"ip_adapter_masks", "ip_hidden_states"} - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters and k not in quiet_attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -class BriaEmbedND(torch.nn.Module): - # modified from https://github.com/black-forest-labs/flux/blob/c00d7c60b085fce8058b9df845e036090873f2ce/src/flux/modules/layers.py#L11 - def __init__(self, theta: int, axes_dim: list[int]): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - - def forward(self, ids: torch.Tensor) -> torch.Tensor: - n_axes = ids.shape[-1] - cos_out = [] - sin_out = [] - pos = ids.float() - freqs_dtype = maybe_adjust_dtype_for_device(torch.float64, ids.device) - for i in range(n_axes): - cos, sin = get_1d_rotary_pos_embed( - self.axes_dim[i], - pos[:, i], - theta=self.theta, - repeat_interleave_real=True, - use_real=True, - freqs_dtype=freqs_dtype, - ) - cos_out.append(cos) - sin_out.append(sin) - freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device) - freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device) - return freqs_cos, freqs_sin - - -class BriaTimesteps(nn.Module): - def __init__( - self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float, scale: int = 1, time_theta=10000 - ): - super().__init__() - self.num_channels = num_channels - self.flip_sin_to_cos = flip_sin_to_cos - self.downscale_freq_shift = downscale_freq_shift - self.scale = scale - self.time_theta = time_theta - - def forward(self, timesteps): - t_emb = get_timestep_embedding( - timesteps, - self.num_channels, - flip_sin_to_cos=self.flip_sin_to_cos, - downscale_freq_shift=self.downscale_freq_shift, - scale=self.scale, - max_period=self.time_theta, - ) - return t_emb - - -class BriaTimestepProjEmbeddings(nn.Module): - def __init__(self, embedding_dim, time_theta): - super().__init__() - - self.time_proj = BriaTimesteps( - num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0, time_theta=time_theta - ) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - def forward(self, timestep, dtype): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=dtype)) # (N, D) - return timesteps_emb - - -class BriaPosEmbed(torch.nn.Module): - # modified from https://github.com/black-forest-labs/flux/blob/c00d7c60b085fce8058b9df845e036090873f2ce/src/flux/modules/layers.py#L11 - def __init__(self, theta: int, axes_dim: list[int]): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - - def forward(self, ids: torch.Tensor) -> torch.Tensor: - n_axes = ids.shape[-1] - cos_out = [] - sin_out = [] - pos = ids.float() - freqs_dtype = maybe_adjust_dtype_for_device(torch.float64, ids.device) - for i in range(n_axes): - cos, sin = get_1d_rotary_pos_embed( - self.axes_dim[i], - pos[:, i], - theta=self.theta, - repeat_interleave_real=True, - use_real=True, - freqs_dtype=freqs_dtype, - ) - cos_out.append(cos) - sin_out.append(sin) - freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device) - freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device) - return freqs_cos, freqs_sin - - -@maybe_allow_in_graph -class BriaTransformerBlock(nn.Module): - def __init__( - self, dim: int, num_attention_heads: int, attention_head_dim: int, qk_norm: str = "rms_norm", eps: float = 1e-6 - ): - super().__init__() - - self.norm1 = AdaLayerNormZero(dim) - self.norm1_context = AdaLayerNormZero(dim) - - self.attn = BriaAttention( - query_dim=dim, - added_kv_proj_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - context_pre_only=False, - bias=True, - processor=BriaAttnProcessor(), - eps=eps, - ) - - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff_context = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - attention_kwargs = attention_kwargs or {} - - # Attention. - attention_outputs = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - **attention_kwargs, - ) - - if len(attention_outputs) == 2: - attn_output, context_attn_output = attention_outputs - elif len(attention_outputs) == 3: - attn_output, context_attn_output, ip_attn_output = attention_outputs - - # Process attention outputs for the `hidden_states`. - attn_output = gate_msa.unsqueeze(1) * attn_output - hidden_states = hidden_states + attn_output - - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - - ff_output = self.ff(norm_hidden_states) - ff_output = gate_mlp.unsqueeze(1) * ff_output - - hidden_states = hidden_states + ff_output - if len(attention_outputs) == 3: - hidden_states = hidden_states + ip_attn_output - - # Process attention outputs for the `encoder_hidden_states`. - context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output - encoder_hidden_states = encoder_hidden_states + context_attn_output - - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - context_ff_output = self.ff_context(norm_encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - if encoder_hidden_states.dtype == torch.float16: - encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) - - return encoder_hidden_states, hidden_states - - -@maybe_allow_in_graph -class BriaSingleTransformerBlock(nn.Module): - def __init__(self, dim: int, num_attention_heads: int, attention_head_dim: int, mlp_ratio: float = 4.0): - super().__init__() - self.mlp_hidden_dim = int(dim * mlp_ratio) - - self.norm = AdaLayerNormZeroSingle(dim) - self.proj_mlp = nn.Linear(dim, self.mlp_hidden_dim) - self.act_mlp = nn.GELU(approximate="tanh") - self.proj_out = nn.Linear(dim + self.mlp_hidden_dim, dim) - - processor = BriaAttnProcessor() - - self.attn = BriaAttention( - query_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - bias=True, - processor=processor, - eps=1e-6, - pre_only=True, - ) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - text_seq_len = encoder_hidden_states.shape[1] - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - residual = hidden_states - norm_hidden_states, gate = self.norm(hidden_states, emb=temb) - mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) - attention_kwargs = attention_kwargs or {} - attn_output = self.attn( - hidden_states=norm_hidden_states, - image_rotary_emb=image_rotary_emb, - **attention_kwargs, - ) - - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) - gate = gate.unsqueeze(1) - hidden_states = gate * self.proj_out(hidden_states) - hidden_states = residual + hidden_states - if hidden_states.dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - encoder_hidden_states, hidden_states = hidden_states[:, :text_seq_len], hidden_states[:, text_seq_len:] - return encoder_hidden_states, hidden_states - - -class BriaTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin): - """ - The Transformer model introduced in Flux. Based on FluxPipeline with several changes: - - no pooled embeddings - - We use zero padding for prompts - - No guidance embedding since this is not a distilled version - Reference: https://blackforestlabs.ai/announcing-black-forest-labs/ - - Parameters: - patch_size (`int`): Patch size to turn the input data into small patches. - in_channels (`int`, *optional*, defaults to 16): The number of channels in the input. - num_layers (`int`, *optional*, defaults to 18): The number of layers of MMDiT blocks to use. - num_single_layers (`int`, *optional*, defaults to 18): The number of layers of single DiT blocks to use. - attention_head_dim (`int`, *optional*, defaults to 64): The number of channels in each head. - num_attention_heads (`int`, *optional*, defaults to 18): The number of heads to use for multi-head attention. - joint_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use. - pooled_projection_dim (`int`): Number of dimensions to use when projecting the `pooled_projections`. - guidance_embeds (`bool`, defaults to False): Whether to use guidance embeddings. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - patch_size: int = 1, - in_channels: int = 64, - num_layers: int = 19, - num_single_layers: int = 38, - attention_head_dim: int = 128, - num_attention_heads: int = 24, - joint_attention_dim: int = 4096, - pooled_projection_dim: int = None, - guidance_embeds: bool = False, - axes_dims_rope: list[int] = [16, 56, 56], - rope_theta=10000, - time_theta=10000, - ): - super().__init__() - self.out_channels = in_channels - self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim - - self.pos_embed = BriaEmbedND(theta=rope_theta, axes_dim=axes_dims_rope) - - self.time_embed = BriaTimestepProjEmbeddings(embedding_dim=self.inner_dim, time_theta=time_theta) - if guidance_embeds: - self.guidance_embed = BriaTimestepProjEmbeddings(embedding_dim=self.inner_dim) - - self.context_embedder = nn.Linear(self.config.joint_attention_dim, self.inner_dim) - self.x_embedder = torch.nn.Linear(self.config.in_channels, self.inner_dim) - - self.transformer_blocks = nn.ModuleList( - [ - BriaTransformerBlock( - dim=self.inner_dim, - num_attention_heads=self.config.num_attention_heads, - attention_head_dim=self.config.attention_head_dim, - ) - for i in range(self.config.num_layers) - ] - ) - - self.single_transformer_blocks = nn.ModuleList( - [ - BriaSingleTransformerBlock( - dim=self.inner_dim, - num_attention_heads=self.config.num_attention_heads, - attention_head_dim=self.config.attention_head_dim, - ) - for i in range(self.config.num_single_layers) - ] - ) - - self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - pooled_projections: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_ids: torch.Tensor = None, - txt_ids: torch.Tensor = None, - guidance: torch.Tensor = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - controlnet_block_samples=None, - controlnet_single_block_samples=None, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`BriaTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.FloatTensor` of shape `(batch size, channel, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.FloatTensor` of shape `(batch size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - pooled_projections (`torch.FloatTensor` of shape `(batch_size, projection_dim)`): Embeddings projected - from the embeddings of input conditions. - timestep ( `torch.LongTensor`): - Used to indicate denoising step. - img_ids (`torch.Tensor`): - Image position ids used to compute the rotary positional embeddings. - txt_ids (`torch.Tensor`): - Text position ids used to compute the rotary positional embeddings. - guidance (`torch.Tensor`, *optional*): - Guidance scale embedding used for guidance-distilled variants of the model. - controlnet_block_samples (`list` of `torch.Tensor`, *optional*): - A list of tensors that if specified are added to the residuals of transformer blocks. - controlnet_single_block_samples (`list` of `torch.Tensor`, *optional*): - A list of tensors that if specified are added to the residuals of single transformer blocks. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - hidden_states = self.x_embedder(hidden_states) - - timestep = timestep.to(hidden_states.dtype) - if guidance is not None: - guidance = guidance.to(hidden_states.dtype) - else: - guidance = None - - temb = self.time_embed(timestep, dtype=hidden_states.dtype) - - if guidance: - temb += self.guidance_embed(guidance, dtype=hidden_states.dtype) - - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - - if len(txt_ids.shape) == 3: - txt_ids = txt_ids[0] - - if len(img_ids.shape) == 3: - img_ids = img_ids[0] - - ids = torch.cat((txt_ids, img_ids), dim=0) - image_rotary_emb = self.pos_embed(ids) - - for index_block, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - attention_kwargs, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - ) - - # controlnet residual - if controlnet_block_samples is not None: - interval_control = len(self.transformer_blocks) / len(controlnet_block_samples) - interval_control = int(np.ceil(interval_control)) - hidden_states = hidden_states + controlnet_block_samples[index_block // interval_control] - - for index_block, block in enumerate(self.single_transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - attention_kwargs, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - ) - - # controlnet residual - if controlnet_single_block_samples is not None: - interval_control = len(self.single_transformer_blocks) / len(controlnet_single_block_samples) - interval_control = int(np.ceil(interval_control)) - hidden_states[:, encoder_hidden_states.shape[1] :, ...] = ( - hidden_states[:, encoder_hidden_states.shape[1] :, ...] - + controlnet_single_block_samples[index_block // interval_control] - ) - - hidden_states = self.norm_out(hidden_states, temb) - output = self.proj_out(hidden_states) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_bria_fibo.py b/diffusers/models/transformers/transformer_bria_fibo.py deleted file mode 100644 index 78545cb7da31d347f9a76df8d0b27a7c893a0cc1..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_bria_fibo.py +++ /dev/null @@ -1,644 +0,0 @@ -# Copyright (c) Bria.ai. All rights reserved. -# -# This file is licensed under the Creative Commons Attribution-NonCommercial 4.0 International Public License (CC-BY-NC-4.0). -# You may obtain a copy of the license at https://creativecommons.org/licenses/by-nc/4.0/ -# -# You are free to share and adapt this material for non-commercial purposes provided you give appropriate credit, -# indicate if changes were made, and do not use the material for commercial purposes. -# -# See the license for further details. -import inspect -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...models.attention_processor import Attention -from ...models.embeddings import TimestepEmbedding, apply_rotary_emb, get_1d_rotary_pos_embed, get_timestep_embedding -from ...models.modeling_outputs import Transformer2DModelOutput -from ...models.modeling_utils import ModelMixin -from ...models.transformers.transformer_bria import BriaAttnProcessor -from ...utils import ( - apply_lora_scale, - logging, -) -from ...utils.torch_utils import maybe_adjust_dtype_for_device, maybe_allow_in_graph -from ..attention import AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..normalization import AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _get_projections(attn: "BriaFiboAttention", hidden_states, encoder_hidden_states=None): - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - encoder_query = encoder_key = encoder_value = None - if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None: - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_fused_projections(attn: "BriaFiboAttention", hidden_states, encoder_hidden_states=None): - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - - encoder_query = encoder_key = encoder_value = (None,) - if encoder_hidden_states is not None and hasattr(attn, "to_added_qkv"): - encoder_query, encoder_key, encoder_value = attn.to_added_qkv(encoder_hidden_states).chunk(3, dim=-1) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_qkv_projections(attn: "BriaFiboAttention", hidden_states, encoder_hidden_states=None): - if attn.fused_projections: - return _get_fused_projections(attn, hidden_states, encoder_hidden_states) - return _get_projections(attn, hidden_states, encoder_hidden_states) - - -# Copied from diffusers.models.transformers.transformer_flux.FluxAttnProcessor with FluxAttnProcessor->BriaFiboAttnProcessor, FluxAttention->BriaFiboAttention -class BriaFiboAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError(f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version.") - - def __call__( - self, - attn: "BriaFiboAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - query, key, value, encoder_query, encoder_key, encoder_value = _get_qkv_projections( - attn, hidden_states, encoder_hidden_states - ) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if attn.added_kv_proj_dim is not None: - encoder_query = encoder_query.unflatten(-1, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(-1, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(-1, (attn.heads, -1)) - - encoder_query = attn.norm_added_q(encoder_query) - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - if encoder_hidden_states is not None: - encoder_hidden_states, hidden_states = hidden_states.split_with_sizes( - [encoder_hidden_states.shape[1], hidden_states.shape[1] - encoder_hidden_states.shape[1]], dim=1 - ) - hidden_states = attn.to_out[0](hidden_states.contiguous()) - hidden_states = attn.to_out[1](hidden_states) - encoder_hidden_states = attn.to_add_out(encoder_hidden_states.contiguous()) - - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -# Based on https://github.com/huggingface/diffusers/blob/55d49d4379007740af20629bb61aba9546c6b053/src/diffusers/models/transformers/transformer_flux.py -class BriaFiboAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = BriaFiboAttnProcessor - _available_processors = [BriaFiboAttnProcessor] - - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - added_kv_proj_dim: int | None = None, - added_proj_bias: bool | None = True, - out_bias: bool = True, - eps: float = 1e-5, - out_dim: int = None, - context_pre_only: bool | None = None, - pre_only: bool = False, - elementwise_affine: bool = True, - processor=None, - ): - super().__init__() - - self.head_dim = dim_head - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.query_dim = query_dim - self.use_bias = bias - self.dropout = dropout - self.out_dim = out_dim if out_dim is not None else query_dim - self.context_pre_only = context_pre_only - self.pre_only = pre_only - self.heads = out_dim // dim_head if out_dim is not None else heads - self.added_kv_proj_dim = added_kv_proj_dim - self.added_proj_bias = added_proj_bias - - self.norm_q = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_v = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - - if not self.pre_only: - self.to_out = torch.nn.ModuleList([]) - self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(torch.nn.Dropout(dropout)) - - if added_kv_proj_dim is not None: - self.norm_added_q = torch.nn.RMSNorm(dim_head, eps=eps) - self.norm_added_k = torch.nn.RMSNorm(dim_head, eps=eps) - self.add_q_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.to_add_out = torch.nn.Linear(self.inner_dim, query_dim, bias=out_bias) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - quiet_attn_parameters = {"ip_adapter_masks", "ip_hidden_states"} - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters and k not in quiet_attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"joint_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -class BriaFiboEmbedND(torch.nn.Module): - # modified from https://github.com/black-forest-labs/flux/blob/c00d7c60b085fce8058b9df845e036090873f2ce/src/flux/modules/layers.py#L11 - def __init__(self, theta: int, axes_dim: list[int]): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - - def forward(self, ids: torch.Tensor) -> torch.Tensor: - n_axes = ids.shape[-1] - cos_out = [] - sin_out = [] - pos = ids.float() - freqs_dtype = maybe_adjust_dtype_for_device(torch.float64, ids.device) - for i in range(n_axes): - cos, sin = get_1d_rotary_pos_embed( - self.axes_dim[i], - pos[:, i], - theta=self.theta, - repeat_interleave_real=True, - use_real=True, - freqs_dtype=freqs_dtype, - ) - cos_out.append(cos) - sin_out.append(sin) - freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device) - freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device) - return freqs_cos, freqs_sin - - -@maybe_allow_in_graph -class BriaFiboSingleTransformerBlock(nn.Module): - def __init__(self, dim: int, num_attention_heads: int, attention_head_dim: int, mlp_ratio: float = 4.0): - super().__init__() - self.mlp_hidden_dim = int(dim * mlp_ratio) - - self.norm = AdaLayerNormZeroSingle(dim) - self.proj_mlp = nn.Linear(dim, self.mlp_hidden_dim) - self.act_mlp = nn.GELU(approximate="tanh") - self.proj_out = nn.Linear(dim + self.mlp_hidden_dim, dim) - - processor = BriaAttnProcessor() - - self.attn = Attention( - query_dim=dim, - cross_attention_dim=None, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - bias=True, - processor=processor, - qk_norm="rms_norm", - eps=1e-6, - pre_only=True, - ) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor: - residual = hidden_states - norm_hidden_states, gate = self.norm(hidden_states, emb=temb) - mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) - joint_attention_kwargs = joint_attention_kwargs or {} - attn_output = self.attn( - hidden_states=norm_hidden_states, - image_rotary_emb=image_rotary_emb, - **joint_attention_kwargs, - ) - - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) - gate = gate.unsqueeze(1) - hidden_states = gate * self.proj_out(hidden_states) - hidden_states = residual + hidden_states - if hidden_states.dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - return hidden_states - - -class BriaFiboTextProjection(nn.Module): - def __init__(self, in_features, hidden_size): - super().__init__() - self.linear = nn.Linear(in_features=in_features, out_features=hidden_size, bias=False) - - def forward(self, caption): - hidden_states = self.linear(caption) - return hidden_states - - -@maybe_allow_in_graph -# Based on from diffusers.models.transformers.transformer_flux.FluxTransformerBlock -class BriaFiboTransformerBlock(nn.Module): - def __init__( - self, dim: int, num_attention_heads: int, attention_head_dim: int, qk_norm: str = "rms_norm", eps: float = 1e-6 - ): - super().__init__() - - self.norm1 = AdaLayerNormZero(dim) - self.norm1_context = AdaLayerNormZero(dim) - - self.attn = BriaFiboAttention( - query_dim=dim, - added_kv_proj_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - context_pre_only=False, - bias=True, - processor=BriaFiboAttnProcessor(), - eps=eps, - ) - - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff_context = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - joint_attention_kwargs = joint_attention_kwargs or {} - - # Attention. - attention_outputs = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - **joint_attention_kwargs, - ) - - if len(attention_outputs) == 2: - attn_output, context_attn_output = attention_outputs - elif len(attention_outputs) == 3: - attn_output, context_attn_output, ip_attn_output = attention_outputs - - # Process attention outputs for the `hidden_states`. - attn_output = gate_msa.unsqueeze(1) * attn_output - hidden_states = hidden_states + attn_output - - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - - ff_output = self.ff(norm_hidden_states) - ff_output = gate_mlp.unsqueeze(1) * ff_output - - hidden_states = hidden_states + ff_output - if len(attention_outputs) == 3: - hidden_states = hidden_states + ip_attn_output - - # Process attention outputs for the `encoder_hidden_states`. - context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output - encoder_hidden_states = encoder_hidden_states + context_attn_output - - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - context_ff_output = self.ff_context(norm_encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - if encoder_hidden_states.dtype == torch.float16: - encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) - - return encoder_hidden_states, hidden_states - - -class BriaFiboTimesteps(nn.Module): - def __init__( - self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float, scale: int = 1, time_theta=10000 - ): - super().__init__() - self.num_channels = num_channels - self.flip_sin_to_cos = flip_sin_to_cos - self.downscale_freq_shift = downscale_freq_shift - self.scale = scale - self.time_theta = time_theta - - def forward(self, timesteps): - t_emb = get_timestep_embedding( - timesteps, - self.num_channels, - flip_sin_to_cos=self.flip_sin_to_cos, - downscale_freq_shift=self.downscale_freq_shift, - scale=self.scale, - max_period=self.time_theta, - ) - return t_emb - - -class BriaFiboTimestepProjEmbeddings(nn.Module): - def __init__(self, embedding_dim, time_theta): - super().__init__() - - self.time_proj = BriaFiboTimesteps( - num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0, time_theta=time_theta - ) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - def forward(self, timestep, dtype): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=dtype)) # (N, D) - return timesteps_emb - - -class BriaFiboTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): - """ - Parameters: - patch_size (`int`): Patch size to turn the input data into small patches. - in_channels (`int`, *optional*, defaults to 16): The number of channels in the input. - num_layers (`int`, *optional*, defaults to 18): The number of layers of MMDiT blocks to use. - num_single_layers (`int`, *optional*, defaults to 18): The number of layers of single DiT blocks to use. - attention_head_dim (`int`, *optional*, defaults to 64): The number of channels in each head. - num_attention_heads (`int`, *optional*, defaults to 18): The number of heads to use for multi-head attention. - joint_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use. - pooled_projection_dim (`int`): Number of dimensions to use when projecting the `pooled_projections`. - guidance_embeds (`bool`, defaults to False): Whether to use guidance embeddings. - ... - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - patch_size: int = 1, - in_channels: int = 64, - num_layers: int = 19, - num_single_layers: int = 38, - attention_head_dim: int = 128, - num_attention_heads: int = 24, - joint_attention_dim: int = 4096, - pooled_projection_dim: int = None, - guidance_embeds: bool = False, - axes_dims_rope: list[int] = [16, 56, 56], - rope_theta=10000, - time_theta=10000, - text_encoder_dim: int = 2048, - ): - super().__init__() - self.out_channels = in_channels - self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim - - self.pos_embed = BriaFiboEmbedND(theta=rope_theta, axes_dim=axes_dims_rope) - - self.time_embed = BriaFiboTimestepProjEmbeddings(embedding_dim=self.inner_dim, time_theta=time_theta) - - if guidance_embeds: - self.guidance_embed = BriaFiboTimestepProjEmbeddings(embedding_dim=self.inner_dim, time_theta=time_theta) - - self.context_embedder = nn.Linear(self.config.joint_attention_dim, self.inner_dim) - self.x_embedder = torch.nn.Linear(self.config.in_channels, self.inner_dim) - - self.transformer_blocks = nn.ModuleList( - [ - BriaFiboTransformerBlock( - dim=self.inner_dim, - num_attention_heads=self.config.num_attention_heads, - attention_head_dim=self.config.attention_head_dim, - ) - for i in range(self.config.num_layers) - ] - ) - - self.single_transformer_blocks = nn.ModuleList( - [ - BriaFiboSingleTransformerBlock( - dim=self.inner_dim, - num_attention_heads=self.config.num_attention_heads, - attention_head_dim=self.config.attention_head_dim, - ) - for i in range(self.config.num_single_layers) - ] - ) - - self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) - - self.gradient_checkpointing = False - - caption_projection = [ - BriaFiboTextProjection(in_features=text_encoder_dim, hidden_size=self.inner_dim // 2) - for i in range(self.config.num_layers + self.config.num_single_layers) - ] - self.caption_projection = nn.ModuleList(caption_projection) - - @apply_lora_scale("joint_attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - text_encoder_layers: list = None, - pooled_projections: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_ids: torch.Tensor = None, - txt_ids: torch.Tensor = None, - guidance: torch.Tensor = None, - joint_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> torch.FloatTensor | Transformer2DModelOutput: - """ - - Args: - hidden_states (`torch.FloatTensor` of shape `(batch size, channel, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.FloatTensor` of shape `(batch size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - text_encoder_layers (`list` of `torch.Tensor`): - Per-block text encoder hidden states, one tensor per transformer block. - pooled_projections (`torch.FloatTensor` of shape `(batch_size, projection_dim)`): Embeddings projected - from the embeddings of input conditions. - timestep ( `torch.LongTensor`): - Used to indicate denoising step. - img_ids (`torch.Tensor`): - Image position ids used to compute the rotary positional embeddings. - txt_ids (`torch.Tensor`): - Text position ids used to compute the rotary positional embeddings. - guidance (`torch.Tensor`, *optional*): - Guidance scale embedding used for guidance-distilled variants of the model. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - - hidden_states = self.x_embedder(hidden_states) - - timestep = timestep.to(hidden_states.dtype) - if guidance is not None: - guidance = guidance.to(hidden_states.dtype) - else: - guidance = None - - temb = self.time_embed(timestep, dtype=hidden_states.dtype) - - if guidance is not None: - temb += self.guidance_embed(guidance, dtype=hidden_states.dtype) - - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - - if len(txt_ids.shape) == 3: - txt_ids = txt_ids[0] - - if len(img_ids.shape) == 3: - img_ids = img_ids[0] - - ids = torch.cat((txt_ids, img_ids), dim=0) - image_rotary_emb = self.pos_embed(ids) - - new_text_encoder_layers = [] - for i, text_encoder_layer in enumerate(text_encoder_layers): - text_encoder_layer = self.caption_projection[i](text_encoder_layer) - new_text_encoder_layers.append(text_encoder_layer) - text_encoder_layers = new_text_encoder_layers - - block_id = 0 - for index_block, block in enumerate(self.transformer_blocks): - current_text_encoder_layer = text_encoder_layers[block_id] - encoder_hidden_states = torch.cat( - [encoder_hidden_states[:, :, : self.inner_dim // 2], current_text_encoder_layer], dim=-1 - ) - block_id += 1 - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - joint_attention_kwargs, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - joint_attention_kwargs=joint_attention_kwargs, - ) - - for index_block, block in enumerate(self.single_transformer_blocks): - current_text_encoder_layer = text_encoder_layers[block_id] - encoder_hidden_states = torch.cat( - [encoder_hidden_states[:, :, : self.inner_dim // 2], current_text_encoder_layer], dim=-1 - ) - block_id += 1 - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - temb, - image_rotary_emb, - joint_attention_kwargs, - ) - - else: - hidden_states = block( - hidden_states=hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - joint_attention_kwargs=joint_attention_kwargs, - ) - - encoder_hidden_states = hidden_states[:, : encoder_hidden_states.shape[1], ...] - hidden_states = hidden_states[:, encoder_hidden_states.shape[1] :, ...] - - hidden_states = self.norm_out(hidden_states, temb) - output = self.proj_out(hidden_states) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_chroma.py b/diffusers/models/transformers/transformer_chroma.py deleted file mode 100644 index 8d7d9d5d6a04e7898718b5827a6451ea6717e1a2..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_chroma.py +++ /dev/null @@ -1,634 +0,0 @@ -# Copyright 2025 Black Forest Labs, The HuggingFace Team and loadstone-rock . All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from typing import Any - -import numpy as np -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FluxTransformer2DLoadersMixin, FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, deprecate, logging -from ...utils.import_utils import is_torch_npu_available -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import AttentionMixin, FeedForward -from ..cache_utils import CacheMixin -from ..embeddings import FluxPosEmbed, PixArtAlphaTextProjection, Timesteps, get_timestep_embedding -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import CombinedTimestepLabelEmbeddings, FP32LayerNorm, RMSNorm -from .transformer_flux import FluxAttention, FluxAttnProcessor - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class ChromaAdaLayerNormZeroPruned(nn.Module): - r""" - Norm layer adaptive layer norm zero (adaLN-Zero). - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - num_embeddings (`int`): The size of the embeddings dictionary. - """ - - def __init__(self, embedding_dim: int, num_embeddings: int | None = None, norm_type="layer_norm", bias=True): - super().__init__() - if num_embeddings is not None: - self.emb = CombinedTimestepLabelEmbeddings(num_embeddings, embedding_dim) - else: - self.emb = None - - if norm_type == "layer_norm": - self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) - elif norm_type == "fp32_layer_norm": - self.norm = FP32LayerNorm(embedding_dim, elementwise_affine=False, bias=False) - else: - raise ValueError( - f"Unsupported `norm_type` ({norm_type}) provided. Supported ones are: 'layer_norm', 'fp32_layer_norm'." - ) - - def forward( - self, - x: torch.Tensor, - timestep: torch.Tensor | None = None, - class_labels: torch.LongTensor | None = None, - hidden_dtype: torch.dtype | None = None, - emb: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - if self.emb is not None: - emb = self.emb(timestep, class_labels, hidden_dtype=hidden_dtype) - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = emb.flatten(1, 2).chunk(6, dim=1) - x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None] - return x, gate_msa, shift_mlp, scale_mlp, gate_mlp - - -class ChromaAdaLayerNormZeroSinglePruned(nn.Module): - r""" - Norm layer adaptive layer norm zero (adaLN-Zero). - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - num_embeddings (`int`): The size of the embeddings dictionary. - """ - - def __init__(self, embedding_dim: int, norm_type="layer_norm", bias=True): - super().__init__() - - if norm_type == "layer_norm": - self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) - else: - raise ValueError( - f"Unsupported `norm_type` ({norm_type}) provided. Supported ones are: 'layer_norm', 'fp32_layer_norm'." - ) - - def forward( - self, - x: torch.Tensor, - emb: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - shift_msa, scale_msa, gate_msa = emb.flatten(1, 2).chunk(3, dim=1) - x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None] - return x, gate_msa - - -class ChromaAdaLayerNormContinuousPruned(nn.Module): - r""" - Adaptive normalization layer with a norm layer (layer_norm or rms_norm). - - Args: - embedding_dim (`int`): Embedding dimension to use during projection. - conditioning_embedding_dim (`int`): Dimension of the input condition. - elementwise_affine (`bool`, defaults to `True`): - Boolean flag to denote if affine transformation should be applied. - eps (`float`, defaults to 1e-5): Epsilon factor. - bias (`bias`, defaults to `True`): Boolean flag to denote if bias should be use. - norm_type (`str`, defaults to `"layer_norm"`): - Normalization layer to use. Values supported: "layer_norm", "rms_norm". - """ - - def __init__( - self, - embedding_dim: int, - conditioning_embedding_dim: int, - # NOTE: It is a bit weird that the norm layer can be configured to have scale and shift parameters - # because the output is immediately scaled and shifted by the projected conditioning embeddings. - # Note that AdaLayerNorm does not let the norm layer have scale and shift parameters. - # However, this is how it was implemented in the original code, and it's rather likely you should - # set `elementwise_affine` to False. - elementwise_affine=True, - eps=1e-5, - bias=True, - norm_type="layer_norm", - ): - super().__init__() - if norm_type == "layer_norm": - self.norm = nn.LayerNorm(embedding_dim, eps, elementwise_affine, bias) - elif norm_type == "rms_norm": - self.norm = RMSNorm(embedding_dim, eps, elementwise_affine) - else: - raise ValueError(f"unknown norm_type {norm_type}") - - def forward(self, x: torch.Tensor, emb: torch.Tensor) -> torch.Tensor: - # convert back to the original dtype in case `conditioning_embedding`` is upcasted to float32 (needed for hunyuanDiT) - shift, scale = torch.chunk(emb.flatten(1, 2).to(x.dtype), 2, dim=1) - x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :] - return x - - -class ChromaCombinedTimestepTextProjEmbeddings(nn.Module): - def __init__(self, num_channels: int, out_dim: int): - super().__init__() - - self.time_proj = Timesteps(num_channels=num_channels, flip_sin_to_cos=True, downscale_freq_shift=0) - self.guidance_proj = Timesteps(num_channels=num_channels, flip_sin_to_cos=True, downscale_freq_shift=0) - - self.register_buffer( - "mod_proj", - get_timestep_embedding( - torch.arange(out_dim) * 1000, 2 * num_channels, flip_sin_to_cos=True, downscale_freq_shift=0 - ), - persistent=False, - ) - - def forward(self, timestep: torch.Tensor) -> torch.Tensor: - mod_index_length = self.mod_proj.shape[0] - batch_size = timestep.shape[0] - - timesteps_proj = self.time_proj(timestep).to(dtype=timestep.dtype) - guidance_proj = self.guidance_proj(torch.tensor([0] * batch_size)).to( - dtype=timestep.dtype, device=timestep.device - ) - - mod_proj = self.mod_proj.to(dtype=timesteps_proj.dtype, device=timesteps_proj.device).repeat(batch_size, 1, 1) - timestep_guidance = ( - torch.cat([timesteps_proj, guidance_proj], dim=1).unsqueeze(1).repeat(1, mod_index_length, 1) - ) - input_vec = torch.cat([timestep_guidance, mod_proj], dim=-1) - return input_vec.to(timestep.dtype) - - -class ChromaApproximator(nn.Module): - def __init__(self, in_dim: int, out_dim: int, hidden_dim: int, n_layers: int = 5): - super().__init__() - self.in_proj = nn.Linear(in_dim, hidden_dim, bias=True) - self.layers = nn.ModuleList( - [PixArtAlphaTextProjection(hidden_dim, hidden_dim, act_fn="silu") for _ in range(n_layers)] - ) - self.norms = nn.ModuleList([nn.RMSNorm(hidden_dim) for _ in range(n_layers)]) - self.out_proj = nn.Linear(hidden_dim, out_dim) - - def forward(self, x): - x = self.in_proj(x) - - for layer, norms in zip(self.layers, self.norms): - x = x + layer(norms(x)) - - return self.out_proj(x) - - -@maybe_allow_in_graph -class ChromaSingleTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float = 4.0, - ): - super().__init__() - self.mlp_hidden_dim = int(dim * mlp_ratio) - self.norm = ChromaAdaLayerNormZeroSinglePruned(dim) - self.proj_mlp = nn.Linear(dim, self.mlp_hidden_dim) - self.act_mlp = nn.GELU(approximate="tanh") - self.proj_out = nn.Linear(dim + self.mlp_hidden_dim, dim) - - if is_torch_npu_available(): - from ..attention_processor import FluxAttnProcessor2_0_NPU - - deprecation_message = ( - "Defaulting to FluxAttnProcessor2_0_NPU for NPU devices will be removed. Attention processors " - "should be set explicitly using the `set_attn_processor` method." - ) - deprecate("npu_processor", "0.34.0", deprecation_message) - processor = FluxAttnProcessor2_0_NPU() - else: - processor = FluxAttnProcessor() - - self.attn = FluxAttention( - query_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - bias=True, - processor=processor, - eps=1e-6, - pre_only=True, - ) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_mask: torch.Tensor | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor: - residual = hidden_states - norm_hidden_states, gate = self.norm(hidden_states, emb=temb) - mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) - joint_attention_kwargs = joint_attention_kwargs or {} - - if attention_mask is not None: - attention_mask = attention_mask[:, None, None, :] * attention_mask[:, None, :, None] - - attn_output = self.attn( - hidden_states=norm_hidden_states, - image_rotary_emb=image_rotary_emb, - attention_mask=attention_mask, - **joint_attention_kwargs, - ) - - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) - gate = gate.unsqueeze(1) - hidden_states = gate * self.proj_out(hidden_states) - hidden_states = residual + hidden_states - if hidden_states.dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - return hidden_states - - -@maybe_allow_in_graph -class ChromaTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - qk_norm: str = "rms_norm", - eps: float = 1e-6, - ): - super().__init__() - self.norm1 = ChromaAdaLayerNormZeroPruned(dim) - self.norm1_context = ChromaAdaLayerNormZeroPruned(dim) - - self.attn = FluxAttention( - query_dim=dim, - added_kv_proj_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - context_pre_only=False, - bias=True, - processor=FluxAttnProcessor(), - eps=eps, - ) - - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff_context = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_mask: torch.Tensor | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - temb_img, temb_txt = temb[:, :6], temb[:, 6:] - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb_img) - - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb_txt - ) - joint_attention_kwargs = joint_attention_kwargs or {} - if attention_mask is not None: - attention_mask = attention_mask[:, None, None, :] * attention_mask[:, None, :, None] - - # Attention. - attention_outputs = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - attention_mask=attention_mask, - **joint_attention_kwargs, - ) - - if len(attention_outputs) == 2: - attn_output, context_attn_output = attention_outputs - elif len(attention_outputs) == 3: - attn_output, context_attn_output, ip_attn_output = attention_outputs - - # Process attention outputs for the `hidden_states`. - attn_output = gate_msa.unsqueeze(1) * attn_output - hidden_states = hidden_states + attn_output - - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - - ff_output = self.ff(norm_hidden_states) - ff_output = gate_mlp.unsqueeze(1) * ff_output - - hidden_states = hidden_states + ff_output - if len(attention_outputs) == 3: - hidden_states = hidden_states + ip_attn_output - - # Process attention outputs for the `encoder_hidden_states`. - - context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output - encoder_hidden_states = encoder_hidden_states + context_attn_output - - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - context_ff_output = self.ff_context(norm_encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - if encoder_hidden_states.dtype == torch.float16: - encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) - - return encoder_hidden_states, hidden_states - - -class ChromaTransformer2DModel( - ModelMixin, - ConfigMixin, - PeftAdapterMixin, - FromOriginalModelMixin, - FluxTransformer2DLoadersMixin, - CacheMixin, - AttentionMixin, -): - """ - The Transformer model introduced in Flux, modified for Chroma. - - Reference: https://huggingface.co/lodestones/Chroma1-HD - - Args: - patch_size (`int`, defaults to `1`): - Patch size to turn the input data into small patches. - in_channels (`int`, defaults to `64`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `None`): - The number of channels in the output. If not specified, it defaults to `in_channels`. - num_layers (`int`, defaults to `19`): - The number of layers of dual stream DiT blocks to use. - num_single_layers (`int`, defaults to `38`): - The number of layers of single stream DiT blocks to use. - attention_head_dim (`int`, defaults to `128`): - The number of dimensions to use for each attention head. - num_attention_heads (`int`, defaults to `24`): - The number of attention heads to use. - joint_attention_dim (`int`, defaults to `4096`): - The number of dimensions to use for the joint attention (embedding/channel dimension of - `encoder_hidden_states`). - axes_dims_rope (`tuple[int]`, defaults to `(16, 56, 56)`): - The dimensions to use for the rotary positional embeddings. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["ChromaTransformerBlock", "ChromaSingleTransformerBlock"] - _repeated_blocks = ["ChromaTransformerBlock", "ChromaSingleTransformerBlock"] - _skip_layerwise_casting_patterns = ["pos_embed", "norm"] - - @register_to_config - def __init__( - self, - patch_size: int = 1, - in_channels: int = 64, - out_channels: int | None = None, - num_layers: int = 19, - num_single_layers: int = 38, - attention_head_dim: int = 128, - num_attention_heads: int = 24, - joint_attention_dim: int = 4096, - axes_dims_rope: tuple[int, ...] = (16, 56, 56), - approximator_num_channels: int = 64, - approximator_hidden_dim: int = 5120, - approximator_layers: int = 5, - ): - super().__init__() - self.out_channels = out_channels or in_channels - self.inner_dim = num_attention_heads * attention_head_dim - - self.pos_embed = FluxPosEmbed(theta=10000, axes_dim=axes_dims_rope) - - self.time_text_embed = ChromaCombinedTimestepTextProjEmbeddings( - num_channels=approximator_num_channels // 4, - out_dim=3 * num_single_layers + 2 * 6 * num_layers + 2, - ) - self.distilled_guidance_layer = ChromaApproximator( - in_dim=approximator_num_channels, - out_dim=self.inner_dim, - hidden_dim=approximator_hidden_dim, - n_layers=approximator_layers, - ) - - self.context_embedder = nn.Linear(joint_attention_dim, self.inner_dim) - self.x_embedder = nn.Linear(in_channels, self.inner_dim) - - self.transformer_blocks = nn.ModuleList( - [ - ChromaTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for _ in range(num_layers) - ] - ) - - self.single_transformer_blocks = nn.ModuleList( - [ - ChromaSingleTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for _ in range(num_single_layers) - ] - ) - - self.norm_out = ChromaAdaLayerNormContinuousPruned( - self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6 - ) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) - - self.gradient_checkpointing = False - - @apply_lora_scale("joint_attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_ids: torch.Tensor = None, - txt_ids: torch.Tensor = None, - attention_mask: torch.Tensor = None, - joint_attention_kwargs: dict[str, Any] | None = None, - controlnet_block_samples=None, - controlnet_single_block_samples=None, - return_dict: bool = True, - controlnet_blocks_repeat: bool = False, - ) -> torch.Tensor | Transformer2DModelOutput: - """ - The [`FluxTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, image_sequence_length, in_channels)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, text_sequence_length, joint_attention_dim)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep ( `torch.LongTensor`): - Used to indicate denoising step. - img_ids (`torch.Tensor`): - Image position ids used to compute the rotary positional embeddings. - txt_ids (`torch.Tensor`): - Text position ids used to compute the rotary positional embeddings. - attention_mask (`torch.Tensor`, *optional*): - Mask applied to `encoder_hidden_states` during attention. - controlnet_block_samples (`list` of `torch.Tensor`, *optional*): - A list of tensors that if specified are added to the residuals of transformer blocks. - controlnet_single_block_samples (`list` of `torch.Tensor`, *optional*): - A list of tensors that if specified are added to the residuals of single transformer blocks. - controlnet_blocks_repeat (`bool`, *optional*, defaults to `False`): - Whether to repeat the controlnet block samples across all transformer blocks. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - - hidden_states = self.x_embedder(hidden_states) - - timestep = timestep.to(hidden_states.dtype) * 1000 - - input_vec = self.time_text_embed(timestep) - pooled_temb = self.distilled_guidance_layer(input_vec) - - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - - if txt_ids.ndim == 3: - logger.warning( - "Passing `txt_ids` 3d torch.Tensor is deprecated." - "Please remove the batch dimension and pass it as a 2d torch Tensor" - ) - txt_ids = txt_ids[0] - if img_ids.ndim == 3: - logger.warning( - "Passing `img_ids` 3d torch.Tensor is deprecated." - "Please remove the batch dimension and pass it as a 2d torch Tensor" - ) - img_ids = img_ids[0] - - ids = torch.cat((txt_ids, img_ids), dim=0) - image_rotary_emb = self.pos_embed(ids) - - if joint_attention_kwargs is not None and "ip_adapter_image_embeds" in joint_attention_kwargs: - ip_adapter_image_embeds = joint_attention_kwargs.pop("ip_adapter_image_embeds") - ip_hidden_states = self.encoder_hid_proj(ip_adapter_image_embeds) - joint_attention_kwargs.update({"ip_hidden_states": ip_hidden_states}) - - for index_block, block in enumerate(self.transformer_blocks): - img_offset = 3 * len(self.single_transformer_blocks) - txt_offset = img_offset + 6 * len(self.transformer_blocks) - img_modulation = img_offset + 6 * index_block - text_modulation = txt_offset + 6 * index_block - temb = torch.cat( - ( - pooled_temb[:, img_modulation : img_modulation + 6], - pooled_temb[:, text_modulation : text_modulation + 6], - ), - dim=1, - ) - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, hidden_states, encoder_hidden_states, temb, image_rotary_emb, attention_mask - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - attention_mask=attention_mask, - joint_attention_kwargs=joint_attention_kwargs, - ) - - # controlnet residual - if controlnet_block_samples is not None: - interval_control = len(self.transformer_blocks) / len(controlnet_block_samples) - interval_control = int(np.ceil(interval_control)) - # For Xlabs ControlNet. - if controlnet_blocks_repeat: - hidden_states = ( - hidden_states + controlnet_block_samples[index_block % len(controlnet_block_samples)] - ) - else: - hidden_states = hidden_states + controlnet_block_samples[index_block // interval_control] - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - for index_block, block in enumerate(self.single_transformer_blocks): - start_idx = 3 * index_block - temb = pooled_temb[:, start_idx : start_idx + 3] - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - temb, - image_rotary_emb, - ) - - else: - hidden_states = block( - hidden_states=hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - attention_mask=attention_mask, - joint_attention_kwargs=joint_attention_kwargs, - ) - - # controlnet residual - if controlnet_single_block_samples is not None: - interval_control = len(self.single_transformer_blocks) / len(controlnet_single_block_samples) - interval_control = int(np.ceil(interval_control)) - hidden_states[:, encoder_hidden_states.shape[1] :, ...] = ( - hidden_states[:, encoder_hidden_states.shape[1] :, ...] - + controlnet_single_block_samples[index_block // interval_control] - ) - - hidden_states = hidden_states[:, encoder_hidden_states.shape[1] :, ...] - - temb = pooled_temb[:, -2:] - hidden_states = self.norm_out(hidden_states, temb) - output = self.proj_out(hidden_states) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_chronoedit.py b/diffusers/models/transformers/transformer_chronoedit.py deleted file mode 100644 index b39a18a98afb0227b62a4d29c2aa3e13a4402a07..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_chronoedit.py +++ /dev/null @@ -1,748 +0,0 @@ -# Copyright 2025 The ChronoEdit Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, deprecate, logging -from ...utils.torch_utils import maybe_allow_in_graph -from .._modeling_parallel import ContextParallelInput, ContextParallelOutput -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps, get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import FP32LayerNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# Copied from diffusers.models.transformers.transformer_wan._get_qkv_projections -def _get_qkv_projections(attn: "WanAttention", hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor): - # encoder_hidden_states is only passed for cross-attention - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - if attn.fused_projections: - if not attn.is_cross_attention: - # In self-attention layers, we can fuse the entire QKV projection into a single linear - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - else: - # In cross-attention layers, we can only fuse the KV projections into a single linear - query = attn.to_q(hidden_states) - key, value = attn.to_kv(encoder_hidden_states).chunk(2, dim=-1) - else: - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - return query, key, value - - -# Copied from diffusers.models.transformers.transformer_wan._get_added_kv_projections -def _get_added_kv_projections(attn: "WanAttention", encoder_hidden_states_img: torch.Tensor): - if attn.fused_projections: - key_img, value_img = attn.to_added_kv(encoder_hidden_states_img).chunk(2, dim=-1) - else: - key_img = attn.add_k_proj(encoder_hidden_states_img) - value_img = attn.add_v_proj(encoder_hidden_states_img) - return key_img, value_img - - -# modified from diffusers.models.transformers.transformer_wan.WanAttnProcessor -class WanAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "WanAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to version 2.0 or higher." - ) - - def __call__( - self, - attn: "WanAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - encoder_hidden_states_img = None - if attn.add_k_proj is not None: - # 512 is the context length of the text encoder, hardcoded for now - image_context_length = encoder_hidden_states.shape[1] - 512 - encoder_hidden_states_img = encoder_hidden_states[:, :image_context_length] - encoder_hidden_states = encoder_hidden_states[:, image_context_length:] - - query, key, value = _get_qkv_projections(attn, hidden_states, encoder_hidden_states) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - if rotary_emb is not None: - - def apply_rotary_emb( - hidden_states: torch.Tensor, - freqs_cos: torch.Tensor, - freqs_sin: torch.Tensor, - ): - x1, x2 = hidden_states.unflatten(-1, (-1, 2)).unbind(-1) - cos = freqs_cos[..., 0::2] - sin = freqs_sin[..., 1::2] - out = torch.empty_like(hidden_states) - out[..., 0::2] = x1 * cos - x2 * sin - out[..., 1::2] = x1 * sin + x2 * cos - return out.type_as(hidden_states) - - query = apply_rotary_emb(query, *rotary_emb) - key = apply_rotary_emb(key, *rotary_emb) - - # I2V task - hidden_states_img = None - if encoder_hidden_states_img is not None: - key_img, value_img = _get_added_kv_projections(attn, encoder_hidden_states_img) - key_img = attn.norm_added_k(key_img) - - key_img = key_img.unflatten(2, (attn.heads, -1)) - value_img = value_img.unflatten(2, (attn.heads, -1)) - - hidden_states_img = dispatch_attention_fn( - query, - key_img, - value_img, - attn_mask=None, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - # Reference: https://github.com/huggingface/diffusers/pull/12660 - parallel_config=None, - ) - hidden_states_img = hidden_states_img.flatten(2, 3) - hidden_states_img = hidden_states_img.type_as(query) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - # Reference: https://github.com/huggingface/diffusers/pull/12660 - parallel_config=(self._parallel_config if encoder_hidden_states is None else None), - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - - if hidden_states_img is not None: - hidden_states = hidden_states + hidden_states_img - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -# Copied from diffusers.models.transformers.transformer_wan.WanAttnProcessor2_0 -class WanAttnProcessor2_0: - def __new__(cls, *args, **kwargs): - deprecation_message = ( - "The WanAttnProcessor2_0 class is deprecated and will be removed in a future version. " - "Please use WanAttnProcessor instead. " - ) - deprecate("WanAttnProcessor2_0", "1.0.0", deprecation_message, standard_warn=False) - return WanAttnProcessor(*args, **kwargs) - - -# Copied from diffusers.models.transformers.transformer_wan.WanAttention -class WanAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = WanAttnProcessor - _available_processors = [WanAttnProcessor] - - def __init__( - self, - dim: int, - heads: int = 8, - dim_head: int = 64, - eps: float = 1e-5, - dropout: float = 0.0, - added_kv_proj_dim: int | None = None, - cross_attention_dim_head: int | None = None, - processor=None, - is_cross_attention=None, - ): - super().__init__() - - self.inner_dim = dim_head * heads - self.heads = heads - self.added_kv_proj_dim = added_kv_proj_dim - self.cross_attention_dim_head = cross_attention_dim_head - self.kv_inner_dim = self.inner_dim if cross_attention_dim_head is None else cross_attention_dim_head * heads - - self.to_q = torch.nn.Linear(dim, self.inner_dim, bias=True) - self.to_k = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) - self.to_v = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) - self.to_out = torch.nn.ModuleList( - [ - torch.nn.Linear(self.inner_dim, dim, bias=True), - torch.nn.Dropout(dropout), - ] - ) - self.norm_q = torch.nn.RMSNorm(dim_head * heads, eps=eps, elementwise_affine=True) - self.norm_k = torch.nn.RMSNorm(dim_head * heads, eps=eps, elementwise_affine=True) - - self.add_k_proj = self.add_v_proj = None - if added_kv_proj_dim is not None: - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=True) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=True) - self.norm_added_k = torch.nn.RMSNorm(dim_head * heads, eps=eps) - - if is_cross_attention is not None: - self.is_cross_attention = is_cross_attention - else: - self.is_cross_attention = cross_attention_dim_head is not None - - self.set_processor(processor) - - def fuse_projections(self): - if getattr(self, "fused_projections", False): - return - - if not self.is_cross_attention: - concatenated_weights = torch.cat([self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data]) - concatenated_bias = torch.cat([self.to_q.bias.data, self.to_k.bias.data, self.to_v.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_qkv = nn.Linear(in_features, out_features, bias=True) - self.to_qkv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - else: - concatenated_weights = torch.cat([self.to_k.weight.data, self.to_v.weight.data]) - concatenated_bias = torch.cat([self.to_k.bias.data, self.to_v.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_kv = nn.Linear(in_features, out_features, bias=True) - self.to_kv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - - if self.added_kv_proj_dim is not None: - concatenated_weights = torch.cat([self.add_k_proj.weight.data, self.add_v_proj.weight.data]) - concatenated_bias = torch.cat([self.add_k_proj.bias.data, self.add_v_proj.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_added_kv = nn.Linear(in_features, out_features, bias=True) - self.to_added_kv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - - self.fused_projections = True - - @torch.no_grad() - def unfuse_projections(self): - if not getattr(self, "fused_projections", False): - return - - if hasattr(self, "to_qkv"): - delattr(self, "to_qkv") - if hasattr(self, "to_kv"): - delattr(self, "to_kv") - if hasattr(self, "to_added_kv"): - delattr(self, "to_added_kv") - - self.fused_projections = False - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - **kwargs, - ) -> torch.Tensor: - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, rotary_emb, **kwargs) - - -# Copied from diffusers.models.transformers.transformer_wan.WanImageEmbedding -class WanImageEmbedding(torch.nn.Module): - def __init__(self, in_features: int, out_features: int, pos_embed_seq_len=None): - super().__init__() - - self.norm1 = FP32LayerNorm(in_features) - self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu") - self.norm2 = FP32LayerNorm(out_features) - if pos_embed_seq_len is not None: - self.pos_embed = nn.Parameter(torch.zeros(1, pos_embed_seq_len, in_features)) - else: - self.pos_embed = None - - def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor: - if self.pos_embed is not None: - batch_size, seq_len, embed_dim = encoder_hidden_states_image.shape - encoder_hidden_states_image = encoder_hidden_states_image.view(-1, 2 * seq_len, embed_dim) - encoder_hidden_states_image = encoder_hidden_states_image + self.pos_embed - - hidden_states = self.norm1(encoder_hidden_states_image) - hidden_states = self.ff(hidden_states) - hidden_states = self.norm2(hidden_states) - return hidden_states - - -# Copied from diffusers.models.transformers.transformer_wan.WanTimeTextImageEmbedding -class WanTimeTextImageEmbedding(nn.Module): - def __init__( - self, - dim: int, - time_freq_dim: int, - time_proj_dim: int, - text_embed_dim: int, - image_embed_dim: int | None = None, - pos_embed_seq_len: int | None = None, - ): - super().__init__() - - self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim) - self.act_fn = nn.SiLU() - self.time_proj = nn.Linear(dim, time_proj_dim) - self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh") - - self.image_embedder = None - if image_embed_dim is not None: - self.image_embedder = WanImageEmbedding(image_embed_dim, dim, pos_embed_seq_len=pos_embed_seq_len) - - def forward( - self, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: torch.Tensor | None = None, - timestep_seq_len: int | None = None, - ): - timestep = self.timesteps_proj(timestep) - if timestep_seq_len is not None: - timestep = timestep.unflatten(0, (-1, timestep_seq_len)) - - time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype - if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8: - timestep = timestep.to(time_embedder_dtype) - temb = self.time_embedder(timestep).type_as(encoder_hidden_states) - timestep_proj = self.time_proj(self.act_fn(temb)) - - encoder_hidden_states = self.text_embedder(encoder_hidden_states) - if encoder_hidden_states_image is not None: - encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image) - - return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image - - -class ChronoEditRotaryPosEmbed(nn.Module): - def __init__( - self, - attention_head_dim: int, - patch_size: tuple[int, int, int], - max_seq_len: int, - theta: float = 10000.0, - temporal_skip_len: int = 8, - ): - super().__init__() - - self.attention_head_dim = attention_head_dim - self.patch_size = patch_size - self.max_seq_len = max_seq_len - self.temporal_skip_len = temporal_skip_len - - h_dim = w_dim = 2 * (attention_head_dim // 6) - t_dim = attention_head_dim - h_dim - w_dim - freqs_dtype = torch.float32 if torch.backends.mps.is_available() else torch.float64 - - freqs_cos = [] - freqs_sin = [] - - for dim in [t_dim, h_dim, w_dim]: - freq_cos, freq_sin = get_1d_rotary_pos_embed( - dim, - max_seq_len, - theta, - use_real=True, - repeat_interleave_real=True, - freqs_dtype=freqs_dtype, - ) - freqs_cos.append(freq_cos) - freqs_sin.append(freq_sin) - - self.register_buffer("freqs_cos", torch.cat(freqs_cos, dim=1), persistent=False) - self.register_buffer("freqs_sin", torch.cat(freqs_sin, dim=1), persistent=False) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.patch_size - ppf, pph, ppw = num_frames // p_t, height // p_h, width // p_w - - split_sizes = [ - self.attention_head_dim - 2 * (self.attention_head_dim // 3), - self.attention_head_dim // 3, - self.attention_head_dim // 3, - ] - - freqs_cos = self.freqs_cos.split(split_sizes, dim=1) - freqs_sin = self.freqs_sin.split(split_sizes, dim=1) - - if num_frames == 2: - freqs_cos_f = freqs_cos[0][: self.temporal_skip_len][[0, -1]].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - else: - freqs_cos_f = freqs_cos[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_cos_h = freqs_cos[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_cos_w = freqs_cos[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - - if num_frames == 2: - freqs_sin_f = freqs_sin[0][: self.temporal_skip_len][[0, -1]].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - else: - freqs_sin_f = freqs_sin[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_sin_h = freqs_sin[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_sin_w = freqs_sin[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - - freqs_cos = torch.cat([freqs_cos_f, freqs_cos_h, freqs_cos_w], dim=-1).reshape(1, ppf * pph * ppw, 1, -1) - freqs_sin = torch.cat([freqs_sin_f, freqs_sin_h, freqs_sin_w], dim=-1).reshape(1, ppf * pph * ppw, 1, -1) - - return freqs_cos, freqs_sin - - -@maybe_allow_in_graph -# Copied from diffusers.models.transformers.transformer_wan.WanTransformerBlock -class WanTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - ffn_dim: int, - num_heads: int, - qk_norm: str = "rms_norm_across_heads", - cross_attn_norm: bool = False, - eps: float = 1e-6, - added_kv_proj_dim: int | None = None, - ): - super().__init__() - - # 1. Self-attention - self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False) - self.attn1 = WanAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - cross_attention_dim_head=None, - processor=WanAttnProcessor(), - ) - - # 2. Cross-attention - self.attn2 = WanAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - added_kv_proj_dim=added_kv_proj_dim, - cross_attention_dim_head=dim // num_heads, - processor=WanAttnProcessor(), - ) - self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity() - - # 3. Feed-forward - self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate") - self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False) - - self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - rotary_emb: torch.Tensor, - ) -> torch.Tensor: - if temb.ndim == 4: - # temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v) - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table.unsqueeze(0) + temb.float() - ).chunk(6, dim=2) - # batch_size, seq_len, 1, inner_dim - shift_msa = shift_msa.squeeze(2) - scale_msa = scale_msa.squeeze(2) - gate_msa = gate_msa.squeeze(2) - c_shift_msa = c_shift_msa.squeeze(2) - c_scale_msa = c_scale_msa.squeeze(2) - c_gate_msa = c_gate_msa.squeeze(2) - else: - # temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B) - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table + temb.float() - ).chunk(6, dim=1) - - # 1. Self-attention - norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states) - attn_output = self.attn1(norm_hidden_states, None, None, rotary_emb) - hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states) - - # 2. Cross-attention - norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states) - attn_output = self.attn2(norm_hidden_states, encoder_hidden_states, None, None) - hidden_states = hidden_states + attn_output - - # 3. Feed-forward - norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as( - hidden_states - ) - ff_output = self.ffn(norm_hidden_states) - hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states) - - return hidden_states - - -# modified from diffusers.models.transformers.transformer_wan.WanTransformer3DModel -class ChronoEditTransformer3DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin, AttentionMixin -): - r""" - A Transformer model for video-like data used in the ChronoEdit model. - - Args: - patch_size (`tuple[int]`, defaults to `(1, 2, 2)`): - 3D patch dimensions for video embedding (t_patch, h_patch, w_patch). - num_attention_heads (`int`, defaults to `40`): - Fixed length for text embeddings. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - text_dim (`int`, defaults to `512`): - Input dimension for text embeddings. - freq_dim (`int`, defaults to `256`): - Dimension for sinusoidal time embeddings. - ffn_dim (`int`, defaults to `13824`): - Intermediate dimension in feed-forward network. - num_layers (`int`, defaults to `40`): - The number of layers of transformer blocks to use. - window_size (`tuple[int]`, defaults to `(-1, -1)`): - Window size for local attention (-1 indicates global attention). - cross_attn_norm (`bool`, defaults to `True`): - Enable cross-attention normalization. - qk_norm (`bool`, defaults to `True`): - Enable query/key normalization. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - add_img_emb (`bool`, defaults to `False`): - Whether to use img_emb. - added_kv_proj_dim (`int`, *optional*, defaults to `None`): - The number of channels to use for the added key and value projections. If `None`, no projection is used. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"] - _no_split_modules = ["WanTransformerBlock"] - _keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"] - _keys_to_ignore_on_load_unexpected = ["norm_added_q"] - _repeated_blocks = ["WanTransformerBlock"] - _cp_plan = { - "rope": { - 0: ContextParallelInput(split_dim=1, expected_dims=4, split_output=True), - 1: ContextParallelInput(split_dim=1, expected_dims=4, split_output=True), - }, - "blocks.0": { - "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - }, - # Reference: https://github.com/huggingface/diffusers/pull/12660 - # We need to disable the splitting of encoder_hidden_states because - # the image_encoder consistently generates 257 tokens for image_embed. This causes - # the shape of encoder_hidden_states—whose token count is always 769 (512 + 257) - # after concatenation—to be indivisible by the number of devices in the CP. - "proj_out": ContextParallelOutput(gather_dim=1, expected_dims=3), - } - - @register_to_config - def __init__( - self, - patch_size: tuple[int] = (1, 2, 2), - num_attention_heads: int = 40, - attention_head_dim: int = 128, - in_channels: int = 16, - out_channels: int = 16, - text_dim: int = 4096, - freq_dim: int = 256, - ffn_dim: int = 13824, - num_layers: int = 40, - cross_attn_norm: bool = True, - qk_norm: str | None = "rms_norm_across_heads", - eps: float = 1e-6, - image_dim: int | None = None, - added_kv_proj_dim: int | None = None, - rope_max_seq_len: int = 1024, - pos_embed_seq_len: int | None = None, - rope_temporal_skip_len: int = 8, - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - # 1. Patch & position embedding - self.rope = ChronoEditRotaryPosEmbed( - attention_head_dim, patch_size, rope_max_seq_len, temporal_skip_len=rope_temporal_skip_len - ) - self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - - # 2. Condition embeddings - # image_embedding_dim=1280 for I2V model - self.condition_embedder = WanTimeTextImageEmbedding( - dim=inner_dim, - time_freq_dim=freq_dim, - time_proj_dim=inner_dim * 6, - text_embed_dim=text_dim, - image_embed_dim=image_dim, - pos_embed_seq_len=pos_embed_seq_len, - ) - - # 3. Transformer blocks - self.blocks = nn.ModuleList( - [ - WanTransformerBlock( - inner_dim, ffn_dim, num_attention_heads, qk_norm, cross_attn_norm, eps, added_kv_proj_dim - ) - for _ in range(num_layers) - ] - ) - - # 4. Output norm & projection - self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size)) - self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: torch.Tensor | None = None, - return_dict: bool = True, - attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor | dict[str, torch.Tensor]: - """ - The [`ChronoEditTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_hidden_states_image (`torch.Tensor`, *optional*): - Conditional image embeddings for image-conditioned generation. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.config.patch_size - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p_h - post_patch_width = width // p_w - - rotary_emb = self.rope(hidden_states) - - hidden_states = self.patch_embedding(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) - - # timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v) - if timestep.ndim == 2: - ts_seq_len = timestep.shape[1] - timestep = timestep.flatten() # batch_size * seq_len - else: - ts_seq_len = None - - temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder( - timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len - ) - if ts_seq_len is not None: - # batch_size, seq_len, 6, inner_dim - timestep_proj = timestep_proj.unflatten(2, (6, -1)) - else: - # batch_size, 6, inner_dim - timestep_proj = timestep_proj.unflatten(1, (6, -1)) - - if encoder_hidden_states_image is not None: - encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) - - # 4. Transformer blocks - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.blocks: - hidden_states = self._gradient_checkpointing_func( - block, hidden_states, encoder_hidden_states, timestep_proj, rotary_emb - ) - else: - for block in self.blocks: - hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb) - - # 5. Output norm, projection & unpatchify - if temb.ndim == 3: - # batch_size, seq_len, inner_dim (wan 2.2 ti2v) - shift, scale = (self.scale_shift_table.unsqueeze(0).to(temb.device) + temb.unsqueeze(2)).chunk(2, dim=2) - shift = shift.squeeze(2) - scale = scale.squeeze(2) - else: - # batch_size, inner_dim - shift, scale = (self.scale_shift_table.to(temb.device) + temb.unsqueeze(1)).chunk(2, dim=1) - - # Move the shift and scale tensors to the same device as hidden_states. - # When using multi-GPU inference via accelerate these will be on the - # first device rather than the last device, which hidden_states ends up - # on. - shift = shift.to(hidden_states.device) - scale = scale.to(hidden_states.device) - - hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states) - hidden_states = self.proj_out(hidden_states) - - hidden_states = hidden_states.reshape( - batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1 - ) - hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) - output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_cogview3plus.py b/diffusers/models/transformers/transformer_cogview3plus.py deleted file mode 100644 index ad6a442acbcc1ef4c92fc8593d2d0f547b0e3f14..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_cogview3plus.py +++ /dev/null @@ -1,308 +0,0 @@ -# Copyright 2025 The CogView team, Tsinghua University & ZhipuAI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ..attention import AttentionMixin, FeedForward -from ..attention_processor import Attention, CogVideoXAttnProcessor2_0 -from ..embeddings import CogView3CombinedTimestepSizeEmbeddings, CogView3PlusPatchEmbed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, CogView3PlusAdaLayerNormZeroTextImage - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class CogView3PlusTransformerBlock(nn.Module): - r""" - Transformer block used in [CogView](https://github.com/THUDM/CogView3) model. - - Args: - dim (`int`): - The number of channels in the input and output. - num_attention_heads (`int`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`): - The number of channels in each head. - time_embed_dim (`int`): - The number of channels in timestep embedding. - """ - - def __init__( - self, - dim: int = 2560, - num_attention_heads: int = 64, - attention_head_dim: int = 40, - time_embed_dim: int = 512, - ): - super().__init__() - - self.norm1 = CogView3PlusAdaLayerNormZeroTextImage(embedding_dim=time_embed_dim, dim=dim) - - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - out_dim=dim, - bias=True, - qk_norm="layer_norm", - elementwise_affine=False, - eps=1e-6, - processor=CogVideoXAttnProcessor2_0(), - ) - - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - - self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - emb: torch.Tensor, - ) -> tuple[torch.Tensor, torch.Tensor]: - text_seq_length = encoder_hidden_states.size(1) - - # norm & modulate - ( - norm_hidden_states, - gate_msa, - shift_mlp, - scale_mlp, - gate_mlp, - norm_encoder_hidden_states, - c_gate_msa, - c_shift_mlp, - c_scale_mlp, - c_gate_mlp, - ) = self.norm1(hidden_states, encoder_hidden_states, emb) - - # attention - attn_hidden_states, attn_encoder_hidden_states = self.attn1( - hidden_states=norm_hidden_states, encoder_hidden_states=norm_encoder_hidden_states - ) - - hidden_states = hidden_states + gate_msa.unsqueeze(1) * attn_hidden_states - encoder_hidden_states = encoder_hidden_states + c_gate_msa.unsqueeze(1) * attn_encoder_hidden_states - - # norm & modulate - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - # feed-forward - norm_hidden_states = torch.cat([norm_encoder_hidden_states, norm_hidden_states], dim=1) - ff_output = self.ff(norm_hidden_states) - - hidden_states = hidden_states + gate_mlp.unsqueeze(1) * ff_output[:, text_seq_length:] - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * ff_output[:, :text_seq_length] - - if hidden_states.dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - if encoder_hidden_states.dtype == torch.float16: - encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) - return hidden_states, encoder_hidden_states - - -class CogView3PlusTransformer2DModel(ModelMixin, AttentionMixin, ConfigMixin): - r""" - The Transformer model introduced in [CogView3: Finer and Faster Text-to-Image Generation via Relay - Diffusion](https://huggingface.co/papers/2403.05121). - - Args: - patch_size (`int`, defaults to `2`): - The size of the patches to use in the patch embedding layer. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - num_layers (`int`, defaults to `30`): - The number of layers of Transformer blocks to use. - attention_head_dim (`int`, defaults to `40`): - The number of channels in each head. - num_attention_heads (`int`, defaults to `64`): - The number of heads to use for multi-head attention. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - text_embed_dim (`int`, defaults to `4096`): - Input dimension of text embeddings from the text encoder. - time_embed_dim (`int`, defaults to `512`): - Output dimension of timestep embeddings. - condition_dim (`int`, defaults to `256`): - The embedding dimension of the input SDXL-style resolution conditions (original_size, target_size, - crop_coords). - pos_embed_max_size (`int`, defaults to `128`): - The maximum resolution of the positional embeddings, from which slices of shape `H x W` are taken and added - to input patched latents, where `H` and `W` are the latent height and width respectively. A value of 128 - means that the maximum supported height and width for image generation is `128 * vae_scale_factor * - patch_size => 128 * 8 * 2 => 2048`. - sample_size (`int`, defaults to `128`): - The base resolution of input latents. If height/width is not provided during generation, this value is used - to determine the resolution as `sample_size * vae_scale_factor => 128 * 8 => 1024` - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["patch_embed", "norm"] - _no_split_modules = ["CogView3PlusTransformerBlock", "CogView3PlusPatchEmbed"] - - @register_to_config - def __init__( - self, - patch_size: int = 2, - in_channels: int = 16, - num_layers: int = 30, - attention_head_dim: int = 40, - num_attention_heads: int = 64, - out_channels: int = 16, - text_embed_dim: int = 4096, - time_embed_dim: int = 512, - condition_dim: int = 256, - pos_embed_max_size: int = 128, - sample_size: int = 128, - ): - super().__init__() - self.out_channels = out_channels - self.inner_dim = num_attention_heads * attention_head_dim - - # CogView3 uses 3 additional SDXL-like conditions - original_size, target_size, crop_coords - # Each of these are sincos embeddings of shape 2 * condition_dim - self.pooled_projection_dim = 3 * 2 * condition_dim - - self.patch_embed = CogView3PlusPatchEmbed( - in_channels=in_channels, - hidden_size=self.inner_dim, - patch_size=patch_size, - text_hidden_size=text_embed_dim, - pos_embed_max_size=pos_embed_max_size, - ) - - self.time_condition_embed = CogView3CombinedTimestepSizeEmbeddings( - embedding_dim=time_embed_dim, - condition_dim=condition_dim, - pooled_projection_dim=self.pooled_projection_dim, - timesteps_dim=self.inner_dim, - ) - - self.transformer_blocks = nn.ModuleList( - [ - CogView3PlusTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - time_embed_dim=time_embed_dim, - ) - for _ in range(num_layers) - ] - ) - - self.norm_out = AdaLayerNormContinuous( - embedding_dim=self.inner_dim, - conditioning_embedding_dim=time_embed_dim, - elementwise_affine=False, - eps=1e-6, - ) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: torch.LongTensor, - original_size: torch.Tensor, - target_size: torch.Tensor, - crop_coords: torch.Tensor, - return_dict: bool = True, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`CogView3PlusTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor`): - Input `hidden_states` of shape `(batch size, channel, height, width)`. - encoder_hidden_states (`torch.Tensor`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) of shape - `(batch_size, sequence_len, text_embed_dim)` - timestep (`torch.LongTensor`): - Used to indicate denoising step. - original_size (`torch.Tensor`): - CogView3 uses SDXL-like micro-conditioning for original image size as explained in section 2.2 of - [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). - target_size (`torch.Tensor`): - CogView3 uses SDXL-like micro-conditioning for target image size as explained in section 2.2 of - [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). - crop_coords (`torch.Tensor`): - CogView3 uses SDXL-like micro-conditioning for crop coordinates as explained in section 2.2 of - [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - `torch.Tensor` or [`~models.transformer_2d.Transformer2DModelOutput`]: - The denoised latents using provided inputs as conditioning. - """ - height, width = hidden_states.shape[-2:] - text_seq_length = encoder_hidden_states.shape[1] - - hidden_states = self.patch_embed( - hidden_states, encoder_hidden_states - ) # takes care of adding positional embeddings too. - emb = self.time_condition_embed(timestep, original_size, target_size, crop_coords, hidden_states.dtype) - - encoder_hidden_states = hidden_states[:, :text_seq_length] - hidden_states = hidden_states[:, text_seq_length:] - - for index_block, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - emb, - ) - else: - hidden_states, encoder_hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - emb=emb, - ) - - hidden_states = self.norm_out(hidden_states, emb) - hidden_states = self.proj_out(hidden_states) # (batch_size, height*width, patch_size*patch_size*out_channels) - - # unpatchify - patch_size = self.config.patch_size - height = height // patch_size - width = width // patch_size - - hidden_states = hidden_states.reshape( - shape=(hidden_states.shape[0], height, width, self.out_channels, patch_size, patch_size) - ) - hidden_states = torch.einsum("nhwcpq->nchpwq", hidden_states) - output = hidden_states.reshape( - shape=(hidden_states.shape[0], self.out_channels, height * patch_size, width * patch_size) - ) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_cogview4.py b/diffusers/models/transformers/transformer_cogview4.py deleted file mode 100644 index 2856fffd2a630879e418655c2e5e206aeb9d15a7..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_cogview4.py +++ /dev/null @@ -1,796 +0,0 @@ -# Copyright 2025 The CogView team, Tsinghua University & ZhipuAI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import FeedForward -from ..attention_processor import Attention -from ..cache_utils import CacheMixin -from ..embeddings import CogView3CombinedTimestepSizeEmbeddings -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import LayerNorm, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class CogView4PatchEmbed(nn.Module): - def __init__( - self, - in_channels: int = 16, - hidden_size: int = 2560, - patch_size: int = 2, - text_hidden_size: int = 4096, - ): - super().__init__() - self.patch_size = patch_size - - self.proj = nn.Linear(in_channels * patch_size**2, hidden_size) - self.text_proj = nn.Linear(text_hidden_size, hidden_size) - - def forward(self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, channel, height, width = hidden_states.shape - post_patch_height = height // self.patch_size - post_patch_width = width // self.patch_size - - hidden_states = hidden_states.reshape( - batch_size, channel, post_patch_height, self.patch_size, post_patch_width, self.patch_size - ) - hidden_states = hidden_states.permute(0, 2, 4, 1, 3, 5).flatten(3, 5).flatten(1, 2) - hidden_states = self.proj(hidden_states) - encoder_hidden_states = self.text_proj(encoder_hidden_states) - - return hidden_states, encoder_hidden_states - - -class CogView4AdaLayerNormZero(nn.Module): - def __init__(self, embedding_dim: int, dim: int) -> None: - super().__init__() - - self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - self.norm_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - self.linear = nn.Linear(embedding_dim, 12 * dim, bias=True) - - def forward( - self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor]: - dtype = hidden_states.dtype - norm_hidden_states = self.norm(hidden_states).to(dtype=dtype) - norm_encoder_hidden_states = self.norm_context(encoder_hidden_states).to(dtype=dtype) - - emb = self.linear(temb) - ( - shift_msa, - c_shift_msa, - scale_msa, - c_scale_msa, - gate_msa, - c_gate_msa, - shift_mlp, - c_shift_mlp, - scale_mlp, - c_scale_mlp, - gate_mlp, - c_gate_mlp, - ) = emb.chunk(12, dim=1) - - hidden_states = norm_hidden_states * (1 + scale_msa.unsqueeze(1)) + shift_msa.unsqueeze(1) - encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_msa.unsqueeze(1)) + c_shift_msa.unsqueeze(1) - - return ( - hidden_states, - gate_msa, - shift_mlp, - scale_mlp, - gate_mlp, - encoder_hidden_states, - c_gate_msa, - c_shift_mlp, - c_scale_mlp, - c_gate_mlp, - ) - - -class CogView4AttnProcessor: - """ - Processor for implementing scaled dot-product attention for the CogView4 model. It applies a rotary embedding on - query and key vectors, but does not include spatial normalization. - - The processor supports passing an attention mask for text tokens. The attention mask should have shape (batch_size, - text_seq_length) where 1 indicates a non-padded token and 0 indicates a padded token. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("CogView4AttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - dtype = encoder_hidden_states.dtype - - batch_size, text_seq_length, embed_dim = encoder_hidden_states.shape - batch_size, image_seq_length, embed_dim = hidden_states.shape - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - # 1. QKV projections - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2) - key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2) - value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2) - - # 2. QK normalization - if attn.norm_q is not None: - query = attn.norm_q(query).to(dtype=dtype) - if attn.norm_k is not None: - key = attn.norm_k(key).to(dtype=dtype) - - # 3. Rotational positional embeddings applied to latent stream - if image_rotary_emb is not None: - from ..embeddings import apply_rotary_emb - - query[:, :, text_seq_length:, :] = apply_rotary_emb( - query[:, :, text_seq_length:, :], image_rotary_emb, use_real_unbind_dim=-2 - ) - key[:, :, text_seq_length:, :] = apply_rotary_emb( - key[:, :, text_seq_length:, :], image_rotary_emb, use_real_unbind_dim=-2 - ) - - # 4. Attention - if attention_mask is not None: - text_attn_mask = attention_mask - assert text_attn_mask.dim() == 2, "the shape of text_attn_mask should be (batch_size, text_seq_length)" - text_attn_mask = text_attn_mask.float().to(query.device) - mix_attn_mask = torch.ones((batch_size, text_seq_length + image_seq_length), device=query.device) - mix_attn_mask[:, :text_seq_length] = text_attn_mask - mix_attn_mask = mix_attn_mask.unsqueeze(2) - attn_mask_matrix = mix_attn_mask @ mix_attn_mask.transpose(1, 2) - attention_mask = (attn_mask_matrix > 0).unsqueeze(1).to(query.dtype) - - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - hidden_states = hidden_states.transpose(1, 2).flatten(2, 3) - hidden_states = hidden_states.type_as(query) - - # 5. Output projection - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - encoder_hidden_states, hidden_states = hidden_states.split( - [text_seq_length, hidden_states.size(1) - text_seq_length], dim=1 - ) - return hidden_states, encoder_hidden_states - - -class CogView4TrainingAttnProcessor: - """ - Training Processor for implementing scaled dot-product attention for the CogView4 model. It applies a rotary - embedding on query and key vectors, but does not include spatial normalization. - - This processor differs from CogView4AttnProcessor in several important ways: - 1. It supports attention masking with variable sequence lengths for multi-resolution training - 2. It unpacks and repacks sequences for efficient training with variable sequence lengths when batch_flag is - provided - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("CogView4AttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - latent_attn_mask: torch.Tensor | None = None, - text_attn_mask: torch.Tensor | None = None, - batch_flag: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]] | None = None, - **kwargs, - ) -> tuple[torch.Tensor, torch.Tensor]: - """ - Args: - attn (`Attention`): - The attention module. - hidden_states (`torch.Tensor`): - The input hidden states. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states for cross-attention. - latent_attn_mask (`torch.Tensor`, *optional*): - Mask for latent tokens where 0 indicates pad token and 1 indicates non-pad token. If None, full - attention is used for all latent tokens. Note: the shape of latent_attn_mask is (batch_size, - num_latent_tokens). - text_attn_mask (`torch.Tensor`, *optional*): - Mask for text tokens where 0 indicates pad token and 1 indicates non-pad token. If None, full attention - is used for all text tokens. - batch_flag (`torch.Tensor`, *optional*): - Values from 0 to n-1 indicating which samples belong to the same batch. Samples with the same - batch_flag are packed together. Example: [0, 1, 1, 2, 2] means sample 0 forms batch0, samples 1-2 form - batch1, and samples 3-4 form batch2. If None, no packing is used. - image_rotary_emb (`tuple[torch.Tensor, torch.Tensor]` or `list[tuple[torch.Tensor, torch.Tensor]]`, *optional*): - The rotary embedding for the image part of the input. - Returns: - `tuple[torch.Tensor, torch.Tensor]`: The processed hidden states for both image and text streams. - """ - - # Get dimensions and device info - batch_size, text_seq_length, embed_dim = encoder_hidden_states.shape - batch_size, image_seq_length, embed_dim = hidden_states.shape - dtype = encoder_hidden_states.dtype - device = encoder_hidden_states.device - latent_hidden_states = hidden_states - # Combine text and image streams for joint processing - mixed_hidden_states = torch.cat([encoder_hidden_states, latent_hidden_states], dim=1) - - # 1. Construct attention mask and maybe packing input - # Create default masks if not provided - if text_attn_mask is None: - text_attn_mask = torch.ones((batch_size, text_seq_length), dtype=torch.int32, device=device) - if latent_attn_mask is None: - latent_attn_mask = torch.ones((batch_size, image_seq_length), dtype=torch.int32, device=device) - - # Validate mask shapes and types - assert text_attn_mask.dim() == 2, "the shape of text_attn_mask should be (batch_size, text_seq_length)" - assert text_attn_mask.dtype == torch.int32, "the dtype of text_attn_mask should be torch.int32" - assert latent_attn_mask.dim() == 2, "the shape of latent_attn_mask should be (batch_size, num_latent_tokens)" - assert latent_attn_mask.dtype == torch.int32, "the dtype of latent_attn_mask should be torch.int32" - - # Create combined mask for text and image tokens - mixed_attn_mask = torch.ones( - (batch_size, text_seq_length + image_seq_length), dtype=torch.int32, device=device - ) - mixed_attn_mask[:, :text_seq_length] = text_attn_mask - mixed_attn_mask[:, text_seq_length:] = latent_attn_mask - - # Convert mask to attention matrix format (where 1 means attend, 0 means don't attend) - mixed_attn_mask_input = mixed_attn_mask.unsqueeze(2).to(dtype=dtype) - attn_mask_matrix = mixed_attn_mask_input @ mixed_attn_mask_input.transpose(1, 2) - - # Handle batch packing if enabled - if batch_flag is not None: - assert batch_flag.dim() == 1 - # Determine packed batch size based on batch_flag - packing_batch_size = torch.max(batch_flag).item() + 1 - - # Calculate actual sequence lengths for each sample based on masks - text_seq_length = torch.sum(text_attn_mask, dim=1) - latent_seq_length = torch.sum(latent_attn_mask, dim=1) - mixed_seq_length = text_seq_length + latent_seq_length - - # Calculate packed sequence lengths for each packed batch - mixed_seq_length_packed = [ - torch.sum(mixed_attn_mask[batch_flag == batch_idx]).item() for batch_idx in range(packing_batch_size) - ] - - assert len(mixed_seq_length_packed) == packing_batch_size - - # Pack sequences by removing padding tokens - mixed_attn_mask_flatten = mixed_attn_mask.flatten(0, 1) - mixed_hidden_states_flatten = mixed_hidden_states.flatten(0, 1) - mixed_hidden_states_unpad = mixed_hidden_states_flatten[mixed_attn_mask_flatten == 1] - assert torch.sum(mixed_seq_length) == mixed_hidden_states_unpad.shape[0] - - # Split the unpadded sequence into packed batches - mixed_hidden_states_packed = torch.split(mixed_hidden_states_unpad, mixed_seq_length_packed) - - # Re-pad to create packed batches with right-side padding - mixed_hidden_states_packed_padded = torch.nn.utils.rnn.pad_sequence( - mixed_hidden_states_packed, - batch_first=True, - padding_value=0.0, - padding_side="right", - ) - - # Create attention mask for packed batches - l = mixed_hidden_states_packed_padded.shape[1] - attn_mask_matrix = torch.zeros( - (packing_batch_size, l, l), - dtype=dtype, - device=device, - ) - - # Fill attention mask with block diagonal matrices - # This ensures that tokens can only attend to other tokens within the same original sample - for idx, mask in enumerate(attn_mask_matrix): - seq_lengths = mixed_seq_length[batch_flag == idx] - offset = 0 - for length in seq_lengths: - # Create a block of 1s for each sample in the packed batch - mask[offset : offset + length, offset : offset + length] = 1 - offset += length - - attn_mask_matrix = attn_mask_matrix.to(dtype=torch.bool) - attn_mask_matrix = attn_mask_matrix.unsqueeze(1) # Add attention head dim - attention_mask = attn_mask_matrix - - # Prepare hidden states for attention computation - if batch_flag is None: - # If no packing, just combine text and image tokens - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - else: - # If packing, use the packed sequence - hidden_states = mixed_hidden_states_packed_padded - - # 2. QKV projections - convert hidden states to query, key, value - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - # Reshape for multi-head attention: [batch, seq_len, heads*dim] -> [batch, heads, seq_len, dim] - query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2) - key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2) - value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2) - - # 3. QK normalization - apply layer norm to queries and keys if configured - if attn.norm_q is not None: - query = attn.norm_q(query).to(dtype=dtype) - if attn.norm_k is not None: - key = attn.norm_k(key).to(dtype=dtype) - - # 4. Apply rotary positional embeddings to image tokens only - if image_rotary_emb is not None: - from ..embeddings import apply_rotary_emb - - if batch_flag is None: - # Apply RoPE only to image tokens (after text tokens) - query[:, :, text_seq_length:, :] = apply_rotary_emb( - query[:, :, text_seq_length:, :], image_rotary_emb, use_real_unbind_dim=-2 - ) - key[:, :, text_seq_length:, :] = apply_rotary_emb( - key[:, :, text_seq_length:, :], image_rotary_emb, use_real_unbind_dim=-2 - ) - else: - # For packed batches, need to carefully apply RoPE to appropriate tokens - assert query.shape[0] == packing_batch_size - assert key.shape[0] == packing_batch_size - assert len(image_rotary_emb) == batch_size - - rope_idx = 0 - for idx in range(packing_batch_size): - offset = 0 - # Get text and image sequence lengths for samples in this packed batch - text_seq_length_bi = text_seq_length[batch_flag == idx] - latent_seq_length_bi = latent_seq_length[batch_flag == idx] - - # Apply RoPE to each image segment in the packed sequence - for tlen, llen in zip(text_seq_length_bi, latent_seq_length_bi): - mlen = tlen + llen - # Apply RoPE only to image tokens (after text tokens) - query[idx, :, offset + tlen : offset + mlen, :] = apply_rotary_emb( - query[idx, :, offset + tlen : offset + mlen, :], - image_rotary_emb[rope_idx], - use_real_unbind_dim=-2, - ) - key[idx, :, offset + tlen : offset + mlen, :] = apply_rotary_emb( - key[idx, :, offset + tlen : offset + mlen, :], - image_rotary_emb[rope_idx], - use_real_unbind_dim=-2, - ) - offset += mlen - rope_idx += 1 - - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - # Reshape back: [batch, heads, seq_len, dim] -> [batch, seq_len, heads*dim] - hidden_states = hidden_states.transpose(1, 2).flatten(2, 3) - hidden_states = hidden_states.type_as(query) - - # 5. Output projection - project attention output to model dimension - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - # Split the output back into text and image streams - if batch_flag is None: - # Simple split for non-packed case - encoder_hidden_states, hidden_states = hidden_states.split( - [text_seq_length, hidden_states.size(1) - text_seq_length], dim=1 - ) - else: - # For packed case: need to unpack, split text/image, then restore to original shapes - # First, unpad the sequence based on the packed sequence lengths - hidden_states_unpad = torch.nn.utils.rnn.unpad_sequence( - hidden_states, - lengths=torch.tensor(mixed_seq_length_packed), - batch_first=True, - ) - # Concatenate all unpadded sequences - hidden_states_flatten = torch.cat(hidden_states_unpad, dim=0) - # Split by original sample sequence lengths - hidden_states_unpack = torch.split(hidden_states_flatten, mixed_seq_length.tolist()) - assert len(hidden_states_unpack) == batch_size - - # Further split each sample's sequence into text and image parts - hidden_states_unpack = [ - torch.split(h, [tlen, llen]) - for h, tlen, llen in zip(hidden_states_unpack, text_seq_length, latent_seq_length) - ] - # Separate text and image sequences - encoder_hidden_states_unpad = [h[0] for h in hidden_states_unpack] - hidden_states_unpad = [h[1] for h in hidden_states_unpack] - - # Update the original tensors with the processed values, respecting the attention masks - for idx in range(batch_size): - # Place unpacked text tokens back in the encoder_hidden_states tensor - encoder_hidden_states[idx][text_attn_mask[idx] == 1] = encoder_hidden_states_unpad[idx] - # Place unpacked image tokens back in the latent_hidden_states tensor - latent_hidden_states[idx][latent_attn_mask[idx] == 1] = hidden_states_unpad[idx] - - # Update the output hidden states - hidden_states = latent_hidden_states - - return hidden_states, encoder_hidden_states - - -@maybe_allow_in_graph -class CogView4TransformerBlock(nn.Module): - def __init__( - self, - dim: int = 2560, - num_attention_heads: int = 64, - attention_head_dim: int = 40, - time_embed_dim: int = 512, - ) -> None: - super().__init__() - - # 1. Attention - self.norm1 = CogView4AdaLayerNormZero(time_embed_dim, dim) - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - out_dim=dim, - bias=True, - qk_norm="layer_norm", - elementwise_affine=False, - eps=1e-5, - processor=CogView4AttnProcessor(), - ) - - # 2. Feedforward - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]] | None = None, - attention_mask: dict[str, torch.Tensor] | None = None, - attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - # 1. Timestep conditioning - ( - norm_hidden_states, - gate_msa, - shift_mlp, - scale_mlp, - gate_mlp, - norm_encoder_hidden_states, - c_gate_msa, - c_shift_mlp, - c_scale_mlp, - c_gate_mlp, - ) = self.norm1(hidden_states, encoder_hidden_states, temb) - - # 2. Attention - if attention_kwargs is None: - attention_kwargs = {} - attn_hidden_states, attn_encoder_hidden_states = self.attn1( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - attention_mask=attention_mask, - **attention_kwargs, - ) - hidden_states = hidden_states + attn_hidden_states * gate_msa.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + attn_encoder_hidden_states * c_gate_msa.unsqueeze(1) - - # 3. Feedforward - norm_hidden_states = self.norm2(hidden_states) * (1 + scale_mlp.unsqueeze(1)) + shift_mlp.unsqueeze(1) - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) * ( - 1 + c_scale_mlp.unsqueeze(1) - ) + c_shift_mlp.unsqueeze(1) - - ff_output = self.ff(norm_hidden_states) - ff_output_context = self.ff(norm_encoder_hidden_states) - hidden_states = hidden_states + ff_output * gate_mlp.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + ff_output_context * c_gate_mlp.unsqueeze(1) - - return hidden_states, encoder_hidden_states - - -class CogView4RotaryPosEmbed(nn.Module): - def __init__(self, dim: int, patch_size: int, rope_axes_dim: tuple[int, int], theta: float = 10000.0) -> None: - super().__init__() - - self.dim = dim - self.patch_size = patch_size - self.rope_axes_dim = rope_axes_dim - self.theta = theta - - def forward(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - batch_size, num_channels, height, width = hidden_states.shape - height, width = height // self.patch_size, width // self.patch_size - - dim_h, dim_w = self.dim // 2, self.dim // 2 - h_inv_freq = 1.0 / ( - self.theta ** (torch.arange(0, dim_h, 2, dtype=torch.float32)[: (dim_h // 2)].float() / dim_h) - ) - w_inv_freq = 1.0 / ( - self.theta ** (torch.arange(0, dim_w, 2, dtype=torch.float32)[: (dim_w // 2)].float() / dim_w) - ) - h_seq = torch.arange(self.rope_axes_dim[0]) - w_seq = torch.arange(self.rope_axes_dim[1]) - freqs_h = torch.outer(h_seq, h_inv_freq) - freqs_w = torch.outer(w_seq, w_inv_freq) - - h_idx = torch.arange(height, device=freqs_h.device) - w_idx = torch.arange(width, device=freqs_w.device) - inner_h_idx = h_idx * self.rope_axes_dim[0] // height - inner_w_idx = w_idx * self.rope_axes_dim[1] // width - - freqs_h = freqs_h[inner_h_idx] - freqs_w = freqs_w[inner_w_idx] - - # Create position matrices for height and width - # [height, 1, dim//4] and [1, width, dim//4] - freqs_h = freqs_h.unsqueeze(1) - freqs_w = freqs_w.unsqueeze(0) - # Broadcast freqs_h and freqs_w to [height, width, dim//4] - freqs_h = freqs_h.expand(height, width, -1) - freqs_w = freqs_w.expand(height, width, -1) - - # Concatenate along last dimension to get [height, width, dim//2] - freqs = torch.cat([freqs_h, freqs_w], dim=-1) - freqs = torch.cat([freqs, freqs], dim=-1) # [height, width, dim] - freqs = freqs.reshape(height * width, -1) - return (freqs.cos(), freqs.sin()) - - -class CogView4AdaLayerNormContinuous(nn.Module): - """ - CogView4-only final AdaLN: LN(x) -> Linear(cond) -> chunk -> affine. Matches Megatron: **no activation** before the - Linear on conditioning embedding. - """ - - def __init__( - self, - embedding_dim: int, - conditioning_embedding_dim: int, - elementwise_affine: bool = True, - eps: float = 1e-5, - bias: bool = True, - norm_type: str = "layer_norm", - ): - super().__init__() - self.linear = nn.Linear(conditioning_embedding_dim, embedding_dim * 2, bias=bias) - if norm_type == "layer_norm": - self.norm = LayerNorm(embedding_dim, eps, elementwise_affine, bias) - elif norm_type == "rms_norm": - self.norm = RMSNorm(embedding_dim, eps, elementwise_affine) - else: - raise ValueError(f"unknown norm_type {norm_type}") - - def forward(self, x: torch.Tensor, conditioning_embedding: torch.Tensor) -> torch.Tensor: - # *** NO SiLU here *** - emb = self.linear(conditioning_embedding.to(x.dtype)) - scale, shift = torch.chunk(emb, 2, dim=1) - x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :] - return x - - -class CogView4Transformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, CacheMixin): - r""" - Args: - patch_size (`int`, defaults to `2`): - The size of the patches to use in the patch embedding layer. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - num_layers (`int`, defaults to `30`): - The number of layers of Transformer blocks to use. - attention_head_dim (`int`, defaults to `40`): - The number of channels in each head. - num_attention_heads (`int`, defaults to `64`): - The number of heads to use for multi-head attention. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - text_embed_dim (`int`, defaults to `4096`): - Input dimension of text embeddings from the text encoder. - time_embed_dim (`int`, defaults to `512`): - Output dimension of timestep embeddings. - condition_dim (`int`, defaults to `256`): - The embedding dimension of the input SDXL-style resolution conditions (original_size, target_size, - crop_coords). - pos_embed_max_size (`int`, defaults to `128`): - The maximum resolution of the positional embeddings, from which slices of shape `H x W` are taken and added - to input patched latents, where `H` and `W` are the latent height and width respectively. A value of 128 - means that the maximum supported height and width for image generation is `128 * vae_scale_factor * - patch_size => 128 * 8 * 2 => 2048`. - sample_size (`int`, defaults to `128`): - The base resolution of input latents. If height/width is not provided during generation, this value is used - to determine the resolution as `sample_size * vae_scale_factor => 128 * 8 => 1024` - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["CogView4TransformerBlock", "CogView4PatchEmbed", "CogView4PatchEmbed"] - _skip_layerwise_casting_patterns = ["patch_embed", "norm", "proj_out"] - - @register_to_config - def __init__( - self, - patch_size: int = 2, - in_channels: int = 16, - out_channels: int = 16, - num_layers: int = 30, - attention_head_dim: int = 40, - num_attention_heads: int = 64, - text_embed_dim: int = 4096, - time_embed_dim: int = 512, - condition_dim: int = 256, - pos_embed_max_size: int = 128, - sample_size: int = 128, - rope_axes_dim: tuple[int, int] = (256, 256), - ): - super().__init__() - - # CogView4 uses 3 additional SDXL-like conditions - original_size, target_size, crop_coords - # Each of these are sincos embeddings of shape 2 * condition_dim - pooled_projection_dim = 3 * 2 * condition_dim - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels - - # 1. RoPE - self.rope = CogView4RotaryPosEmbed(attention_head_dim, patch_size, rope_axes_dim, theta=10000.0) - - # 2. Patch & Text-timestep embedding - self.patch_embed = CogView4PatchEmbed(in_channels, inner_dim, patch_size, text_embed_dim) - - self.time_condition_embed = CogView3CombinedTimestepSizeEmbeddings( - embedding_dim=time_embed_dim, - condition_dim=condition_dim, - pooled_projection_dim=pooled_projection_dim, - timesteps_dim=inner_dim, - ) - - # 3. Transformer blocks - self.transformer_blocks = nn.ModuleList( - [ - CogView4TransformerBlock(inner_dim, num_attention_heads, attention_head_dim, time_embed_dim) - for _ in range(num_layers) - ] - ) - - # 4. Output projection - self.norm_out = CogView4AdaLayerNormContinuous(inner_dim, time_embed_dim, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels, bias=True) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: torch.LongTensor, - original_size: torch.Tensor, - target_size: torch.Tensor, - crop_coords: torch.Tensor, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]] | None = None, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`CogView4Transformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, in_channels, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - original_size (`torch.Tensor`): - Original image size conditioning. - target_size (`torch.Tensor`): - Target image size conditioning. - crop_coords (`torch.Tensor`): - Crop coordinates conditioning. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - attention_mask (`torch.Tensor`, *optional*): - Mask applied to attention scores. - image_rotary_emb (`tuple` of `torch.Tensor`, *optional*): - Pre-computed rotary positional embeddings. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_channels, height, width = hidden_states.shape - - # 1. RoPE - if image_rotary_emb is None: - image_rotary_emb = self.rope(hidden_states) - - # 2. Patch & Timestep embeddings - p = self.config.patch_size - post_patch_height = height // p - post_patch_width = width // p - - hidden_states, encoder_hidden_states = self.patch_embed(hidden_states, encoder_hidden_states) - - temb = self.time_condition_embed(timestep, original_size, target_size, crop_coords, hidden_states.dtype) - temb = F.silu(temb) - - # 3. Transformer blocks - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - attention_mask, - attention_kwargs, - ) - else: - hidden_states, encoder_hidden_states = block( - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - attention_mask, - attention_kwargs, - ) - - # 4. Output norm & projection - hidden_states = self.norm_out(hidden_states, temb) - hidden_states = self.proj_out(hidden_states) - - # 5. Unpatchify - hidden_states = hidden_states.reshape(batch_size, post_patch_height, post_patch_width, -1, p, p) - output = hidden_states.permute(0, 3, 1, 4, 2, 5).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_cosmos.py b/diffusers/models/transformers/transformer_cosmos.py deleted file mode 100644 index d901bb5809de47251ae5ab63c8721f721656ece1..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_cosmos.py +++ /dev/null @@ -1,840 +0,0 @@ -# Copyright 2025 The NVIDIA Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import numpy as np -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import is_torchvision_available -from ..attention import FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..attention_processor import Attention -from ..embeddings import Timesteps -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import RMSNorm - - -if is_torchvision_available(): - from torchvision import transforms - - -class CosmosPatchEmbed(nn.Module): - def __init__( - self, in_channels: int, out_channels: int, patch_size: tuple[int, int, int], bias: bool = True - ) -> None: - super().__init__() - self.patch_size = patch_size - - self.proj = nn.Linear(in_channels * patch_size[0] * patch_size[1] * patch_size[2], out_channels, bias=bias) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.patch_size - hidden_states = hidden_states.reshape( - batch_size, num_channels, num_frames // p_t, p_t, height // p_h, p_h, width // p_w, p_w - ) - hidden_states = hidden_states.permute(0, 2, 4, 6, 1, 3, 5, 7).flatten(4, 7) - hidden_states = self.proj(hidden_states) - return hidden_states - - -class CosmosTimestepEmbedding(nn.Module): - def __init__(self, in_features: int, out_features: int) -> None: - super().__init__() - self.linear_1 = nn.Linear(in_features, out_features, bias=False) - self.activation = nn.SiLU() - self.linear_2 = nn.Linear(out_features, 3 * out_features, bias=False) - - def forward(self, timesteps: torch.Tensor) -> torch.Tensor: - emb = self.linear_1(timesteps) - emb = self.activation(emb) - emb = self.linear_2(emb) - return emb - - -class CosmosEmbedding(nn.Module): - def __init__(self, embedding_dim: int, condition_dim: int) -> None: - super().__init__() - - self.time_proj = Timesteps(embedding_dim, flip_sin_to_cos=True, downscale_freq_shift=0.0) - self.t_embedder = CosmosTimestepEmbedding(embedding_dim, condition_dim) - self.norm = RMSNorm(embedding_dim, eps=1e-6, elementwise_affine=True) - - def forward(self, hidden_states: torch.Tensor, timestep: torch.LongTensor) -> torch.Tensor: - timesteps_proj = self.time_proj(timestep).type_as(hidden_states) - temb = self.t_embedder(timesteps_proj) - embedded_timestep = self.norm(timesteps_proj) - return temb, embedded_timestep - - -class CosmosAdaLayerNorm(nn.Module): - def __init__(self, in_features: int, hidden_features: int) -> None: - super().__init__() - self.embedding_dim = in_features - - self.activation = nn.SiLU() - self.norm = nn.LayerNorm(in_features, elementwise_affine=False, eps=1e-6) - self.linear_1 = nn.Linear(in_features, hidden_features, bias=False) - self.linear_2 = nn.Linear(hidden_features, 2 * in_features, bias=False) - - def forward( - self, hidden_states: torch.Tensor, embedded_timestep: torch.Tensor, temb: torch.Tensor | None = None - ) -> torch.Tensor: - embedded_timestep = self.activation(embedded_timestep) - embedded_timestep = self.linear_1(embedded_timestep) - embedded_timestep = self.linear_2(embedded_timestep) - - if temb is not None: - embedded_timestep = embedded_timestep + temb[..., : 2 * self.embedding_dim] - - shift, scale = embedded_timestep.chunk(2, dim=-1) - hidden_states = self.norm(hidden_states) - - if embedded_timestep.ndim == 2: - shift, scale = (x.unsqueeze(1) for x in (shift, scale)) - - hidden_states = hidden_states * (1 + scale) + shift - return hidden_states - - -class CosmosAdaLayerNormZero(nn.Module): - def __init__(self, in_features: int, hidden_features: int | None = None) -> None: - super().__init__() - - self.norm = nn.LayerNorm(in_features, elementwise_affine=False, eps=1e-6) - self.activation = nn.SiLU() - - if hidden_features is None: - self.linear_1 = nn.Identity() - else: - self.linear_1 = nn.Linear(in_features, hidden_features, bias=False) - - self.linear_2 = nn.Linear(hidden_features, 3 * in_features, bias=False) - - def forward( - self, - hidden_states: torch.Tensor, - embedded_timestep: torch.Tensor, - temb: torch.Tensor | None = None, - ) -> torch.Tensor: - embedded_timestep = self.activation(embedded_timestep) - embedded_timestep = self.linear_1(embedded_timestep) - embedded_timestep = self.linear_2(embedded_timestep) - - if temb is not None: - embedded_timestep = embedded_timestep + temb - - shift, scale, gate = embedded_timestep.chunk(3, dim=-1) - hidden_states = self.norm(hidden_states) - - if embedded_timestep.ndim == 2: - shift, scale, gate = (x.unsqueeze(1) for x in (shift, scale, gate)) - - hidden_states = hidden_states * (1 + scale) + shift - return hidden_states, gate - - -class CosmosAttnProcessor2_0: - def __init__(self): - if not hasattr(torch.nn.functional, "scaled_dot_product_attention"): - raise ImportError("CosmosAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - # 1. QKV projections - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2) - key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2) - value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2) - - # 2. QK normalization - query = attn.norm_q(query) - key = attn.norm_k(key) - - # 3. Apply RoPE - if image_rotary_emb is not None: - from ..embeddings import apply_rotary_emb - - query = apply_rotary_emb(query, image_rotary_emb, use_real=True, use_real_unbind_dim=-2) - key = apply_rotary_emb(key, image_rotary_emb, use_real=True, use_real_unbind_dim=-2) - - # 4. Prepare for GQA - if torch.onnx.is_in_onnx_export(): - query_idx = torch.tensor(query.size(3), device=query.device) - key_idx = torch.tensor(key.size(3), device=key.device) - value_idx = torch.tensor(value.size(3), device=value.device) - else: - query_idx = query.size(3) - key_idx = key.size(3) - value_idx = value.size(3) - key = key.repeat_interleave(query_idx // key_idx, dim=3) - value = value.repeat_interleave(query_idx // value_idx, dim=3) - - # 5. Attention - hidden_states = dispatch_attention_fn( - query.transpose(1, 2), - key.transpose(1, 2), - value.transpose(1, 2), - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - ) - hidden_states = hidden_states.flatten(2, 3).type_as(query) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - return hidden_states - - -class CosmosAttnProcessor2_5: - def __init__(self): - if not hasattr(torch.nn.functional, "scaled_dot_product_attention"): - raise ImportError("CosmosAttnProcessor2_5 requires PyTorch 2.0. Please upgrade PyTorch to 2.0 or newer.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: tuple[torch.Tensor, torch.Tensor], - attention_mask: tuple[torch.Tensor, torch.Tensor], - image_rotary_emb=None, - ) -> torch.Tensor: - if not isinstance(encoder_hidden_states, tuple): - raise ValueError("Expected encoder_hidden_states as (text_context, img_context) tuple.") - - text_context, img_context = encoder_hidden_states if encoder_hidden_states else (None, None) - text_mask, img_mask = attention_mask if attention_mask else (None, None) - - if text_context is None: - text_context = hidden_states - - query = attn.to_q(hidden_states) - key = attn.to_k(text_context) - value = attn.to_v(text_context) - - query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2) - key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2) - value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if image_rotary_emb is not None: - from ..embeddings import apply_rotary_emb - - query = apply_rotary_emb(query, image_rotary_emb, use_real=True, use_real_unbind_dim=-2) - key = apply_rotary_emb(key, image_rotary_emb, use_real=True, use_real_unbind_dim=-2) - - if torch.onnx.is_in_onnx_export(): - query_idx = torch.tensor(query.size(3), device=query.device) - key_idx = torch.tensor(key.size(3), device=key.device) - value_idx = torch.tensor(value.size(3), device=value.device) - else: - query_idx = query.size(3) - key_idx = key.size(3) - value_idx = value.size(3) - key = key.repeat_interleave(query_idx // key_idx, dim=3) - value = value.repeat_interleave(query_idx // value_idx, dim=3) - - attn_out = dispatch_attention_fn( - query.transpose(1, 2), - key.transpose(1, 2), - value.transpose(1, 2), - attn_mask=text_mask, - dropout_p=0.0, - is_causal=False, - ) - attn_out = attn_out.flatten(2, 3).type_as(query) - - if img_context is not None: - q_img = attn.q_img(hidden_states) - k_img = attn.k_img(img_context) - v_img = attn.v_img(img_context) - - batch_size = hidden_states.shape[0] - dim_head = attn.out_dim // attn.heads - - q_img = q_img.view(batch_size, -1, attn.heads, dim_head).transpose(1, 2) - k_img = k_img.view(batch_size, -1, attn.heads, dim_head).transpose(1, 2) - v_img = v_img.view(batch_size, -1, attn.heads, dim_head).transpose(1, 2) - - q_img = attn.q_img_norm(q_img) - k_img = attn.k_img_norm(k_img) - - q_img_idx = q_img.size(3) - k_img_idx = k_img.size(3) - v_img_idx = v_img.size(3) - k_img = k_img.repeat_interleave(q_img_idx // k_img_idx, dim=3) - v_img = v_img.repeat_interleave(q_img_idx // v_img_idx, dim=3) - - img_out = dispatch_attention_fn( - q_img.transpose(1, 2), - k_img.transpose(1, 2), - v_img.transpose(1, 2), - attn_mask=img_mask, - dropout_p=0.0, - is_causal=False, - ) - img_out = img_out.flatten(2, 3).type_as(q_img) - hidden_states = attn_out + img_out - else: - hidden_states = attn_out - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class CosmosAttention(Attention): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - - # add parameters for image q/k/v - inner_dim = self.heads * self.to_q.out_features // self.heads - self.q_img = nn.Linear(self.query_dim, inner_dim, bias=False) - self.k_img = nn.Linear(self.query_dim, inner_dim, bias=False) - self.v_img = nn.Linear(self.query_dim, inner_dim, bias=False) - self.q_img_norm = RMSNorm(self.to_q.out_features // self.heads, eps=1e-6, elementwise_affine=True) - self.k_img_norm = RMSNorm(self.to_k.out_features // self.heads, eps=1e-6, elementwise_affine=True) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: tuple[torch.Tensor, torch.Tensor], - attention_mask: torch.Tensor | None = None, - **cross_attention_kwargs, - ) -> torch.Tensor: - return super().forward( - hidden_states=hidden_states, - # NOTE: type-hint in base class can be ignored - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - **cross_attention_kwargs, - ) - - -class CosmosTransformerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - cross_attention_dim: int, - mlp_ratio: float = 4.0, - adaln_lora_dim: int = 256, - qk_norm: str = "rms_norm", - out_bias: bool = False, - img_context: bool = False, - before_proj: bool = False, - after_proj: bool = False, - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.norm1 = CosmosAdaLayerNormZero(in_features=hidden_size, hidden_features=adaln_lora_dim) - self.img_context = img_context - self.attn1 = Attention( - query_dim=hidden_size, - cross_attention_dim=None, - heads=num_attention_heads, - dim_head=attention_head_dim, - qk_norm=qk_norm, - elementwise_affine=True, - out_bias=out_bias, - processor=CosmosAttnProcessor2_0(), - ) - - self.norm2 = CosmosAdaLayerNormZero(in_features=hidden_size, hidden_features=adaln_lora_dim) - if img_context: - self.attn2 = CosmosAttention( - query_dim=hidden_size, - cross_attention_dim=cross_attention_dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - qk_norm=qk_norm, - elementwise_affine=True, - out_bias=out_bias, - processor=CosmosAttnProcessor2_5(), - ) - else: - self.attn2 = Attention( - query_dim=hidden_size, - cross_attention_dim=cross_attention_dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - qk_norm=qk_norm, - elementwise_affine=True, - out_bias=out_bias, - processor=CosmosAttnProcessor2_0(), - ) - - self.norm3 = CosmosAdaLayerNormZero(in_features=hidden_size, hidden_features=adaln_lora_dim) - self.ff = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu", bias=out_bias) - - # NOTE: zero conv for CosmosControlNet - self.before_proj = None - self.after_proj = None - if before_proj: - self.before_proj = nn.Linear(hidden_size, hidden_size) - if after_proj: - self.after_proj = nn.Linear(hidden_size, hidden_size) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None | tuple[torch.Tensor | None, torch.Tensor | None], - embedded_timestep: torch.Tensor, - temb: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - extra_pos_emb: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - controlnet_residual: torch.Tensor | None = None, - latents: torch.Tensor | None = None, - block_idx: int | None = None, - ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: - if self.before_proj is not None: - hidden_states = self.before_proj(hidden_states) + latents - - if extra_pos_emb is not None: - hidden_states = hidden_states + extra_pos_emb - - # 1. Self Attention - norm_hidden_states, gate = self.norm1(hidden_states, embedded_timestep, temb) - attn_output = self.attn1(norm_hidden_states, image_rotary_emb=image_rotary_emb) - hidden_states = hidden_states + gate * attn_output - - # 2. Cross Attention - norm_hidden_states, gate = self.norm2(hidden_states, embedded_timestep, temb) - attn_output = self.attn2( - norm_hidden_states, encoder_hidden_states=encoder_hidden_states, attention_mask=attention_mask - ) - hidden_states = hidden_states + gate * attn_output - - # 3. Feed Forward - norm_hidden_states, gate = self.norm3(hidden_states, embedded_timestep, temb) - ff_output = self.ff(norm_hidden_states) - hidden_states = hidden_states + gate * ff_output - - if controlnet_residual is not None: - assert self.after_proj is None - # NOTE: this is assumed to be scaled by the controlnet - hidden_states += controlnet_residual - - if self.after_proj is not None: - assert controlnet_residual is None - hs_proj = self.after_proj(hidden_states) - return hidden_states, hs_proj - - return hidden_states - - -class CosmosRotaryPosEmbed(nn.Module): - def __init__( - self, - hidden_size: int, - max_size: tuple[int, int, int] = (128, 240, 240), - patch_size: tuple[int, int, int] = (1, 2, 2), - base_fps: int = 24, - rope_scale: tuple[float, float, float] = (2.0, 1.0, 1.0), - ) -> None: - super().__init__() - - self.max_size = [size // patch for size, patch in zip(max_size, patch_size)] - self.patch_size = patch_size - self.base_fps = base_fps - - self.dim_h = hidden_size // 6 * 2 - self.dim_w = hidden_size // 6 * 2 - self.dim_t = hidden_size - self.dim_h - self.dim_w - - self.h_ntk_factor = rope_scale[1] ** (self.dim_h / (self.dim_h - 2)) - self.w_ntk_factor = rope_scale[2] ** (self.dim_w / (self.dim_w - 2)) - self.t_ntk_factor = rope_scale[0] ** (self.dim_t / (self.dim_t - 2)) - - def forward(self, hidden_states: torch.Tensor, fps: int | None = None) -> tuple[torch.Tensor, torch.Tensor]: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - pe_size = [num_frames // self.patch_size[0], height // self.patch_size[1], width // self.patch_size[2]] - device = hidden_states.device - - h_theta = 10000.0 * self.h_ntk_factor - w_theta = 10000.0 * self.w_ntk_factor - t_theta = 10000.0 * self.t_ntk_factor - - seq = torch.arange(max(self.max_size), device=device, dtype=torch.float32) - dim_h_range = ( - torch.arange(0, self.dim_h, 2, device=device, dtype=torch.float32)[: (self.dim_h // 2)] / self.dim_h - ) - dim_w_range = ( - torch.arange(0, self.dim_w, 2, device=device, dtype=torch.float32)[: (self.dim_w // 2)] / self.dim_w - ) - dim_t_range = ( - torch.arange(0, self.dim_t, 2, device=device, dtype=torch.float32)[: (self.dim_t // 2)] / self.dim_t - ) - h_spatial_freqs = 1.0 / (h_theta**dim_h_range) - w_spatial_freqs = 1.0 / (w_theta**dim_w_range) - temporal_freqs = 1.0 / (t_theta**dim_t_range) - - emb_h = torch.outer(seq[: pe_size[1]], h_spatial_freqs)[None, :, None, :].repeat(pe_size[0], 1, pe_size[2], 1) - emb_w = torch.outer(seq[: pe_size[2]], w_spatial_freqs)[None, None, :, :].repeat(pe_size[0], pe_size[1], 1, 1) - - # Apply sequence scaling in temporal dimension - if fps is None: - # Images - emb_t = torch.outer(seq[: pe_size[0]], temporal_freqs) - else: - # Videos - emb_t = torch.outer(seq[: pe_size[0]] / fps * self.base_fps, temporal_freqs) - - emb_t = emb_t[:, None, None, :].repeat(1, pe_size[1], pe_size[2], 1) - freqs = torch.cat([emb_t, emb_h, emb_w] * 2, dim=-1).flatten(0, 2).float() - cos = torch.cos(freqs) - sin = torch.sin(freqs) - return cos, sin - - -class CosmosLearnablePositionalEmbed(nn.Module): - def __init__( - self, - hidden_size: int, - max_size: tuple[int, int, int], - patch_size: tuple[int, int, int], - eps: float = 1e-6, - ) -> None: - super().__init__() - - self.max_size = [size // patch for size, patch in zip(max_size, patch_size)] - self.patch_size = patch_size - self.eps = eps - - self.pos_emb_t = nn.Parameter(torch.zeros(self.max_size[0], hidden_size)) - self.pos_emb_h = nn.Parameter(torch.zeros(self.max_size[1], hidden_size)) - self.pos_emb_w = nn.Parameter(torch.zeros(self.max_size[2], hidden_size)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - pe_size = [num_frames // self.patch_size[0], height // self.patch_size[1], width // self.patch_size[2]] - - emb_t = self.pos_emb_t[: pe_size[0]][None, :, None, None, :].repeat(batch_size, 1, pe_size[1], pe_size[2], 1) - emb_h = self.pos_emb_h[: pe_size[1]][None, None, :, None, :].repeat(batch_size, pe_size[0], 1, pe_size[2], 1) - emb_w = self.pos_emb_w[: pe_size[2]][None, None, None, :, :].repeat(batch_size, pe_size[0], pe_size[1], 1, 1) - emb = emb_t + emb_h + emb_w - emb = emb.flatten(1, 3) - - norm = torch.linalg.vector_norm(emb, dim=-1, keepdim=True, dtype=torch.float32) - norm = torch.add(self.eps, norm, alpha=np.sqrt(norm.numel() / emb.numel())) - return (emb / norm).type_as(hidden_states) - - -class CosmosTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin, PeftAdapterMixin): - r""" - A Transformer model for video-like data used in [Cosmos](https://github.com/NVIDIA/Cosmos). - - Args: - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - num_attention_heads (`int`, defaults to `32`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each attention head. - num_layers (`int`, defaults to `28`): - The number of layers of transformer blocks to use. - mlp_ratio (`float`, defaults to `4.0`): - The ratio of the hidden layer size to the input size in the feedforward network. - text_embed_dim (`int`, defaults to `4096`): - Input dimension of text embeddings from the text encoder. - adaln_lora_dim (`int`, defaults to `256`): - The hidden dimension of the Adaptive LayerNorm LoRA layer. - max_size (`tuple[int, int, int]`, defaults to `(128, 240, 240)`): - The maximum size of the input latent tensors in the temporal, height, and width dimensions. - patch_size (`tuple[int, int, int]`, defaults to `(1, 2, 2)`): - The patch size to use for patchifying the input latent tensors in the temporal, height, and width - dimensions. - rope_scale (`tuple[float, float, float]`, defaults to `(2.0, 1.0, 1.0)`): - The scaling factor to use for RoPE in the temporal, height, and width dimensions. - concat_padding_mask (`bool`, defaults to `True`): - Whether to concatenate the padding mask to the input latent tensors. - extra_pos_embed_type (`str`, *optional*, defaults to `learnable`): - The type of extra positional embeddings to use. Can be one of `None` or `learnable`. - controlnet_block_every_n (`int`, *optional*): - Interval between transformer blocks that should receive control residuals (for example, `7` to inject after - every seventh block). Required for Cosmos Transfer2.5. - img_context_dim_in (`int`, *optional*): - The dimension of the input image context feature vector, i.e. it is the D in [B, N, D]. - img_context_num_tokens (`int`): - The number of tokens in the image context feature vector, i.e. it is the N in [B, N, D]. If - `img_context_dim_in` is not provided, then this parameter is ignored. - img_context_dim_out (`int`): - The output dimension of the image context projection layer. If `img_context_dim_in` is not provided, then - this parameter is ignored. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["patch_embed", "final_layer", "norm"] - _no_split_modules = ["CosmosTransformerBlock"] - _keep_in_fp32_modules = ["learnable_pos_embed"] - - @register_to_config - def __init__( - self, - in_channels: int = 16, - out_channels: int = 16, - num_attention_heads: int = 32, - attention_head_dim: int = 128, - num_layers: int = 28, - mlp_ratio: float = 4.0, - text_embed_dim: int = 1024, - adaln_lora_dim: int = 256, - max_size: tuple[int, int, int] = (128, 240, 240), - patch_size: tuple[int, int, int] = (1, 2, 2), - rope_scale: tuple[float, float, float] = (2.0, 1.0, 1.0), - concat_padding_mask: bool = True, - extra_pos_embed_type: str | None = "learnable", - use_crossattn_projection: bool = False, - crossattn_proj_in_channels: int = 1024, - encoder_hidden_states_channels: int = 1024, - controlnet_block_every_n: int | None = None, - img_context_dim_in: int | None = None, - img_context_num_tokens: int = 256, - img_context_dim_out: int = 2048, - ) -> None: - super().__init__() - hidden_size = num_attention_heads * attention_head_dim - - # 1. Patch Embedding - patch_embed_in_channels = in_channels + 1 if concat_padding_mask else in_channels - self.patch_embed = CosmosPatchEmbed(patch_embed_in_channels, hidden_size, patch_size, bias=False) - - # 2. Positional Embedding - self.rope = CosmosRotaryPosEmbed( - hidden_size=attention_head_dim, max_size=max_size, patch_size=patch_size, rope_scale=rope_scale - ) - - self.learnable_pos_embed = None - if extra_pos_embed_type == "learnable": - self.learnable_pos_embed = CosmosLearnablePositionalEmbed( - hidden_size=hidden_size, - max_size=max_size, - patch_size=patch_size, - ) - - # 3. Time Embedding - self.time_embed = CosmosEmbedding(hidden_size, hidden_size) - - # 4. Transformer Blocks - self.transformer_blocks = nn.ModuleList( - [ - CosmosTransformerBlock( - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - cross_attention_dim=text_embed_dim, - mlp_ratio=mlp_ratio, - adaln_lora_dim=adaln_lora_dim, - qk_norm="rms_norm", - out_bias=False, - img_context=self.config.img_context_dim_in is not None and self.config.img_context_dim_in > 0, - ) - for _ in range(num_layers) - ] - ) - - # 5. Output norm & projection - self.norm_out = CosmosAdaLayerNorm(hidden_size, adaln_lora_dim) - self.proj_out = nn.Linear( - hidden_size, patch_size[0] * patch_size[1] * patch_size[2] * out_channels, bias=False - ) - - if self.config.use_crossattn_projection: - self.crossattn_proj = nn.Sequential( - nn.Linear(crossattn_proj_in_channels, encoder_hidden_states_channels, bias=True), - nn.GELU(), - ) - - self.gradient_checkpointing = False - - if self.config.img_context_dim_in: - self.img_context_proj = nn.Sequential( - nn.Linear(self.config.img_context_dim_in, self.config.img_context_dim_out, bias=True), - nn.GELU(), - ) - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - block_controlnet_hidden_states: list[torch.Tensor] | None = None, - attention_mask: torch.Tensor | None = None, - fps: int | None = None, - condition_mask: torch.Tensor | None = None, - padding_mask: torch.Tensor | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`CosmosTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - block_controlnet_hidden_states (`list` of `torch.Tensor`, *optional*): - A list of tensors that if specified are added to the residuals of transformer blocks. - attention_mask (`torch.Tensor`, *optional*): - Mask applied to `encoder_hidden_states` during attention. - fps (`int`, *optional*): - Frames per second of the input video used to compute the rotary positional embeddings. - condition_mask (`torch.Tensor`, *optional*): - Mask channel concatenated to `hidden_states` to indicate the conditioning region. - padding_mask (`torch.Tensor`, *optional*): - Padding mask concatenated to `hidden_states` when `concat_padding_mask` is enabled. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_channels, num_frames, height, width = hidden_states.shape - - # 1. Concatenate padding mask if needed & prepare attention mask - if condition_mask is not None: - hidden_states = torch.cat([hidden_states, condition_mask], dim=1) - - if self.config.concat_padding_mask: - padding_mask_resized = transforms.functional.resize( - padding_mask, list(hidden_states.shape[-2:]), interpolation=transforms.InterpolationMode.NEAREST - ) - hidden_states = torch.cat( - [hidden_states, padding_mask_resized.unsqueeze(2).repeat(batch_size, 1, num_frames, 1, 1)], dim=1 - ) - - if attention_mask is not None: - attention_mask = attention_mask.unsqueeze(1).unsqueeze(1) # [B, 1, 1, S] - - # 2. Generate positional embeddings - image_rotary_emb = self.rope(hidden_states, fps=fps) - extra_pos_emb = self.learnable_pos_embed(hidden_states) if self.config.extra_pos_embed_type else None - - # 3. Patchify input - p_t, p_h, p_w = self.config.patch_size - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p_h - post_patch_width = width // p_w - - hidden_states = self.patch_embed(hidden_states) - hidden_states = hidden_states.flatten(1, 3) # [B, T, H, W, C] -> [B, THW, C] - - # 4. Timestep embeddings - if timestep.ndim == 1: - temb, embedded_timestep = self.time_embed(hidden_states, timestep) - elif timestep.ndim == 5: - assert timestep.shape == (batch_size, 1, num_frames, 1, 1), ( - f"Expected timestep to have shape [B, 1, T, 1, 1], but got {timestep.shape}" - ) - timestep = timestep.flatten() - temb, embedded_timestep = self.time_embed(hidden_states, timestep) - # We can do this because num_frames == post_patch_num_frames, as p_t is 1 - temb, embedded_timestep = ( - x.view(batch_size, post_patch_num_frames, 1, 1, -1) - .expand(-1, -1, post_patch_height, post_patch_width, -1) - .flatten(1, 3) - for x in (temb, embedded_timestep) - ) # [BT, C] -> [B, T, 1, 1, C] -> [B, T, H, W, C] -> [B, THW, C] - else: - raise ValueError(f"Expected timestep to have shape [B, 1, T, 1, 1] or [T], but got {timestep.shape}") - - # 5. Process encoder hidden states - text_context, img_context = ( - encoder_hidden_states if isinstance(encoder_hidden_states, tuple) else (encoder_hidden_states, None) - ) - if self.config.use_crossattn_projection: - text_context = self.crossattn_proj(text_context) - - if img_context is not None and self.config.img_context_dim_in: - img_context = self.img_context_proj(img_context) - - processed_encoder_hidden_states = ( - (text_context, img_context) if isinstance(encoder_hidden_states, tuple) else text_context - ) - - # 6. Build controlnet block index map - controlnet_block_index_map = {} - if block_controlnet_hidden_states is not None: - n_blocks = len(self.transformer_blocks) - controlnet_block_index_map = { - block_idx: block_controlnet_hidden_states[idx] - for idx, block_idx in list(enumerate(range(0, n_blocks, self.config.controlnet_block_every_n))) - } - - # 7. Transformer blocks - for block_idx, block in enumerate(self.transformer_blocks): - controlnet_residual = controlnet_block_index_map.get(block_idx) - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - processed_encoder_hidden_states, - embedded_timestep, - temb, - image_rotary_emb, - extra_pos_emb, - attention_mask, - controlnet_residual, - ) - else: - hidden_states = block( - hidden_states, - processed_encoder_hidden_states, - embedded_timestep, - temb, - image_rotary_emb, - extra_pos_emb, - attention_mask, - controlnet_residual, - ) - - # 8. Output norm & projection & unpatchify - hidden_states = self.norm_out(hidden_states, embedded_timestep, temb) - hidden_states = self.proj_out(hidden_states) - hidden_states = hidden_states.unflatten(2, (p_h, p_w, p_t, -1)) - hidden_states = hidden_states.unflatten(1, (post_patch_num_frames, post_patch_height, post_patch_width)) - # NOTE: The permutation order here is not the inverse operation of what happens when patching as usually expected. - # It might be a source of confusion to the reader, but this is correct - hidden_states = hidden_states.permute(0, 7, 1, 6, 2, 4, 3, 5) - hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (hidden_states,) - - return Transformer2DModelOutput(sample=hidden_states) diff --git a/diffusers/models/transformers/transformer_cosmos3.py b/diffusers/models/transformers/transformer_cosmos3.py deleted file mode 100644 index f7cfc317bc7922c5e7aebe91f693c1bab6343e02..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_cosmos3.py +++ /dev/null @@ -1,851 +0,0 @@ -# Copyright 2025 The NVIDIA Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from dataclasses import dataclass - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import BaseOutput -from ..attention import AttentionMixin, AttentionModuleMixin -from ..attention_dispatch import dispatch_attention_fn -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from ..normalization import RMSNorm - - -@dataclass -class Cosmos3OmniTransformerOutput(BaseOutput): - """Output of [`Cosmos3OmniTransformer`]. - - Args: - sample (`list[torch.Tensor]`): - Per-item vision velocity predictions. - sound (`list[torch.Tensor]`, *optional*): - Per-item sound velocity predictions when sound generation is enabled. - action (`list[torch.Tensor]`, *optional*): - Per-item action velocity predictions when action generation is enabled. - """ - - sample: list[torch.Tensor] - sound: list[torch.Tensor] | None = None - action: list[torch.Tensor] | None = None - - -class Cosmos3AttnProcessor: - """Dual-pathway attention processor for Cosmos3. - - Projects, normalizes, applies rotary position embeddings, then runs separate causal (understanding) and full - (generation) attention pathways. The generation pathway cross-attends to both und and gen keys/values. - """ - - _attention_backend = None - _parallel_config = None - - def __call__( - self, - attn: "Cosmos3PackedMoTAttention", - und_seq: torch.Tensor, - gen_seq: torch.Tensor, - rotary_emb: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor], - ) -> tuple[torch.Tensor, torch.Tensor]: - # Per-pathway projections - q_und = attn.to_q(und_seq).view(-1, attn.num_attention_heads, attn.head_dim) - k_und = attn.to_k(und_seq).view(-1, attn.num_key_value_heads, attn.head_dim) - v_und = attn.to_v(und_seq).view(-1, attn.num_key_value_heads, attn.head_dim) - q_gen = attn.add_q_proj(gen_seq).view(-1, attn.num_attention_heads, attn.head_dim) - k_gen = attn.add_k_proj(gen_seq).view(-1, attn.num_key_value_heads, attn.head_dim) - v_gen = attn.add_v_proj(gen_seq).view(-1, attn.num_key_value_heads, attn.head_dim) - - q_und = attn.norm_q(q_und) - k_und = attn.norm_k(k_und) - k_und_for_gen = attn.k_norm_und_for_gen(k_und) if attn.k_norm_und_for_gen is not None else k_und - q_gen = attn.norm_added_q(q_gen) - k_gen = attn.norm_added_k(k_gen) - - # Apply rotary position embeddings per pathway - cos_und, sin_und, cos_gen, sin_gen = rotary_emb - cos_und = cos_und.unsqueeze(1) - sin_und = sin_und.unsqueeze(1) - q_und = q_und * cos_und + _rotate_half(q_und) * sin_und - k_und = k_und * cos_und + _rotate_half(k_und) * sin_und - k_und_for_gen = k_und_for_gen * cos_und + _rotate_half(k_und_for_gen) * sin_und - cos_gen = cos_gen.unsqueeze(1) - sin_gen = sin_gen.unsqueeze(1) - q_gen = q_gen * cos_gen + _rotate_half(q_gen) * sin_gen - k_gen = k_gen * cos_gen + _rotate_half(k_gen) * sin_gen - - # Causal pathway (understanding): und tokens self-attend with causal masking. - causal_out = dispatch_attention_fn( - q_und.unsqueeze(0), - k_und.unsqueeze(0), - v_und.unsqueeze(0), - is_causal=True, - enable_gqa=True, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - causal_out = causal_out.squeeze(0).flatten(-2, -1) - - # Full pathway (generation): gen tokens cross-attend to all (und + gen) keys/values. - all_k = torch.cat([k_und_for_gen, k_gen], dim=0) - all_v = torch.cat([v_und, v_gen], dim=0) - full_out = dispatch_attention_fn( - q_gen.unsqueeze(0), - all_k.unsqueeze(0), - all_v.unsqueeze(0), - is_causal=False, - enable_gqa=True, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - full_out = full_out.squeeze(0).flatten(-2, -1) - - # Per-pathway output projection - und_out = attn.to_out(causal_out) - gen_out = attn.to_add_out(full_out) - return und_out, gen_out - - -def _rotate_half(x: torch.Tensor) -> torch.Tensor: - half = x.shape[-1] // 2 - return torch.cat((-x[..., half:], x[..., :half]), dim=-1) - - -class Cosmos3VLTextRotaryEmbedding(nn.Module): - def __init__(self, head_dim: int, rope_theta: float, rope_axes_dim: tuple[int, int, int]): - super().__init__() - inv_freq = 1.0 / (rope_theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)) - self.register_buffer("inv_freq", inv_freq, persistent=False) - self.rope_axes_dim = rope_axes_dim - - def apply_interleaved_mrope(self, freqs, rope_axes_dim): - """Reorganize chunked [TTT...HHH...WWW] frequency layout into interleaved - [THTHWHTHW...TT], preserving frequency continuity across the 3 grids.""" - freqs_t = freqs[0] - for dim, offset in enumerate((1, 2), start=1): # H, W - length = rope_axes_dim[dim] * 3 - idx = slice(offset, length, 3) - freqs_t[..., idx] = freqs[dim, ..., idx] - return freqs_t - - def forward(self, position_ids, device, dtype): - if position_ids.ndim == 2: - position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1) # [3,B,N] - inv_freq_expanded = ( - self.inv_freq[None, None, :, None].float().expand(3, position_ids.shape[1], -1, 1).to(device) - ) # [3,B,head_dim//2,1] - position_ids_expanded = position_ids[:, :, None, :].float() # [3,B,1,N] - # Disable autocast so the position-id matmul runs in float32: under an ambient autocast it would run in - # bfloat16, which cannot represent consecutive integers past 256, collapsing positions onto the same - # frequency and degrading the rotary embedding. - with torch.autocast(device_type=position_ids.device.type, enabled=False): - freqs = inv_freq_expanded @ position_ids_expanded - freqs = freqs.transpose(2, 3) # [3,B,N,head_dim//2] - freqs = self.apply_interleaved_mrope(freqs, self.rope_axes_dim) # [B,N,head_dim//2] - emb = torch.cat((freqs, freqs), dim=-1) # [B,N,head_dim] - return emb.cos().to(dtype=dtype), emb.sin().to(dtype=dtype) # each: [B,N,head_dim] - - -class Cosmos3NemotronRMSNorm(nn.Module): - def __init__(self, dim: int, eps: float): - super().__init__() - self.eps = eps - self.weight = nn.Parameter(torch.ones(dim)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - input_dtype = hidden_states.dtype - hidden_states = hidden_states.float() - variance = hidden_states.pow(2).mean(-1, keepdim=True) - hidden_states = hidden_states * torch.rsqrt(variance + self.eps) - return (self.weight.float() * hidden_states).to(input_dtype) - - -class Cosmos3VLTextMLP(nn.Module): - def __init__(self, hidden_size: int, intermediate_size: int, hidden_act: str = "silu"): - super().__init__() - if hidden_act not in ("relu2", "silu"): - raise ValueError(f"Cosmos3 only supports `hidden_act` values 'relu2' and 'silu', got {hidden_act!r}.") - self.hidden_act = hidden_act - if hidden_act == "silu": - self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False) - self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False) - self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False) - self.act_fn = nn.SiLU() if hidden_act == "silu" else None - - def forward(self, x): - if self.hidden_act == "relu2": - return self.down_proj(torch.relu(self.up_proj(x)).square()) - return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) - - -class DomainAwareLinear(nn.Module): - """Linear projection with one weight/bias pair per embodiment domain.""" - - def __init__(self, input_size: int, output_size: int, num_domains: int) -> None: - super().__init__() - self.input_size = input_size - self.output_size = output_size - self.num_domains = num_domains - self.fc = nn.Embedding(self.num_domains, self.output_size * self.input_size) - self.bias = nn.Embedding(self.num_domains, self.output_size) - - def forward(self, x: torch.Tensor, domain_id: torch.Tensor) -> torch.Tensor: - if domain_id.ndim == 0: - domain_id = domain_id.unsqueeze(0) - domain_id = domain_id.to(device=x.device, dtype=torch.long).reshape(-1) - if x.shape[0] != domain_id.shape[0]: - raise ValueError( - "Cosmos3 action domain_id batch size must match action tokens: " - f"tokens={x.shape[0]}, domain_id={domain_id.shape[0]}." - ) - if torch.any((domain_id < 0) | (domain_id >= self.num_domains)): - raise ValueError(f"Cosmos3 action domain_id must be in [0, {self.num_domains}), got {domain_id.tolist()}.") - weight = self.fc(domain_id).view(domain_id.shape[0], self.input_size, self.output_size) - bias = self.bias(domain_id).view(domain_id.shape[0], self.output_size) - if x.ndim == 2: - return torch.bmm(x.unsqueeze(1), weight).squeeze(1) + bias - if x.ndim == 3: - return torch.bmm(x, weight) + bias.unsqueeze(1) - raise ValueError(f"Cosmos3 DomainAwareLinear expected rank-2 or rank-3 input, got {tuple(x.shape)}.") - - -class Cosmos3PackedMoTAttention(nn.Module, AttentionModuleMixin): - """Dual-pathway packed attention with separate projections for the understanding and generation token streams.""" - - _default_processor_cls = Cosmos3AttnProcessor - _available_processors = [Cosmos3AttnProcessor] - _supports_qkv_fusion = False - - def __init__( - self, - hidden_size: int, - head_dim: int, - num_attention_heads: int, - num_key_value_heads: int, - attention_bias: bool, - rms_norm_eps: float, - qk_norm_for_text: bool = True, - use_und_k_norm_for_gen: bool = False, - norm_type: str = "rms_norm", - processor=None, - ): - super().__init__() - self.hidden_size = hidden_size - self.head_dim = head_dim - self.num_attention_heads = num_attention_heads - self.num_key_value_heads = num_key_value_heads - self.num_key_value_groups = num_attention_heads // num_key_value_heads - - # Understanding pathway. norm_q / norm_k are applied per-head (only on - # head_dim), so no reshape is needed after them. - self.to_q = nn.Linear(hidden_size, num_attention_heads * head_dim, bias=attention_bias) - self.to_k = nn.Linear(hidden_size, num_key_value_heads * head_dim, bias=attention_bias) - self.to_v = nn.Linear(hidden_size, num_key_value_heads * head_dim, bias=attention_bias) - self.to_out = nn.Linear(num_attention_heads * head_dim, hidden_size, bias=attention_bias) - if not qk_norm_for_text: - self.norm_q = nn.Identity() - self.norm_k = nn.Identity() - elif norm_type == "nemotron_rms_norm": - self.norm_q = Cosmos3NemotronRMSNorm(head_dim, eps=rms_norm_eps) - self.norm_k = Cosmos3NemotronRMSNorm(head_dim, eps=rms_norm_eps) - else: - self.norm_q = RMSNorm(head_dim, eps=rms_norm_eps, elementwise_affine=True, bias=False) - self.norm_k = RMSNorm(head_dim, eps=rms_norm_eps, elementwise_affine=True, bias=False) - - if use_und_k_norm_for_gen and not qk_norm_for_text: - if norm_type == "nemotron_rms_norm": - self.k_norm_und_for_gen = Cosmos3NemotronRMSNorm(head_dim, eps=rms_norm_eps) - else: - self.k_norm_und_for_gen = RMSNorm(head_dim, eps=rms_norm_eps, elementwise_affine=True, bias=False) - else: - self.k_norm_und_for_gen = None - - # Generation pathway - self.add_q_proj = nn.Linear(hidden_size, num_attention_heads * head_dim, bias=attention_bias) - self.add_k_proj = nn.Linear(hidden_size, num_key_value_heads * head_dim, bias=attention_bias) - self.add_v_proj = nn.Linear(hidden_size, num_key_value_heads * head_dim, bias=attention_bias) - self.to_add_out = nn.Linear(num_attention_heads * head_dim, hidden_size, bias=attention_bias) - if norm_type == "nemotron_rms_norm": - self.norm_added_q = Cosmos3NemotronRMSNorm(head_dim, eps=rms_norm_eps) - self.norm_added_k = Cosmos3NemotronRMSNorm(head_dim, eps=rms_norm_eps) - else: - self.norm_added_q = RMSNorm(head_dim, eps=rms_norm_eps, elementwise_affine=True, bias=False) - self.norm_added_k = RMSNorm(head_dim, eps=rms_norm_eps, elementwise_affine=True, bias=False) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - und_seq: torch.Tensor, - gen_seq: torch.Tensor, - rotary_emb: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor], - ) -> tuple[torch.Tensor, torch.Tensor]: - return self.processor(self, und_seq, gen_seq, rotary_emb) - - -class Cosmos3VLTextMoTDecoderLayer(nn.Module): - """Cosmos3 text MoT decoder layer for the Qwen3 and Nemotron dense backbones.""" - - def __init__( - self, - hidden_size: int, - head_dim: int, - num_attention_heads: int, - num_key_value_heads: int, - intermediate_size: int, - attention_bias: bool, - rms_norm_eps: float, - hidden_act: str = "silu", - qk_norm_for_text: bool = True, - use_und_k_norm_for_gen: bool = False, - ): - super().__init__() - self.hidden_size = hidden_size - norm_type = "nemotron_rms_norm" if hidden_act == "relu2" else "rms_norm" - self.self_attn = Cosmos3PackedMoTAttention( - hidden_size=hidden_size, - head_dim=head_dim, - num_attention_heads=num_attention_heads, - num_key_value_heads=num_key_value_heads, - attention_bias=attention_bias, - rms_norm_eps=rms_norm_eps, - qk_norm_for_text=qk_norm_for_text, - use_und_k_norm_for_gen=use_und_k_norm_for_gen, - norm_type=norm_type, - ) - - self.mlp = Cosmos3VLTextMLP( - hidden_size=hidden_size, intermediate_size=intermediate_size, hidden_act=hidden_act - ) - self.mlp_moe_gen = Cosmos3VLTextMLP( - hidden_size=hidden_size, intermediate_size=intermediate_size, hidden_act=hidden_act - ) - - if norm_type == "nemotron_rms_norm": - self.input_layernorm = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps) - self.input_layernorm_moe_gen = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps) - self.post_attention_layernorm = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps) - self.post_attention_layernorm_moe_gen = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps) - else: - self.input_layernorm = RMSNorm(hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False) - self.input_layernorm_moe_gen = RMSNorm(hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False) - self.post_attention_layernorm = RMSNorm(hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False) - self.post_attention_layernorm_moe_gen = RMSNorm( - hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False - ) - - def forward( - self, - und_seq: torch.Tensor, - gen_seq: torch.Tensor, - rotary_emb: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor], - ) -> tuple[torch.Tensor, torch.Tensor]: - und_norm = self.input_layernorm(und_seq) - gen_norm = self.input_layernorm_moe_gen(gen_seq) - - und_attn_out, gen_attn_out = self.self_attn(und_norm, gen_norm, rotary_emb) - residual_und = und_seq + und_attn_out - residual_gen = gen_seq + gen_attn_out - - mlp_out_und = self.mlp(self.post_attention_layernorm(residual_und)) - mlp_out_gen = self.mlp_moe_gen(self.post_attention_layernorm_moe_gen(residual_gen)) - - return residual_und + mlp_out_und, residual_gen + mlp_out_gen - - -class Cosmos3OmniTransformer(ModelMixin, ConfigMixin, PeftAdapterMixin, AttentionMixin): - _supports_gradient_checkpointing = True - _no_split_modules = ["Cosmos3VLTextMoTDecoderLayer"] - _repeated_blocks = ["Cosmos3VLTextMoTDecoderLayer"] - _skip_layerwise_casting_patterns = ["embed_tokens", "time_embedder", "norm"] - _keep_in_fp32_modules = ["time_embedder"] - # Optional context-parallelism seams. They default to ``None`` (no-op) so the - # model itself carries no CP logic. `forward` applies `_cp_shard_fn` to the - # per-pathway hidden states + rotary embeddings before the decoder layers, and - # `_cp_gather_fn` to the per-pathway outputs after the final norm. An external - # helper (see `examples/cosmos3/cosmos_parallel.py`) sets these to - # shard/gather across a device mesh and installs a context-parallel attention - # processor — the packed dual-pathway + GQA + ragged-length structure cannot be - # expressed as diffusers' declarative `_cp_plan`, so CP lives outside the model. - _cp_shard_fn = None - _cp_gather_fn = None - # `dtype` is injected into init_dict by ModelMixin.from_pretrained (configuration_utils.py:289), - # so __init__ must accept it. Excluding it here keeps save_pretrained from writing it into - # config.json — the value is a load-time runtime hint, not part of the model architecture. - ignore_for_config = ["dtype"] - - @register_to_config - def __init__( - self, - attention_bias: bool = False, - attention_dropout: float = 0.0, - dtype: str = "bfloat16", # required by the loader (see `ignore_for_config` above); not read here - head_dim: int = 128, - hidden_size: int = 4096, - intermediate_size: int = 12288, - base_fps: int = 24, - enable_fps_modulation: bool = True, - latent_channel: int = 48, - unified_3d_mrope_reset_spatial_ids: bool = True, - unified_3d_mrope_temporal_modality_margin: int = 15000, - latent_patch_size: int = 2, - num_attention_heads: int = 32, - num_hidden_layers: int = 36, - num_key_value_heads: int = 8, - patch_latent_dim: int = 192, - rms_norm_eps: float = 1e-6, - rope_scaling: dict | None = None, - rope_theta: float = 5000000.0, - action_dim: int | None = None, - action_gen: bool = False, - num_embodiment_domains: int = 32, - sound_dim: int | None = None, - sound_gen: bool = False, - sound_latent_fps: float = 25.0, - timestep_scale: float = 0.001, - vocab_size: int = 151936, - hidden_act: str = "silu", - qk_norm_for_text: bool = True, - use_und_k_norm_for_gen: bool = False, - rope_axes_dim: tuple[int, int, int] | list[int] | None = None, - ): - super().__init__() - - if rope_axes_dim is None: - rope_axes_dim = ( - rope_scaling.get("mrope_section", [24, 20, 20]) if rope_scaling is not None else [24, 20, 20] - ) - self.register_to_config(rope_axes_dim=rope_axes_dim) - - # Text-model layers live directly on the transformer (flat layout). The published - # checkpoint must be re-keyed with the leading `model.` prefix stripped — see - # scripts/build_flat_layout_repo.py for the rewrite. - self.embed_tokens = nn.Embedding(vocab_size, hidden_size) - self.layers = nn.ModuleList( - [ - Cosmos3VLTextMoTDecoderLayer( - hidden_size=hidden_size, - head_dim=head_dim, - num_attention_heads=num_attention_heads, - num_key_value_heads=num_key_value_heads, - intermediate_size=intermediate_size, - attention_bias=attention_bias, - rms_norm_eps=rms_norm_eps, - hidden_act=hidden_act, - qk_norm_for_text=qk_norm_for_text, - use_und_k_norm_for_gen=use_und_k_norm_for_gen, - ) - for _ in range(num_hidden_layers) - ] - ) - if hidden_act == "relu2": - self.norm = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps) - self.norm_moe_gen = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps) - else: - self.norm = RMSNorm(hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False) - self.norm_moe_gen = RMSNorm(hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False) - self.rotary_emb = Cosmos3VLTextRotaryEmbedding( - head_dim=head_dim, rope_theta=rope_theta, rope_axes_dim=rope_axes_dim - ) - - # Modality projection heads + timestep embedding. - self.vocab_size = vocab_size - self.lm_head = nn.Linear(hidden_size, vocab_size, bias=False) - self.proj_in = nn.Linear(patch_latent_dim, hidden_size, bias=True) - self.proj_out = nn.Linear(hidden_size, patch_latent_dim, bias=True) - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.time_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=hidden_size) - self.action_gen = action_gen - self.action_dim = action_dim - self.num_embodiment_domains = num_embodiment_domains - if action_gen: - if self.action_dim is None: - raise ValueError("`action_dim` must be provided when `action_gen=True`.") - self.action_proj_in = DomainAwareLinear(self.action_dim, hidden_size, self.num_embodiment_domains) - self.action_proj_out = DomainAwareLinear(hidden_size, self.action_dim, self.num_embodiment_domains) - self.action_modality_embed = nn.Parameter(torch.zeros(hidden_size)) - if sound_gen: - if sound_dim is None: - raise ValueError("`sound_dim` must be provided when `sound_gen=True`.") - self.audio_proj_in = nn.Linear(sound_dim, hidden_size, bias=True) - self.audio_proj_out = nn.Linear(hidden_size, sound_dim, bias=True) - self.audio_modality_embed = nn.Parameter(torch.zeros(hidden_size)) - - self.gradient_checkpointing = False - - # ------------------------------------------------------------------------- - # Pure-tensor packing/unpacking helpers (no layer state). - # ------------------------------------------------------------------------- - - def _apply_timestep_embeds_to_noisy_tokens( - self, - packed_tokens: torch.Tensor, - packed_timestep_embeds: torch.Tensor, - noisy_frame_indexes: list[torch.Tensor], - token_shapes: list[tuple[int, ...]], - ) -> torch.Tensor: - start_noisy_index = 0 - flattened_noisy_frame_indexes: list[torch.Tensor] = [] - for noisy_indexes_i, token_shape_i in zip(noisy_frame_indexes, token_shapes): - spatial_numel_i = math.prod(token_shape_i[1:]) - spatial_indexes_i = torch.arange(spatial_numel_i, device=packed_tokens.device) - # Broadcast [N, 1] + [spatial_numel_i] → [N, spatial_numel_i] - frame_offsets = (noisy_indexes_i * spatial_numel_i).unsqueeze(-1) + spatial_indexes_i + start_noisy_index - flattened_noisy_frame_indexes.append(frame_offsets.flatten()) - start_noisy_index += token_shape_i[0] * spatial_numel_i - flattened = torch.cat(flattened_noisy_frame_indexes, dim=0).unsqueeze(-1).expand(-1, packed_tokens.shape[1]) - return packed_tokens.scatter_add(dim=0, index=flattened, src=packed_timestep_embeds) - - def _patchify_and_pack_latents( - self, - tokens_vision: list[torch.Tensor], - ) -> tuple[torch.Tensor, list[tuple[int, int, int]]]: - p = self.config.latent_patch_size - latent_channel = self.config.latent_channel - packed_latent: list[torch.Tensor] = [] - original_latent_shapes: list[tuple[int, int, int]] = [] - for latent in tokens_vision: - latent = latent.squeeze(0) # [C, T, H, W] - _, t_actual, h_actual, w_actual = latent.shape - original_latent_shapes.append((t_actual, h_actual, w_actual)) - h_padded = ((h_actual + p - 1) // p) * p - w_padded = ((w_actual + p - 1) // p) * p - if h_padded != h_actual or w_padded != w_actual: - padded = torch.zeros( - (latent_channel, t_actual, h_padded, w_padded), - device=latent.device, - dtype=latent.dtype, - ) - padded[:, :, :h_actual, :w_actual] = latent - latent = padded - h_patches = h_padded // p - w_patches = w_padded // p - latent = latent.reshape(latent_channel, t_actual, h_patches, p, w_patches, p) - latent = torch.einsum("cthpwq->thwpqc", latent).reshape(-1, p * p * latent_channel) - packed_latent.append(latent) - return torch.cat(packed_latent, dim=0), original_latent_shapes - - def _unpatchify_and_unpack_latents( - self, - packed_mse_preds: torch.Tensor, - token_shapes_vision: list[tuple[int, int, int]], - noisy_frame_indexes_vision: list[torch.Tensor], - original_latent_shapes: list[tuple[int, int, int]], - ) -> list[torch.Tensor]: - p = self.config.latent_patch_size - latent_channel = self.config.latent_channel - unpatchified_latents: list[torch.Tensor] = [] - start_idx = 0 - for token_shape, noisy_frame_indexes, original_shape in zip( - token_shapes_vision, noisy_frame_indexes_vision, original_latent_shapes - ): - t_c = token_shape[0] - _, h_orig, w_orig = original_shape - h_padded = ((h_orig + p - 1) // p) * p - w_padded = ((w_orig + p - 1) // p) * p - h_patches = h_padded // p - w_patches = w_padded // p - t_n = len(noisy_frame_indexes) - output_tensor = torch.zeros( - (latent_channel, t_c, h_orig, w_orig), - device=packed_mse_preds.device, - dtype=packed_mse_preds.dtype, - ) - num_patches = t_n * h_patches * w_patches - if num_patches > 0: - end_idx = start_idx + num_patches - latent_patches = packed_mse_preds[start_idx:end_idx] - latent_patches = latent_patches.reshape(t_n, h_patches, w_patches, p, p, latent_channel) - latent = torch.einsum("thwpqc->cthpwq", latent_patches) - latent = latent.reshape(latent_channel, t_n, h_patches * p, w_patches * p) - latent = latent[:, :, :h_orig, :w_orig] - output_tensor[:, noisy_frame_indexes] = latent - start_idx = end_idx - unpatchified_latents.append(output_tensor.unsqueeze(0)) - return unpatchified_latents - - def _pack_sound_latents( - self, - tokens_sound: list[torch.Tensor], - token_shapes_sound: list[tuple[int, int, int]], - ) -> torch.Tensor: - """List of ``[C, T]`` tensors → packed ``[total_T, C]`` tensor.""" - return torch.cat( - [sound[:, : shape[0]].permute(1, 0) for sound, shape in zip(tokens_sound, token_shapes_sound)], - dim=0, - ) - - def _unpack_sound_latents( - self, - packed_preds: torch.Tensor, - token_shapes_sound: list[tuple[int, int, int]], - noisy_frame_indexes_sound: list[torch.Tensor], - ) -> list[torch.Tensor]: - """Packed ``[total_noisy_T, C]`` predictions → list of ``[C, T]`` tensors (zeros at conditioned positions).""" - sound_dim = self.config.sound_dim - unpacked: list[torch.Tensor] = [] - start_idx = 0 - for shape, noisy_idxs in zip(token_shapes_sound, noisy_frame_indexes_sound): - T = shape[0] - output = torch.zeros((sound_dim, T), device=packed_preds.device, dtype=packed_preds.dtype) - t_n = len(noisy_idxs) - if t_n > 0: - output[:, noisy_idxs] = packed_preds[start_idx : start_idx + t_n].T - start_idx += t_n - unpacked.append(output) - return unpacked - - def _pack_action_latents( - self, - tokens_action: list[torch.Tensor], - token_shapes_action: list[tuple[int, int, int]], - domain_ids_action: list[torch.Tensor], - ) -> tuple[torch.Tensor, torch.Tensor]: - """List of ``[T, D]`` tensors → packed ``[total_T, D]`` plus per-token domain ids.""" - packed: list[torch.Tensor] = [] - domain_ids: list[torch.Tensor] = [] - for action, shape, domain_id in zip(tokens_action, token_shapes_action, domain_ids_action): - token_count = shape[0] - packed.append(action[:token_count]) - domain_ids.append(domain_id.reshape(1).expand(token_count)) - return torch.cat(packed, dim=0), torch.cat(domain_ids, dim=0) - - def _unpack_action_latents( - self, - packed_preds: torch.Tensor, - token_shapes_action: list[tuple[int, int, int]], - noisy_frame_indexes_action: list[torch.Tensor], - ) -> list[torch.Tensor]: - """Packed ``[total_noisy_T, D]`` predictions → list of ``[T, D]`` tensors.""" - unpacked: list[torch.Tensor] = [] - start_idx = 0 - for shape, noisy_idxs in zip(token_shapes_action, noisy_frame_indexes_action): - T = shape[0] - output = torch.zeros((T, self.action_dim), device=packed_preds.device, dtype=packed_preds.dtype) - t_n = len(noisy_idxs) - if t_n > 0: - output[noisy_idxs] = packed_preds[start_idx : start_idx + t_n] - start_idx += t_n - unpacked.append(output) - return unpacked - - # ------------------------------------------------------------------------- - # forward: full per-step pass — encode text/vision/sound/action → run layers → - # decode vision/sound/action. Pipeline calls this once per CFG pass. - # ------------------------------------------------------------------------- - - def forward( - self, - input_ids: torch.Tensor, - text_indexes: torch.Tensor, - position_ids: torch.Tensor, - und_len: int, - sequence_length: int, - vision_tokens: list[torch.Tensor], - vision_token_shapes: list[tuple[int, int, int]], - vision_sequence_indexes: torch.Tensor, - vision_mse_loss_indexes: torch.Tensor, - vision_timesteps: torch.Tensor, - vision_noisy_frame_indexes: list[torch.Tensor], - sound_tokens: list[torch.Tensor] | None = None, - sound_token_shapes: list[tuple[int, int, int]] | None = None, - sound_sequence_indexes: torch.Tensor | None = None, - sound_mse_loss_indexes: torch.Tensor | None = None, - sound_timesteps: torch.Tensor | None = None, - sound_noisy_frame_indexes: list[torch.Tensor] | None = None, - action_tokens: list[torch.Tensor] | None = None, - action_token_shapes: list[tuple[int, int, int]] | None = None, - action_sequence_indexes: torch.Tensor | None = None, - action_mse_loss_indexes: torch.Tensor | None = None, - action_timesteps: torch.Tensor | None = None, - action_noisy_frame_indexes: list[torch.Tensor] | None = None, - action_domain_ids: list[torch.Tensor] | None = None, - return_dict: bool = True, - ) -> ( - Cosmos3OmniTransformerOutput | tuple[list[torch.Tensor], list[torch.Tensor] | None, list[torch.Tensor] | None] - ): - """Run a full denoising-step forward pass. - - Args: - input_ids: Text token IDs placed at ``text_indexes`` in the joint sequence. - text_indexes: Indices of text tokens in the joint sequence. - position_ids: ``[3, sequence_length]`` mRoPE position IDs for the full joint sequence. - und_len: Length of the causal text (understanding) prefix; generation tokens follow. - sequence_length: Total length of the joint packed sequence. - vision_tokens: Per-item vision latent tensors before patchify. - vision_token_shapes: Patch grid shapes ``(T, H, W)`` per vision item. - vision_sequence_indexes: Indices of vision tokens in the joint sequence. - vision_mse_loss_indexes: Indices used to read vision predictions after the backbone. - vision_timesteps: Per-patch diffusion timesteps for vision tokens. - vision_noisy_frame_indexes: Noisy frame indices per vision item. - sound_tokens: Optional sound latent tensors before packing. - sound_token_shapes: Optional patch grid shapes for sound items. - sound_sequence_indexes: Optional indices of sound tokens in the joint sequence. - sound_mse_loss_indexes: Optional indices used to read sound predictions. - sound_timesteps: Optional per-token diffusion timesteps for sound. - sound_noisy_frame_indexes: Optional noisy frame indices per sound item. - action_tokens: Optional action latent tensors before packing. - action_token_shapes: Optional patch grid shapes ``(T, H, W)`` per action item. - action_sequence_indexes: Optional indices of action tokens in the joint sequence. - action_mse_loss_indexes: Optional indices used to read action predictions after the backbone. - action_timesteps: Optional per-token diffusion timesteps for action tokens. - action_noisy_frame_indexes: Optional noisy frame indices per action item. - action_domain_ids: Optional per-item domain IDs selecting the action head weights. - return_dict: Whether to return a [`Cosmos3OmniTransformerOutput`] instead of a tuple. - - Returns: - A [`Cosmos3OmniTransformerOutput`] or a tuple of per-modality prediction lists. Optional modalities return - ``None`` when their inputs are omitted. - """ - has_sound = sound_tokens is not None and sound_sequence_indexes is not None - has_action = action_tokens is not None and action_sequence_indexes is not None - - # Embed text tokens into the joint hidden_states buffer at their sequence positions. - packed_text_embedding = self.embed_tokens(input_ids) - target_dtype = packed_text_embedding.dtype - hidden_states = packed_text_embedding.new_zeros(size=(sequence_length, self.config.hidden_size)) - hidden_states[text_indexes] = packed_text_embedding - - # Patchify + project vision latents, then add timestep embeddings to noisy frames. - packed_tokens_vision, original_latent_shapes = self._patchify_and_pack_latents(vision_tokens) - packed_tokens_vision = self.proj_in(packed_tokens_vision) - timesteps_vision = vision_timesteps * self.config.timestep_scale - time_embedder_dtype = next(self.time_embedder.parameters()).dtype - packed_timestep_embeds_vision = self.time_embedder(self.time_proj(timesteps_vision).to(time_embedder_dtype)) - packed_timestep_embeds_vision = packed_timestep_embeds_vision.to(target_dtype) - packed_tokens_vision = self._apply_timestep_embeds_to_noisy_tokens( - packed_tokens=packed_tokens_vision, - packed_timestep_embeds=packed_timestep_embeds_vision, - noisy_frame_indexes=vision_noisy_frame_indexes, - token_shapes=vision_token_shapes, - ) - hidden_states[vision_sequence_indexes] = packed_tokens_vision - - # Pack + project sound latents (when present); all sound frames are noisy. - if has_sound: - packed_tokens_sound = self._pack_sound_latents(sound_tokens, sound_token_shapes).to(target_dtype) - packed_tokens_sound = self.audio_proj_in(packed_tokens_sound) + self.audio_modality_embed - timesteps_sound = sound_timesteps * self.config.timestep_scale - packed_timestep_embeds_sound = self.time_embedder(self.time_proj(timesteps_sound).to(time_embedder_dtype)) - packed_timestep_embeds_sound = packed_timestep_embeds_sound.to(target_dtype) - packed_tokens_sound = self._apply_timestep_embeds_to_noisy_tokens( - packed_tokens=packed_tokens_sound, - packed_timestep_embeds=packed_timestep_embeds_sound, - noisy_frame_indexes=sound_noisy_frame_indexes, - token_shapes=sound_token_shapes, - ) - hidden_states[sound_sequence_indexes] = packed_tokens_sound - - # Pack + project action latents (when present). Domain ids select the action head weights. - if has_action: - packed_tokens_action, per_token_domain_ids = self._pack_action_latents( - action_tokens, action_token_shapes, action_domain_ids - ) - packed_tokens_action = packed_tokens_action.to(target_dtype) - per_token_domain_ids = per_token_domain_ids.to(device=packed_tokens_action.device) - packed_tokens_action = self.action_proj_in(packed_tokens_action, per_token_domain_ids) - packed_tokens_action = packed_tokens_action + self.action_modality_embed - if action_mse_loss_indexes.numel() > 0: - timesteps_action = action_timesteps * self.config.timestep_scale - packed_timestep_embeds_action = self.time_embedder( - self.time_proj(timesteps_action).to(time_embedder_dtype) - ) - packed_timestep_embeds_action = packed_timestep_embeds_action.to(target_dtype) - packed_tokens_action = self._apply_timestep_embeds_to_noisy_tokens( - packed_tokens=packed_tokens_action, - packed_timestep_embeds=packed_timestep_embeds_action, - noisy_frame_indexes=action_noisy_frame_indexes, - token_shapes=action_token_shapes, - ) - hidden_states[action_sequence_indexes] = packed_tokens_action - - # Compute rotary embeddings once for the joint sequence, then slice into und/gen halves. - _meta_tensor = torch.tensor([], dtype=hidden_states.dtype, device=hidden_states.device) - cos, sin = self.rotary_emb( - position_ids=position_ids.unsqueeze(0) if position_ids.ndim == 1 else position_ids.unsqueeze(1), - device=hidden_states.device, - dtype=hidden_states.dtype, - ) - # cos, sin: [1, N, head_dim] (1-D pos_ids) or [3, 1, N, head_dim] (mrope pos_ids) - cos = cos.squeeze(0) - sin = sin.squeeze(0) - - und_seq = hidden_states[:und_len] - gen_seq = hidden_states[und_len:] - rotary_emb = (cos[:und_len], sin[:und_len], cos[und_len:], sin[und_len:]) - - # Optional context-parallelism shard seam (no-op unless set by an external - # helper, e.g. `examples/cosmos3/cosmos_parallel.py`). When set, it - # shards each pathway's sequence and rotary embeddings across a device mesh, so - # the decoder layers below run on local sequence shards. - if self._cp_shard_fn is not None: - und_seq, gen_seq, rotary_emb = self._cp_shard_fn(und_seq, gen_seq, rotary_emb) - - for decoder_layer in self.layers: - if torch.is_grad_enabled() and self.gradient_checkpointing: - und_seq, gen_seq = self._gradient_checkpointing_func( - decoder_layer.__call__, und_seq, gen_seq, rotary_emb - ) - else: - und_seq, gen_seq = decoder_layer(und_seq, gen_seq, rotary_emb) - und_out = self.norm(und_seq) - gen_out = self.norm_moe_gen(gen_seq) - - # Optional context-parallelism gather seam: re-gather the full per-pathway - # sequence on every rank (and drop the padding) before the global-index decode - # below, since the downstream indexes address positions in the unpadded joint - # sequence. No-op unless `_cp_shard_fn`'s counterpart is set. - if self._cp_gather_fn is not None: - und_out, gen_out = self._cp_gather_fn(und_out, gen_out) - - last_hidden_state = torch.cat([und_out, gen_out], dim=0) - - # Decode vision predictions from the joint hidden state. - preds_vision_packed = self.proj_out(last_hidden_state[vision_mse_loss_indexes]) - preds_vision = self._unpatchify_and_unpack_latents( - preds_vision_packed, - token_shapes_vision=vision_token_shapes, - noisy_frame_indexes_vision=vision_noisy_frame_indexes, - original_latent_shapes=original_latent_shapes, - ) - - preds_sound: list[torch.Tensor] | None = None - if has_sound: - preds_sound_packed = self.audio_proj_out(last_hidden_state[sound_mse_loss_indexes]) - preds_sound = self._unpack_sound_latents(preds_sound_packed, sound_token_shapes, sound_noisy_frame_indexes) - - preds_action: list[torch.Tensor] | None = None - if has_action: - per_noisy_domain_ids = [ - domain_id.reshape(1).expand(len(noisy_idxs)) - for domain_id, noisy_idxs in zip(action_domain_ids, action_noisy_frame_indexes) - ] - per_noisy_domain_ids = torch.cat(per_noisy_domain_ids, dim=0).to(device=last_hidden_state.device) - preds_action_packed = self.action_proj_out( - last_hidden_state[action_mse_loss_indexes], per_noisy_domain_ids - ) - preds_action = self._unpack_action_latents( - preds_action_packed, action_token_shapes, action_noisy_frame_indexes - ) - - if not return_dict: - return preds_vision, preds_sound, preds_action - - return Cosmos3OmniTransformerOutput(sample=preds_vision, sound=preds_sound, action=preds_action) diff --git a/diffusers/models/transformers/transformer_easyanimate.py b/diffusers/models/transformers/transformer_easyanimate.py deleted file mode 100644 index 24c874ad40ef1a1ddb3241010c5697a579706bdc..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_easyanimate.py +++ /dev/null @@ -1,552 +0,0 @@ -# Copyright 2025 The EasyAnimate team and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -import torch.nn.functional as F -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import Attention, FeedForward -from ..embeddings import TimestepEmbedding, Timesteps, get_3d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNorm, FP32LayerNorm, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class EasyAnimateLayerNormZero(nn.Module): - def __init__( - self, - conditioning_dim: int, - embedding_dim: int, - elementwise_affine: bool = True, - eps: float = 1e-5, - bias: bool = True, - norm_type: str = "fp32_layer_norm", - ) -> None: - super().__init__() - - self.silu = nn.SiLU() - self.linear = nn.Linear(conditioning_dim, 6 * embedding_dim, bias=bias) - - if norm_type == "layer_norm": - self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=elementwise_affine, eps=eps) - elif norm_type == "fp32_layer_norm": - self.norm = FP32LayerNorm(embedding_dim, elementwise_affine=elementwise_affine, eps=eps) - else: - raise ValueError( - f"Unsupported `norm_type` ({norm_type}) provided. Supported ones are: 'layer_norm', 'fp32_layer_norm'." - ) - - def forward( - self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - shift, scale, gate, enc_shift, enc_scale, enc_gate = self.linear(self.silu(temb)).chunk(6, dim=1) - hidden_states = self.norm(hidden_states) * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) - encoder_hidden_states = self.norm(encoder_hidden_states) * (1 + enc_scale.unsqueeze(1)) + enc_shift.unsqueeze( - 1 - ) - return hidden_states, encoder_hidden_states, gate, enc_gate - - -class EasyAnimateRotaryPosEmbed(nn.Module): - def __init__(self, patch_size: int, rope_dim: list[int]) -> None: - super().__init__() - - self.patch_size = patch_size - self.rope_dim = rope_dim - - def get_resize_crop_region_for_grid(self, src, tgt_width, tgt_height): - tw = tgt_width - th = tgt_height - h, w = src - r = h / w - if r > (th / tw): - resize_height = th - resize_width = int(round(th / h * w)) - else: - resize_width = tw - resize_height = int(round(tw / w * h)) - - crop_top = int(round((th - resize_height) / 2.0)) - crop_left = int(round((tw - resize_width) / 2.0)) - - return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - bs, c, num_frames, grid_height, grid_width = hidden_states.size() - grid_height = grid_height // self.patch_size - grid_width = grid_width // self.patch_size - base_size_width = 90 // self.patch_size - base_size_height = 60 // self.patch_size - - grid_crops_coords = self.get_resize_crop_region_for_grid( - (grid_height, grid_width), base_size_width, base_size_height - ) - image_rotary_emb = get_3d_rotary_pos_embed( - self.rope_dim, - grid_crops_coords, - grid_size=(grid_height, grid_width), - temporal_size=hidden_states.size(2), - use_real=True, - ) - return image_rotary_emb - - -class EasyAnimateAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is - used in the EasyAnimateTransformer3DModel model. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "EasyAnimateAttnProcessor2_0 requires PyTorch 2.0 or above. To use it, please install PyTorch 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - if attn.add_q_proj is None and encoder_hidden_states is not None: - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - # 1. QKV projections - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2) - key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2) - value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2) - - # 2. QK normalization - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # 3. Encoder condition QKV projection and normalization - if attn.add_q_proj is not None and encoder_hidden_states is not None: - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - encoder_query = encoder_query.unflatten(2, (attn.heads, -1)).transpose(1, 2) - encoder_key = encoder_key.unflatten(2, (attn.heads, -1)).transpose(1, 2) - encoder_value = encoder_value.unflatten(2, (attn.heads, -1)).transpose(1, 2) - - if attn.norm_added_q is not None: - encoder_query = attn.norm_added_q(encoder_query) - if attn.norm_added_k is not None: - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([encoder_query, query], dim=2) - key = torch.cat([encoder_key, key], dim=2) - value = torch.cat([encoder_value, value], dim=2) - - if image_rotary_emb is not None: - from ..embeddings import apply_rotary_emb - - query[:, :, encoder_hidden_states.shape[1] :] = apply_rotary_emb( - query[:, :, encoder_hidden_states.shape[1] :], image_rotary_emb - ) - if not attn.is_cross_attention: - key[:, :, encoder_hidden_states.shape[1] :] = apply_rotary_emb( - key[:, :, encoder_hidden_states.shape[1] :], image_rotary_emb - ) - - # 5. Attention - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - hidden_states = hidden_states.transpose(1, 2).flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - # 6. Output projection - if encoder_hidden_states is not None: - encoder_hidden_states, hidden_states = ( - hidden_states[:, : encoder_hidden_states.shape[1]], - hidden_states[:, encoder_hidden_states.shape[1] :], - ) - - if getattr(attn, "to_out", None) is not None: - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - if getattr(attn, "to_add_out", None) is not None: - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - else: - if getattr(attn, "to_out", None) is not None: - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - return hidden_states, encoder_hidden_states - - -@maybe_allow_in_graph -class EasyAnimateTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - time_embed_dim: int, - dropout: float = 0.0, - activation_fn: str = "gelu-approximate", - norm_elementwise_affine: bool = True, - norm_eps: float = 1e-6, - final_dropout: bool = True, - ff_inner_dim: int | None = None, - ff_bias: bool = True, - qk_norm: bool = True, - after_norm: bool = False, - norm_type: str = "fp32_layer_norm", - is_mmdit_block: bool = True, - ): - super().__init__() - - # Attention Part - self.norm1 = EasyAnimateLayerNormZero( - time_embed_dim, dim, norm_elementwise_affine, norm_eps, norm_type=norm_type, bias=True - ) - - self.attn1 = Attention( - query_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - qk_norm="layer_norm" if qk_norm else None, - eps=1e-6, - bias=True, - added_proj_bias=True, - added_kv_proj_dim=dim if is_mmdit_block else None, - context_pre_only=False if is_mmdit_block else None, - processor=EasyAnimateAttnProcessor2_0(), - ) - - # FFN Part - self.norm2 = EasyAnimateLayerNormZero( - time_embed_dim, dim, norm_elementwise_affine, norm_eps, norm_type=norm_type, bias=True - ) - self.ff = FeedForward( - dim, - dropout=dropout, - activation_fn=activation_fn, - final_dropout=final_dropout, - inner_dim=ff_inner_dim, - bias=ff_bias, - ) - - self.txt_ff = None - if is_mmdit_block: - self.txt_ff = FeedForward( - dim, - dropout=dropout, - activation_fn=activation_fn, - final_dropout=final_dropout, - inner_dim=ff_inner_dim, - bias=ff_bias, - ) - - self.norm3 = None - if after_norm: - self.norm3 = FP32LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - # 1. Attention - norm_hidden_states, norm_encoder_hidden_states, gate_msa, enc_gate_msa = self.norm1( - hidden_states, encoder_hidden_states, temb - ) - attn_hidden_states, attn_encoder_hidden_states = self.attn1( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - ) - hidden_states = hidden_states + gate_msa.unsqueeze(1) * attn_hidden_states - encoder_hidden_states = encoder_hidden_states + enc_gate_msa.unsqueeze(1) * attn_encoder_hidden_states - - # 2. Feed-forward - norm_hidden_states, norm_encoder_hidden_states, gate_ff, enc_gate_ff = self.norm2( - hidden_states, encoder_hidden_states, temb - ) - if self.norm3 is not None: - norm_hidden_states = self.norm3(self.ff(norm_hidden_states)) - if self.txt_ff is not None: - norm_encoder_hidden_states = self.norm3(self.txt_ff(norm_encoder_hidden_states)) - else: - norm_encoder_hidden_states = self.norm3(self.ff(norm_encoder_hidden_states)) - else: - norm_hidden_states = self.ff(norm_hidden_states) - if self.txt_ff is not None: - norm_encoder_hidden_states = self.txt_ff(norm_encoder_hidden_states) - else: - norm_encoder_hidden_states = self.ff(norm_encoder_hidden_states) - hidden_states = hidden_states + gate_ff.unsqueeze(1) * norm_hidden_states - encoder_hidden_states = encoder_hidden_states + enc_gate_ff.unsqueeze(1) * norm_encoder_hidden_states - return hidden_states, encoder_hidden_states - - -class EasyAnimateTransformer3DModel(ModelMixin, ConfigMixin): - """ - A Transformer model for video-like data in [EasyAnimate](https://github.com/aigc-apps/EasyAnimate). - - Parameters: - num_attention_heads (`int`, defaults to `48`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `64`): - The number of channels in each head. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `16`): - The number of channels in the output. - patch_size (`int`, defaults to `2`): - The size of the patches to use in the patch embedding layer. - sample_width (`int`, defaults to `90`): - The width of the input latents. - sample_height (`int`, defaults to `60`): - The height of the input latents. - activation_fn (`str`, defaults to `"gelu-approximate"`): - Activation function to use in feed-forward. - timestep_activation_fn (`str`, defaults to `"silu"`): - Activation function to use when generating the timestep embeddings. - num_layers (`int`, defaults to `30`): - The number of layers of Transformer blocks to use. - mmdit_layers (`int`, defaults to `1000`): - The number of layers of Multi Modal Transformer blocks to use. - dropout (`float`, defaults to `0.0`): - The dropout probability to use. - time_embed_dim (`int`, defaults to `512`): - Output dimension of timestep embeddings. - text_embed_dim (`int`, defaults to `4096`): - Input dimension of text embeddings from the text encoder. - norm_eps (`float`, defaults to `1e-5`): - The epsilon value to use in normalization layers. - norm_elementwise_affine (`bool`, defaults to `True`): - Whether to use elementwise affine in normalization layers. - flip_sin_to_cos (`bool`, defaults to `True`): - Whether to flip the sin to cos in the time embedding. - time_position_encoding_type (`str`, defaults to `3d_rope`): - Type of time position encoding. - after_norm (`bool`, defaults to `False`): - Flag to apply normalization after. - resize_inpaint_mask_directly (`bool`, defaults to `True`): - Flag to resize inpaint mask directly. - enable_text_attention_mask (`bool`, defaults to `True`): - Flag to enable text attention mask. - add_noise_in_inpaint_model (`bool`, defaults to `False`): - Flag to add noise in inpaint model. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["EasyAnimateTransformerBlock"] - _skip_layerwise_casting_patterns = ["^proj$", "norm", "^proj_out$"] - - @register_to_config - def __init__( - self, - num_attention_heads: int = 48, - attention_head_dim: int = 64, - in_channels: int | None = None, - out_channels: int | None = None, - patch_size: int | None = None, - sample_width: int = 90, - sample_height: int = 60, - activation_fn: str = "gelu-approximate", - timestep_activation_fn: str = "silu", - freq_shift: int = 0, - num_layers: int = 48, - mmdit_layers: int = 48, - dropout: float = 0.0, - time_embed_dim: int = 512, - add_norm_text_encoder: bool = False, - text_embed_dim: int = 3584, - text_embed_dim_t5: int = None, - norm_eps: float = 1e-5, - norm_elementwise_affine: bool = True, - flip_sin_to_cos: bool = True, - time_position_encoding_type: str = "3d_rope", - after_norm=False, - resize_inpaint_mask_directly: bool = True, - enable_text_attention_mask: bool = True, - add_noise_in_inpaint_model: bool = True, - ): - super().__init__() - inner_dim = num_attention_heads * attention_head_dim - - # 1. Timestep embedding - self.time_proj = Timesteps(inner_dim, flip_sin_to_cos, freq_shift) - self.time_embedding = TimestepEmbedding(inner_dim, time_embed_dim, timestep_activation_fn) - self.rope_embedding = EasyAnimateRotaryPosEmbed(patch_size, attention_head_dim) - - # 2. Patch embedding - self.proj = nn.Conv2d( - in_channels, inner_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=True - ) - - # 3. Text refined embedding - self.text_proj = None - self.text_proj_t5 = None - if not add_norm_text_encoder: - self.text_proj = nn.Linear(text_embed_dim, inner_dim) - if text_embed_dim_t5 is not None: - self.text_proj_t5 = nn.Linear(text_embed_dim_t5, inner_dim) - else: - self.text_proj = nn.Sequential( - RMSNorm(text_embed_dim, 1e-6, elementwise_affine=True), nn.Linear(text_embed_dim, inner_dim) - ) - if text_embed_dim_t5 is not None: - self.text_proj_t5 = nn.Sequential( - RMSNorm(text_embed_dim, 1e-6, elementwise_affine=True), nn.Linear(text_embed_dim_t5, inner_dim) - ) - - # 4. Transformer blocks - self.transformer_blocks = nn.ModuleList( - [ - EasyAnimateTransformerBlock( - dim=inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - time_embed_dim=time_embed_dim, - dropout=dropout, - activation_fn=activation_fn, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - after_norm=after_norm, - is_mmdit_block=True if _ < mmdit_layers else False, - ) - for _ in range(num_layers) - ] - ) - self.norm_final = nn.LayerNorm(inner_dim, norm_eps, norm_elementwise_affine) - - # 5. Output norm & projection - self.norm_out = AdaLayerNorm( - embedding_dim=time_embed_dim, - output_dim=2 * inner_dim, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - chunk_dim=1, - ) - self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - timestep_cond: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - encoder_hidden_states_t5: torch.Tensor | None = None, - inpaint_latents: torch.Tensor | None = None, - control_latents: torch.Tensor | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`EasyAnimateTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - timestep_cond (`torch.Tensor`, *optional*): - Conditional embeddings for timestep. If provided, the embeddings will be summed with the samples passed - through the `self.time_embedding` layer to obtain the final timestep embeddings. - encoder_hidden_states (`torch.Tensor`, *optional*): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_hidden_states_t5 (`torch.Tensor`, *optional*): - Additional conditional embeddings computed from a T5 text encoder. - inpaint_latents (`torch.Tensor`, *optional*): - Latents concatenated to `hidden_states` for inpainting variants of the model. - control_latents (`torch.Tensor`, *optional*): - Latents concatenated to `hidden_states` for control variants of the model. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, channels, video_length, height, width = hidden_states.size() - p = self.config.patch_size - post_patch_height = height // p - post_patch_width = width // p - - # 1. Time embedding - temb = self.time_proj(timestep).to(dtype=hidden_states.dtype) - temb = self.time_embedding(temb, timestep_cond) - image_rotary_emb = self.rope_embedding(hidden_states) - - # 2. Patch embedding - if inpaint_latents is not None: - hidden_states = torch.concat([hidden_states, inpaint_latents], 1) - if control_latents is not None: - hidden_states = torch.concat([hidden_states, control_latents], 1) - - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) # [B, C, F, H, W] -> [BF, C, H, W] - hidden_states = self.proj(hidden_states) - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute( - 0, 2, 1, 3, 4 - ) # [BF, C, H, W] -> [B, F, C, H, W] - hidden_states = hidden_states.flatten(2, 4).transpose(1, 2) # [B, F, C, H, W] -> [B, FHW, C] - - # 3. Text embedding - encoder_hidden_states = self.text_proj(encoder_hidden_states) - if encoder_hidden_states_t5 is not None: - encoder_hidden_states_t5 = self.text_proj_t5(encoder_hidden_states_t5) - encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states_t5], dim=1).contiguous() - - # 4. Transformer blocks - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, hidden_states, encoder_hidden_states, temb, image_rotary_emb - ) - else: - hidden_states, encoder_hidden_states = block( - hidden_states, encoder_hidden_states, temb, image_rotary_emb - ) - - hidden_states = self.norm_final(hidden_states) - - # 5. Output norm & projection - hidden_states = self.norm_out(hidden_states, temb=temb) - hidden_states = self.proj_out(hidden_states) - - # 6. Unpatchify - p = self.config.patch_size - output = hidden_states.reshape(batch_size, video_length, post_patch_height, post_patch_width, channels, p, p) - output = output.permute(0, 4, 1, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_ernie_image.py b/diffusers/models/transformers/transformer_ernie_image.py deleted file mode 100644 index 0abc5d254bb2a7014f2fc878ab97df1c004bf9f0..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_ernie_image.py +++ /dev/null @@ -1,453 +0,0 @@ -# Copyright 2025 Baidu ERNIE-Image Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -""" -Ernie-Image Transformer2DModel for HuggingFace Diffusers. -""" - -import inspect -from dataclasses import dataclass -from typing import Tuple - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import BaseOutput, logging -from ..attention import AttentionModuleMixin -from ..attention_dispatch import dispatch_attention_fn -from ..attention_processor import Attention -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from ..normalization import RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class ErnieImageTransformer2DModelOutput(BaseOutput): - sample: torch.Tensor - - -def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor: - assert dim % 2 == 0 - scale = torch.arange(0, dim, 2, dtype=torch.float32, device=pos.device) / dim - omega = 1.0 / (theta**scale) - # Disable autocast so the position-id einsum runs in float32: under an ambient autocast it would run in - # bfloat16, which cannot represent consecutive integers past 256, so position ids beyond that point would - # collapse onto the same frequency and degrade the rotary embedding. - with torch.autocast(device_type=pos.device.type, enabled=False): - out = torch.einsum("...n,d->...nd", pos, omega) - return out.float() - - -class ErnieImageEmbedND3(nn.Module): - def __init__(self, dim: int, theta: int, axes_dim: Tuple[int, int, int]): - super().__init__() - self.dim = dim - self.theta = theta - self.axes_dim = list(axes_dim) - - def forward(self, ids: torch.Tensor) -> torch.Tensor: - emb = torch.cat([rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(3)], dim=-1) - emb = emb.unsqueeze(2) # [B, S, 1, head_dim//2] - return torch.stack([emb, emb], dim=-1).reshape(*emb.shape[:-1], -1) # [B, S, 1, head_dim] - - -class ErnieImagePatchEmbedDynamic(nn.Module): - def __init__(self, in_channels: int, embed_dim: int, patch_size: int): - super().__init__() - self.patch_size = patch_size - self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size, bias=True) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = self.proj(x) - batch_size, dim, height, width = x.shape - return x.reshape(batch_size, dim, height * width).transpose(1, 2).contiguous() - - -class ErnieImageSingleStreamAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "ErnieImageSingleStreamAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to version 2.0 or higher." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - freqs_cis: torch.Tensor | None = None, - ) -> torch.Tensor: - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - # Apply Norms - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE: same rotate_half logic as Megatron _apply_rotary_pos_emb_bshd (rotary_interleaved=False) - # x_in: [B, S, heads, head_dim], freqs_cis: [B, S, 1, head_dim] with angles [θ0,θ0,θ1,θ1,...] - def apply_rotary_emb(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: - rot_dim = freqs_cis.shape[-1] - x, x_pass = x_in[..., :rot_dim], x_in[..., rot_dim:] - cos_ = torch.cos(freqs_cis).to(x.dtype) - sin_ = torch.sin(freqs_cis).to(x.dtype) - # Non-interleaved rotate_half: [-x2, x1] - x1, x2 = x.chunk(2, dim=-1) - x_rotated = torch.cat((-x2, x1), dim=-1) - return torch.cat((x * cos_ + x_rotated * sin_, x_pass), dim=-1) - - if freqs_cis is not None: - query = apply_rotary_emb(query, freqs_cis) - key = apply_rotary_emb(key, freqs_cis) - - # Cast to correct dtype - dtype = query.dtype - query, key = query.to(dtype), key.to(dtype) - - # From [batch, seq_len] to [batch, 1, 1, seq_len] -> broadcast to [batch, heads, seq_len, seq_len] - if attention_mask is not None and attention_mask.ndim == 2: - attention_mask = attention_mask[:, None, None, :] - - # Compute joint attention - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - # Reshape back - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(dtype) - output = attn.to_out[0](hidden_states) - - return output - - -class ErnieImageAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = ErnieImageSingleStreamAttnProcessor - - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - qk_norm: str = "rms_norm", - added_proj_bias: bool | None = True, - out_bias: bool = True, - eps: float = 1e-5, - out_dim: int = None, - elementwise_affine: bool = True, - processor=None, - ): - super().__init__() - - self.head_dim = dim_head - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.query_dim = query_dim - self.out_dim = out_dim if out_dim is not None else query_dim - self.heads = out_dim // dim_head if out_dim is not None else heads - - self.use_bias = bias - self.dropout = dropout - - self.added_proj_bias = added_proj_bias - - self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_v = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - - # QK Norm - if qk_norm == "layer_norm": - self.norm_q = torch.nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = torch.nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - elif qk_norm == "rms_norm": - self.norm_q = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - else: - raise ValueError( - f"unknown qk_norm: {qk_norm}. Should be one of None, 'layer_norm', 'fp32_layer_norm', 'layer_norm_across_heads', 'rms_norm', 'rms_norm_across_heads', 'l2'." - ) - - self.to_out = torch.nn.ModuleList([]) - self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"joint_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -class ErnieImageFeedForward(nn.Module): - def __init__(self, hidden_size: int, ffn_hidden_size: int): - super().__init__() - # Separate gate and up projections (matches converted weights) - self.gate_proj = nn.Linear(hidden_size, ffn_hidden_size, bias=False) - self.up_proj = nn.Linear(hidden_size, ffn_hidden_size, bias=False) - self.linear_fc2 = nn.Linear(ffn_hidden_size, hidden_size, bias=False) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.linear_fc2(self.up_proj(x) * F.gelu(self.gate_proj(x))) - - -class ErnieImageSharedAdaLNBlock(nn.Module): - def __init__( - self, hidden_size: int, num_heads: int, ffn_hidden_size: int, eps: float = 1e-6, qk_layernorm: bool = True - ): - super().__init__() - self.adaLN_sa_ln = RMSNorm(hidden_size, eps=eps) - self.self_attention = ErnieImageAttention( - query_dim=hidden_size, - dim_head=hidden_size // num_heads, - heads=num_heads, - qk_norm="rms_norm" if qk_layernorm else None, - eps=eps, - bias=False, - out_bias=False, - processor=ErnieImageSingleStreamAttnProcessor(), - ) - self.adaLN_mlp_ln = RMSNorm(hidden_size, eps=eps) - self.mlp = ErnieImageFeedForward(hidden_size, ffn_hidden_size) - - def forward( - self, - x, - rotary_pos_emb, - temb: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor], - attention_mask: torch.Tensor | None = None, - ): - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = temb - residual = x - x = self.adaLN_sa_ln(x) - x = (x.float() * (1 + scale_msa.float()) + shift_msa.float()).to(x.dtype) - x_bsh = x.permute(1, 0, 2) # [S, B, H] → [B, S, H] for diffusers Attention (batch-first) - attn_out = self.self_attention(x_bsh, attention_mask=attention_mask, image_rotary_emb=rotary_pos_emb) - attn_out = attn_out.permute(1, 0, 2) # [B, S, H] → [S, B, H] - x = residual + (gate_msa.float() * attn_out.float()).to(x.dtype) - residual = x - x = self.adaLN_mlp_ln(x) - x = (x.float() * (1 + scale_mlp.float()) + shift_mlp.float()).to(x.dtype) - return residual + (gate_mlp.float() * self.mlp(x).float()).to(x.dtype) - - -class ErnieImageAdaLNContinuous(nn.Module): - def __init__(self, hidden_size: int, eps: float = 1e-6): - super().__init__() - self.norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=eps) - self.linear = nn.Linear(hidden_size, hidden_size * 2) - - def forward(self, x: torch.Tensor, conditioning: torch.Tensor) -> torch.Tensor: - scale, shift = self.linear(conditioning).chunk(2, dim=-1) - x = self.norm(x) - # Broadcast conditioning to sequence dimension - x = x * (1 + scale.unsqueeze(0)) + shift.unsqueeze(0) - return x - - -class ErnieImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): - _supports_gradient_checkpointing = True - _repeated_blocks = ["ErnieImageSharedAdaLNBlock"] - - @register_to_config - def __init__( - self, - hidden_size: int = 3072, - num_attention_heads: int = 24, - num_layers: int = 24, - ffn_hidden_size: int = 8192, - in_channels: int = 128, - out_channels: int = 128, - patch_size: int = 1, - text_in_dim: int = 2560, - rope_theta: int = 256, - rope_axes_dim: Tuple[int, int, int] = (32, 48, 48), - eps: float = 1e-6, - qk_layernorm: bool = True, - ): - super().__init__() - self.hidden_size = hidden_size - self.num_heads = num_attention_heads - self.head_dim = hidden_size // num_attention_heads - self.num_layers = num_layers - self.patch_size = patch_size - self.in_channels = in_channels - self.out_channels = out_channels - self.text_in_dim = text_in_dim - - self.x_embedder = ErnieImagePatchEmbedDynamic(in_channels, hidden_size, patch_size) - self.text_proj = nn.Linear(text_in_dim, hidden_size, bias=False) if text_in_dim != hidden_size else None - self.time_proj = Timesteps(hidden_size, flip_sin_to_cos=False, downscale_freq_shift=0) - self.time_embedding = TimestepEmbedding(hidden_size, hidden_size) - self.pos_embed = ErnieImageEmbedND3(dim=self.head_dim, theta=rope_theta, axes_dim=rope_axes_dim) - self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size)) - nn.init.zeros_(self.adaLN_modulation[-1].weight) - nn.init.zeros_(self.adaLN_modulation[-1].bias) - self.layers = nn.ModuleList( - [ - ErnieImageSharedAdaLNBlock( - hidden_size, num_attention_heads, ffn_hidden_size, eps, qk_layernorm=qk_layernorm - ) - for _ in range(num_layers) - ] - ) - self.final_norm = ErnieImageAdaLNContinuous(hidden_size, eps) - self.final_linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels) - nn.init.zeros_(self.final_linear.weight) - nn.init.zeros_(self.final_linear.bias) - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - # encoder_hidden_states: List[torch.Tensor], - text_bth: torch.Tensor, - text_lens: torch.Tensor, - return_dict: bool = True, - ): - """ - The [`ErnieImageTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, in_channels, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - text_bth (`torch.Tensor`): - Conditional text embeddings (embeddings computed from the input conditions such as prompts) to use, - shaped `(batch_size, text_length, embed_dims)`. - text_lens (`torch.Tensor`): - Per-sample text sequence lengths used to build the attention mask. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - """ - device, dtype = hidden_states.device, hidden_states.dtype - B, C, H, W = hidden_states.shape - p, Hp, Wp = self.patch_size, H // self.patch_size, W // self.patch_size - N_img = Hp * Wp - - img_sbh = self.x_embedder(hidden_states).transpose(0, 1).contiguous() - # text_bth, text_lens = self._pad_text(encoder_hidden_states, device, dtype) - if self.text_proj is not None and text_bth.numel() > 0: - text_bth = self.text_proj(text_bth) - Tmax = text_bth.shape[1] - text_sbh = text_bth.transpose(0, 1).contiguous() - - x = torch.cat([img_sbh, text_sbh], dim=0) - S = x.shape[0] - - # Position IDs - text_ids = ( - torch.cat( - [ - torch.arange(Tmax, device=device, dtype=torch.float32).view(1, Tmax, 1).expand(B, -1, -1), - torch.zeros((B, Tmax, 2), device=device), - ], - dim=-1, - ) - if Tmax > 0 - else torch.zeros((B, 0, 3), device=device) - ) - grid_yx = torch.stack( - torch.meshgrid( - torch.arange(Hp, device=device, dtype=torch.float32), - torch.arange(Wp, device=device, dtype=torch.float32), - indexing="ij", - ), - dim=-1, - ).reshape(-1, 2) - image_ids = torch.cat( - [text_lens.float().view(B, 1, 1).expand(-1, N_img, -1), grid_yx.view(1, N_img, 2).expand(B, -1, -1)], - dim=-1, - ) - rotary_pos_emb = self.pos_embed(torch.cat([image_ids, text_ids], dim=1)) - - # Attention mask: True = valid (attend), False = padding (mask out), matches sdpa bool convention - valid_text = ( - torch.arange(Tmax, device=device).view(1, Tmax) < text_lens.view(B, 1) - if Tmax > 0 - else torch.zeros((B, 0), device=device, dtype=torch.bool) - ) - attention_mask = torch.cat([torch.ones((B, N_img), device=device, dtype=torch.bool), valid_text], dim=1)[ - :, None, None, : - ] - - # AdaLN - sample = self.time_proj(timestep) - sample = sample.to(dtype=dtype) - c = self.time_embedding(sample) - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = [ - t.unsqueeze(0).expand(S, -1, -1).contiguous() for t in self.adaLN_modulation(c).chunk(6, dim=-1) - ] - for layer in self.layers: - temb = [shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp] - if torch.is_grad_enabled() and self.gradient_checkpointing: - x = self._gradient_checkpointing_func( - layer, - x, - rotary_pos_emb, - temb, - attention_mask, - ) - else: - x = layer(x, rotary_pos_emb, temb, attention_mask) - x = self.final_norm(x, c).type_as(x) - patches = self.final_linear(x)[:N_img].transpose(0, 1).contiguous() - output = ( - patches.view(B, Hp, Wp, p, p, self.out_channels) - .permute(0, 5, 1, 3, 2, 4) - .contiguous() - .view(B, self.out_channels, H, W) - ) - - return ErnieImageTransformer2DModelOutput(sample=output) if return_dict else (output,) diff --git a/diffusers/models/transformers/transformer_flux.py b/diffusers/models/transformers/transformer_flux.py deleted file mode 100644 index 94857dffacb29cf591c1b9404e1c68ed412fbfcb..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_flux.py +++ /dev/null @@ -1,786 +0,0 @@ -# Copyright 2025 Black Forest Labs, The HuggingFace Team and The InstantX Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -from typing import Any - -import numpy as np -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FluxTransformer2DLoadersMixin, FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device, maybe_allow_in_graph -from .._modeling_parallel import ContextParallelInput, ContextParallelOutput -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import ( - CombinedTimestepGuidanceTextProjEmbeddings, - CombinedTimestepTextProjEmbeddings, - apply_rotary_emb, - get_1d_rotary_pos_embed, -) -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _get_projections(attn: "FluxAttention", hidden_states, encoder_hidden_states=None): - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - encoder_query = encoder_key = encoder_value = None - if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None: - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_fused_projections(attn: "FluxAttention", hidden_states, encoder_hidden_states=None): - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - - encoder_query = encoder_key = encoder_value = (None,) - if encoder_hidden_states is not None and hasattr(attn, "to_added_qkv"): - encoder_query, encoder_key, encoder_value = attn.to_added_qkv(encoder_hidden_states).chunk(3, dim=-1) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_qkv_projections(attn: "FluxAttention", hidden_states, encoder_hidden_states=None): - if attn.fused_projections: - return _get_fused_projections(attn, hidden_states, encoder_hidden_states) - return _get_projections(attn, hidden_states, encoder_hidden_states) - - -class FluxAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError(f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version.") - - def __call__( - self, - attn: "FluxAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - query, key, value, encoder_query, encoder_key, encoder_value = _get_qkv_projections( - attn, hidden_states, encoder_hidden_states - ) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if attn.added_kv_proj_dim is not None: - encoder_query = encoder_query.unflatten(-1, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(-1, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(-1, (attn.heads, -1)) - - encoder_query = attn.norm_added_q(encoder_query) - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - if encoder_hidden_states is not None: - encoder_hidden_states, hidden_states = hidden_states.split_with_sizes( - [encoder_hidden_states.shape[1], hidden_states.shape[1] - encoder_hidden_states.shape[1]], dim=1 - ) - hidden_states = attn.to_out[0](hidden_states.contiguous()) - hidden_states = attn.to_out[1](hidden_states) - encoder_hidden_states = attn.to_add_out(encoder_hidden_states.contiguous()) - - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -class FluxIPAdapterAttnProcessor(torch.nn.Module): - """Flux Attention processor for IP-Adapter.""" - - _attention_backend = None - _parallel_config = None - - def __init__( - self, hidden_size: int, cross_attention_dim: int, num_tokens=(4,), scale=1.0, device=None, dtype=None - ): - super().__init__() - - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - f"{self.__class__.__name__} requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - self.hidden_size = hidden_size - self.cross_attention_dim = cross_attention_dim - - if not isinstance(num_tokens, (tuple, list)): - num_tokens = [num_tokens] - - if not isinstance(scale, list): - scale = [scale] * len(num_tokens) - if len(scale) != len(num_tokens): - raise ValueError("`scale` should be a list of integers with the same length as `num_tokens`.") - self.scale = scale - - self.to_k_ip = nn.ModuleList( - [ - nn.Linear(cross_attention_dim, hidden_size, bias=True, device=device, dtype=dtype) - for _ in range(len(num_tokens)) - ] - ) - self.to_v_ip = nn.ModuleList( - [ - nn.Linear(cross_attention_dim, hidden_size, bias=True, device=device, dtype=dtype) - for _ in range(len(num_tokens)) - ] - ) - - def __call__( - self, - attn: "FluxAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ip_hidden_states: list[torch.Tensor] | None = None, - ip_adapter_masks: torch.Tensor | None = None, - ) -> torch.Tensor: - batch_size = hidden_states.shape[0] - - query, key, value, encoder_query, encoder_key, encoder_value = _get_qkv_projections( - attn, hidden_states, encoder_hidden_states - ) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - ip_query = query - - if encoder_hidden_states is not None: - encoder_query = encoder_query.unflatten(-1, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(-1, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(-1, (attn.heads, -1)) - - encoder_query = attn.norm_added_q(encoder_query) - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - if encoder_hidden_states is not None: - encoder_hidden_states, hidden_states = hidden_states.split_with_sizes( - [encoder_hidden_states.shape[1], hidden_states.shape[1] - encoder_hidden_states.shape[1]], dim=1 - ) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - # IP-adapter - ip_attn_output = torch.zeros_like(hidden_states) - - for current_ip_hidden_states, scale, to_k_ip, to_v_ip in zip( - ip_hidden_states, self.scale, self.to_k_ip, self.to_v_ip - ): - ip_key = to_k_ip(current_ip_hidden_states) - ip_value = to_v_ip(current_ip_hidden_states) - - ip_key = ip_key.view(batch_size, -1, attn.heads, attn.head_dim) - ip_value = ip_value.view(batch_size, -1, attn.heads, attn.head_dim) - - current_ip_hidden_states = dispatch_attention_fn( - ip_query, - ip_key, - ip_value, - attn_mask=None, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - current_ip_hidden_states = current_ip_hidden_states.reshape(batch_size, -1, attn.heads * attn.head_dim) - current_ip_hidden_states = current_ip_hidden_states.to(ip_query.dtype) - ip_attn_output += scale * current_ip_hidden_states - - return hidden_states, encoder_hidden_states, ip_attn_output - else: - return hidden_states - - -class FluxAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = FluxAttnProcessor - _available_processors = [ - FluxAttnProcessor, - FluxIPAdapterAttnProcessor, - ] - - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - added_kv_proj_dim: int | None = None, - added_proj_bias: bool | None = True, - out_bias: bool = True, - eps: float = 1e-5, - out_dim: int = None, - context_pre_only: bool | None = None, - pre_only: bool = False, - elementwise_affine: bool = True, - processor=None, - ): - super().__init__() - - self.head_dim = dim_head - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.query_dim = query_dim - self.use_bias = bias - self.dropout = dropout - self.out_dim = out_dim if out_dim is not None else query_dim - self.context_pre_only = context_pre_only - self.pre_only = pre_only - self.heads = out_dim // dim_head if out_dim is not None else heads - self.added_kv_proj_dim = added_kv_proj_dim - self.added_proj_bias = added_proj_bias - - self.norm_q = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_v = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - - if not self.pre_only: - self.to_out = torch.nn.ModuleList([]) - self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(torch.nn.Dropout(dropout)) - - if added_kv_proj_dim is not None: - self.norm_added_q = torch.nn.RMSNorm(dim_head, eps=eps) - self.norm_added_k = torch.nn.RMSNorm(dim_head, eps=eps) - self.add_q_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.to_add_out = torch.nn.Linear(self.inner_dim, query_dim, bias=out_bias) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - quiet_attn_parameters = {"ip_adapter_masks", "ip_hidden_states"} - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters and k not in quiet_attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"joint_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -@maybe_allow_in_graph -class FluxSingleTransformerBlock(nn.Module): - def __init__(self, dim: int, num_attention_heads: int, attention_head_dim: int, mlp_ratio: float = 4.0): - super().__init__() - self.mlp_hidden_dim = int(dim * mlp_ratio) - - self.norm = AdaLayerNormZeroSingle(dim) - self.proj_mlp = nn.Linear(dim, self.mlp_hidden_dim) - self.act_mlp = nn.GELU(approximate="tanh") - self.proj_out = nn.Linear(dim + self.mlp_hidden_dim, dim) - - self.attn = FluxAttention( - query_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - bias=True, - processor=FluxAttnProcessor(), - eps=1e-6, - pre_only=True, - ) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - text_seq_len = encoder_hidden_states.shape[1] - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - residual = hidden_states - norm_hidden_states, gate = self.norm(hidden_states, emb=temb) - mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) - joint_attention_kwargs = joint_attention_kwargs or {} - attn_output = self.attn( - hidden_states=norm_hidden_states, - image_rotary_emb=image_rotary_emb, - **joint_attention_kwargs, - ) - - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) - gate = gate.unsqueeze(1) - hidden_states = gate * self.proj_out(hidden_states) - hidden_states = residual + hidden_states - if hidden_states.dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - encoder_hidden_states, hidden_states = hidden_states[:, :text_seq_len], hidden_states[:, text_seq_len:] - return encoder_hidden_states, hidden_states - - -@maybe_allow_in_graph -class FluxTransformerBlock(nn.Module): - def __init__( - self, dim: int, num_attention_heads: int, attention_head_dim: int, qk_norm: str = "rms_norm", eps: float = 1e-6 - ): - super().__init__() - - self.norm1 = AdaLayerNormZero(dim) - self.norm1_context = AdaLayerNormZero(dim) - - self.attn = FluxAttention( - query_dim=dim, - added_kv_proj_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - context_pre_only=False, - bias=True, - processor=FluxAttnProcessor(), - eps=eps, - ) - - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff_context = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - joint_attention_kwargs = joint_attention_kwargs or {} - - # Attention. - attention_outputs = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - **joint_attention_kwargs, - ) - - if len(attention_outputs) == 2: - attn_output, context_attn_output = attention_outputs - elif len(attention_outputs) == 3: - attn_output, context_attn_output, ip_attn_output = attention_outputs - - # Process attention outputs for the `hidden_states`. - attn_output = gate_msa.unsqueeze(1) * attn_output - hidden_states = hidden_states + attn_output - - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - - ff_output = self.ff(norm_hidden_states) - ff_output = gate_mlp.unsqueeze(1) * ff_output - - hidden_states = hidden_states + ff_output - if len(attention_outputs) == 3: - hidden_states = hidden_states + ip_attn_output - - # Process attention outputs for the `encoder_hidden_states`. - context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output - encoder_hidden_states = encoder_hidden_states + context_attn_output - - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - context_ff_output = self.ff_context(norm_encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - if encoder_hidden_states.dtype == torch.float16: - encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) - - return encoder_hidden_states, hidden_states - - -class FluxPosEmbed(nn.Module): - # modified from https://github.com/black-forest-labs/flux/blob/c00d7c60b085fce8058b9df845e036090873f2ce/src/flux/modules/layers.py#L11 - def __init__(self, theta: int, axes_dim: list[int]): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - - def forward(self, ids: torch.Tensor) -> torch.Tensor: - n_axes = ids.shape[-1] - cos_out = [] - sin_out = [] - pos = ids.float() - freqs_dtype = maybe_adjust_dtype_for_device(torch.float64, ids.device) - for i in range(n_axes): - cos, sin = get_1d_rotary_pos_embed( - self.axes_dim[i], - pos[:, i], - theta=self.theta, - repeat_interleave_real=True, - use_real=True, - freqs_dtype=freqs_dtype, - ) - cos_out.append(cos) - sin_out.append(sin) - freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device) - freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device) - return freqs_cos, freqs_sin - - -class FluxTransformer2DModel( - ModelMixin, - ConfigMixin, - PeftAdapterMixin, - FromOriginalModelMixin, - FluxTransformer2DLoadersMixin, - CacheMixin, - AttentionMixin, -): - """ - The Transformer model introduced in Flux. - - Reference: https://blackforestlabs.ai/announcing-black-forest-labs/ - - Args: - patch_size (`int`, defaults to `1`): - Patch size to turn the input data into small patches. - in_channels (`int`, defaults to `64`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `None`): - The number of channels in the output. If not specified, it defaults to `in_channels`. - num_layers (`int`, defaults to `19`): - The number of layers of dual stream DiT blocks to use. - num_single_layers (`int`, defaults to `38`): - The number of layers of single stream DiT blocks to use. - attention_head_dim (`int`, defaults to `128`): - The number of dimensions to use for each attention head. - num_attention_heads (`int`, defaults to `24`): - The number of attention heads to use. - joint_attention_dim (`int`, defaults to `4096`): - The number of dimensions to use for the joint attention (embedding/channel dimension of - `encoder_hidden_states`). - pooled_projection_dim (`int`, defaults to `768`): - The number of dimensions to use for the pooled projection. - guidance_embeds (`bool`, defaults to `False`): - Whether to use guidance embeddings for guidance-distilled variant of the model. - axes_dims_rope (`tuple[int]`, defaults to `(16, 56, 56)`): - The dimensions to use for the rotary positional embeddings. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["FluxTransformerBlock", "FluxSingleTransformerBlock"] - _skip_layerwise_casting_patterns = ["pos_embed", "norm"] - _repeated_blocks = ["FluxTransformerBlock", "FluxSingleTransformerBlock"] - _cp_plan = { - "": { - "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - "encoder_hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - "img_ids": ContextParallelInput(split_dim=0, expected_dims=2, split_output=False), - "txt_ids": ContextParallelInput(split_dim=0, expected_dims=2, split_output=False), - }, - "proj_out": ContextParallelOutput(gather_dim=1, expected_dims=3), - } - - @register_to_config - def __init__( - self, - patch_size: int = 1, - in_channels: int = 64, - out_channels: int | None = None, - num_layers: int = 19, - num_single_layers: int = 38, - attention_head_dim: int = 128, - num_attention_heads: int = 24, - joint_attention_dim: int = 4096, - pooled_projection_dim: int = 768, - guidance_embeds: bool = False, - axes_dims_rope: tuple[int, int, int] = (16, 56, 56), - ): - super().__init__() - self.out_channels = out_channels or in_channels - self.inner_dim = num_attention_heads * attention_head_dim - - self.pos_embed = FluxPosEmbed(theta=10000, axes_dim=axes_dims_rope) - - text_time_guidance_cls = ( - CombinedTimestepGuidanceTextProjEmbeddings if guidance_embeds else CombinedTimestepTextProjEmbeddings - ) - self.time_text_embed = text_time_guidance_cls( - embedding_dim=self.inner_dim, pooled_projection_dim=pooled_projection_dim - ) - - self.context_embedder = nn.Linear(joint_attention_dim, self.inner_dim) - self.x_embedder = nn.Linear(in_channels, self.inner_dim) - - self.transformer_blocks = nn.ModuleList( - [ - FluxTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for _ in range(num_layers) - ] - ) - - self.single_transformer_blocks = nn.ModuleList( - [ - FluxSingleTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for _ in range(num_single_layers) - ] - ) - - self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) - - self.gradient_checkpointing = False - - @apply_lora_scale("joint_attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - pooled_projections: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_ids: torch.Tensor = None, - txt_ids: torch.Tensor = None, - guidance: torch.Tensor = None, - joint_attention_kwargs: dict[str, Any] | None = None, - controlnet_block_samples=None, - controlnet_single_block_samples=None, - return_dict: bool = True, - controlnet_blocks_repeat: bool = False, - ) -> torch.Tensor | Transformer2DModelOutput: - """ - The [`FluxTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, image_sequence_length, in_channels)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, text_sequence_length, joint_attention_dim)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - pooled_projections (`torch.Tensor` of shape `(batch_size, projection_dim)`): Embeddings projected - from the embeddings of input conditions. - timestep ( `torch.LongTensor`): - Used to indicate denoising step. - img_ids (`torch.Tensor`): - Image position ids used to compute the rotary positional embeddings. - txt_ids (`torch.Tensor`): - Text position ids used to compute the rotary positional embeddings. - guidance (`torch.Tensor`, *optional*): - Guidance scale embedding used for guidance-distilled variants of the model. - controlnet_block_samples (`list` of `torch.Tensor`, *optional*): - A list of tensors that if specified are added to the residuals of transformer blocks. - controlnet_single_block_samples (`list` of `torch.Tensor`, *optional*): - A list of tensors that if specified are added to the residuals of single transformer blocks. - controlnet_blocks_repeat (`bool`, *optional*, defaults to `False`): - Whether to repeat the controlnet block samples across all transformer blocks. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - - hidden_states = self.x_embedder(hidden_states) - - timestep = timestep.to(hidden_states.dtype) * 1000 - if guidance is not None: - guidance = guidance.to(hidden_states.dtype) * 1000 - - temb = ( - self.time_text_embed(timestep, pooled_projections) - if guidance is None - else self.time_text_embed(timestep, guidance, pooled_projections) - ) - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - - if txt_ids.ndim == 3: - logger.warning( - "Passing `txt_ids` 3d torch.Tensor is deprecated." - "Please remove the batch dimension and pass it as a 2d torch Tensor" - ) - txt_ids = txt_ids[0] - if img_ids.ndim == 3: - logger.warning( - "Passing `img_ids` 3d torch.Tensor is deprecated." - "Please remove the batch dimension and pass it as a 2d torch Tensor" - ) - img_ids = img_ids[0] - - ids = torch.cat((txt_ids, img_ids), dim=0) - image_rotary_emb = self.pos_embed(ids) - - if joint_attention_kwargs is not None and "ip_adapter_image_embeds" in joint_attention_kwargs: - ip_adapter_image_embeds = joint_attention_kwargs.pop("ip_adapter_image_embeds") - ip_hidden_states = self.encoder_hid_proj(ip_adapter_image_embeds) - joint_attention_kwargs.update({"ip_hidden_states": ip_hidden_states}) - - for index_block, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - joint_attention_kwargs, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - joint_attention_kwargs=joint_attention_kwargs, - ) - - # controlnet residual - if controlnet_block_samples is not None: - interval_control = len(self.transformer_blocks) / len(controlnet_block_samples) - interval_control = int(np.ceil(interval_control)) - # For Xlabs ControlNet. - if controlnet_blocks_repeat: - hidden_states = ( - hidden_states + controlnet_block_samples[index_block % len(controlnet_block_samples)] - ) - else: - hidden_states = hidden_states + controlnet_block_samples[index_block // interval_control] - - for index_block, block in enumerate(self.single_transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - joint_attention_kwargs, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - joint_attention_kwargs=joint_attention_kwargs, - ) - - # controlnet residual - if controlnet_single_block_samples is not None: - interval_control = len(self.single_transformer_blocks) / len(controlnet_single_block_samples) - interval_control = int(np.ceil(interval_control)) - hidden_states = hidden_states + controlnet_single_block_samples[index_block // interval_control] - - hidden_states = self.norm_out(hidden_states, temb) - output = self.proj_out(hidden_states) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_flux2.py b/diffusers/models/transformers/transformer_flux2.py deleted file mode 100644 index 17c8bd0ffd525228484a7dba25d654b1b3d2a9d5..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_flux2.py +++ /dev/null @@ -1,1386 +0,0 @@ -# Copyright 2025 Black Forest Labs, The HuggingFace Team and The InstantX Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -from dataclasses import dataclass -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FluxTransformer2DLoadersMixin, FromOriginalModelMixin, PeftAdapterMixin -from ...utils import BaseOutput, apply_lora_scale, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from .._modeling_parallel import ContextParallelInput, ContextParallelOutput -from ..attention import AttentionMixin, AttentionModuleMixin -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import ( - TimestepEmbedding, - Timesteps, - apply_rotary_emb, - get_1d_rotary_pos_embed, -) -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class Flux2Transformer2DModelOutput(BaseOutput): - """ - The output of [`Flux2Transformer2DModel`]. - - Args: - sample (`torch.Tensor` of shape `(batch_size, num_channels, height, width)`): - The hidden states output conditioned on the `encoder_hidden_states` input. - kv_cache (`Flux2KVCache`, *optional*): - The populated KV cache for reference image tokens. Only returned when `kv_cache_mode="extract"`. - """ - - sample: "torch.Tensor" # noqa: F821 - kv_cache: "Flux2KVCache | None" = None - - -class Flux2KVLayerCache: - """Per-layer KV cache for reference image tokens in the Flux2 Klein KV model. - - Stores the K and V projections (post-RoPE) for reference tokens extracted during the first denoising step. Tensor - format: (batch_size, num_ref_tokens, num_heads, head_dim). - """ - - def __init__(self): - self.k_ref: torch.Tensor | None = None - self.v_ref: torch.Tensor | None = None - - def store(self, k_ref: torch.Tensor, v_ref: torch.Tensor): - """Store reference token K/V.""" - self.k_ref = k_ref - self.v_ref = v_ref - - def get(self) -> tuple[torch.Tensor, torch.Tensor]: - """Retrieve cached reference token K/V.""" - if self.k_ref is None: - raise RuntimeError("KV cache has not been populated yet.") - return self.k_ref, self.v_ref - - def clear(self): - self.k_ref = None - self.v_ref = None - - -class Flux2KVCache: - """Container for all layers' reference-token KV caches. - - Holds separate cache lists for double-stream and single-stream transformer blocks. - """ - - def __init__(self, num_double_layers: int, num_single_layers: int): - self.double_block_caches = [Flux2KVLayerCache() for _ in range(num_double_layers)] - self.single_block_caches = [Flux2KVLayerCache() for _ in range(num_single_layers)] - self.num_ref_tokens: int = 0 - - def get_double(self, layer_idx: int) -> Flux2KVLayerCache: - return self.double_block_caches[layer_idx] - - def get_single(self, layer_idx: int) -> Flux2KVLayerCache: - return self.single_block_caches[layer_idx] - - def clear(self): - for cache in self.double_block_caches: - cache.clear() - for cache in self.single_block_caches: - cache.clear() - self.num_ref_tokens = 0 - - -def _flux2_kv_causal_attention( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - num_txt_tokens: int, - num_ref_tokens: int, - kv_cache: Flux2KVLayerCache | None = None, - backend=None, -) -> torch.Tensor: - """Causal attention for KV caching where reference tokens only self-attend. - - All tensors use the diffusers convention: (batch_size, seq_len, num_heads, head_dim). - - Without cache (extract mode): sequence layout is [txt, ref, img]. txt+img tokens attend to all tokens, ref tokens - only attend to themselves. With cache (cached mode): sequence layout is [txt, img]. Cached ref K/V are injected - between txt and img. - """ - # No ref tokens and no cache — standard full attention - if num_ref_tokens == 0 and kv_cache is None: - return dispatch_attention_fn(query, key, value, backend=backend) - - if kv_cache is not None: - # Cached mode: inject ref K/V between txt and img - k_ref, v_ref = kv_cache.get() - - k_all = torch.cat([key[:, :num_txt_tokens], k_ref, key[:, num_txt_tokens:]], dim=1) - v_all = torch.cat([value[:, :num_txt_tokens], v_ref, value[:, num_txt_tokens:]], dim=1) - - return dispatch_attention_fn(query, k_all, v_all, backend=backend) - - # Extract mode: ref tokens self-attend, txt+img attend to all - ref_start = num_txt_tokens - ref_end = num_txt_tokens + num_ref_tokens - - q_txt = query[:, :ref_start] - q_ref = query[:, ref_start:ref_end] - q_img = query[:, ref_end:] - - k_txt = key[:, :ref_start] - k_ref = key[:, ref_start:ref_end] - k_img = key[:, ref_end:] - - v_txt = value[:, :ref_start] - v_ref = value[:, ref_start:ref_end] - v_img = value[:, ref_end:] - - # txt+img attend to all tokens - q_txt_img = torch.cat([q_txt, q_img], dim=1) - k_all = torch.cat([k_txt, k_ref, k_img], dim=1) - v_all = torch.cat([v_txt, v_ref, v_img], dim=1) - attn_txt_img = dispatch_attention_fn(q_txt_img, k_all, v_all, backend=backend) - attn_txt = attn_txt_img[:, :ref_start] - attn_img = attn_txt_img[:, ref_start:] - - # ref tokens self-attend only - attn_ref = dispatch_attention_fn(q_ref, k_ref, v_ref, backend=backend) - - return torch.cat([attn_txt, attn_ref, attn_img], dim=1) - - -def _blend_mod_params( - img_params: tuple[torch.Tensor, ...], - ref_params: tuple[torch.Tensor, ...], - num_ref: int, - seq_len: int, -) -> tuple[torch.Tensor, ...]: - """Blend modulation parameters so that the first `num_ref` positions use `ref_params`.""" - blended = [] - for im, rm in zip(img_params, ref_params): - if im.ndim == 2: - im = im.unsqueeze(1) - rm = rm.unsqueeze(1) - B = im.shape[0] - blended.append( - torch.cat( - [rm.expand(B, num_ref, -1), im.expand(B, seq_len, -1)[:, num_ref:, :]], - dim=1, - ) - ) - return tuple(blended) - - -def _blend_double_block_mods( - img_mod: torch.Tensor, - ref_mod: torch.Tensor, - num_ref: int, - seq_len: int, -) -> torch.Tensor: - """Blend double-block image-stream modulations for a [ref, img] sequence layout. - - Takes raw modulation tensors (before `Flux2Modulation.split`) and returns a blended raw tensor that is compatible - with `Flux2Modulation.split(mod, 2)`. - """ - if img_mod.ndim == 2: - img_mod = img_mod.unsqueeze(1) - ref_mod = ref_mod.unsqueeze(1) - img_chunks = torch.chunk(img_mod, 6, dim=-1) - ref_chunks = torch.chunk(ref_mod, 6, dim=-1) - img_mods = (img_chunks[0:3], img_chunks[3:6]) - ref_mods = (ref_chunks[0:3], ref_chunks[3:6]) - - all_params = [] - for img_set, ref_set in zip(img_mods, ref_mods): - blended = _blend_mod_params(img_set, ref_set, num_ref, seq_len) - all_params.extend(blended) - return torch.cat(all_params, dim=-1) - - -def _blend_single_block_mods( - single_mod: torch.Tensor, - ref_mod: torch.Tensor, - num_txt: int, - num_ref: int, - seq_len: int, -) -> torch.Tensor: - """Blend single-block modulations for a [txt, ref, img] sequence layout. - - Takes raw modulation tensors and returns a blended raw tensor compatible with `Flux2Modulation.split(mod, 1)`. - """ - if single_mod.ndim == 2: - single_mod = single_mod.unsqueeze(1) - ref_mod = ref_mod.unsqueeze(1) - img_params = torch.chunk(single_mod, 3, dim=-1) - ref_params = torch.chunk(ref_mod, 3, dim=-1) - - blended = [] - for im, rm in zip(img_params, ref_params): - if im.ndim == 2: - im = im.unsqueeze(1) - rm = rm.unsqueeze(1) - B = im.shape[0] - im_expanded = im.expand(B, seq_len, -1) - rm_expanded = rm.expand(B, num_ref, -1) - blended.append( - torch.cat( - [im_expanded[:, :num_txt, :], rm_expanded, im_expanded[:, num_txt + num_ref :, :]], - dim=1, - ) - ) - return torch.cat(blended, dim=-1) - - -def _get_projections(attn: "Flux2Attention", hidden_states, encoder_hidden_states=None): - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - encoder_query = encoder_key = encoder_value = None - if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None: - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_fused_projections(attn: "Flux2Attention", hidden_states, encoder_hidden_states=None): - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - - encoder_query = encoder_key = encoder_value = (None,) - if encoder_hidden_states is not None and hasattr(attn, "to_added_qkv"): - encoder_query, encoder_key, encoder_value = attn.to_added_qkv(encoder_hidden_states).chunk(3, dim=-1) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_qkv_projections(attn: "Flux2Attention", hidden_states, encoder_hidden_states=None): - if attn.fused_projections: - return _get_fused_projections(attn, hidden_states, encoder_hidden_states) - return _get_projections(attn, hidden_states, encoder_hidden_states) - - -class Flux2SwiGLU(nn.Module): - """ - Flux 2 uses a SwiGLU-style activation in the transformer feedforward sub-blocks, but with the linear projection - layer fused into the first linear layer of the FF sub-block. Thus, this module has no trainable parameters. - """ - - def __init__(self): - super().__init__() - self.gate_fn = nn.SiLU() - - def forward(self, x: torch.Tensor) -> torch.Tensor: - half = x.shape[-1] // 2 - x = self.gate_fn(x[..., :half]) * x[..., half:] - return x - - -class Flux2FeedForward(nn.Module): - def __init__( - self, - dim: int, - dim_out: int | None = None, - mult: float = 3.0, - inner_dim: int | None = None, - bias: bool = False, - ): - super().__init__() - if inner_dim is None: - inner_dim = int(dim * mult) - dim_out = dim_out or dim - - # Flux2SwiGLU will reduce the dimension by half - self.linear_in = nn.Linear(dim, inner_dim * 2, bias=bias) - self.act_fn = Flux2SwiGLU() - self.linear_out = nn.Linear(inner_dim, dim_out, bias=bias) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = self.linear_in(x) - x = self.act_fn(x) - x = self.linear_out(x) - return x - - -class Flux2AttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError(f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version.") - - def __call__( - self, - attn: "Flux2Attention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - query, key, value, encoder_query, encoder_key, encoder_value = _get_qkv_projections( - attn, hidden_states, encoder_hidden_states - ) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if attn.added_kv_proj_dim is not None: - encoder_query = encoder_query.unflatten(-1, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(-1, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(-1, (attn.heads, -1)) - - encoder_query = attn.norm_added_q(encoder_query) - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - if encoder_hidden_states is not None: - encoder_hidden_states, hidden_states = hidden_states.split_with_sizes( - [encoder_hidden_states.shape[1], hidden_states.shape[1] - encoder_hidden_states.shape[1]], dim=1 - ) - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - if encoder_hidden_states is not None: - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -class Flux2KVAttnProcessor: - """ - Attention processor for Flux2 double-stream blocks with KV caching support for reference image tokens. - - When `kv_cache_mode` is "extract", reference token K/V are stored in the cache after RoPE and causal attention is - used (ref tokens self-attend only, txt+img attend to all). When `kv_cache_mode` is "cached", cached ref K/V are - injected during attention. When no KV args are provided, behaves identically to `Flux2AttnProcessor`. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError(f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version.") - - def __call__( - self, - attn: "Flux2Attention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - kv_cache: Flux2KVLayerCache | None = None, - kv_cache_mode: str | None = None, - num_ref_tokens: int = 0, - ) -> torch.Tensor: - query, key, value, encoder_query, encoder_key, encoder_value = _get_qkv_projections( - attn, hidden_states, encoder_hidden_states - ) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if attn.added_kv_proj_dim is not None: - encoder_query = encoder_query.unflatten(-1, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(-1, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(-1, (attn.heads, -1)) - - encoder_query = attn.norm_added_q(encoder_query) - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - num_txt_tokens = encoder_hidden_states.shape[1] if encoder_hidden_states is not None else 0 - - # Extract ref K/V from the combined sequence - if kv_cache_mode == "extract" and kv_cache is not None and num_ref_tokens > 0: - ref_start = num_txt_tokens - ref_end = num_txt_tokens + num_ref_tokens - kv_cache.store(key[:, ref_start:ref_end].clone(), value[:, ref_start:ref_end].clone()) - - # Dispatch attention - if kv_cache_mode == "extract" and num_ref_tokens > 0: - hidden_states = _flux2_kv_causal_attention( - query, key, value, num_txt_tokens, num_ref_tokens, backend=self._attention_backend - ) - elif kv_cache_mode == "cached" and kv_cache is not None: - hidden_states = _flux2_kv_causal_attention( - query, key, value, num_txt_tokens, 0, kv_cache=kv_cache, backend=self._attention_backend - ) - else: - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - if encoder_hidden_states is not None: - encoder_hidden_states, hidden_states = hidden_states.split_with_sizes( - [encoder_hidden_states.shape[1], hidden_states.shape[1] - encoder_hidden_states.shape[1]], dim=1 - ) - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - if encoder_hidden_states is not None: - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -class Flux2Attention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = Flux2AttnProcessor - _available_processors = [Flux2AttnProcessor, Flux2KVAttnProcessor] - - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - added_kv_proj_dim: int | None = None, - added_proj_bias: bool | None = True, - out_bias: bool = True, - eps: float = 1e-5, - out_dim: int = None, - elementwise_affine: bool = True, - processor=None, - ): - super().__init__() - - self.head_dim = dim_head - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.query_dim = query_dim - self.out_dim = out_dim if out_dim is not None else query_dim - self.heads = out_dim // dim_head if out_dim is not None else heads - - self.use_bias = bias - self.dropout = dropout - - self.added_kv_proj_dim = added_kv_proj_dim - self.added_proj_bias = added_proj_bias - - self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_v = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - - # QK Norm - self.norm_q = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - - self.to_out = torch.nn.ModuleList([]) - self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(torch.nn.Dropout(dropout)) - - if added_kv_proj_dim is not None: - self.norm_added_q = torch.nn.RMSNorm(dim_head, eps=eps) - self.norm_added_k = torch.nn.RMSNorm(dim_head, eps=eps) - self.add_q_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.to_add_out = torch.nn.Linear(self.inner_dim, query_dim, bias=out_bias) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"joint_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -class Flux2ParallelSelfAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError(f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version.") - - def __call__( - self, - attn: "Flux2ParallelSelfAttention", - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - # Parallel in (QKV + MLP in) projection - hidden_states = attn.to_qkv_mlp_proj(hidden_states) - qkv, mlp_hidden_states = torch.split( - hidden_states, [3 * attn.inner_dim, attn.mlp_hidden_dim * attn.mlp_mult_factor], dim=-1 - ) - - # Handle the attention logic - query, key, value = qkv.chunk(3, dim=-1) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - # Handle the feedforward (FF) logic - mlp_hidden_states = attn.mlp_act_fn(mlp_hidden_states) - - # Concatenate and parallel output projection - hidden_states = torch.cat([hidden_states, mlp_hidden_states], dim=-1) - hidden_states = attn.to_out(hidden_states) - - return hidden_states - - -class Flux2KVParallelSelfAttnProcessor: - """ - Attention processor for Flux2 single-stream blocks with KV caching support for reference image tokens. - - When `kv_cache_mode` is "extract", reference token K/V are stored and causal attention is used. When - `kv_cache_mode` is "cached", cached ref K/V are injected during attention. When no KV args are provided, behaves - identically to `Flux2ParallelSelfAttnProcessor`. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError(f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version.") - - def __call__( - self, - attn: "Flux2ParallelSelfAttention", - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - kv_cache: Flux2KVLayerCache | None = None, - kv_cache_mode: str | None = None, - num_txt_tokens: int = 0, - num_ref_tokens: int = 0, - ) -> torch.Tensor: - # Parallel in (QKV + MLP in) projection - hidden_states_proj = attn.to_qkv_mlp_proj(hidden_states) - qkv, mlp_hidden_states = torch.split( - hidden_states_proj, [3 * attn.inner_dim, attn.mlp_hidden_dim * attn.mlp_mult_factor], dim=-1 - ) - - query, key, value = qkv.chunk(3, dim=-1) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - # Extract ref K/V from the combined sequence - if kv_cache_mode == "extract" and kv_cache is not None and num_ref_tokens > 0: - ref_start = num_txt_tokens - ref_end = num_txt_tokens + num_ref_tokens - kv_cache.store(key[:, ref_start:ref_end].clone(), value[:, ref_start:ref_end].clone()) - - # Dispatch attention - if kv_cache_mode == "extract" and num_ref_tokens > 0: - attn_output = _flux2_kv_causal_attention( - query, key, value, num_txt_tokens, num_ref_tokens, backend=self._attention_backend - ) - elif kv_cache_mode == "cached" and kv_cache is not None: - attn_output = _flux2_kv_causal_attention( - query, key, value, num_txt_tokens, 0, kv_cache=kv_cache, backend=self._attention_backend - ) - else: - attn_output = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - attn_output = attn_output.flatten(2, 3) - attn_output = attn_output.to(query.dtype) - - # Handle the feedforward (FF) logic - mlp_hidden_states = attn.mlp_act_fn(mlp_hidden_states) - - # Concatenate and parallel output projection - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=-1) - hidden_states = attn.to_out(hidden_states) - - return hidden_states - - -class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin): - """ - Flux 2 parallel self-attention for the Flux 2 single-stream transformer blocks. - - This implements a parallel transformer block, where the attention QKV projections are fused to the feedforward (FF) - input projections, and the attention output projections are fused to the FF output projections. See the [ViT-22B - paper](https://arxiv.org/abs/2302.05442) for a visual depiction of this type of transformer block. - """ - - _default_processor_cls = Flux2ParallelSelfAttnProcessor - _available_processors = [Flux2ParallelSelfAttnProcessor, Flux2KVParallelSelfAttnProcessor] - # Does not support QKV fusion as the QKV projections are always fused - _supports_qkv_fusion = False - - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - out_bias: bool = True, - eps: float = 1e-5, - out_dim: int = None, - elementwise_affine: bool = True, - mlp_ratio: float = 4.0, - mlp_mult_factor: int = 2, - processor=None, - ): - super().__init__() - - self.head_dim = dim_head - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.query_dim = query_dim - self.out_dim = out_dim if out_dim is not None else query_dim - self.heads = out_dim // dim_head if out_dim is not None else heads - - self.use_bias = bias - self.dropout = dropout - - self.mlp_ratio = mlp_ratio - self.mlp_hidden_dim = int(query_dim * self.mlp_ratio) - self.mlp_mult_factor = mlp_mult_factor - - # Fused QKV projections + MLP input projection - self.to_qkv_mlp_proj = torch.nn.Linear( - self.query_dim, self.inner_dim * 3 + self.mlp_hidden_dim * self.mlp_mult_factor, bias=bias - ) - self.mlp_act_fn = Flux2SwiGLU() - - # QK Norm - self.norm_q = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - - # Fused attention output projection + MLP output projection - self.to_out = torch.nn.Linear(self.inner_dim + self.mlp_hidden_dim, self.out_dim, bias=out_bias) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"joint_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -class Flux2SingleTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float = 3.0, - eps: float = 1e-6, - bias: bool = False, - ): - super().__init__() - - self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=eps) - - # Note that the MLP in/out linear layers are fused with the attention QKV/out projections, respectively; this - # is often called a "parallel" transformer block. See the [ViT-22B paper](https://arxiv.org/abs/2302.05442) - # for a visual depiction of this type of transformer block. - self.attn = Flux2ParallelSelfAttention( - query_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - bias=bias, - out_bias=bias, - eps=eps, - mlp_ratio=mlp_ratio, - mlp_mult_factor=2, - processor=Flux2ParallelSelfAttnProcessor(), - ) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None, - temb_mod: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - split_hidden_states: bool = False, - text_seq_len: int | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - # If encoder_hidden_states is None, hidden_states is assumed to have encoder_hidden_states already - # concatenated - if encoder_hidden_states is not None: - text_seq_len = encoder_hidden_states.shape[1] - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - mod_shift, mod_scale, mod_gate = Flux2Modulation.split(temb_mod, 1)[0] - - norm_hidden_states = self.norm(hidden_states) - norm_hidden_states = (1 + mod_scale) * norm_hidden_states + mod_shift - - joint_attention_kwargs = joint_attention_kwargs or {} - attn_output = self.attn( - hidden_states=norm_hidden_states, - image_rotary_emb=image_rotary_emb, - **joint_attention_kwargs, - ) - - hidden_states = hidden_states + mod_gate * attn_output - if hidden_states.dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - if split_hidden_states: - encoder_hidden_states, hidden_states = hidden_states[:, :text_seq_len], hidden_states[:, text_seq_len:] - return encoder_hidden_states, hidden_states - else: - return hidden_states - - -class Flux2TransformerBlock(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float = 3.0, - eps: float = 1e-6, - bias: bool = False, - ): - super().__init__() - self.mlp_hidden_dim = int(dim * mlp_ratio) - - self.norm1 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps) - self.norm1_context = nn.LayerNorm(dim, elementwise_affine=False, eps=eps) - - self.attn = Flux2Attention( - query_dim=dim, - added_kv_proj_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - bias=bias, - added_proj_bias=bias, - out_bias=bias, - eps=eps, - processor=Flux2AttnProcessor(), - ) - - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps) - self.ff = Flux2FeedForward(dim=dim, dim_out=dim, mult=mlp_ratio, bias=bias) - - self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=eps) - self.ff_context = Flux2FeedForward(dim=dim, dim_out=dim, mult=mlp_ratio, bias=bias) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb_mod_img: torch.Tensor, - temb_mod_txt: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - joint_attention_kwargs = joint_attention_kwargs or {} - - # Modulation parameters shape: [1, 1, self.dim] - (shift_msa, scale_msa, gate_msa), (shift_mlp, scale_mlp, gate_mlp) = Flux2Modulation.split(temb_mod_img, 2) - (c_shift_msa, c_scale_msa, c_gate_msa), (c_shift_mlp, c_scale_mlp, c_gate_mlp) = Flux2Modulation.split( - temb_mod_txt, 2 - ) - - # Img stream - norm_hidden_states = self.norm1(hidden_states) - norm_hidden_states = (1 + scale_msa) * norm_hidden_states + shift_msa - - # Conditioning txt stream - norm_encoder_hidden_states = self.norm1_context(encoder_hidden_states) - norm_encoder_hidden_states = (1 + c_scale_msa) * norm_encoder_hidden_states + c_shift_msa - - # Attention on concatenated img + txt stream - attention_outputs = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - **joint_attention_kwargs, - ) - - attn_output, context_attn_output = attention_outputs - - # Process attention outputs for the image stream (`hidden_states`). - attn_output = gate_msa * attn_output - hidden_states = hidden_states + attn_output - - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp - - ff_output = self.ff(norm_hidden_states) - hidden_states = hidden_states + gate_mlp * ff_output - - # Process attention outputs for the text stream (`encoder_hidden_states`). - context_attn_output = c_gate_msa * context_attn_output - encoder_hidden_states = encoder_hidden_states + context_attn_output - - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp) + c_shift_mlp - - context_ff_output = self.ff_context(norm_encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states + c_gate_mlp * context_ff_output - if encoder_hidden_states.dtype == torch.float16: - encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) - - return encoder_hidden_states, hidden_states - - -class Flux2PosEmbed(nn.Module): - # modified from https://github.com/black-forest-labs/flux/blob/c00d7c60b085fce8058b9df845e036090873f2ce/src/flux/modules/layers.py#L11 - def __init__(self, theta: int, axes_dim: list[int]): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - - def forward(self, ids: torch.Tensor) -> torch.Tensor: - # Expected ids shape: [S, len(self.axes_dim)] - cos_out = [] - sin_out = [] - pos = ids.float() - freqs_dtype = maybe_adjust_dtype_for_device(torch.float64, ids.device) - # Unlike Flux 1, loop over len(self.axes_dim) rather than ids.shape[-1] - for i in range(len(self.axes_dim)): - cos, sin = get_1d_rotary_pos_embed( - self.axes_dim[i], - pos[..., i], - theta=self.theta, - repeat_interleave_real=True, - use_real=True, - freqs_dtype=freqs_dtype, - ) - cos_out.append(cos) - sin_out.append(sin) - freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device) - freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device) - return freqs_cos, freqs_sin - - -class Flux2TimestepGuidanceEmbeddings(nn.Module): - def __init__( - self, - in_channels: int = 256, - embedding_dim: int = 6144, - bias: bool = False, - guidance_embeds: bool = True, - ): - super().__init__() - - self.time_proj = Timesteps(num_channels=in_channels, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding( - in_channels=in_channels, time_embed_dim=embedding_dim, sample_proj_bias=bias - ) - - if guidance_embeds: - self.guidance_embedder = TimestepEmbedding( - in_channels=in_channels, time_embed_dim=embedding_dim, sample_proj_bias=bias - ) - else: - self.guidance_embedder = None - - def forward(self, timestep: torch.Tensor, guidance: torch.Tensor) -> torch.Tensor: - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(timestep.dtype)) # (N, D) - - if guidance is not None and self.guidance_embedder is not None: - guidance_proj = self.time_proj(guidance) - guidance_emb = self.guidance_embedder(guidance_proj.to(guidance.dtype)) # (N, D) - time_guidance_emb = timesteps_emb + guidance_emb - return time_guidance_emb - else: - return timesteps_emb - - -class Flux2Modulation(nn.Module): - def __init__(self, dim: int, mod_param_sets: int = 2, bias: bool = False): - super().__init__() - self.mod_param_sets = mod_param_sets - - self.linear = nn.Linear(dim, dim * 3 * self.mod_param_sets, bias=bias) - self.act_fn = nn.SiLU() - - def forward(self, temb: torch.Tensor) -> torch.Tensor: - mod = self.act_fn(temb) - mod = self.linear(mod) - return mod - - @staticmethod - # split inside the transformer blocks, to avoid passing tuples into checkpoints https://github.com/huggingface/diffusers/issues/12776 - def split(mod: torch.Tensor, mod_param_sets: int) -> tuple[tuple[torch.Tensor, torch.Tensor, torch.Tensor], ...]: - if mod.ndim == 2: - mod = mod.unsqueeze(1) - mod_params = torch.chunk(mod, 3 * mod_param_sets, dim=-1) - # Return tuple of 3-tuples of modulation params shift/scale/gate - return tuple(mod_params[3 * i : 3 * (i + 1)] for i in range(mod_param_sets)) - - -class Flux2Transformer2DModel( - ModelMixin, - ConfigMixin, - PeftAdapterMixin, - FromOriginalModelMixin, - FluxTransformer2DLoadersMixin, - CacheMixin, - AttentionMixin, -): - """ - The Transformer model introduced in Flux 2. - - Reference: https://blackforestlabs.ai/announcing-black-forest-labs/ - - Args: - patch_size (`int`, defaults to `1`): - Patch size to turn the input data into small patches. - in_channels (`int`, defaults to `128`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `None`): - The number of channels in the output. If not specified, it defaults to `in_channels`. - num_layers (`int`, defaults to `8`): - The number of layers of dual stream DiT blocks to use. - num_single_layers (`int`, defaults to `48`): - The number of layers of single stream DiT blocks to use. - attention_head_dim (`int`, defaults to `128`): - The number of dimensions to use for each attention head. - num_attention_heads (`int`, defaults to `48`): - The number of attention heads to use. - joint_attention_dim (`int`, defaults to `15360`): - The number of dimensions to use for the joint attention (embedding/channel dimension of - `encoder_hidden_states`). - pooled_projection_dim (`int`, defaults to `768`): - The number of dimensions to use for the pooled projection. - guidance_embeds (`bool`, defaults to `True`): - Whether to use guidance embeddings for guidance-distilled variant of the model. - axes_dims_rope (`tuple[int]`, defaults to `(32, 32, 32, 32)`): - The dimensions to use for the rotary positional embeddings. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["Flux2TransformerBlock", "Flux2SingleTransformerBlock"] - _skip_layerwise_casting_patterns = ["pos_embed", "norm"] - _repeated_blocks = ["Flux2TransformerBlock", "Flux2SingleTransformerBlock"] - _cp_plan = { - "": { - "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - "encoder_hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - "img_ids": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - "txt_ids": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - }, - "proj_out": ContextParallelOutput(gather_dim=1, expected_dims=3), - } - - @register_to_config - def __init__( - self, - patch_size: int = 1, - in_channels: int = 128, - out_channels: int | None = None, - num_layers: int = 8, - num_single_layers: int = 48, - attention_head_dim: int = 128, - num_attention_heads: int = 48, - joint_attention_dim: int = 15360, - timestep_guidance_channels: int = 256, - mlp_ratio: float = 3.0, - axes_dims_rope: tuple[int, ...] = (32, 32, 32, 32), - rope_theta: int = 2000, - eps: float = 1e-6, - guidance_embeds: bool = True, - ): - super().__init__() - self.out_channels = out_channels or in_channels - self.inner_dim = num_attention_heads * attention_head_dim - - # 1. Sinusoidal positional embedding for RoPE on image and text tokens - self.pos_embed = Flux2PosEmbed(theta=rope_theta, axes_dim=axes_dims_rope) - - # 2. Combined timestep + guidance embedding - self.time_guidance_embed = Flux2TimestepGuidanceEmbeddings( - in_channels=timestep_guidance_channels, - embedding_dim=self.inner_dim, - bias=False, - guidance_embeds=guidance_embeds, - ) - - # 3. Modulation (double stream and single stream blocks share modulation parameters, resp.) - # Two sets of shift/scale/gate modulation parameters for the double stream attn and FF sub-blocks - self.double_stream_modulation_img = Flux2Modulation(self.inner_dim, mod_param_sets=2, bias=False) - self.double_stream_modulation_txt = Flux2Modulation(self.inner_dim, mod_param_sets=2, bias=False) - # Only one set of modulation parameters as the attn and FF sub-blocks are run in parallel for single stream - self.single_stream_modulation = Flux2Modulation(self.inner_dim, mod_param_sets=1, bias=False) - - # 4. Input projections - self.x_embedder = nn.Linear(in_channels, self.inner_dim, bias=False) - self.context_embedder = nn.Linear(joint_attention_dim, self.inner_dim, bias=False) - - # 5. Double Stream Transformer Blocks - self.transformer_blocks = nn.ModuleList( - [ - Flux2TransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - mlp_ratio=mlp_ratio, - eps=eps, - bias=False, - ) - for _ in range(num_layers) - ] - ) - - # 6. Single Stream Transformer Blocks - self.single_transformer_blocks = nn.ModuleList( - [ - Flux2SingleTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - mlp_ratio=mlp_ratio, - eps=eps, - bias=False, - ) - for _ in range(num_single_layers) - ] - ) - - # 7. Output layers - self.norm_out = AdaLayerNormContinuous( - self.inner_dim, self.inner_dim, elementwise_affine=False, eps=eps, bias=False - ) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=False) - - self.gradient_checkpointing = False - - _skip_keys = ["kv_cache"] - - @apply_lora_scale("joint_attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_ids: torch.Tensor = None, - txt_ids: torch.Tensor = None, - guidance: torch.Tensor = None, - joint_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - kv_cache: "Flux2KVCache | None" = None, - kv_cache_mode: str | None = None, - num_ref_tokens: int = 0, - ref_fixed_timestep: float = 0.0, - ) -> torch.Tensor | Flux2Transformer2DModelOutput: - """ - The [`Flux2Transformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, image_sequence_length, in_channels)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, text_sequence_length, joint_attention_dim)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - img_ids (`torch.Tensor`): - Image position ids used to compute the rotary positional embeddings. - txt_ids (`torch.Tensor`): - Text position ids used to compute the rotary positional embeddings. - guidance (`torch.Tensor`, *optional*): - Guidance scale embedding used for guidance-distilled variants of the model. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - kv_cache (`Flux2KVCache`, *optional*): - KV cache for reference image tokens. When `kv_cache_mode` is "extract", a new cache is created and - returned. When "cached", the provided cache is used to inject ref K/V during attention. - kv_cache_mode (`str`, *optional*): - One of "extract" (first step with ref tokens) or "cached" (subsequent steps using cached ref K/V). When - `None`, standard forward pass without KV caching. - num_ref_tokens (`int`, defaults to `0`): - Number of reference image tokens prepended to `hidden_states` (only used when - `kv_cache_mode="extract"`). - ref_fixed_timestep (`float`, defaults to `0.0`): - Fixed timestep for reference token modulation (only used when `kv_cache_mode="extract"`). - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. When `kv_cache_mode="extract"`, also returns the - populated `Flux2KVCache`. - """ - num_txt_tokens = encoder_hidden_states.shape[1] - - # 1. Calculate timestep embedding and modulation parameters - timestep = timestep.to(hidden_states.dtype) * 1000 - - if guidance is not None: - guidance = guidance.to(hidden_states.dtype) * 1000 - - temb = self.time_guidance_embed(timestep, guidance) - - double_stream_mod_img = self.double_stream_modulation_img(temb) - double_stream_mod_txt = self.double_stream_modulation_txt(temb) - single_stream_mod = self.single_stream_modulation(temb) - - # KV extract mode: create cache and blend modulations for ref tokens - if kv_cache_mode == "extract" and num_ref_tokens > 0: - num_img_tokens = hidden_states.shape[1] # includes ref tokens - - kv_cache = Flux2KVCache( - num_double_layers=len(self.transformer_blocks), - num_single_layers=len(self.single_transformer_blocks), - ) - kv_cache.num_ref_tokens = num_ref_tokens - - # Ref tokens use a fixed timestep for modulation - ref_timestep = torch.full_like(timestep, ref_fixed_timestep * 1000) - ref_temb = self.time_guidance_embed(ref_timestep, guidance) - - ref_double_mod_img = self.double_stream_modulation_img(ref_temb) - ref_single_mod = self.single_stream_modulation(ref_temb) - - # Blend double block img modulation: [ref_mod, img_mod] - double_stream_mod_img = _blend_double_block_mods( - double_stream_mod_img, ref_double_mod_img, num_ref_tokens, num_img_tokens - ) - - # 2. Input projection for image (hidden_states) and conditioning text (encoder_hidden_states) - hidden_states = self.x_embedder(hidden_states) - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - - # 3. Calculate RoPE embeddings from image and text tokens - if img_ids.ndim == 3: - img_ids = img_ids[0] - if txt_ids.ndim == 3: - txt_ids = txt_ids[0] - - image_rotary_emb = self.pos_embed(img_ids) - text_rotary_emb = self.pos_embed(txt_ids) - concat_rotary_emb = ( - torch.cat([text_rotary_emb[0], image_rotary_emb[0]], dim=0), - torch.cat([text_rotary_emb[1], image_rotary_emb[1]], dim=0), - ) - - # 4. Build joint_attention_kwargs with KV cache info - if kv_cache_mode == "extract": - kv_attn_kwargs = { - **(joint_attention_kwargs or {}), - "kv_cache": None, - "kv_cache_mode": "extract", - "num_ref_tokens": num_ref_tokens, - } - elif kv_cache_mode == "cached" and kv_cache is not None: - kv_attn_kwargs = { - **(joint_attention_kwargs or {}), - "kv_cache": None, - "kv_cache_mode": "cached", - "num_ref_tokens": kv_cache.num_ref_tokens, - } - else: - kv_attn_kwargs = joint_attention_kwargs - - # 5. Double Stream Transformer Blocks - for index_block, block in enumerate(self.transformer_blocks): - if kv_cache_mode is not None and kv_cache is not None: - kv_attn_kwargs["kv_cache"] = kv_cache.get_double(index_block) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - double_stream_mod_img, - double_stream_mod_txt, - concat_rotary_emb, - kv_attn_kwargs, - ) - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb_mod_img=double_stream_mod_img, - temb_mod_txt=double_stream_mod_txt, - image_rotary_emb=concat_rotary_emb, - joint_attention_kwargs=kv_attn_kwargs, - ) - - # Concatenate text and image streams for single-block inference - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - # Blend single block modulation for extract mode: [txt_mod, ref_mod, img_mod] - if kv_cache_mode == "extract" and num_ref_tokens > 0: - total_single_len = hidden_states.shape[1] - single_stream_mod = _blend_single_block_mods( - single_stream_mod, ref_single_mod, num_txt_tokens, num_ref_tokens, total_single_len - ) - - # Build single-block KV kwargs (single blocks need num_txt_tokens) - if kv_cache_mode is not None: - kv_attn_kwargs_single = {**kv_attn_kwargs, "num_txt_tokens": num_txt_tokens} - else: - kv_attn_kwargs_single = kv_attn_kwargs - - # 6. Single Stream Transformer Blocks - for index_block, block in enumerate(self.single_transformer_blocks): - if kv_cache_mode is not None and kv_cache is not None: - kv_attn_kwargs_single["kv_cache"] = kv_cache.get_single(index_block) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - None, - single_stream_mod, - concat_rotary_emb, - kv_attn_kwargs_single, - ) - else: - hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=None, - temb_mod=single_stream_mod, - image_rotary_emb=concat_rotary_emb, - joint_attention_kwargs=kv_attn_kwargs_single, - ) - - # Remove text tokens (and ref tokens in extract mode) from concatenated stream - if kv_cache_mode == "extract" and num_ref_tokens > 0: - hidden_states = hidden_states[:, num_txt_tokens + num_ref_tokens :, ...] - else: - hidden_states = hidden_states[:, num_txt_tokens:, ...] - - # 7. Output layers - hidden_states = self.norm_out(hidden_states, temb) - output = self.proj_out(hidden_states) - - if kv_cache_mode == "extract": - if not return_dict: - return (output, kv_cache) - return Flux2Transformer2DModelOutput(sample=output, kv_cache=kv_cache) - - if not return_dict: - return (output,) - - return Flux2Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_glm_image.py b/diffusers/models/transformers/transformer_glm_image.py deleted file mode 100644 index e2d883d2fecdd3c580565ceff570d2f9de1726be..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_glm_image.py +++ /dev/null @@ -1,705 +0,0 @@ -# Copyright 2025 The CogView team, Tsinghua University & ZhipuAI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..attention_processor import Attention -from ..cache_utils import CacheMixin -from ..embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import LayerNorm, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class GlmImageCombinedTimestepSizeEmbeddings(nn.Module): - def __init__(self, embedding_dim: int, condition_dim: int, pooled_projection_dim: int, timesteps_dim: int = 256): - super().__init__() - - self.time_proj = Timesteps(num_channels=timesteps_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.condition_proj = Timesteps(num_channels=condition_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=timesteps_dim, time_embed_dim=embedding_dim) - self.condition_embedder = PixArtAlphaTextProjection(pooled_projection_dim, embedding_dim, act_fn="silu") - - def forward( - self, - timestep: torch.Tensor, - target_size: torch.Tensor, - crop_coords: torch.Tensor, - hidden_dtype: torch.dtype, - ) -> torch.Tensor: - timesteps_proj = self.time_proj(timestep) - - crop_coords_proj = self.condition_proj(crop_coords.flatten()).view(crop_coords.size(0), -1) - target_size_proj = self.condition_proj(target_size.flatten()).view(target_size.size(0), -1) - - # (B, 2 * condition_dim) - condition_proj = torch.cat([crop_coords_proj, target_size_proj], dim=1) - - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (B, embedding_dim) - condition_emb = self.condition_embedder(condition_proj.to(dtype=hidden_dtype)) # (B, embedding_dim) - - conditioning = timesteps_emb + condition_emb - conditioning = F.silu(conditioning) - - return conditioning - - -class GlmImageImageProjector(nn.Module): - def __init__( - self, - in_channels: int = 16, - hidden_size: int = 2560, - patch_size: int = 2, - ): - super().__init__() - self.patch_size = patch_size - - self.proj = nn.Linear(in_channels * patch_size**2, hidden_size) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, channel, height, width = hidden_states.shape - post_patch_height = height // self.patch_size - post_patch_width = width // self.patch_size - - hidden_states = hidden_states.reshape( - batch_size, channel, post_patch_height, self.patch_size, post_patch_width, self.patch_size - ) - hidden_states = hidden_states.permute(0, 2, 4, 1, 3, 5).flatten(3, 5).flatten(1, 2) - hidden_states = self.proj(hidden_states) - - return hidden_states - - -class GlmImageAdaLayerNormZero(nn.Module): - def __init__(self, embedding_dim: int, dim: int) -> None: - super().__init__() - - self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - self.norm_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - self.linear = nn.Linear(embedding_dim, 12 * dim, bias=True) - - def forward( - self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor]: - dtype = hidden_states.dtype - norm_hidden_states = self.norm(hidden_states).to(dtype=dtype) - norm_encoder_hidden_states = self.norm_context(encoder_hidden_states).to(dtype=dtype) - - emb = self.linear(temb) - ( - shift_msa, - c_shift_msa, - scale_msa, - c_scale_msa, - gate_msa, - c_gate_msa, - shift_mlp, - c_shift_mlp, - scale_mlp, - c_scale_mlp, - gate_mlp, - c_gate_mlp, - ) = emb.chunk(12, dim=1) - - hidden_states = norm_hidden_states * (1 + scale_msa.unsqueeze(1)) + shift_msa.unsqueeze(1) - encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_msa.unsqueeze(1)) + c_shift_msa.unsqueeze(1) - - return ( - hidden_states, - gate_msa, - shift_mlp, - scale_mlp, - gate_mlp, - encoder_hidden_states, - c_gate_msa, - c_shift_mlp, - c_scale_mlp, - c_gate_mlp, - ) - - -class GlmImageLayerKVCache: - """KV cache for GlmImage model. - Supports per-sample caching for batch processing where each sample may have different condition images. - """ - - def __init__(self): - self.k_caches: list[torch.Tensor | None] = [] - self.v_caches: list[torch.Tensor | None] = [] - self.mode: str | None = None # "write", "read", "skip" - self.current_sample_idx: int = 0 # Current sample index for writing - - def store(self, k: torch.Tensor, v: torch.Tensor): - """Store KV cache for the current sample.""" - # k, v shape: (1, seq_len, num_heads, head_dim) - if len(self.k_caches) <= self.current_sample_idx: - # First time storing for this sample - self.k_caches.append(k) - self.v_caches.append(v) - else: - # Append to existing cache for this sample (multiple condition images) - self.k_caches[self.current_sample_idx] = torch.cat([self.k_caches[self.current_sample_idx], k], dim=1) - self.v_caches[self.current_sample_idx] = torch.cat([self.v_caches[self.current_sample_idx], v], dim=1) - - def get(self, k: torch.Tensor, v: torch.Tensor): - """Get combined KV cache for all samples in the batch. - - Args: - k: Current key tensor, shape (batch_size, seq_len, num_heads, head_dim) - v: Current value tensor, shape (batch_size, seq_len, num_heads, head_dim) - Returns: - Combined key and value tensors with cached values prepended. - """ - batch_size = k.shape[0] - num_cached_samples = len(self.k_caches) - if num_cached_samples == 0: - return k, v - if num_cached_samples == 1: - # Single cache, expand for all batch samples (shared condition images) - k_cache_expanded = self.k_caches[0].expand(batch_size, -1, -1, -1) - v_cache_expanded = self.v_caches[0].expand(batch_size, -1, -1, -1) - elif num_cached_samples == batch_size: - # Per-sample cache, concatenate along batch dimension - k_cache_expanded = torch.cat(self.k_caches, dim=0) - v_cache_expanded = torch.cat(self.v_caches, dim=0) - else: - # Mismatch: try to handle by repeating the caches - # This handles cases like num_images_per_prompt > 1 - repeat_factor = batch_size // num_cached_samples - if batch_size % num_cached_samples == 0: - k_cache_list = [] - v_cache_list = [] - for i in range(num_cached_samples): - k_cache_list.append(self.k_caches[i].expand(repeat_factor, -1, -1, -1)) - v_cache_list.append(self.v_caches[i].expand(repeat_factor, -1, -1, -1)) - k_cache_expanded = torch.cat(k_cache_list, dim=0) - v_cache_expanded = torch.cat(v_cache_list, dim=0) - else: - raise ValueError( - f"Cannot match {num_cached_samples} cached samples to batch size {batch_size}. " - f"Batch size must be a multiple of the number of cached samples." - ) - - k_combined = torch.cat([k_cache_expanded, k], dim=1) - v_combined = torch.cat([v_cache_expanded, v], dim=1) - return k_combined, v_combined - - def clear(self): - self.k_caches = [] - self.v_caches = [] - self.mode = None - self.current_sample_idx = 0 - - def next_sample(self): - """Move to the next sample for writing.""" - self.current_sample_idx += 1 - - -class GlmImageKVCache: - """Container for all layers' KV caches. - Supports per-sample caching for batch processing where each sample may have different condition images. - """ - - def __init__(self, num_layers: int): - self.num_layers = num_layers - self.caches = [GlmImageLayerKVCache() for _ in range(num_layers)] - - def __getitem__(self, layer_idx: int) -> GlmImageLayerKVCache: - return self.caches[layer_idx] - - def set_mode(self, mode: str): - if mode is not None and mode not in ["write", "read", "skip"]: - raise ValueError(f"Invalid mode: {mode}, must be one of 'write', 'read', 'skip'") - for cache in self.caches: - cache.mode = mode - - def next_sample(self): - """Move to the next sample for writing. Call this after processing - all condition images for one batch sample.""" - for cache in self.caches: - cache.next_sample() - - def clear(self): - for cache in self.caches: - cache.clear() - - -class GlmImageAttnProcessor: - """ - Processor for implementing scaled dot-product attention for the GlmImage model. It applies a rotary embedding on - query and key vectors, but does not include spatial normalization. - - The processor supports passing an attention mask for text tokens. The attention mask should have shape (batch_size, - text_seq_length) where 1 indicates a non-padded token and 0 indicates a padded token. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("GlmImageAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - kv_cache: GlmImageLayerKVCache | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - dtype = encoder_hidden_states.dtype - - batch_size, text_seq_length, embed_dim = encoder_hidden_states.shape - batch_size, image_seq_length, embed_dim = hidden_states.shape - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - # 1. QKV projections - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - # 2. QK normalization - if attn.norm_q is not None: - query = attn.norm_q(query).to(dtype=dtype) - if attn.norm_k is not None: - key = attn.norm_k(key).to(dtype=dtype) - - # 3. Rotational positional embeddings applied to latent stream - if image_rotary_emb is not None: - from ..embeddings import apply_rotary_emb - - query[:, text_seq_length:, :, :] = apply_rotary_emb( - query[:, text_seq_length:, :, :], image_rotary_emb, sequence_dim=1, use_real_unbind_dim=-2 - ) - key[:, text_seq_length:, :, :] = apply_rotary_emb( - key[:, text_seq_length:, :, :], image_rotary_emb, sequence_dim=1, use_real_unbind_dim=-2 - ) - - if kv_cache is not None: - if kv_cache.mode == "write": - kv_cache.store(key, value) - elif kv_cache.mode == "read": - key, value = kv_cache.get(key, value) - elif kv_cache.mode == "skip": - pass - - # 4. Attention - if attention_mask is not None: - text_attn_mask = attention_mask - assert text_attn_mask.dim() == 2, "the shape of text_attn_mask should be (batch_size, text_seq_length)" - text_attn_mask = text_attn_mask.float().to(query.device) - mix_attn_mask = torch.ones((batch_size, text_seq_length + image_seq_length), device=query.device) - mix_attn_mask[:, :text_seq_length] = text_attn_mask - mix_attn_mask = mix_attn_mask.unsqueeze(2) - attn_mask_matrix = mix_attn_mask @ mix_attn_mask.transpose(1, 2) - attention_mask = (attn_mask_matrix > 0).unsqueeze(1).to(query.dtype) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - # 5. Output projection - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - encoder_hidden_states, hidden_states = hidden_states.split( - [text_seq_length, hidden_states.size(1) - text_seq_length], dim=1 - ) - return hidden_states, encoder_hidden_states - - -@maybe_allow_in_graph -class GlmImageTransformerBlock(nn.Module): - def __init__( - self, - dim: int = 2560, - num_attention_heads: int = 64, - attention_head_dim: int = 40, - time_embed_dim: int = 512, - ) -> None: - super().__init__() - - # 1. Attention - self.norm1 = GlmImageAdaLayerNormZero(time_embed_dim, dim) - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - out_dim=dim, - bias=True, - qk_norm="layer_norm", - elementwise_affine=False, - eps=1e-5, - processor=GlmImageAttnProcessor(), - ) - - # 2. Feedforward - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5) - self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]] | None = None, - attention_mask: dict[str, torch.Tensor] | None = None, - attention_kwargs: dict[str, Any] | None = None, - kv_cache: GlmImageLayerKVCache | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - # 1. Timestep conditioning - ( - norm_hidden_states, - gate_msa, - shift_mlp, - scale_mlp, - gate_mlp, - norm_encoder_hidden_states, - c_gate_msa, - c_shift_mlp, - c_scale_mlp, - c_gate_mlp, - ) = self.norm1(hidden_states, encoder_hidden_states, temb) - - # 2. Attention - attention_kwargs = attention_kwargs or {} - - attn_hidden_states, attn_encoder_hidden_states = self.attn1( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - attention_mask=attention_mask, - kv_cache=kv_cache, - **attention_kwargs, - ) - hidden_states = hidden_states + attn_hidden_states * gate_msa.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + attn_encoder_hidden_states * c_gate_msa.unsqueeze(1) - - # 3. Feedforward - norm_hidden_states = self.norm2(hidden_states) * (1 + scale_mlp.unsqueeze(1)) + shift_mlp.unsqueeze(1) - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) * ( - 1 + c_scale_mlp.unsqueeze(1) - ) + c_shift_mlp.unsqueeze(1) - - ff_output = self.ff(norm_hidden_states) - ff_output_context = self.ff(norm_encoder_hidden_states) - hidden_states = hidden_states + ff_output * gate_mlp.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + ff_output_context * c_gate_mlp.unsqueeze(1) - - return hidden_states, encoder_hidden_states - - -class GlmImageRotaryPosEmbed(nn.Module): - def __init__(self, dim: int, patch_size: int, theta: float = 10000.0) -> None: - super().__init__() - - self.dim = dim - self.patch_size = patch_size - self.theta = theta - - def forward(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - batch_size, num_channels, height, width = hidden_states.shape - height, width = height // self.patch_size, width // self.patch_size - - dim_h, dim_w = self.dim // 2, self.dim // 2 - h_inv_freq = 1.0 / ( - self.theta ** (torch.arange(0, dim_h, 2, dtype=torch.float32)[: (dim_h // 2)].float() / dim_h) - ) - w_inv_freq = 1.0 / ( - self.theta ** (torch.arange(0, dim_w, 2, dtype=torch.float32)[: (dim_w // 2)].float() / dim_w) - ) - h_seq = torch.arange(height) - w_seq = torch.arange(width) - freqs_h = torch.outer(h_seq, h_inv_freq) - freqs_w = torch.outer(w_seq, w_inv_freq) - - # Create position matrices for height and width - # [height, 1, dim//4] and [1, width, dim//4] - freqs_h = freqs_h.unsqueeze(1) - freqs_w = freqs_w.unsqueeze(0) - # Broadcast freqs_h and freqs_w to [height, width, dim//4] - freqs_h = freqs_h.expand(height, width, -1) - freqs_w = freqs_w.expand(height, width, -1) - - # Concatenate along last dimension to get [height, width, dim//2] - freqs = torch.cat([freqs_h, freqs_w], dim=-1) - freqs = torch.cat([freqs, freqs], dim=-1) # [height, width, dim] - freqs = freqs.reshape(height * width, -1) - return (freqs.cos(), freqs.sin()) - - -class GlmImageAdaLayerNormContinuous(nn.Module): - """ - GlmImage-only final AdaLN: LN(x) -> Linear(cond) -> chunk -> affine. Matches Megatron: **no activation** before the - Linear on conditioning embedding. - """ - - def __init__( - self, - embedding_dim: int, - conditioning_embedding_dim: int, - elementwise_affine: bool = True, - eps: float = 1e-5, - bias: bool = True, - norm_type: str = "layer_norm", - ): - super().__init__() - self.linear = nn.Linear(conditioning_embedding_dim, embedding_dim * 2, bias=bias) - if norm_type == "layer_norm": - self.norm = LayerNorm(embedding_dim, eps, elementwise_affine, bias) - elif norm_type == "rms_norm": - self.norm = RMSNorm(embedding_dim, eps, elementwise_affine) - else: - raise ValueError(f"unknown norm_type {norm_type}") - - def forward(self, x: torch.Tensor, conditioning_embedding: torch.Tensor) -> torch.Tensor: - # *** NO SiLU here *** - emb = self.linear(conditioning_embedding.to(x.dtype)) - scale, shift = torch.chunk(emb, 2, dim=1) - x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :] - return x - - -class GlmImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, CacheMixin): - r""" - Args: - patch_size (`int`, defaults to `2`): - The size of the patches to use in the patch embedding layer. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - num_layers (`int`, defaults to `30`): - The number of layers of Transformer blocks to use. - attention_head_dim (`int`, defaults to `40`): - The number of channels in each head. - num_attention_heads (`int`, defaults to `64`): - The number of heads to use for multi-head attention. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - text_embed_dim (`int`, defaults to `1472`): - Input dimension of text embeddings from the text encoder. - time_embed_dim (`int`, defaults to `512`): - Output dimension of timestep embeddings. - condition_dim (`int`, defaults to `256`): - The embedding dimension of the input SDXL-style resolution conditions (original_size, target_size, - crop_coords). - pos_embed_max_size (`int`, defaults to `128`): - The maximum resolution of the positional embeddings, from which slices of shape `H x W` are taken and added - to input patched latents, where `H` and `W` are the latent height and width respectively. A value of 128 - means that the maximum supported height and width for image generation is `128 * vae_scale_factor * - patch_size => 128 * 8 * 2 => 2048`. - sample_size (`int`, defaults to `128`): - The base resolution of input latents. If height/width is not provided during generation, this value is used - to determine the resolution as `sample_size * vae_scale_factor => 128 * 8 => 1024` - """ - - _supports_gradient_checkpointing = True - _repeated_blocks = ["GlmImageTransformerBlock"] - _no_split_modules = [ - "GlmImageTransformerBlock", - "GlmImageImageProjector", - "GlmImageCombinedTimestepSizeEmbeddings", - ] - _skip_layerwise_casting_patterns = ["patch_embed", "norm", "proj_out"] - _skip_keys = ["kv_caches"] - - @register_to_config - def __init__( - self, - patch_size: int = 2, - in_channels: int = 16, - out_channels: int = 16, - num_layers: int = 30, - attention_head_dim: int = 40, - num_attention_heads: int = 64, - text_embed_dim: int = 1472, - time_embed_dim: int = 512, - condition_dim: int = 256, - prior_vq_quantizer_codebook_size: int = 16384, - ): - super().__init__() - - # GlmImage uses 2 additional SDXL-like conditions - target_size, crop_coords - # Each of these are sincos embeddings of shape 2 * condition_dim - pooled_projection_dim = 2 * 2 * condition_dim - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels - - # 1. RoPE - self.rope = GlmImageRotaryPosEmbed(attention_head_dim, patch_size, theta=10000.0) - - # 2. Patch & Text-timestep embedding - self.image_projector = GlmImageImageProjector(in_channels, inner_dim, patch_size) - self.glyph_projector = FeedForward(text_embed_dim, inner_dim, inner_dim=inner_dim, activation_fn="gelu") - self.prior_token_embedding = nn.Embedding(prior_vq_quantizer_codebook_size, inner_dim) - self.prior_projector = FeedForward(inner_dim, inner_dim, inner_dim=inner_dim, activation_fn="linear-silu") - - self.time_condition_embed = GlmImageCombinedTimestepSizeEmbeddings( - embedding_dim=time_embed_dim, - condition_dim=condition_dim, - pooled_projection_dim=pooled_projection_dim, - timesteps_dim=time_embed_dim, - ) - - # 3. Transformer blocks - self.transformer_blocks = nn.ModuleList( - [ - GlmImageTransformerBlock(inner_dim, num_attention_heads, attention_head_dim, time_embed_dim) - for _ in range(num_layers) - ] - ) - - # 4. Output projection - self.norm_out = GlmImageAdaLayerNormContinuous(inner_dim, time_embed_dim, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels, bias=True) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - prior_token_id: torch.Tensor, - prior_token_drop: torch.Tensor, - timestep: torch.LongTensor, - target_size: torch.Tensor, - crop_coords: torch.Tensor, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - attention_mask: torch.Tensor | None = None, - kv_caches: GlmImageKVCache | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]] | None = None, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`GlmImageTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, in_channels, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - prior_token_id (`torch.Tensor`): - Token ids for the prior embedding lookup. - prior_token_drop (`torch.Tensor`): - Boolean mask indicating which prior embeddings should be dropped (zeroed out). - timestep (`torch.LongTensor`): - Used to indicate denoising step. - target_size (`torch.Tensor`): - Target image size conditioning. - crop_coords (`torch.Tensor`): - Crop coordinates conditioning. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - attention_mask (`torch.Tensor`, *optional*): - Mask applied to attention scores. - kv_caches (`GlmImageKVCache`, *optional*): - Pre-computed key/value caches used to speed up inference. - image_rotary_emb (`tuple` of `torch.Tensor`, *optional*): - Pre-computed rotary positional embeddings. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_channels, height, width = hidden_states.shape - - # 1. RoPE - if image_rotary_emb is None: - image_rotary_emb = self.rope(hidden_states) - - # 2. Patch & Timestep embeddings - p = self.config.patch_size - post_patch_height = height // p - post_patch_width = width // p - - hidden_states = self.image_projector(hidden_states) - encoder_hidden_states = self.glyph_projector(encoder_hidden_states) - prior_embedding = self.prior_token_embedding(prior_token_id) - prior_embedding[prior_token_drop] *= 0.0 - prior_hidden_states = self.prior_projector(prior_embedding) - - hidden_states = hidden_states + prior_hidden_states - - temb = self.time_condition_embed(timestep, target_size, crop_coords, hidden_states.dtype) - - # 3. Transformer blocks - for idx, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - attention_mask, - attention_kwargs, - kv_caches[idx] if kv_caches is not None else None, - ) - else: - hidden_states, encoder_hidden_states = block( - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - attention_mask, - attention_kwargs, - kv_cache=kv_caches[idx] if kv_caches is not None else None, - ) - - # 4. Output norm & projection - hidden_states = self.norm_out(hidden_states, temb) - hidden_states = self.proj_out(hidden_states) - - # 5. Unpatchify - hidden_states = hidden_states.reshape(batch_size, post_patch_height, post_patch_width, -1, p, p) - - # Rearrange tensor from (B, H_p, W_p, C, p, p) to (B, C, H_p * p, W_p * p) - output = hidden_states.permute(0, 3, 1, 4, 2, 5).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_helios.py b/diffusers/models/transformers/transformer_helios.py deleted file mode 100644 index b99ab1e3f34fe2f528361a3609e0f7216fba9a37..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_helios.py +++ /dev/null @@ -1,859 +0,0 @@ -# Copyright 2025 The Helios Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_allow_in_graph -from .._modeling_parallel import ContextParallelInput, ContextParallelOutput -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import FP32LayerNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def pad_for_3d_conv(x, kernel_size): - b, c, t, h, w = x.shape - pt, ph, pw = kernel_size - pad_t = (pt - (t % pt)) % pt - pad_h = (ph - (h % ph)) % ph - pad_w = (pw - (w % pw)) % pw - return torch.nn.functional.pad(x, (0, pad_w, 0, pad_h, 0, pad_t), mode="replicate") - - -def center_down_sample_3d(x, kernel_size): - return torch.nn.functional.avg_pool3d(x, kernel_size, stride=kernel_size) - - -def apply_rotary_emb_transposed( - hidden_states: torch.Tensor, - freqs_cis: torch.Tensor, -): - x_1, x_2 = hidden_states.unflatten(-1, (-1, 2)).unbind(-1) - cos, sin = freqs_cis.unsqueeze(-2).chunk(2, dim=-1) - out = torch.empty_like(hidden_states) - out[..., 0::2] = x_1 * cos[..., 0::2] - x_2 * sin[..., 1::2] - out[..., 1::2] = x_1 * sin[..., 1::2] + x_2 * cos[..., 0::2] - return out.type_as(hidden_states) - - -def _get_qkv_projections(attn: "HeliosAttention", hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor): - # encoder_hidden_states is only passed for cross-attention - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - if attn.fused_projections: - if not attn.is_cross_attention: - # In self-attention layers, we can fuse the entire QKV projection into a single linear - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - else: - # In cross-attention layers, we can only fuse the KV projections into a single linear - query = attn.to_q(hidden_states) - key, value = attn.to_kv(encoder_hidden_states).chunk(2, dim=-1) - else: - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - return query, key, value - - -class HeliosOutputNorm(nn.Module): - def __init__(self, dim: int, eps: float = 1e-6, elementwise_affine: bool = False): - super().__init__() - self.scale_shift_table = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5) - self.norm = FP32LayerNorm(dim, eps, elementwise_affine=False) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor, original_context_length: int): - temb = temb[:, -original_context_length:, :] - shift, scale = (self.scale_shift_table.unsqueeze(0).to(temb.device) + temb.unsqueeze(2)).chunk(2, dim=2) - shift, scale = shift.squeeze(2).to(hidden_states.device), scale.squeeze(2).to(hidden_states.device) - hidden_states = hidden_states[:, -original_context_length:, :] - hidden_states = (self.norm(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states) - return hidden_states - - -class HeliosAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "HeliosAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to version 2.0 or higher." - ) - - def __call__( - self, - attn: "HeliosAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - original_context_length: int = None, - ) -> torch.Tensor: - query, key, value = _get_qkv_projections(attn, hidden_states, encoder_hidden_states) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - if rotary_emb is not None: - query = apply_rotary_emb_transposed(query, rotary_emb) - key = apply_rotary_emb_transposed(key, rotary_emb) - - if not attn.is_cross_attention and attn.is_amplify_history: - history_seq_len = hidden_states.shape[1] - original_context_length - - if history_seq_len > 0: - scale_key = 1.0 + torch.sigmoid(attn.history_key_scale) * (attn.max_scale - 1.0) - if attn.history_scale_mode == "per_head": - scale_key = scale_key.view(1, 1, -1, 1) - key = torch.cat([key[:, :history_seq_len] * scale_key, key[:, history_seq_len:]], dim=1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - # Reference: https://github.com/huggingface/diffusers/pull/12909 - parallel_config=(self._parallel_config if encoder_hidden_states is None else None), - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class HeliosAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = HeliosAttnProcessor - _available_processors = [HeliosAttnProcessor] - - def __init__( - self, - dim: int, - heads: int = 8, - dim_head: int = 64, - eps: float = 1e-5, - dropout: float = 0.0, - added_kv_proj_dim: int | None = None, - cross_attention_dim_head: int | None = None, - processor=None, - is_cross_attention=None, - is_amplify_history=False, - history_scale_mode="per_head", # [scalar, per_head] - ): - super().__init__() - - self.inner_dim = dim_head * heads - self.heads = heads - self.added_kv_proj_dim = added_kv_proj_dim - self.cross_attention_dim_head = cross_attention_dim_head - self.kv_inner_dim = self.inner_dim if cross_attention_dim_head is None else cross_attention_dim_head * heads - - self.to_q = torch.nn.Linear(dim, self.inner_dim, bias=True) - self.to_k = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) - self.to_v = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) - self.to_out = torch.nn.ModuleList( - [ - torch.nn.Linear(self.inner_dim, dim, bias=True), - torch.nn.Dropout(dropout), - ] - ) - self.norm_q = torch.nn.RMSNorm(dim_head * heads, eps=eps, elementwise_affine=True) - self.norm_k = torch.nn.RMSNorm(dim_head * heads, eps=eps, elementwise_affine=True) - - self.add_k_proj = self.add_v_proj = None - if added_kv_proj_dim is not None: - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=True) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=True) - self.norm_added_k = torch.nn.RMSNorm(dim_head * heads, eps=eps) - - if is_cross_attention is not None: - self.is_cross_attention = is_cross_attention - else: - self.is_cross_attention = cross_attention_dim_head is not None - - self.set_processor(processor) - - self.is_amplify_history = is_amplify_history - if is_amplify_history: - if history_scale_mode == "scalar": - self.history_key_scale = nn.Parameter(torch.ones(1)) - elif history_scale_mode == "per_head": - self.history_key_scale = nn.Parameter(torch.ones(heads)) - else: - raise ValueError(f"Unknown history_scale_mode: {history_scale_mode}") - self.history_scale_mode = history_scale_mode - self.max_scale = 10.0 - - def fuse_projections(self): - if getattr(self, "fused_projections", False): - return - - if not self.is_cross_attention: - concatenated_weights = torch.cat([self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data]) - concatenated_bias = torch.cat([self.to_q.bias.data, self.to_k.bias.data, self.to_v.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_qkv = nn.Linear(in_features, out_features, bias=True) - self.to_qkv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - else: - concatenated_weights = torch.cat([self.to_k.weight.data, self.to_v.weight.data]) - concatenated_bias = torch.cat([self.to_k.bias.data, self.to_v.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_kv = nn.Linear(in_features, out_features, bias=True) - self.to_kv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - - if self.added_kv_proj_dim is not None: - concatenated_weights = torch.cat([self.add_k_proj.weight.data, self.add_v_proj.weight.data]) - concatenated_bias = torch.cat([self.add_k_proj.bias.data, self.add_v_proj.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_added_kv = nn.Linear(in_features, out_features, bias=True) - self.to_added_kv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - - self.fused_projections = True - - @torch.no_grad() - def unfuse_projections(self): - if not getattr(self, "fused_projections", False): - return - - if hasattr(self, "to_qkv"): - delattr(self, "to_qkv") - if hasattr(self, "to_kv"): - delattr(self, "to_kv") - if hasattr(self, "to_added_kv"): - delattr(self, "to_added_kv") - - self.fused_projections = False - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - original_context_length: int = None, - **kwargs, - ) -> torch.Tensor: - return self.processor( - self, - hidden_states, - encoder_hidden_states, - attention_mask, - rotary_emb, - original_context_length, - **kwargs, - ) - - -class HeliosTimeTextEmbedding(nn.Module): - def __init__( - self, - dim: int, - time_freq_dim: int, - time_proj_dim: int, - text_embed_dim: int, - ): - super().__init__() - - self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim) - self.act_fn = nn.SiLU() - self.time_proj = nn.Linear(dim, time_proj_dim) - self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh") - - def forward( - self, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - is_return_encoder_hidden_states: bool = True, - ): - timestep = self.timesteps_proj(timestep) - - time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype - if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8: - timestep = timestep.to(time_embedder_dtype) - temb = self.time_embedder(timestep).type_as(encoder_hidden_states) - timestep_proj = self.time_proj(self.act_fn(temb)) - - if encoder_hidden_states is not None and is_return_encoder_hidden_states: - encoder_hidden_states = self.text_embedder(encoder_hidden_states) - - return temb, timestep_proj, encoder_hidden_states - - -class HeliosRotaryPosEmbed(nn.Module): - def __init__(self, rope_dim, theta): - super().__init__() - self.DT, self.DY, self.DX = rope_dim - self.theta = theta - self.register_buffer("freqs_base_t", self._get_freqs_base(self.DT), persistent=False) - self.register_buffer("freqs_base_y", self._get_freqs_base(self.DY), persistent=False) - self.register_buffer("freqs_base_x", self._get_freqs_base(self.DX), persistent=False) - - def _get_freqs_base(self, dim): - return 1.0 / (self.theta ** (torch.arange(0, dim, 2, dtype=torch.float32)[: (dim // 2)] / dim)) - - @torch.no_grad() - def get_frequency_batched(self, freqs_base, pos): - # Disable autocast so the position-grid einsum runs in float32: under an ambient autocast it would run - # in bfloat16, which cannot represent consecutive integers past 256, so positions beyond that point - # would collapse onto the same frequency and degrade the rotary embedding. - with torch.autocast(device_type=pos.device.type, enabled=False): - freqs = torch.einsum("d,bthw->dbthw", freqs_base, pos) - freqs = freqs.repeat_interleave(2, dim=0) - return freqs.cos(), freqs.sin() - - @torch.no_grad() - def _get_spatial_meshgrid(self, height, width, device_str): - device = torch.device(device_str) - grid_y_coords = torch.arange(height, device=device, dtype=torch.float32) - grid_x_coords = torch.arange(width, device=device, dtype=torch.float32) - grid_y, grid_x = torch.meshgrid(grid_y_coords, grid_x_coords, indexing="ij") - return grid_y, grid_x - - @torch.no_grad() - def forward(self, frame_indices, height, width, device): - batch_size = frame_indices.shape[0] - num_frames = frame_indices.shape[1] - - frame_indices = frame_indices.to(device=device, dtype=torch.float32) - grid_y, grid_x = self._get_spatial_meshgrid(height, width, str(device)) - - grid_t = frame_indices[:, :, None, None].expand(batch_size, num_frames, height, width) - grid_y_batch = grid_y[None, None, :, :].expand(batch_size, num_frames, -1, -1) - grid_x_batch = grid_x[None, None, :, :].expand(batch_size, num_frames, -1, -1) - - freqs_cos_t, freqs_sin_t = self.get_frequency_batched(self.freqs_base_t, grid_t) - freqs_cos_y, freqs_sin_y = self.get_frequency_batched(self.freqs_base_y, grid_y_batch) - freqs_cos_x, freqs_sin_x = self.get_frequency_batched(self.freqs_base_x, grid_x_batch) - - result = torch.cat([freqs_cos_t, freqs_cos_y, freqs_cos_x, freqs_sin_t, freqs_sin_y, freqs_sin_x], dim=0) - - return result.permute(1, 0, 2, 3, 4) - - -@maybe_allow_in_graph -class HeliosTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - ffn_dim: int, - num_heads: int, - qk_norm: str = "rms_norm_across_heads", - cross_attn_norm: bool = False, - eps: float = 1e-6, - added_kv_proj_dim: int | None = None, - guidance_cross_attn: bool = False, - is_amplify_history: bool = False, - history_scale_mode: str = "per_head", # [scalar, per_head] - ): - super().__init__() - - # 1. Self-attention - self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False) - self.attn1 = HeliosAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - cross_attention_dim_head=None, - processor=HeliosAttnProcessor(), - is_amplify_history=is_amplify_history, - history_scale_mode=history_scale_mode, - ) - - # 2. Cross-attention - self.attn2 = HeliosAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - added_kv_proj_dim=added_kv_proj_dim, - cross_attention_dim_head=dim // num_heads, - processor=HeliosAttnProcessor(), - ) - self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity() - - # 3. Feed-forward - self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate") - self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False) - - self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) - - # 4. Guidance cross-attention - self.guidance_cross_attn = guidance_cross_attn - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - rotary_emb: torch.Tensor, - original_context_length: int = None, - ) -> torch.Tensor: - if temb.ndim == 4: - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table.unsqueeze(0) + temb.float() - ).chunk(6, dim=2) - # batch_size, seq_len, 1, inner_dim - shift_msa = shift_msa.squeeze(2) - scale_msa = scale_msa.squeeze(2) - gate_msa = gate_msa.squeeze(2) - c_shift_msa = c_shift_msa.squeeze(2) - c_scale_msa = c_scale_msa.squeeze(2) - c_gate_msa = c_gate_msa.squeeze(2) - else: - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table + temb.float() - ).chunk(6, dim=1) - - # 1. Self-attention - norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states) - attn_output = self.attn1( - norm_hidden_states, - None, - None, - rotary_emb, - original_context_length, - ) - hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states) - - # 2. Cross-attention - if self.guidance_cross_attn: - history_seq_len = hidden_states.shape[1] - original_context_length - - history_hidden_states, hidden_states = torch.split( - hidden_states, [history_seq_len, original_context_length], dim=1 - ) - norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states) - attn_output = self.attn2( - norm_hidden_states, - encoder_hidden_states, - None, - None, - original_context_length, - ) - hidden_states = hidden_states + attn_output - hidden_states = torch.cat([history_hidden_states, hidden_states], dim=1) - else: - norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states) - attn_output = self.attn2( - norm_hidden_states, - encoder_hidden_states, - None, - None, - original_context_length, - ) - hidden_states = hidden_states + attn_output - - # 3. Feed-forward - norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as( - hidden_states - ) - ff_output = self.ffn(norm_hidden_states) - hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states) - - return hidden_states - - -class HeliosTransformer3DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin, AttentionMixin -): - r""" - A Transformer model for video-like data used in the Helios model. - - Args: - patch_size (`tuple[int]`, defaults to `(1, 2, 2)`): - 3D patch dimensions for video embedding (t_patch, h_patch, w_patch). - num_attention_heads (`int`, defaults to `40`): - Fixed length for text embeddings. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - text_dim (`int`, defaults to `512`): - Input dimension for text embeddings. - freq_dim (`int`, defaults to `256`): - Dimension for sinusoidal time embeddings. - ffn_dim (`int`, defaults to `13824`): - Intermediate dimension in feed-forward network. - num_layers (`int`, defaults to `40`): - The number of layers of transformer blocks to use. - window_size (`tuple[int]`, defaults to `(-1, -1)`): - Window size for local attention (-1 indicates global attention). - cross_attn_norm (`bool`, defaults to `True`): - Enable cross-attention normalization. - qk_norm (`bool`, defaults to `True`): - Enable query/key normalization. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - add_img_emb (`bool`, defaults to `False`): - Whether to use img_emb. - added_kv_proj_dim (`int`, *optional*, defaults to `None`): - The number of channels to use for the added key and value projections. If `None`, no projection is used. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = [ - "patch_embedding", - "patch_short", - "patch_mid", - "patch_long", - "condition_embedder", - "norm", - ] - _no_split_modules = ["HeliosTransformerBlock", "HeliosOutputNorm"] - _keep_in_fp32_modules = [ - "time_embedder", - "scale_shift_table", - "norm1", - "norm2", - "norm3", - "history_key_scale", - ] - _keys_to_ignore_on_load_unexpected = ["norm_added_q"] - _repeated_blocks = ["HeliosTransformerBlock"] - _cp_plan = { - # Input split at attn level and ffn level. - "blocks.*.attn1": { - "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - "rotary_emb": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - }, - "blocks.*.attn2": { - "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - }, - "blocks.*.ffn": { - "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - }, - # Output gather at attn level and ffn level. - **{f"blocks.{i}.attn1": ContextParallelOutput(gather_dim=1, expected_dims=3) for i in range(40)}, - **{f"blocks.{i}.attn2": ContextParallelOutput(gather_dim=1, expected_dims=3) for i in range(40)}, - **{f"blocks.{i}.ffn": ContextParallelOutput(gather_dim=1, expected_dims=3) for i in range(40)}, - } - - @register_to_config - def __init__( - self, - patch_size: tuple[int, ...] = (1, 2, 2), - num_attention_heads: int = 40, - attention_head_dim: int = 128, - in_channels: int = 16, - out_channels: int = 16, - text_dim: int = 4096, - freq_dim: int = 256, - ffn_dim: int = 13824, - num_layers: int = 40, - cross_attn_norm: bool = True, - qk_norm: str | None = "rms_norm_across_heads", - eps: float = 1e-6, - added_kv_proj_dim: int | None = None, - rope_dim: tuple[int, ...] = (44, 42, 42), - rope_theta: float = 10000.0, - guidance_cross_attn: bool = True, - zero_history_timestep: bool = True, - has_multi_term_memory_patch: bool = True, - is_amplify_history: bool = False, - history_scale_mode: str = "per_head", # [scalar, per_head] - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - # 1. Patch & position embedding - self.rope = HeliosRotaryPosEmbed(rope_dim=rope_dim, theta=rope_theta) - self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - - # 2. Initial Multi Term Memory Patch - self.zero_history_timestep = zero_history_timestep - if has_multi_term_memory_patch: - self.patch_short = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - self.patch_mid = nn.Conv3d( - in_channels, - inner_dim, - kernel_size=tuple(2 * p for p in patch_size), - stride=tuple(2 * p for p in patch_size), - ) - self.patch_long = nn.Conv3d( - in_channels, - inner_dim, - kernel_size=tuple(4 * p for p in patch_size), - stride=tuple(4 * p for p in patch_size), - ) - - # 3. Condition embeddings - self.condition_embedder = HeliosTimeTextEmbedding( - dim=inner_dim, - time_freq_dim=freq_dim, - time_proj_dim=inner_dim * 6, - text_embed_dim=text_dim, - ) - - # 4. Transformer blocks - self.blocks = nn.ModuleList( - [ - HeliosTransformerBlock( - inner_dim, - ffn_dim, - num_attention_heads, - qk_norm, - cross_attn_norm, - eps, - added_kv_proj_dim, - guidance_cross_attn=guidance_cross_attn, - is_amplify_history=is_amplify_history, - history_scale_mode=history_scale_mode, - ) - for _ in range(num_layers) - ] - ) - - # 5. Output norm & projection - self.norm_out = HeliosOutputNorm(inner_dim, eps, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size)) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - # ------------ Stage 1 ------------ - indices_hidden_states=None, - indices_latents_history_short=None, - indices_latents_history_mid=None, - indices_latents_history_long=None, - latents_history_short=None, - latents_history_mid=None, - latents_history_long=None, - return_dict: bool = True, - attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor | dict[str, torch.Tensor]: - """ - The [`HeliosTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - indices_hidden_states (`torch.Tensor`, *optional*): - Frame indices for `hidden_states` used to compute the rotary positional embeddings. - indices_latents_history_short (`torch.Tensor`, *optional*): - Frame indices for the short history latents. - indices_latents_history_mid (`torch.Tensor`, *optional*): - Frame indices for the mid history latents. - indices_latents_history_long (`torch.Tensor`, *optional*): - Frame indices for the long history latents. - latents_history_short (`torch.Tensor`, *optional*): - Short history latents conditioning. - latents_history_mid (`torch.Tensor`, *optional*): - Mid history latents conditioning. - latents_history_long (`torch.Tensor`, *optional*): - Long history latents conditioning. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - # 1. Input - batch_size = hidden_states.shape[0] - p_t, p_h, p_w = self.config.patch_size - - # 2. Process noisy latents - hidden_states = self.patch_embedding(hidden_states) - _, _, post_patch_num_frames, post_patch_height, post_patch_width = hidden_states.shape - - if indices_hidden_states is None: - indices_hidden_states = torch.arange(0, post_patch_num_frames).unsqueeze(0).expand(batch_size, -1) - - hidden_states = hidden_states.flatten(2).transpose(1, 2) - rotary_emb = self.rope( - frame_indices=indices_hidden_states, - height=post_patch_height, - width=post_patch_width, - device=hidden_states.device, - ) - rotary_emb = rotary_emb.flatten(2).transpose(1, 2) - original_context_length = hidden_states.shape[1] - - # 3. Process short history latents - if latents_history_short is not None and indices_latents_history_short is not None: - latents_history_short = self.patch_short(latents_history_short) - _, _, _, H1, W1 = latents_history_short.shape - latents_history_short = latents_history_short.flatten(2).transpose(1, 2) - - rotary_emb_history_short = self.rope( - frame_indices=indices_latents_history_short, - height=H1, - width=W1, - device=latents_history_short.device, - ) - rotary_emb_history_short = rotary_emb_history_short.flatten(2).transpose(1, 2) - - hidden_states = torch.cat([latents_history_short, hidden_states], dim=1) - rotary_emb = torch.cat([rotary_emb_history_short, rotary_emb], dim=1) - - # 4. Process mid history latents - if latents_history_mid is not None and indices_latents_history_mid is not None: - latents_history_mid = pad_for_3d_conv(latents_history_mid, (2, 4, 4)) - latents_history_mid = self.patch_mid(latents_history_mid) - latents_history_mid = latents_history_mid.flatten(2).transpose(1, 2) - - rotary_emb_history_mid = self.rope( - frame_indices=indices_latents_history_mid, - height=H1, - width=W1, - device=latents_history_mid.device, - ) - rotary_emb_history_mid = pad_for_3d_conv(rotary_emb_history_mid, (2, 2, 2)) - rotary_emb_history_mid = center_down_sample_3d(rotary_emb_history_mid, (2, 2, 2)) - rotary_emb_history_mid = rotary_emb_history_mid.flatten(2).transpose(1, 2) - - hidden_states = torch.cat([latents_history_mid, hidden_states], dim=1) - rotary_emb = torch.cat([rotary_emb_history_mid, rotary_emb], dim=1) - - # 5. Process long history latents - if latents_history_long is not None and indices_latents_history_long is not None: - latents_history_long = pad_for_3d_conv(latents_history_long, (4, 8, 8)) - latents_history_long = self.patch_long(latents_history_long) - latents_history_long = latents_history_long.flatten(2).transpose(1, 2) - - rotary_emb_history_long = self.rope( - frame_indices=indices_latents_history_long, - height=H1, - width=W1, - device=latents_history_long.device, - ) - rotary_emb_history_long = pad_for_3d_conv(rotary_emb_history_long, (4, 4, 4)) - rotary_emb_history_long = center_down_sample_3d(rotary_emb_history_long, (4, 4, 4)) - rotary_emb_history_long = rotary_emb_history_long.flatten(2).transpose(1, 2) - - hidden_states = torch.cat([latents_history_long, hidden_states], dim=1) - rotary_emb = torch.cat([rotary_emb_history_long, rotary_emb], dim=1) - - history_context_length = hidden_states.shape[1] - original_context_length - - if indices_hidden_states is not None and self.zero_history_timestep: - timestep_t0 = torch.zeros((1), dtype=timestep.dtype, device=timestep.device) - temb_t0, timestep_proj_t0, _ = self.condition_embedder( - timestep_t0, encoder_hidden_states, is_return_encoder_hidden_states=False - ) - temb_t0 = temb_t0.unsqueeze(1).expand(batch_size, history_context_length, -1) - timestep_proj_t0 = ( - timestep_proj_t0.unflatten(-1, (6, -1)) - .view(1, 6, 1, -1) - .expand(batch_size, -1, history_context_length, -1) - ) - - temb, timestep_proj, encoder_hidden_states = self.condition_embedder(timestep, encoder_hidden_states) - timestep_proj = timestep_proj.unflatten(-1, (6, -1)) - - if indices_hidden_states is not None and not self.zero_history_timestep: - main_repeat_size = hidden_states.shape[1] - else: - main_repeat_size = original_context_length - temb = temb.view(batch_size, 1, -1).expand(batch_size, main_repeat_size, -1) - timestep_proj = timestep_proj.view(batch_size, 6, 1, -1).expand(batch_size, 6, main_repeat_size, -1) - - if indices_hidden_states is not None and self.zero_history_timestep: - temb = torch.cat([temb_t0, temb], dim=1) - timestep_proj = torch.cat([timestep_proj_t0, timestep_proj], dim=2) - - if timestep_proj.ndim == 4: - timestep_proj = timestep_proj.permute(0, 2, 1, 3) - - # 6. Transformer blocks - hidden_states = hidden_states.contiguous() - encoder_hidden_states = encoder_hidden_states.contiguous() - rotary_emb = rotary_emb.contiguous() - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.blocks: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - timestep_proj, - rotary_emb, - original_context_length, - ) - else: - for block in self.blocks: - hidden_states = block( - hidden_states, - encoder_hidden_states, - timestep_proj, - rotary_emb, - original_context_length, - ) - - # 7. Normalization - hidden_states = self.norm_out(hidden_states, temb, original_context_length) - hidden_states = self.proj_out(hidden_states) - - # 8. Unpatchify - hidden_states = hidden_states.reshape( - batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1 - ) - hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) - output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_hidream_image.py b/diffusers/models/transformers/transformer_hidream_image.py deleted file mode 100644 index bd69d5de68cab381cff5c39a1adfc1d99c7e24d6..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_hidream_image.py +++ /dev/null @@ -1,959 +0,0 @@ -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...models.modeling_outputs import Transformer2DModelOutput -from ...models.modeling_utils import ModelMixin -from ...utils import apply_lora_scale, deprecate, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device, maybe_allow_in_graph -from ..attention import Attention -from ..embeddings import TimestepEmbedding, Timesteps - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class HiDreamImageFeedForwardSwiGLU(nn.Module): - def __init__( - self, - dim: int, - hidden_dim: int, - multiple_of: int = 256, - ffn_dim_multiplier: float | None = None, - ): - super().__init__() - hidden_dim = int(2 * hidden_dim / 3) - # custom dim factor multiplier - if ffn_dim_multiplier is not None: - hidden_dim = int(ffn_dim_multiplier * hidden_dim) - hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of) - - self.w1 = nn.Linear(dim, hidden_dim, bias=False) - self.w2 = nn.Linear(hidden_dim, dim, bias=False) - self.w3 = nn.Linear(dim, hidden_dim, bias=False) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.w2(torch.nn.functional.silu(self.w1(x)) * self.w3(x)) - - -class HiDreamImagePooledEmbed(nn.Module): - def __init__(self, text_emb_dim, hidden_size): - super().__init__() - self.pooled_embedder = TimestepEmbedding(in_channels=text_emb_dim, time_embed_dim=hidden_size) - - def forward(self, pooled_embed: torch.Tensor) -> torch.Tensor: - return self.pooled_embedder(pooled_embed) - - -class HiDreamImageTimestepEmbed(nn.Module): - def __init__(self, hidden_size, frequency_embedding_size=256): - super().__init__() - self.time_proj = Timesteps(num_channels=frequency_embedding_size, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=frequency_embedding_size, time_embed_dim=hidden_size) - - def forward(self, timesteps: torch.Tensor, wdtype: torch.dtype | None = None) -> torch.Tensor: - t_emb = self.time_proj(timesteps).to(dtype=wdtype) - t_emb = self.timestep_embedder(t_emb) - return t_emb - - -class HiDreamImageOutEmbed(nn.Module): - def __init__(self, hidden_size, patch_size, out_channels): - super().__init__() - self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) - self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True)) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor) -> torch.Tensor: - shift, scale = self.adaLN_modulation(temb).chunk(2, dim=1) - hidden_states = self.norm_final(hidden_states) * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) - hidden_states = self.linear(hidden_states) - return hidden_states - - -class HiDreamImagePatchEmbed(nn.Module): - def __init__( - self, - patch_size=2, - in_channels=4, - out_channels=1024, - ): - super().__init__() - self.patch_size = patch_size - self.out_channels = out_channels - self.proj = nn.Linear(in_channels * patch_size * patch_size, out_channels, bias=True) - - def forward(self, latent) -> torch.Tensor: - latent = self.proj(latent) - return latent - - -def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor: - assert dim % 2 == 0, "The dimension must be even." - - dtype = maybe_adjust_dtype_for_device(torch.float64, pos.device) - - scale = torch.arange(0, dim, 2, dtype=dtype, device=pos.device) / dim - omega = 1.0 / (theta**scale) - - batch_size, seq_length = pos.shape - out = torch.einsum("...n,d->...nd", pos, omega) - cos_out = torch.cos(out) - sin_out = torch.sin(out) - - stacked_out = torch.stack([cos_out, -sin_out, sin_out, cos_out], dim=-1) - out = stacked_out.view(batch_size, -1, dim // 2, 2, 2) - return out.float() - - -class HiDreamImageEmbedND(nn.Module): - def __init__(self, theta: int, axes_dim: list[int]): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - - def forward(self, ids: torch.Tensor) -> torch.Tensor: - n_axes = ids.shape[-1] - emb = torch.cat( - [rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(n_axes)], - dim=-3, - ) - return emb.unsqueeze(2) - - -def apply_rope(xq: torch.Tensor, xk: torch.Tensor, freqs_cis: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2) - xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2) - xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1] - xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1] - return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk) - - -@maybe_allow_in_graph -class HiDreamAttention(Attention): - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - upcast_attention: bool = False, - upcast_softmax: bool = False, - scale_qk: bool = True, - eps: float = 1e-5, - processor=None, - out_dim: int = None, - single: bool = False, - ): - super(Attention, self).__init__() - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.query_dim = query_dim - self.upcast_attention = upcast_attention - self.upcast_softmax = upcast_softmax - self.out_dim = out_dim if out_dim is not None else query_dim - - self.scale_qk = scale_qk - self.scale = dim_head**-0.5 if self.scale_qk else 1.0 - - self.heads = out_dim // dim_head if out_dim is not None else heads - self.sliceable_head_dim = heads - self.single = single - - self.to_q = nn.Linear(query_dim, self.inner_dim) - self.to_k = nn.Linear(self.inner_dim, self.inner_dim) - self.to_v = nn.Linear(self.inner_dim, self.inner_dim) - self.to_out = nn.Linear(self.inner_dim, self.out_dim) - self.q_rms_norm = nn.RMSNorm(self.inner_dim, eps) - self.k_rms_norm = nn.RMSNorm(self.inner_dim, eps) - - if not single: - self.to_q_t = nn.Linear(query_dim, self.inner_dim) - self.to_k_t = nn.Linear(self.inner_dim, self.inner_dim) - self.to_v_t = nn.Linear(self.inner_dim, self.inner_dim) - self.to_out_t = nn.Linear(self.inner_dim, self.out_dim) - self.q_rms_norm_t = nn.RMSNorm(self.inner_dim, eps) - self.k_rms_norm_t = nn.RMSNorm(self.inner_dim, eps) - - self.set_processor(processor) - - def forward( - self, - norm_hidden_states: torch.Tensor, - hidden_states_masks: torch.Tensor = None, - norm_encoder_hidden_states: torch.Tensor = None, - image_rotary_emb: torch.Tensor = None, - ) -> torch.Tensor: - return self.processor( - self, - hidden_states=norm_hidden_states, - hidden_states_masks=hidden_states_masks, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - ) - - -class HiDreamAttnProcessor: - """Attention processor used typically in processing the SD3-like self-attention projections.""" - - def __call__( - self, - attn: HiDreamAttention, - hidden_states: torch.Tensor, - hidden_states_masks: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor = None, - *args, - **kwargs, - ) -> torch.Tensor: - dtype = hidden_states.dtype - batch_size = hidden_states.shape[0] - - query_i = attn.q_rms_norm(attn.to_q(hidden_states)).to(dtype=dtype) - key_i = attn.k_rms_norm(attn.to_k(hidden_states)).to(dtype=dtype) - value_i = attn.to_v(hidden_states) - - inner_dim = key_i.shape[-1] - head_dim = inner_dim // attn.heads - - query_i = query_i.view(batch_size, -1, attn.heads, head_dim) - key_i = key_i.view(batch_size, -1, attn.heads, head_dim) - value_i = value_i.view(batch_size, -1, attn.heads, head_dim) - if hidden_states_masks is not None: - key_i = key_i * hidden_states_masks.view(batch_size, -1, 1, 1) - - if not attn.single: - query_t = attn.q_rms_norm_t(attn.to_q_t(encoder_hidden_states)).to(dtype=dtype) - key_t = attn.k_rms_norm_t(attn.to_k_t(encoder_hidden_states)).to(dtype=dtype) - value_t = attn.to_v_t(encoder_hidden_states) - - query_t = query_t.view(batch_size, -1, attn.heads, head_dim) - key_t = key_t.view(batch_size, -1, attn.heads, head_dim) - value_t = value_t.view(batch_size, -1, attn.heads, head_dim) - - num_image_tokens = query_i.shape[1] - num_text_tokens = query_t.shape[1] - query = torch.cat([query_i, query_t], dim=1) - key = torch.cat([key_i, key_t], dim=1) - value = torch.cat([value_i, value_t], dim=1) - else: - query = query_i - key = key_i - value = value_i - - if query.shape[-1] == image_rotary_emb.shape[-3] * 2: - query, key = apply_rope(query, key, image_rotary_emb) - - else: - query_1, query_2 = query.chunk(2, dim=-1) - key_1, key_2 = key.chunk(2, dim=-1) - query_1, key_1 = apply_rope(query_1, key_1, image_rotary_emb) - query = torch.cat([query_1, query_2], dim=-1) - key = torch.cat([key_1, key_2], dim=-1) - - hidden_states = F.scaled_dot_product_attention( - query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2), dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - if not attn.single: - hidden_states_i, hidden_states_t = torch.split(hidden_states, [num_image_tokens, num_text_tokens], dim=1) - hidden_states_i = attn.to_out(hidden_states_i) - hidden_states_t = attn.to_out_t(hidden_states_t) - return hidden_states_i, hidden_states_t - else: - hidden_states = attn.to_out(hidden_states) - return hidden_states - - -# Modified from https://github.com/deepseek-ai/DeepSeek-V3/blob/main/inference/model.py -class MoEGate(nn.Module): - def __init__( - self, - embed_dim, - num_routed_experts=4, - num_activated_experts=2, - aux_loss_alpha=0.01, - _force_inference_output=False, - ): - super().__init__() - self.top_k = num_activated_experts - self.n_routed_experts = num_routed_experts - - self.scoring_func = "softmax" - self.alpha = aux_loss_alpha - self.seq_aux = False - - # topk selection algorithm - self.norm_topk_prob = False - self.gating_dim = embed_dim - self.weight = nn.Parameter(torch.randn(self.n_routed_experts, self.gating_dim) / embed_dim**0.5) - - self._force_inference_output = _force_inference_output - - def forward(self, hidden_states): - bsz, seq_len, h = hidden_states.shape - ### compute gating score - hidden_states = hidden_states.view(-1, h) - logits = F.linear(hidden_states, self.weight, None) - if self.scoring_func == "softmax": - scores = logits.softmax(dim=-1) - else: - raise NotImplementedError(f"insupportable scoring function for MoE gating: {self.scoring_func}") - - ### select top-k experts - topk_weight, topk_idx = torch.topk(scores, k=self.top_k, dim=-1, sorted=False) - - ### norm gate to sum 1 - if self.top_k > 1 and self.norm_topk_prob: - denominator = topk_weight.sum(dim=-1, keepdim=True) + 1e-20 - topk_weight = topk_weight / denominator - - ### expert-level computation auxiliary loss - if self.training and self.alpha > 0.0 and not self._force_inference_output: - scores_for_aux = scores - aux_topk = self.top_k - # always compute aux loss based on the naive greedy topk method - topk_idx_for_aux_loss = topk_idx.view(bsz, -1) - if self.seq_aux: - scores_for_seq_aux = scores_for_aux.view(bsz, seq_len, -1) - ce = torch.zeros(bsz, self.n_routed_experts, device=hidden_states.device) - ce.scatter_add_( - 1, topk_idx_for_aux_loss, torch.ones(bsz, seq_len * aux_topk, device=hidden_states.device) - ).div_(seq_len * aux_topk / self.n_routed_experts) - aux_loss = (ce * scores_for_seq_aux.mean(dim=1)).sum(dim=1).mean() * self.alpha - else: - mask_ce = F.one_hot(topk_idx_for_aux_loss.view(-1), num_classes=self.n_routed_experts) - ce = mask_ce.float().mean(0) - - Pi = scores_for_aux.mean(0) - fi = ce * self.n_routed_experts - aux_loss = (Pi * fi).sum() * self.alpha - else: - aux_loss = None - return topk_idx, topk_weight, aux_loss - - -# Modified from https://github.com/deepseek-ai/DeepSeek-V3/blob/main/inference/model.py -class MOEFeedForwardSwiGLU(nn.Module): - def __init__( - self, - dim: int, - hidden_dim: int, - num_routed_experts: int, - num_activated_experts: int, - _force_inference_output: bool = False, - ): - super().__init__() - self.shared_experts = HiDreamImageFeedForwardSwiGLU(dim, hidden_dim // 2) - self.experts = nn.ModuleList( - [HiDreamImageFeedForwardSwiGLU(dim, hidden_dim) for i in range(num_routed_experts)] - ) - self._force_inference_output = _force_inference_output - self.gate = MoEGate( - embed_dim=dim, - num_routed_experts=num_routed_experts, - num_activated_experts=num_activated_experts, - _force_inference_output=_force_inference_output, - ) - self.num_activated_experts = num_activated_experts - - def forward(self, x): - wtype = x.dtype - identity = x - orig_shape = x.shape - topk_idx, topk_weight, aux_loss = self.gate(x) - x = x.view(-1, x.shape[-1]) - flat_topk_idx = topk_idx.view(-1) - if self.training and not self._force_inference_output: - x = x.repeat_interleave(self.num_activated_experts, dim=0) - y = torch.empty_like(x, dtype=wtype) - for i, expert in enumerate(self.experts): - y[flat_topk_idx == i] = expert(x[flat_topk_idx == i]).to(dtype=wtype) - y = (y.view(*topk_weight.shape, -1) * topk_weight.unsqueeze(-1)).sum(dim=1) - y = y.view(*orig_shape).to(dtype=wtype) - # y = AddAuxiliaryLoss.apply(y, aux_loss) - else: - y = self.moe_infer(x, flat_topk_idx, topk_weight.view(-1, 1)).view(*orig_shape) - y = y + self.shared_experts(identity) - return y - - @torch.no_grad() - def moe_infer(self, x, flat_expert_indices, flat_expert_weights): - expert_cache = torch.zeros_like(x) - idxs = flat_expert_indices.argsort() - tokens_per_expert = flat_expert_indices.bincount().cpu().numpy().cumsum(0) - token_idxs = idxs // self.num_activated_experts - for i, end_idx in enumerate(tokens_per_expert): - start_idx = 0 if i == 0 else tokens_per_expert[i - 1] - if start_idx == end_idx: - continue - expert = self.experts[i] - exp_token_idx = token_idxs[start_idx:end_idx] - expert_tokens = x[exp_token_idx] - expert_out = expert(expert_tokens) - expert_out.mul_(flat_expert_weights[idxs[start_idx:end_idx]]) - - # for fp16 and other dtype - expert_cache = expert_cache.to(expert_out.dtype) - expert_cache.scatter_reduce_(0, exp_token_idx.view(-1, 1).repeat(1, x.shape[-1]), expert_out, reduce="sum") - return expert_cache - - -class TextProjection(nn.Module): - def __init__(self, in_features, hidden_size): - super().__init__() - self.linear = nn.Linear(in_features=in_features, out_features=hidden_size, bias=False) - - def forward(self, caption): - hidden_states = self.linear(caption) - return hidden_states - - -@maybe_allow_in_graph -class HiDreamImageSingleTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - num_routed_experts: int = 4, - num_activated_experts: int = 2, - _force_inference_output: bool = False, - ): - super().__init__() - self.num_attention_heads = num_attention_heads - self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(dim, 6 * dim, bias=True)) - - # 1. Attention - self.norm1_i = nn.LayerNorm(dim, eps=1e-06, elementwise_affine=False) - self.attn1 = HiDreamAttention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - processor=HiDreamAttnProcessor(), - single=True, - ) - - # 3. Feed-forward - self.norm3_i = nn.LayerNorm(dim, eps=1e-06, elementwise_affine=False) - if num_routed_experts > 0: - self.ff_i = MOEFeedForwardSwiGLU( - dim=dim, - hidden_dim=4 * dim, - num_routed_experts=num_routed_experts, - num_activated_experts=num_activated_experts, - _force_inference_output=_force_inference_output, - ) - else: - self.ff_i = HiDreamImageFeedForwardSwiGLU(dim=dim, hidden_dim=4 * dim) - - def forward( - self, - hidden_states: torch.Tensor, - hidden_states_masks: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor = None, - ) -> torch.Tensor: - wtype = hidden_states.dtype - shift_msa_i, scale_msa_i, gate_msa_i, shift_mlp_i, scale_mlp_i, gate_mlp_i = self.adaLN_modulation(temb)[ - :, None - ].chunk(6, dim=-1) - - # 1. MM-Attention - norm_hidden_states = self.norm1_i(hidden_states).to(dtype=wtype) - norm_hidden_states = norm_hidden_states * (1 + scale_msa_i) + shift_msa_i - attn_output_i = self.attn1( - norm_hidden_states, - hidden_states_masks, - image_rotary_emb=image_rotary_emb, - ) - hidden_states = gate_msa_i * attn_output_i + hidden_states - - # 2. Feed-forward - norm_hidden_states = self.norm3_i(hidden_states).to(dtype=wtype) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp_i) + shift_mlp_i - ff_output_i = gate_mlp_i * self.ff_i(norm_hidden_states.to(dtype=wtype)) - hidden_states = ff_output_i + hidden_states - return hidden_states - - -@maybe_allow_in_graph -class HiDreamImageTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - num_routed_experts: int = 4, - num_activated_experts: int = 2, - _force_inference_output: bool = False, - ): - super().__init__() - self.num_attention_heads = num_attention_heads - self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(dim, 12 * dim, bias=True)) - - # 1. Attention - self.norm1_i = nn.LayerNorm(dim, eps=1e-06, elementwise_affine=False) - self.norm1_t = nn.LayerNorm(dim, eps=1e-06, elementwise_affine=False) - self.attn1 = HiDreamAttention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - processor=HiDreamAttnProcessor(), - single=False, - ) - - # 3. Feed-forward - self.norm3_i = nn.LayerNorm(dim, eps=1e-06, elementwise_affine=False) - if num_routed_experts > 0: - self.ff_i = MOEFeedForwardSwiGLU( - dim=dim, - hidden_dim=4 * dim, - num_routed_experts=num_routed_experts, - num_activated_experts=num_activated_experts, - _force_inference_output=_force_inference_output, - ) - else: - self.ff_i = HiDreamImageFeedForwardSwiGLU(dim=dim, hidden_dim=4 * dim) - self.norm3_t = nn.LayerNorm(dim, eps=1e-06, elementwise_affine=False) - self.ff_t = HiDreamImageFeedForwardSwiGLU(dim=dim, hidden_dim=4 * dim) - - def forward( - self, - hidden_states: torch.Tensor, - hidden_states_masks: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - wtype = hidden_states.dtype - ( - shift_msa_i, - scale_msa_i, - gate_msa_i, - shift_mlp_i, - scale_mlp_i, - gate_mlp_i, - shift_msa_t, - scale_msa_t, - gate_msa_t, - shift_mlp_t, - scale_mlp_t, - gate_mlp_t, - ) = self.adaLN_modulation(temb)[:, None].chunk(12, dim=-1) - - # 1. MM-Attention - norm_hidden_states = self.norm1_i(hidden_states).to(dtype=wtype) - norm_hidden_states = norm_hidden_states * (1 + scale_msa_i) + shift_msa_i - norm_encoder_hidden_states = self.norm1_t(encoder_hidden_states).to(dtype=wtype) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + scale_msa_t) + shift_msa_t - - attn_output_i, attn_output_t = self.attn1( - norm_hidden_states, - hidden_states_masks, - norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - ) - - hidden_states = gate_msa_i * attn_output_i + hidden_states - encoder_hidden_states = gate_msa_t * attn_output_t + encoder_hidden_states - - # 2. Feed-forward - norm_hidden_states = self.norm3_i(hidden_states).to(dtype=wtype) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp_i) + shift_mlp_i - norm_encoder_hidden_states = self.norm3_t(encoder_hidden_states).to(dtype=wtype) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + scale_mlp_t) + shift_mlp_t - - ff_output_i = gate_mlp_i * self.ff_i(norm_hidden_states) - ff_output_t = gate_mlp_t * self.ff_t(norm_encoder_hidden_states) - hidden_states = ff_output_i + hidden_states - encoder_hidden_states = ff_output_t + encoder_hidden_states - return hidden_states, encoder_hidden_states - - -class HiDreamBlock(nn.Module): - def __init__(self, block: HiDreamImageTransformerBlock | HiDreamImageSingleTransformerBlock): - super().__init__() - self.block = block - - def forward( - self, - hidden_states: torch.Tensor, - hidden_states_masks: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor = None, - ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: - return self.block( - hidden_states=hidden_states, - hidden_states_masks=hidden_states_masks, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - ) - - -class HiDreamImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): - _supports_gradient_checkpointing = True - _no_split_modules = ["HiDreamImageTransformerBlock", "HiDreamImageSingleTransformerBlock"] - - @register_to_config - def __init__( - self, - patch_size: int | None = None, - in_channels: int = 64, - out_channels: int | None = None, - num_layers: int = 16, - num_single_layers: int = 32, - attention_head_dim: int = 128, - num_attention_heads: int = 20, - caption_channels: list[int] = None, - text_emb_dim: int = 2048, - num_routed_experts: int = 4, - num_activated_experts: int = 2, - axes_dims_rope: tuple[int, int] = (32, 32), - max_resolution: tuple[int, int] = (128, 128), - llama_layers: list[int] = None, - force_inference_output: bool = False, - ): - super().__init__() - self.out_channels = out_channels or in_channels - self.inner_dim = num_attention_heads * attention_head_dim - - self.t_embedder = HiDreamImageTimestepEmbed(self.inner_dim) - self.p_embedder = HiDreamImagePooledEmbed(text_emb_dim, self.inner_dim) - self.x_embedder = HiDreamImagePatchEmbed( - patch_size=patch_size, - in_channels=in_channels, - out_channels=self.inner_dim, - ) - self.pe_embedder = HiDreamImageEmbedND(theta=10000, axes_dim=axes_dims_rope) - - self.double_stream_blocks = nn.ModuleList( - [ - HiDreamBlock( - HiDreamImageTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - num_routed_experts=num_routed_experts, - num_activated_experts=num_activated_experts, - _force_inference_output=force_inference_output, - ) - ) - for _ in range(num_layers) - ] - ) - - self.single_stream_blocks = nn.ModuleList( - [ - HiDreamBlock( - HiDreamImageSingleTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - num_routed_experts=num_routed_experts, - num_activated_experts=num_activated_experts, - _force_inference_output=force_inference_output, - ) - ) - for _ in range(num_single_layers) - ] - ) - - self.final_layer = HiDreamImageOutEmbed(self.inner_dim, patch_size, self.out_channels) - - caption_channels = [caption_channels[1]] * (num_layers + num_single_layers) + [caption_channels[0]] - caption_projection = [] - for caption_channel in caption_channels: - caption_projection.append(TextProjection(in_features=caption_channel, hidden_size=self.inner_dim)) - self.caption_projection = nn.ModuleList(caption_projection) - self.max_seq = max_resolution[0] * max_resolution[1] // (patch_size * patch_size) - - self.gradient_checkpointing = False - - def unpatchify(self, x: torch.Tensor, img_sizes: list[tuple[int, int]], is_training: bool) -> list[torch.Tensor]: - if is_training and not self.config.force_inference_output: - B, S, F = x.shape - C = F // (self.config.patch_size * self.config.patch_size) - x = ( - x.reshape(B, S, self.config.patch_size, self.config.patch_size, C) - .permute(0, 4, 1, 2, 3) - .reshape(B, C, S, self.config.patch_size * self.config.patch_size) - ) - else: - x_arr = [] - p1 = self.config.patch_size - p2 = self.config.patch_size - for i, img_size in enumerate(img_sizes): - pH, pW = img_size - t = x[i, : pH * pW].reshape(1, pH, pW, -1) - F_token = t.shape[-1] - C = F_token // (p1 * p2) - t = t.reshape(1, pH, pW, p1, p2, C) - t = t.permute(0, 5, 1, 3, 2, 4) - t = t.reshape(1, C, pH * p1, pW * p2) - x_arr.append(t) - x = torch.cat(x_arr, dim=0) - return x - - def patchify(self, hidden_states): - batch_size, channels, height, width = hidden_states.shape - patch_size = self.config.patch_size - patch_height, patch_width = height // patch_size, width // patch_size - device = hidden_states.device - dtype = hidden_states.dtype - - # create img_sizes - img_sizes = torch.tensor([patch_height, patch_width], dtype=torch.int64, device=device).reshape(-1) - img_sizes = img_sizes.unsqueeze(0).repeat(batch_size, 1) - - # create hidden_states_masks - if hidden_states.shape[-2] != hidden_states.shape[-1]: - hidden_states_masks = torch.zeros((batch_size, self.max_seq), dtype=dtype, device=device) - hidden_states_masks[:, : patch_height * patch_width] = 1.0 - else: - hidden_states_masks = None - - # create img_ids - img_ids = torch.zeros(patch_height, patch_width, 3, device=device) - row_indices = torch.arange(patch_height, device=device)[:, None] - col_indices = torch.arange(patch_width, device=device)[None, :] - img_ids[..., 1] = img_ids[..., 1] + row_indices - img_ids[..., 2] = img_ids[..., 2] + col_indices - img_ids = img_ids.reshape(patch_height * patch_width, -1) - - if hidden_states.shape[-2] != hidden_states.shape[-1]: - # Handle non-square latents - img_ids_pad = torch.zeros(self.max_seq, 3, device=device) - img_ids_pad[: patch_height * patch_width, :] = img_ids - img_ids = img_ids_pad.unsqueeze(0).repeat(batch_size, 1, 1) - else: - img_ids = img_ids.unsqueeze(0).repeat(batch_size, 1, 1) - - # patchify hidden_states - if hidden_states.shape[-2] != hidden_states.shape[-1]: - # Handle non-square latents - out = torch.zeros( - (batch_size, channels, self.max_seq, patch_size * patch_size), - dtype=dtype, - device=device, - ) - hidden_states = hidden_states.reshape( - batch_size, channels, patch_height, patch_size, patch_width, patch_size - ) - hidden_states = hidden_states.permute(0, 1, 2, 4, 3, 5) - hidden_states = hidden_states.reshape( - batch_size, channels, patch_height * patch_width, patch_size * patch_size - ) - out[:, :, 0 : patch_height * patch_width] = hidden_states - hidden_states = out - hidden_states = hidden_states.permute(0, 2, 3, 1).reshape( - batch_size, self.max_seq, patch_size * patch_size * channels - ) - - else: - # Handle square latents - hidden_states = hidden_states.reshape( - batch_size, channels, patch_height, patch_size, patch_width, patch_size - ) - hidden_states = hidden_states.permute(0, 2, 4, 3, 5, 1) - hidden_states = hidden_states.reshape( - batch_size, patch_height * patch_width, patch_size * patch_size * channels - ) - - return hidden_states, hidden_states_masks, img_sizes, img_ids - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timesteps: torch.LongTensor = None, - encoder_hidden_states_t5: torch.Tensor = None, - encoder_hidden_states_llama3: torch.Tensor = None, - pooled_embeds: torch.Tensor = None, - img_ids: torch.Tensor | None = None, - img_sizes: list[tuple[int, int]] | None = None, - hidden_states_masks: torch.Tensor | None = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - **kwargs, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`HiDreamImageTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, in_channels, height, width)` or `(batch_size, patch_height * patch_width, patch_size * patch_size * channels)`): - Input `hidden_states`. - timesteps (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states_t5 (`torch.Tensor`): - Conditional embeddings computed from the T5 text encoder. - encoder_hidden_states_llama3 (`torch.Tensor`): - Conditional embeddings computed from the Llama3 text encoder. - pooled_embeds (`torch.Tensor`): - Pooled text embeddings used for additional conditioning. - img_ids (`torch.Tensor`, *optional*): - Image position ids for the patched hidden states. - img_sizes (`list` of `tuple` of `int`, *optional*): - Per-sample patch grid sizes used to unpatchify the output. - hidden_states_masks (`torch.Tensor`, *optional*): - Mask over patched `hidden_states`. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - encoder_hidden_states = kwargs.get("encoder_hidden_states", None) - - if encoder_hidden_states is not None: - deprecation_message = "The `encoder_hidden_states` argument is deprecated. Please use `encoder_hidden_states_t5` and `encoder_hidden_states_llama3` instead." - deprecate("encoder_hidden_states", "0.35.0", deprecation_message) - encoder_hidden_states_t5 = encoder_hidden_states[0] - encoder_hidden_states_llama3 = encoder_hidden_states[1] - - if img_ids is not None and img_sizes is not None and hidden_states_masks is None: - deprecation_message = ( - "Passing `img_ids` and `img_sizes` with unpachified `hidden_states` is deprecated and will be ignored." - ) - deprecate("img_ids", "0.35.0", deprecation_message) - - if hidden_states_masks is not None and (img_ids is None or img_sizes is None): - raise ValueError("if `hidden_states_masks` is passed, `img_ids` and `img_sizes` must also be passed.") - elif hidden_states_masks is not None and hidden_states.ndim != 3: - raise ValueError( - "if `hidden_states_masks` is passed, `hidden_states` must be a 3D tensors with shape (batch_size, patch_height * patch_width, patch_size * patch_size * channels)" - ) - - # spatial forward - batch_size = hidden_states.shape[0] - hidden_states_type = hidden_states.dtype - - # Patchify the input - if hidden_states_masks is None: - hidden_states, hidden_states_masks, img_sizes, img_ids = self.patchify(hidden_states) - - # Embed the hidden states - hidden_states = self.x_embedder(hidden_states) - - # 0. time - timesteps = self.t_embedder(timesteps, hidden_states_type) - p_embedder = self.p_embedder(pooled_embeds) - temb = timesteps + p_embedder - - encoder_hidden_states = [encoder_hidden_states_llama3[k] for k in self.config.llama_layers] - - if self.caption_projection is not None: - new_encoder_hidden_states = [] - for i, enc_hidden_state in enumerate(encoder_hidden_states): - enc_hidden_state = self.caption_projection[i](enc_hidden_state) - enc_hidden_state = enc_hidden_state.view(batch_size, -1, hidden_states.shape[-1]) - new_encoder_hidden_states.append(enc_hidden_state) - encoder_hidden_states = new_encoder_hidden_states - encoder_hidden_states_t5 = self.caption_projection[-1](encoder_hidden_states_t5) - encoder_hidden_states_t5 = encoder_hidden_states_t5.view(batch_size, -1, hidden_states.shape[-1]) - encoder_hidden_states.append(encoder_hidden_states_t5) - - txt_ids = torch.zeros( - batch_size, - encoder_hidden_states[-1].shape[1] - + encoder_hidden_states[-2].shape[1] - + encoder_hidden_states[0].shape[1], - 3, - device=img_ids.device, - dtype=img_ids.dtype, - ) - ids = torch.cat((img_ids, txt_ids), dim=1) - image_rotary_emb = self.pe_embedder(ids) - - # 2. Blocks - block_id = 0 - initial_encoder_hidden_states = torch.cat( - [ - encoder_hidden_states[-1].to(hidden_states.device), - encoder_hidden_states[-2].to(hidden_states.device), - ], - dim=1, - ) - initial_encoder_hidden_states_seq_len = initial_encoder_hidden_states.shape[1] - for bid, block in enumerate(self.double_stream_blocks): - cur_llama31_encoder_hidden_states = encoder_hidden_states[block_id].to(hidden_states.device) - cur_encoder_hidden_states = torch.cat( - [initial_encoder_hidden_states, cur_llama31_encoder_hidden_states], dim=1 - ) - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, initial_encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - hidden_states_masks, - cur_encoder_hidden_states, - temb, - image_rotary_emb, - ) - else: - hidden_states, initial_encoder_hidden_states = block( - hidden_states=hidden_states, - hidden_states_masks=hidden_states_masks, - encoder_hidden_states=cur_encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - ) - initial_encoder_hidden_states = initial_encoder_hidden_states[:, :initial_encoder_hidden_states_seq_len] - block_id += 1 - - image_tokens_seq_len = hidden_states.shape[1] - hidden_states = torch.cat([hidden_states, initial_encoder_hidden_states], dim=1) - hidden_states_seq_len = hidden_states.shape[1] - if hidden_states_masks is not None: - encoder_attention_mask_ones = torch.ones( - (batch_size, initial_encoder_hidden_states.shape[1] + cur_llama31_encoder_hidden_states.shape[1]), - device=hidden_states_masks.device, - dtype=hidden_states_masks.dtype, - ) - hidden_states_masks = torch.cat([hidden_states_masks, encoder_attention_mask_ones], dim=1) - - for bid, block in enumerate(self.single_stream_blocks): - cur_llama31_encoder_hidden_states = encoder_hidden_states[block_id].to(hidden_states.device) - hidden_states = torch.cat([hidden_states, cur_llama31_encoder_hidden_states], dim=1) - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - hidden_states_masks, - None, - temb, - image_rotary_emb, - ) - else: - hidden_states = block( - hidden_states=hidden_states, - hidden_states_masks=hidden_states_masks, - encoder_hidden_states=None, - temb=temb, - image_rotary_emb=image_rotary_emb, - ) - hidden_states = hidden_states[:, :hidden_states_seq_len] - block_id += 1 - - hidden_states = hidden_states[:, :image_tokens_seq_len, ...] - output = self.final_layer(hidden_states, temb) - output = self.unpatchify(output, img_sizes, self.training) - if hidden_states_masks is not None: - hidden_states_masks = hidden_states_masks[:, :image_tokens_seq_len] - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_hunyuan_video.py b/diffusers/models/transformers/transformer_hunyuan_video.py deleted file mode 100644 index 3730cc8ffa569760965c585ab0ddac829387ca65..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_hunyuan_video.py +++ /dev/null @@ -1,1126 +0,0 @@ -# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from diffusers.loaders import FromOriginalModelMixin - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ..attention import AttentionMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..attention_processor import Attention -from ..cache_utils import CacheMixin -from ..embeddings import ( - CombinedTimestepTextProjEmbeddings, - PixArtAlphaTextProjection, - TimestepEmbedding, - Timesteps, - get_1d_rotary_pos_embed, -) -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle, FP32LayerNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class HunyuanVideoAttnProcessor2_0: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "HunyuanVideoAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - if attn.add_q_proj is None and encoder_hidden_states is not None: - hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1) - - # 1. QKV projections - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - # 2. QK normalization - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # 3. Rotational positional embeddings applied to latent stream - if image_rotary_emb is not None: - from ..embeddings import apply_rotary_emb - - if attn.add_q_proj is None and encoder_hidden_states is not None: - query = torch.cat( - [ - apply_rotary_emb( - query[:, : -encoder_hidden_states.shape[1]], - image_rotary_emb, - sequence_dim=1, - ), - query[:, -encoder_hidden_states.shape[1] :], - ], - dim=1, - ) - key = torch.cat( - [ - apply_rotary_emb( - key[:, : -encoder_hidden_states.shape[1]], - image_rotary_emb, - sequence_dim=1, - ), - key[:, -encoder_hidden_states.shape[1] :], - ], - dim=1, - ) - else: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - # 4. Encoder condition QKV projection and normalization - if attn.add_q_proj is not None and encoder_hidden_states is not None: - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - encoder_query = encoder_query.unflatten(2, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(2, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(2, (attn.heads, -1)) - - if attn.norm_added_q is not None: - encoder_query = attn.norm_added_q(encoder_query) - if attn.norm_added_k is not None: - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([query, encoder_query], dim=1) - key = torch.cat([key, encoder_key], dim=1) - value = torch.cat([value, encoder_value], dim=1) - - # 5. Attention - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - # 6. Output projection - if encoder_hidden_states is not None: - hidden_states, encoder_hidden_states = ( - hidden_states[:, : -encoder_hidden_states.shape[1]], - hidden_states[:, -encoder_hidden_states.shape[1] :], - ) - - if getattr(attn, "to_out", None) is not None: - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - if getattr(attn, "to_add_out", None) is not None: - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - return hidden_states, encoder_hidden_states - - -class HunyuanVideoPatchEmbed(nn.Module): - def __init__( - self, - patch_size: int | tuple[int, int, int] = 16, - in_chans: int = 3, - embed_dim: int = 768, - ) -> None: - super().__init__() - - patch_size = (patch_size, patch_size, patch_size) if isinstance(patch_size, int) else patch_size - self.proj = nn.Conv3d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.proj(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) # BCFHW -> BNC - return hidden_states - - -class HunyuanVideoAdaNorm(nn.Module): - def __init__(self, in_features: int, out_features: int | None = None) -> None: - super().__init__() - - out_features = out_features or 2 * in_features - self.linear = nn.Linear(in_features, out_features) - self.nonlinearity = nn.SiLU() - - def forward( - self, temb: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - temb = self.linear(self.nonlinearity(temb)) - gate_msa, gate_mlp = temb.chunk(2, dim=1) - gate_msa, gate_mlp = gate_msa.unsqueeze(1), gate_mlp.unsqueeze(1) - return gate_msa, gate_mlp - - -class HunyuanVideoTokenReplaceAdaLayerNormZero(nn.Module): - def __init__(self, embedding_dim: int, norm_type: str = "layer_norm", bias: bool = True): - super().__init__() - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=bias) - - if norm_type == "layer_norm": - self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) - elif norm_type == "fp32_layer_norm": - self.norm = FP32LayerNorm(embedding_dim, elementwise_affine=False, bias=False) - else: - raise ValueError( - f"Unsupported `norm_type` ({norm_type}) provided. Supported ones are: 'layer_norm', 'fp32_layer_norm'." - ) - - def forward( - self, - hidden_states: torch.Tensor, - emb: torch.Tensor, - token_replace_emb: torch.Tensor, - first_frame_num_tokens: int, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - emb = self.linear(self.silu(emb)) - token_replace_emb = self.linear(self.silu(token_replace_emb)) - - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = emb.chunk(6, dim=1) - tr_shift_msa, tr_scale_msa, tr_gate_msa, tr_shift_mlp, tr_scale_mlp, tr_gate_mlp = token_replace_emb.chunk( - 6, dim=1 - ) - - norm_hidden_states = self.norm(hidden_states) - hidden_states_zero = ( - norm_hidden_states[:, :first_frame_num_tokens] * (1 + tr_scale_msa[:, None]) + tr_shift_msa[:, None] - ) - hidden_states_orig = ( - norm_hidden_states[:, first_frame_num_tokens:] * (1 + scale_msa[:, None]) + shift_msa[:, None] - ) - hidden_states = torch.cat([hidden_states_zero, hidden_states_orig], dim=1) - - return ( - hidden_states, - gate_msa, - shift_mlp, - scale_mlp, - gate_mlp, - tr_gate_msa, - tr_shift_mlp, - tr_scale_mlp, - tr_gate_mlp, - ) - - -class HunyuanVideoTokenReplaceAdaLayerNormZeroSingle(nn.Module): - def __init__(self, embedding_dim: int, norm_type: str = "layer_norm", bias: bool = True): - super().__init__() - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, 3 * embedding_dim, bias=bias) - - if norm_type == "layer_norm": - self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) - else: - raise ValueError( - f"Unsupported `norm_type` ({norm_type}) provided. Supported ones are: 'layer_norm', 'fp32_layer_norm'." - ) - - def forward( - self, - hidden_states: torch.Tensor, - emb: torch.Tensor, - token_replace_emb: torch.Tensor, - first_frame_num_tokens: int, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - emb = self.linear(self.silu(emb)) - token_replace_emb = self.linear(self.silu(token_replace_emb)) - - shift_msa, scale_msa, gate_msa = emb.chunk(3, dim=1) - tr_shift_msa, tr_scale_msa, tr_gate_msa = token_replace_emb.chunk(3, dim=1) - - norm_hidden_states = self.norm(hidden_states) - hidden_states_zero = ( - norm_hidden_states[:, :first_frame_num_tokens] * (1 + tr_scale_msa[:, None]) + tr_shift_msa[:, None] - ) - hidden_states_orig = ( - norm_hidden_states[:, first_frame_num_tokens:] * (1 + scale_msa[:, None]) + shift_msa[:, None] - ) - hidden_states = torch.cat([hidden_states_zero, hidden_states_orig], dim=1) - - return hidden_states, gate_msa, tr_gate_msa - - -class HunyuanVideoConditionEmbedding(nn.Module): - def __init__( - self, - embedding_dim: int, - pooled_projection_dim: int, - guidance_embeds: bool, - image_condition_type: str | None = None, - ): - super().__init__() - - self.image_condition_type = image_condition_type - - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - self.text_embedder = PixArtAlphaTextProjection(pooled_projection_dim, embedding_dim, act_fn="silu") - - self.guidance_embedder = None - if guidance_embeds: - self.guidance_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - def forward( - self, timestep: torch.Tensor, pooled_projection: torch.Tensor, guidance: torch.Tensor | None = None - ) -> tuple[torch.Tensor, torch.Tensor]: - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D) - pooled_projections = self.text_embedder(pooled_projection) - - token_replace_emb = None - if self.image_condition_type == "token_replace": - token_replace_timestep = torch.zeros_like(timestep) - token_replace_proj = self.time_proj(token_replace_timestep) - token_replace_emb = self.timestep_embedder(token_replace_proj.to(dtype=pooled_projection.dtype)) - token_replace_emb = token_replace_emb + pooled_projections - - if self.guidance_embedder is not None: - guidance_proj = self.time_proj(guidance) - guidance_emb = self.guidance_embedder(guidance_proj.to(dtype=pooled_projection.dtype)) - conditioning = timesteps_emb + guidance_emb + pooled_projections - else: - conditioning = timesteps_emb + pooled_projections - return conditioning, token_replace_emb - - -class HunyuanVideoIndividualTokenRefinerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - mlp_width_ratio: str = 4.0, - mlp_drop_rate: float = 0.0, - attention_bias: bool = True, - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6) - self.attn = Attention( - query_dim=hidden_size, - cross_attention_dim=None, - heads=num_attention_heads, - dim_head=attention_head_dim, - bias=attention_bias, - ) - - self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6) - self.ff = FeedForward(hidden_size, mult=mlp_width_ratio, activation_fn="linear-silu", dropout=mlp_drop_rate) - - self.norm_out = HunyuanVideoAdaNorm(hidden_size, 2 * hidden_size) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - norm_hidden_states = self.norm1(hidden_states) - - attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=None, - attention_mask=attention_mask, - ) - - gate_msa, gate_mlp = self.norm_out(temb) - hidden_states = hidden_states + attn_output * gate_msa - - ff_output = self.ff(self.norm2(hidden_states)) - hidden_states = hidden_states + ff_output * gate_mlp - - return hidden_states - - -class HunyuanVideoIndividualTokenRefiner(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - num_layers: int, - mlp_width_ratio: float = 4.0, - mlp_drop_rate: float = 0.0, - attention_bias: bool = True, - ) -> None: - super().__init__() - - self.refiner_blocks = nn.ModuleList( - [ - HunyuanVideoIndividualTokenRefinerBlock( - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - mlp_width_ratio=mlp_width_ratio, - mlp_drop_rate=mlp_drop_rate, - attention_bias=attention_bias, - ) - for _ in range(num_layers) - ] - ) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - ) -> None: - self_attn_mask = None - if attention_mask is not None: - batch_size = attention_mask.shape[0] - seq_len = attention_mask.shape[1] - attention_mask = attention_mask.to(hidden_states.device).bool() - self_attn_mask_1 = attention_mask.view(batch_size, 1, 1, seq_len).repeat(1, 1, seq_len, 1) - self_attn_mask_2 = self_attn_mask_1.transpose(2, 3) - self_attn_mask = (self_attn_mask_1 & self_attn_mask_2).bool() - self_attn_mask[:, :, :, 0] = True - - for block in self.refiner_blocks: - hidden_states = block(hidden_states, temb, self_attn_mask) - - return hidden_states - - -class HunyuanVideoTokenRefiner(nn.Module): - def __init__( - self, - in_channels: int, - num_attention_heads: int, - attention_head_dim: int, - num_layers: int, - mlp_ratio: float = 4.0, - mlp_drop_rate: float = 0.0, - attention_bias: bool = True, - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.time_text_embed = CombinedTimestepTextProjEmbeddings( - embedding_dim=hidden_size, pooled_projection_dim=in_channels - ) - self.proj_in = nn.Linear(in_channels, hidden_size, bias=True) - self.token_refiner = HunyuanVideoIndividualTokenRefiner( - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - num_layers=num_layers, - mlp_width_ratio=mlp_ratio, - mlp_drop_rate=mlp_drop_rate, - attention_bias=attention_bias, - ) - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - attention_mask: torch.LongTensor | None = None, - ) -> torch.Tensor: - if attention_mask is None: - pooled_projections = hidden_states.mean(dim=1) - else: - original_dtype = hidden_states.dtype - mask_float = attention_mask.float().unsqueeze(-1) - pooled_projections = (hidden_states * mask_float).sum(dim=1) / mask_float.sum(dim=1) - pooled_projections = pooled_projections.to(original_dtype) - - temb = self.time_text_embed(timestep, pooled_projections) - hidden_states = self.proj_in(hidden_states) - hidden_states = self.token_refiner(hidden_states, temb, attention_mask) - - return hidden_states - - -class HunyuanVideoRotaryPosEmbed(nn.Module): - def __init__(self, patch_size: int, patch_size_t: int, rope_dim: list[int], theta: float = 256.0) -> None: - super().__init__() - - self.patch_size = patch_size - self.patch_size_t = patch_size_t - self.rope_dim = rope_dim - self.theta = theta - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - rope_sizes = [num_frames // self.patch_size_t, height // self.patch_size, width // self.patch_size] - - axes_grids = [] - for i in range(3): - # Note: The following line diverges from original behaviour. We create the grid on the device, whereas - # original implementation creates it on CPU and then moves it to device. This results in numerical - # differences in layerwise debugging outputs, but visually it is the same. - grid = torch.arange(0, rope_sizes[i], device=hidden_states.device, dtype=torch.float32) - axes_grids.append(grid) - grid = torch.meshgrid(*axes_grids, indexing="ij") # [W, H, T] - grid = torch.stack(grid, dim=0) # [3, W, H, T] - - freqs = [] - for i in range(3): - freq = get_1d_rotary_pos_embed(self.rope_dim[i], grid[i].reshape(-1), self.theta, use_real=True) - freqs.append(freq) - - freqs_cos = torch.cat([f[0] for f in freqs], dim=1) # (W * H * T, D / 2) - freqs_sin = torch.cat([f[1] for f in freqs], dim=1) # (W * H * T, D / 2) - return freqs_cos, freqs_sin - - -class HunyuanVideoSingleTransformerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float = 4.0, - qk_norm: str = "rms_norm", - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - mlp_dim = int(hidden_size * mlp_ratio) - - self.attn = Attention( - query_dim=hidden_size, - cross_attention_dim=None, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=hidden_size, - bias=True, - processor=HunyuanVideoAttnProcessor2_0(), - qk_norm=qk_norm, - eps=1e-6, - pre_only=True, - ) - - self.norm = AdaLayerNormZeroSingle(hidden_size, norm_type="layer_norm") - self.proj_mlp = nn.Linear(hidden_size, mlp_dim) - self.act_mlp = nn.GELU(approximate="tanh") - self.proj_out = nn.Linear(hidden_size + mlp_dim, hidden_size) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - *args, - **kwargs, - ) -> tuple[torch.Tensor, torch.Tensor]: - text_seq_length = encoder_hidden_states.shape[1] - hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1) - - residual = hidden_states - - # 1. Input normalization - norm_hidden_states, gate = self.norm(hidden_states, emb=temb) - mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) - - norm_hidden_states, norm_encoder_hidden_states = ( - norm_hidden_states[:, :-text_seq_length, :], - norm_hidden_states[:, -text_seq_length:, :], - ) - - # 2. Attention - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - attn_output = torch.cat([attn_output, context_attn_output], dim=1) - - # 3. Modulation and residual connection - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) - hidden_states = gate.unsqueeze(1) * self.proj_out(hidden_states) - hidden_states = hidden_states + residual - - hidden_states, encoder_hidden_states = ( - hidden_states[:, :-text_seq_length, :], - hidden_states[:, -text_seq_length:, :], - ) - return hidden_states, encoder_hidden_states - - -class HunyuanVideoTransformerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float, - qk_norm: str = "rms_norm", - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.norm1 = AdaLayerNormZero(hidden_size, norm_type="layer_norm") - self.norm1_context = AdaLayerNormZero(hidden_size, norm_type="layer_norm") - - self.attn = Attention( - query_dim=hidden_size, - cross_attention_dim=None, - added_kv_proj_dim=hidden_size, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=hidden_size, - context_pre_only=False, - bias=True, - processor=HunyuanVideoAttnProcessor2_0(), - qk_norm=qk_norm, - eps=1e-6, - ) - - self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate") - - self.norm2_context = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.ff_context = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None, - *args, - **kwargs, - ) -> tuple[torch.Tensor, torch.Tensor]: - # 1. Input normalization - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - - # 2. Joint attention - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=freqs_cis, - ) - - # 3. Modulation and residual connection - hidden_states = hidden_states + attn_output * gate_msa.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + context_attn_output * c_gate_msa.unsqueeze(1) - - norm_hidden_states = self.norm2(hidden_states) - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - # 4. Feed-forward - ff_output = self.ff(norm_hidden_states) - context_ff_output = self.ff_context(norm_encoder_hidden_states) - - hidden_states = hidden_states + gate_mlp.unsqueeze(1) * ff_output - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - - return hidden_states, encoder_hidden_states - - -class HunyuanVideoTokenReplaceSingleTransformerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float = 4.0, - qk_norm: str = "rms_norm", - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - mlp_dim = int(hidden_size * mlp_ratio) - - self.attn = Attention( - query_dim=hidden_size, - cross_attention_dim=None, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=hidden_size, - bias=True, - processor=HunyuanVideoAttnProcessor2_0(), - qk_norm=qk_norm, - eps=1e-6, - pre_only=True, - ) - - self.norm = HunyuanVideoTokenReplaceAdaLayerNormZeroSingle(hidden_size, norm_type="layer_norm") - self.proj_mlp = nn.Linear(hidden_size, mlp_dim) - self.act_mlp = nn.GELU(approximate="tanh") - self.proj_out = nn.Linear(hidden_size + mlp_dim, hidden_size) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - token_replace_emb: torch.Tensor = None, - num_tokens: int = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - text_seq_length = encoder_hidden_states.shape[1] - hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1) - - residual = hidden_states - - # 1. Input normalization - norm_hidden_states, gate, tr_gate = self.norm(hidden_states, temb, token_replace_emb, num_tokens) - mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) - - norm_hidden_states, norm_encoder_hidden_states = ( - norm_hidden_states[:, :-text_seq_length, :], - norm_hidden_states[:, -text_seq_length:, :], - ) - - # 2. Attention - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - attn_output = torch.cat([attn_output, context_attn_output], dim=1) - - # 3. Modulation and residual connection - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) - - proj_output = self.proj_out(hidden_states) - hidden_states_zero = proj_output[:, :num_tokens] * tr_gate.unsqueeze(1) - hidden_states_orig = proj_output[:, num_tokens:] * gate.unsqueeze(1) - hidden_states = torch.cat([hidden_states_zero, hidden_states_orig], dim=1) - hidden_states = hidden_states + residual - - hidden_states, encoder_hidden_states = ( - hidden_states[:, :-text_seq_length, :], - hidden_states[:, -text_seq_length:, :], - ) - return hidden_states, encoder_hidden_states - - -class HunyuanVideoTokenReplaceTransformerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float, - qk_norm: str = "rms_norm", - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.norm1 = HunyuanVideoTokenReplaceAdaLayerNormZero(hidden_size, norm_type="layer_norm") - self.norm1_context = AdaLayerNormZero(hidden_size, norm_type="layer_norm") - - self.attn = Attention( - query_dim=hidden_size, - cross_attention_dim=None, - added_kv_proj_dim=hidden_size, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=hidden_size, - context_pre_only=False, - bias=True, - processor=HunyuanVideoAttnProcessor2_0(), - qk_norm=qk_norm, - eps=1e-6, - ) - - self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate") - - self.norm2_context = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.ff_context = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None, - token_replace_emb: torch.Tensor = None, - num_tokens: int = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - # 1. Input normalization - ( - norm_hidden_states, - gate_msa, - shift_mlp, - scale_mlp, - gate_mlp, - tr_gate_msa, - tr_shift_mlp, - tr_scale_mlp, - tr_gate_mlp, - ) = self.norm1(hidden_states, temb, token_replace_emb, num_tokens) - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - - # 2. Joint attention - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=freqs_cis, - ) - - # 3. Modulation and residual connection - hidden_states_zero = hidden_states[:, :num_tokens] + attn_output[:, :num_tokens] * tr_gate_msa.unsqueeze(1) - hidden_states_orig = hidden_states[:, num_tokens:] + attn_output[:, num_tokens:] * gate_msa.unsqueeze(1) - hidden_states = torch.cat([hidden_states_zero, hidden_states_orig], dim=1) - encoder_hidden_states = encoder_hidden_states + context_attn_output * c_gate_msa.unsqueeze(1) - - norm_hidden_states = self.norm2(hidden_states) - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - - hidden_states_zero = norm_hidden_states[:, :num_tokens] * (1 + tr_scale_mlp[:, None]) + tr_shift_mlp[:, None] - hidden_states_orig = norm_hidden_states[:, num_tokens:] * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - norm_hidden_states = torch.cat([hidden_states_zero, hidden_states_orig], dim=1) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - # 4. Feed-forward - ff_output = self.ff(norm_hidden_states) - context_ff_output = self.ff_context(norm_encoder_hidden_states) - - hidden_states_zero = hidden_states[:, :num_tokens] + ff_output[:, :num_tokens] * tr_gate_mlp.unsqueeze(1) - hidden_states_orig = hidden_states[:, num_tokens:] + ff_output[:, num_tokens:] * gate_mlp.unsqueeze(1) - hidden_states = torch.cat([hidden_states_zero, hidden_states_orig], dim=1) - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - - return hidden_states, encoder_hidden_states - - -class HunyuanVideoTransformer3DModel( - ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin -): - r""" - A Transformer model for video-like data used in [HunyuanVideo](https://huggingface.co/tencent/HunyuanVideo). - - Args: - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - num_attention_heads (`int`, defaults to `24`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - num_layers (`int`, defaults to `20`): - The number of layers of dual-stream blocks to use. - num_single_layers (`int`, defaults to `40`): - The number of layers of single-stream blocks to use. - num_refiner_layers (`int`, defaults to `2`): - The number of layers of refiner blocks to use. - mlp_ratio (`float`, defaults to `4.0`): - The ratio of the hidden layer size to the input size in the feedforward network. - patch_size (`int`, defaults to `2`): - The size of the spatial patches to use in the patch embedding layer. - patch_size_t (`int`, defaults to `1`): - The size of the tmeporal patches to use in the patch embedding layer. - qk_norm (`str`, defaults to `rms_norm`): - The normalization to use for the query and key projections in the attention layers. - guidance_embeds (`bool`, defaults to `True`): - Whether to use guidance embeddings in the model. - text_embed_dim (`int`, defaults to `4096`): - Input dimension of text embeddings from the text encoder. - pooled_projection_dim (`int`, defaults to `768`): - The dimension of the pooled projection of the text embeddings. - rope_theta (`float`, defaults to `256.0`): - The value of theta to use in the RoPE layer. - rope_axes_dim (`tuple[int]`, defaults to `(16, 56, 56)`): - The dimensions of the axes to use in the RoPE layer. - image_condition_type (`str`, *optional*, defaults to `None`): - The type of image conditioning to use. If `None`, no image conditioning is used. If `latent_concat`, the - image is concatenated to the latent stream. If `token_replace`, the image is used to replace first-frame - tokens in the latent stream and apply conditioning. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["x_embedder", "context_embedder", "norm"] - _no_split_modules = [ - "HunyuanVideoTransformerBlock", - "HunyuanVideoSingleTransformerBlock", - "HunyuanVideoTokenReplaceTransformerBlock", - "HunyuanVideoTokenReplaceSingleTransformerBlock", - "HunyuanVideoPatchEmbed", - "HunyuanVideoTokenRefiner", - ] - _repeated_blocks = [ - "HunyuanVideoTransformerBlock", - "HunyuanVideoSingleTransformerBlock", - "HunyuanVideoPatchEmbed", - "HunyuanVideoTokenRefiner", - ] - - @register_to_config - def __init__( - self, - in_channels: int = 16, - out_channels: int = 16, - num_attention_heads: int = 24, - attention_head_dim: int = 128, - num_layers: int = 20, - num_single_layers: int = 40, - num_refiner_layers: int = 2, - mlp_ratio: float = 4.0, - patch_size: int = 2, - patch_size_t: int = 1, - qk_norm: str = "rms_norm", - guidance_embeds: bool = True, - text_embed_dim: int = 4096, - pooled_projection_dim: int = 768, - rope_theta: float = 256.0, - rope_axes_dim: tuple[int, ...] = (16, 56, 56), - image_condition_type: str | None = None, - ) -> None: - super().__init__() - - supported_image_condition_types = ["latent_concat", "token_replace"] - if image_condition_type is not None and image_condition_type not in supported_image_condition_types: - raise ValueError( - f"Invalid `image_condition_type` ({image_condition_type}). Supported ones are: {supported_image_condition_types}" - ) - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - # 1. Latent and condition embedders - self.x_embedder = HunyuanVideoPatchEmbed((patch_size_t, patch_size, patch_size), in_channels, inner_dim) - self.context_embedder = HunyuanVideoTokenRefiner( - text_embed_dim, num_attention_heads, attention_head_dim, num_layers=num_refiner_layers - ) - - self.time_text_embed = HunyuanVideoConditionEmbedding( - inner_dim, pooled_projection_dim, guidance_embeds, image_condition_type - ) - - # 2. RoPE - self.rope = HunyuanVideoRotaryPosEmbed(patch_size, patch_size_t, rope_axes_dim, rope_theta) - - # 3. Dual stream transformer blocks - if image_condition_type == "token_replace": - self.transformer_blocks = nn.ModuleList( - [ - HunyuanVideoTokenReplaceTransformerBlock( - num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm - ) - for _ in range(num_layers) - ] - ) - else: - self.transformer_blocks = nn.ModuleList( - [ - HunyuanVideoTransformerBlock( - num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm - ) - for _ in range(num_layers) - ] - ) - - # 4. Single stream transformer blocks - if image_condition_type == "token_replace": - self.single_transformer_blocks = nn.ModuleList( - [ - HunyuanVideoTokenReplaceSingleTransformerBlock( - num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm - ) - for _ in range(num_single_layers) - ] - ) - else: - self.single_transformer_blocks = nn.ModuleList( - [ - HunyuanVideoSingleTransformerBlock( - num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm - ) - for _ in range(num_single_layers) - ] - ) - - # 5. Output projection - self.norm_out = AdaLayerNormContinuous(inner_dim, inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(inner_dim, patch_size_t * patch_size * patch_size * out_channels) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_attention_mask: torch.Tensor, - pooled_projections: torch.Tensor, - guidance: torch.Tensor = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`HunyuanVideoTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_attention_mask (`torch.Tensor`): - Mask applied to `encoder_hidden_states` during attention. - pooled_projections (`torch.Tensor` of shape `(batch_size, projection_dim)`): - Embeddings projected from the embeddings of input conditions. - guidance (`torch.Tensor`, *optional*): - Guidance scale embedding used for guidance-distilled variants of the model. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p, p_t = self.config.patch_size, self.config.patch_size_t - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p - post_patch_width = width // p - first_frame_num_tokens = 1 * post_patch_height * post_patch_width - - # 1. RoPE - image_rotary_emb = self.rope(hidden_states) - - # 2. Conditional embeddings - temb, token_replace_emb = self.time_text_embed(timestep, pooled_projections, guidance) - - hidden_states = self.x_embedder(hidden_states) - encoder_hidden_states = self.context_embedder(encoder_hidden_states, timestep, encoder_attention_mask) - - # 3. Attention mask preparation - latent_sequence_length = hidden_states.shape[1] - condition_sequence_length = encoder_hidden_states.shape[1] - sequence_length = latent_sequence_length + condition_sequence_length - attention_mask = torch.ones( - batch_size, sequence_length, device=hidden_states.device, dtype=torch.bool - ) # [B, N] - effective_condition_sequence_length = encoder_attention_mask.sum(dim=1, dtype=torch.int) # [B,] - effective_sequence_length = latent_sequence_length + effective_condition_sequence_length - indices = torch.arange(sequence_length, device=hidden_states.device).unsqueeze(0) # [1, N] - mask_indices = indices >= effective_sequence_length.unsqueeze(1) # [B, N] - attention_mask = attention_mask.masked_fill(mask_indices, False) - attention_mask = attention_mask.unsqueeze(1).unsqueeze(1) # [B, 1, 1, N] - - # 4. Transformer blocks - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.transformer_blocks: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - attention_mask, - image_rotary_emb, - token_replace_emb, - first_frame_num_tokens, - ) - - for block in self.single_transformer_blocks: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - attention_mask, - image_rotary_emb, - token_replace_emb, - first_frame_num_tokens, - ) - - else: - for block in self.transformer_blocks: - hidden_states, encoder_hidden_states = block( - hidden_states, - encoder_hidden_states, - temb, - attention_mask, - image_rotary_emb, - token_replace_emb, - first_frame_num_tokens, - ) - - for block in self.single_transformer_blocks: - hidden_states, encoder_hidden_states = block( - hidden_states, - encoder_hidden_states, - temb, - attention_mask, - image_rotary_emb, - token_replace_emb, - first_frame_num_tokens, - ) - - # 5. Output projection - hidden_states = self.norm_out(hidden_states, temb) - hidden_states = self.proj_out(hidden_states) - - hidden_states = hidden_states.reshape( - batch_size, post_patch_num_frames, post_patch_height, post_patch_width, -1, p_t, p, p - ) - hidden_states = hidden_states.permute(0, 4, 1, 5, 2, 6, 3, 7) - hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (hidden_states,) - - return Transformer2DModelOutput(sample=hidden_states) diff --git a/diffusers/models/transformers/transformer_hunyuan_video15.py b/diffusers/models/transformers/transformer_hunyuan_video15.py deleted file mode 100644 index 64c18e541d7ce63513d5d5ada0f6074fda4da06c..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_hunyuan_video15.py +++ /dev/null @@ -1,807 +0,0 @@ -# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from diffusers.loaders import FromOriginalModelMixin - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ..attention import AttentionMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..attention_processor import Attention -from ..cache_utils import CacheMixin -from ..embeddings import ( - CombinedTimestepTextProjEmbeddings, - TimestepEmbedding, - Timesteps, - get_1d_rotary_pos_embed, -) -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, AdaLayerNormZero - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class HunyuanVideo15AttnProcessor2_0: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "HunyuanVideo15AttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - # 1. QKV projections - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - # 2. QK normalization - query = attn.norm_q(query) - key = attn.norm_k(key) - - # 3. Rotational positional embeddings applied to latent stream - if image_rotary_emb is not None: - from ..embeddings import apply_rotary_emb - - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - # 4. Encoder condition QKV projection and normalization - if encoder_hidden_states is not None: - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - encoder_query = encoder_query.unflatten(2, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(2, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(2, (attn.heads, -1)) - - if attn.norm_added_q is not None: - encoder_query = attn.norm_added_q(encoder_query) - if attn.norm_added_k is not None: - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([query, encoder_query], dim=1) - key = torch.cat([key, encoder_key], dim=1) - value = torch.cat([value, encoder_value], dim=1) - - batch_size, seq_len, heads, dim = query.shape - attention_mask = F.pad(attention_mask, (seq_len - attention_mask.shape[1], 0), value=True) - attention_mask = attention_mask.bool() - self_attn_mask_1 = attention_mask.view(batch_size, 1, 1, seq_len).repeat(1, 1, seq_len, 1) - self_attn_mask_2 = self_attn_mask_1.transpose(2, 3) - attention_mask = (self_attn_mask_1 & self_attn_mask_2).bool() - - # 5. Attention - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - # 6. Output projection - if encoder_hidden_states is not None: - hidden_states, encoder_hidden_states = ( - hidden_states[:, : -encoder_hidden_states.shape[1]], - hidden_states[:, -encoder_hidden_states.shape[1] :], - ) - - if getattr(attn, "to_out", None) is not None: - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - if getattr(attn, "to_add_out", None) is not None: - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - return hidden_states, encoder_hidden_states - - -class HunyuanVideo15PatchEmbed(nn.Module): - def __init__( - self, - patch_size: int | tuple[int, int, int] = 16, - in_chans: int = 3, - embed_dim: int = 768, - ) -> None: - super().__init__() - - patch_size = (patch_size, patch_size, patch_size) if isinstance(patch_size, int) else patch_size - self.proj = nn.Conv3d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.proj(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) # BCFHW -> BNC - return hidden_states - - -class HunyuanVideo15AdaNorm(nn.Module): - def __init__(self, in_features: int, out_features: int | None = None) -> None: - super().__init__() - - out_features = out_features or 2 * in_features - self.linear = nn.Linear(in_features, out_features) - self.nonlinearity = nn.SiLU() - - def forward( - self, temb: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - temb = self.linear(self.nonlinearity(temb)) - gate_msa, gate_mlp = temb.chunk(2, dim=1) - gate_msa, gate_mlp = gate_msa.unsqueeze(1), gate_mlp.unsqueeze(1) - return gate_msa, gate_mlp - - -class HunyuanVideo15TimeEmbedding(nn.Module): - r""" - Time embedding for HunyuanVideo 1.5. - - Supports standard timestep embedding and optional reference timestep embedding for MeanFlow-based super-resolution - models. - - Args: - embedding_dim (`int`): - The dimension of the output embedding. - """ - - def __init__(self, embedding_dim: int, use_meanflow: bool = False): - super().__init__() - - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - self.use_meanflow = use_meanflow - self.time_proj_r = None - self.timestep_embedder_r = None - if use_meanflow: - self.time_proj_r = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder_r = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - def forward( - self, - timestep: torch.Tensor, - timestep_r: torch.Tensor | None = None, - ) -> torch.Tensor: - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=timestep.dtype)) - - if timestep_r is not None: - timesteps_proj_r = self.time_proj_r(timestep_r) - timesteps_emb_r = self.timestep_embedder_r(timesteps_proj_r.to(dtype=timestep.dtype)) - timesteps_emb = timesteps_emb + timesteps_emb_r - - return timesteps_emb - - -class HunyuanVideo15IndividualTokenRefinerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - mlp_width_ratio: str = 4.0, - mlp_drop_rate: float = 0.0, - attention_bias: bool = True, - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6) - self.attn = Attention( - query_dim=hidden_size, - cross_attention_dim=None, - heads=num_attention_heads, - dim_head=attention_head_dim, - bias=attention_bias, - ) - - self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6) - self.ff = FeedForward(hidden_size, mult=mlp_width_ratio, activation_fn="linear-silu", dropout=mlp_drop_rate) - - self.norm_out = HunyuanVideo15AdaNorm(hidden_size, 2 * hidden_size) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - norm_hidden_states = self.norm1(hidden_states) - - attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=None, - attention_mask=attention_mask, - ) - - gate_msa, gate_mlp = self.norm_out(temb) - hidden_states = hidden_states + attn_output * gate_msa - - ff_output = self.ff(self.norm2(hidden_states)) - hidden_states = hidden_states + ff_output * gate_mlp - - return hidden_states - - -class HunyuanVideo15IndividualTokenRefiner(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - num_layers: int, - mlp_width_ratio: float = 4.0, - mlp_drop_rate: float = 0.0, - attention_bias: bool = True, - ) -> None: - super().__init__() - - self.refiner_blocks = nn.ModuleList( - [ - HunyuanVideo15IndividualTokenRefinerBlock( - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - mlp_width_ratio=mlp_width_ratio, - mlp_drop_rate=mlp_drop_rate, - attention_bias=attention_bias, - ) - for _ in range(num_layers) - ] - ) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - ) -> None: - self_attn_mask = None - if attention_mask is not None: - batch_size = attention_mask.shape[0] - seq_len = attention_mask.shape[1] - attention_mask = attention_mask.to(hidden_states.device).bool() - self_attn_mask_1 = attention_mask.view(batch_size, 1, 1, seq_len).repeat(1, 1, seq_len, 1) - self_attn_mask_2 = self_attn_mask_1.transpose(2, 3) - self_attn_mask = (self_attn_mask_1 & self_attn_mask_2).bool() - - for block in self.refiner_blocks: - hidden_states = block(hidden_states, temb, self_attn_mask) - - return hidden_states - - -class HunyuanVideo15TokenRefiner(nn.Module): - def __init__( - self, - in_channels: int, - num_attention_heads: int, - attention_head_dim: int, - num_layers: int, - mlp_ratio: float = 4.0, - mlp_drop_rate: float = 0.0, - attention_bias: bool = True, - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.time_text_embed = CombinedTimestepTextProjEmbeddings( - embedding_dim=hidden_size, pooled_projection_dim=in_channels - ) - self.proj_in = nn.Linear(in_channels, hidden_size, bias=True) - self.token_refiner = HunyuanVideo15IndividualTokenRefiner( - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - num_layers=num_layers, - mlp_width_ratio=mlp_ratio, - mlp_drop_rate=mlp_drop_rate, - attention_bias=attention_bias, - ) - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - attention_mask: torch.LongTensor | None = None, - ) -> torch.Tensor: - if attention_mask is None: - pooled_projections = hidden_states.mean(dim=1) - else: - original_dtype = hidden_states.dtype - mask_float = attention_mask.float().unsqueeze(-1) - pooled_projections = (hidden_states * mask_float).sum(dim=1) / mask_float.sum(dim=1) - pooled_projections = pooled_projections.to(original_dtype) - - temb = self.time_text_embed(timestep, pooled_projections) - hidden_states = self.proj_in(hidden_states) - hidden_states = self.token_refiner(hidden_states, temb, attention_mask) - - return hidden_states - - -class HunyuanVideo15RotaryPosEmbed(nn.Module): - def __init__(self, patch_size: int, patch_size_t: int, rope_dim: list[int], theta: float = 256.0) -> None: - super().__init__() - - self.patch_size = patch_size - self.patch_size_t = patch_size_t - self.rope_dim = rope_dim - self.theta = theta - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - rope_sizes = [num_frames // self.patch_size_t, height // self.patch_size, width // self.patch_size] - - axes_grids = [] - for i in range(len(rope_sizes)): - # Note: The following line diverges from original behaviour. We create the grid on the device, whereas - # original implementation creates it on CPU and then moves it to device. This results in numerical - # differences in layerwise debugging outputs, but visually it is the same. - grid = torch.arange(0, rope_sizes[i], device=hidden_states.device, dtype=torch.float32) - axes_grids.append(grid) - grid = torch.meshgrid(*axes_grids, indexing="ij") # [W, H, T] - grid = torch.stack(grid, dim=0) # [3, W, H, T] - - freqs = [] - for i in range(3): - freq = get_1d_rotary_pos_embed(self.rope_dim[i], grid[i].reshape(-1), self.theta, use_real=True) - freqs.append(freq) - - freqs_cos = torch.cat([f[0] for f in freqs], dim=1) # (W * H * T, D / 2) - freqs_sin = torch.cat([f[1] for f in freqs], dim=1) # (W * H * T, D / 2) - return freqs_cos, freqs_sin - - -class HunyuanVideo15ByT5TextProjection(nn.Module): - def __init__(self, in_features: int, hidden_size: int, out_features: int): - super().__init__() - self.norm = nn.LayerNorm(in_features) - self.linear_1 = nn.Linear(in_features, hidden_size) - self.linear_2 = nn.Linear(hidden_size, hidden_size) - self.linear_3 = nn.Linear(hidden_size, out_features) - self.act_fn = nn.GELU() - - def forward(self, encoder_hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.norm(encoder_hidden_states) - hidden_states = self.linear_1(hidden_states) - hidden_states = self.act_fn(hidden_states) - hidden_states = self.linear_2(hidden_states) - hidden_states = self.act_fn(hidden_states) - hidden_states = self.linear_3(hidden_states) - return hidden_states - - -class HunyuanVideo15ImageProjection(nn.Module): - def __init__(self, in_channels: int, hidden_size: int): - super().__init__() - self.norm_in = nn.LayerNorm(in_channels) - self.linear_1 = nn.Linear(in_channels, in_channels) - self.act_fn = nn.GELU() - self.linear_2 = nn.Linear(in_channels, hidden_size) - self.norm_out = nn.LayerNorm(hidden_size) - - def forward(self, image_embeds: torch.Tensor) -> torch.Tensor: - hidden_states = self.norm_in(image_embeds) - hidden_states = self.linear_1(hidden_states) - hidden_states = self.act_fn(hidden_states) - hidden_states = self.linear_2(hidden_states) - hidden_states = self.norm_out(hidden_states) - return hidden_states - - -class HunyuanVideo15TransformerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float, - qk_norm: str = "rms_norm", - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.norm1 = AdaLayerNormZero(hidden_size, norm_type="layer_norm") - self.norm1_context = AdaLayerNormZero(hidden_size, norm_type="layer_norm") - - self.attn = Attention( - query_dim=hidden_size, - cross_attention_dim=None, - added_kv_proj_dim=hidden_size, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=hidden_size, - context_pre_only=False, - bias=True, - processor=HunyuanVideo15AttnProcessor2_0(), - qk_norm=qk_norm, - eps=1e-6, - ) - - self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate") - - self.norm2_context = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.ff_context = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None, - *args, - **kwargs, - ) -> tuple[torch.Tensor, torch.Tensor]: - # 1. Input normalization - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - - # 2. Joint attention - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=freqs_cis, - ) - - # 3. Modulation and residual connection - hidden_states = hidden_states + attn_output * gate_msa.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + context_attn_output * c_gate_msa.unsqueeze(1) - - norm_hidden_states = self.norm2(hidden_states) - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - # 4. Feed-forward - ff_output = self.ff(norm_hidden_states) - context_ff_output = self.ff_context(norm_encoder_hidden_states) - - hidden_states = hidden_states + gate_mlp.unsqueeze(1) * ff_output - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - - return hidden_states, encoder_hidden_states - - -class HunyuanVideo15Transformer3DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin, AttentionMixin -): - r""" - A Transformer model for video-like data used in [HunyuanVideo1.5](https://huggingface.co/tencent/HunyuanVideo1.5). - - Args: - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - num_attention_heads (`int`, defaults to `24`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - num_layers (`int`, defaults to `20`): - The number of layers of dual-stream blocks to use. - num_refiner_layers (`int`, defaults to `2`): - The number of layers of refiner blocks to use. - mlp_ratio (`float`, defaults to `4.0`): - The ratio of the hidden layer size to the input size in the feedforward network. - patch_size (`int`, defaults to `2`): - The size of the spatial patches to use in the patch embedding layer. - patch_size_t (`int`, defaults to `1`): - The size of the tmeporal patches to use in the patch embedding layer. - qk_norm (`str`, defaults to `rms_norm`): - The normalization to use for the query and key projections in the attention layers. - guidance_embeds (`bool`, defaults to `True`): - Whether to use guidance embeddings in the model. - text_embed_dim (`int`, defaults to `4096`): - Input dimension of text embeddings from the text encoder. - pooled_projection_dim (`int`, defaults to `768`): - The dimension of the pooled projection of the text embeddings. - rope_theta (`float`, defaults to `256.0`): - The value of theta to use in the RoPE layer. - rope_axes_dim (`tuple[int]`, defaults to `(16, 56, 56)`): - The dimensions of the axes to use in the RoPE layer. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["x_embedder", "context_embedder", "norm"] - _no_split_modules = [ - "HunyuanVideo15TransformerBlock", - "HunyuanVideo15PatchEmbed", - "HunyuanVideo15TokenRefiner", - ] - _repeated_blocks = [ - "HunyuanVideo15TransformerBlock", - "HunyuanVideo15PatchEmbed", - "HunyuanVideo15TokenRefiner", - ] - - @register_to_config - def __init__( - self, - in_channels: int = 65, - out_channels: int = 32, - num_attention_heads: int = 16, - attention_head_dim: int = 128, - num_layers: int = 54, - num_refiner_layers: int = 2, - mlp_ratio: float = 4.0, - patch_size: int = 1, - patch_size_t: int = 1, - qk_norm: str = "rms_norm", - text_embed_dim: int = 3584, - text_embed_2_dim: int = 1472, - image_embed_dim: int = 1152, - rope_theta: float = 256.0, - rope_axes_dim: tuple[int, ...] = (16, 56, 56), - # YiYi Notes: config based on target_size_config https://github.com/yiyixuxu/hy15/blob/main/hyvideo/pipelines/hunyuan_video_pipeline.py#L205 - target_size: int = 640, # did not name sample_size since it is in pixel spaces - task_type: str = "i2v", - use_meanflow: bool = False, - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - # 1. Latent and condition embedders - self.x_embedder = HunyuanVideo15PatchEmbed((patch_size_t, patch_size, patch_size), in_channels, inner_dim) - self.image_embedder = HunyuanVideo15ImageProjection(image_embed_dim, inner_dim) - - self.context_embedder = HunyuanVideo15TokenRefiner( - text_embed_dim, num_attention_heads, attention_head_dim, num_layers=num_refiner_layers - ) - self.context_embedder_2 = HunyuanVideo15ByT5TextProjection(text_embed_2_dim, 2048, inner_dim) - - self.time_embed = HunyuanVideo15TimeEmbedding(inner_dim, use_meanflow=use_meanflow) - - self.cond_type_embed = nn.Embedding(3, inner_dim) - - # 2. RoPE - self.rope = HunyuanVideo15RotaryPosEmbed(patch_size, patch_size_t, rope_axes_dim, rope_theta) - - # 3. Dual stream transformer blocks - - self.transformer_blocks = nn.ModuleList( - [ - HunyuanVideo15TransformerBlock( - num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm - ) - for _ in range(num_layers) - ] - ) - - # 5. Output projection - self.norm_out = AdaLayerNormContinuous(inner_dim, inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(inner_dim, patch_size_t * patch_size * patch_size * out_channels) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_attention_mask: torch.Tensor, - timestep_r: torch.LongTensor | None = None, - encoder_hidden_states_2: torch.Tensor | None = None, - encoder_attention_mask_2: torch.Tensor | None = None, - image_embeds: torch.Tensor | None = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`HunyuanVideo15Transformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_attention_mask (`torch.Tensor`): - Mask applied to `encoder_hidden_states` during attention. - timestep_r (`torch.LongTensor`, *optional*): - Refiner timestep conditioning. - encoder_hidden_states_2 (`torch.Tensor`, *optional*): - Additional conditional embeddings computed from a second text encoder (ByT5). - encoder_attention_mask_2 (`torch.Tensor`, *optional*): - Mask applied to `encoder_hidden_states_2` during attention. - image_embeds (`torch.Tensor`, *optional*): - Image embeddings for image-conditioned generation. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.config.patch_size_t, self.config.patch_size, self.config.patch_size - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p_h - post_patch_width = width // p_w - - # 1. RoPE - image_rotary_emb = self.rope(hidden_states) - - # 2. Conditional embeddings - temb = self.time_embed(timestep, timestep_r=timestep_r) - - hidden_states = self.x_embedder(hidden_states) - - # qwen text embedding - encoder_hidden_states = self.context_embedder(encoder_hidden_states, timestep, encoder_attention_mask) - - encoder_hidden_states_cond_emb = self.cond_type_embed( - torch.zeros_like(encoder_hidden_states[:, :, 0], dtype=torch.long) - ) - encoder_hidden_states = encoder_hidden_states + encoder_hidden_states_cond_emb - - # byt5 text embedding - encoder_hidden_states_2 = self.context_embedder_2(encoder_hidden_states_2) - - encoder_hidden_states_2_cond_emb = self.cond_type_embed( - torch.ones_like(encoder_hidden_states_2[:, :, 0], dtype=torch.long) - ) - encoder_hidden_states_2 = encoder_hidden_states_2 + encoder_hidden_states_2_cond_emb - - # image embed - encoder_hidden_states_3 = self.image_embedder(image_embeds) - is_t2v = torch.all(image_embeds == 0) - if is_t2v: - encoder_hidden_states_3 = encoder_hidden_states_3 * 0.0 - encoder_attention_mask_3 = torch.zeros( - (batch_size, encoder_hidden_states_3.shape[1]), - dtype=encoder_attention_mask.dtype, - device=encoder_attention_mask.device, - ) - else: - encoder_attention_mask_3 = torch.ones( - (batch_size, encoder_hidden_states_3.shape[1]), - dtype=encoder_attention_mask.dtype, - device=encoder_attention_mask.device, - ) - encoder_hidden_states_3_cond_emb = self.cond_type_embed( - 2 - * torch.ones_like( - encoder_hidden_states_3[:, :, 0], - dtype=torch.long, - ) - ) - encoder_hidden_states_3 = encoder_hidden_states_3 + encoder_hidden_states_3_cond_emb - - # reorder and combine text tokens: combine valid tokens first, then padding - encoder_attention_mask = encoder_attention_mask.bool() - encoder_attention_mask_2 = encoder_attention_mask_2.bool() - encoder_attention_mask_3 = encoder_attention_mask_3.bool() - new_encoder_hidden_states = [] - new_encoder_attention_mask = [] - - for text, text_mask, text_2, text_mask_2, image, image_mask in zip( - encoder_hidden_states, - encoder_attention_mask, - encoder_hidden_states_2, - encoder_attention_mask_2, - encoder_hidden_states_3, - encoder_attention_mask_3, - ): - # Concatenate: [valid_image, valid_byt5, valid_mllm, invalid_image, invalid_byt5, invalid_mllm] - new_encoder_hidden_states.append( - torch.cat( - [ - image[image_mask], # valid image - text_2[text_mask_2], # valid byt5 - text[text_mask], # valid mllm - image[~image_mask], # invalid image - torch.zeros_like(text_2[~text_mask_2]), # invalid byt5 (zeroed) - torch.zeros_like(text[~text_mask]), # invalid mllm (zeroed) - ], - dim=0, - ) - ) - - # Apply same reordering to attention masks - new_encoder_attention_mask.append( - torch.cat( - [ - image_mask[image_mask], - text_mask_2[text_mask_2], - text_mask[text_mask], - image_mask[~image_mask], - text_mask_2[~text_mask_2], - text_mask[~text_mask], - ], - dim=0, - ) - ) - - encoder_hidden_states = torch.stack(new_encoder_hidden_states) - encoder_attention_mask = torch.stack(new_encoder_attention_mask) - - # 4. Transformer blocks - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.transformer_blocks: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - encoder_attention_mask, - image_rotary_emb, - ) - - else: - for block in self.transformer_blocks: - hidden_states, encoder_hidden_states = block( - hidden_states, - encoder_hidden_states, - temb, - encoder_attention_mask, - image_rotary_emb, - ) - - # 5. Output projection - hidden_states = self.norm_out(hidden_states, temb) - hidden_states = self.proj_out(hidden_states) - - hidden_states = hidden_states.reshape( - batch_size, post_patch_num_frames, post_patch_height, post_patch_width, -1, p_t, p_h, p_w - ) - hidden_states = hidden_states.permute(0, 4, 1, 5, 2, 6, 3, 7) - hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (hidden_states,) - - return Transformer2DModelOutput(sample=hidden_states) diff --git a/diffusers/models/transformers/transformer_hunyuan_video_framepack.py b/diffusers/models/transformers/transformer_hunyuan_video_framepack.py deleted file mode 100644 index 9a3dbc00f4ec9a637969e5fea8e1a8fa7c3787d8..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_hunyuan_video_framepack.py +++ /dev/null @@ -1,442 +0,0 @@ -# Copyright 2025 The Framepack Team, The Hunyuan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, get_logger -from ..cache_utils import CacheMixin -from ..embeddings import get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous -from .transformer_hunyuan_video import ( - HunyuanVideoConditionEmbedding, - HunyuanVideoPatchEmbed, - HunyuanVideoSingleTransformerBlock, - HunyuanVideoTokenRefiner, - HunyuanVideoTransformerBlock, -) - - -logger = get_logger(__name__) # pylint: disable=invalid-name - - -class HunyuanVideoFramepackRotaryPosEmbed(nn.Module): - def __init__(self, patch_size: int, patch_size_t: int, rope_dim: list[int], theta: float = 256.0) -> None: - super().__init__() - - self.patch_size = patch_size - self.patch_size_t = patch_size_t - self.rope_dim = rope_dim - self.theta = theta - - def forward(self, frame_indices: torch.Tensor, height: int, width: int, device: torch.device): - height = height // self.patch_size - width = width // self.patch_size - grid = torch.meshgrid( - frame_indices.to(device=device, dtype=torch.float32), - torch.arange(0, height, device=device, dtype=torch.float32), - torch.arange(0, width, device=device, dtype=torch.float32), - indexing="ij", - ) # 3 * [W, H, T] - grid = torch.stack(grid, dim=0) # [3, W, H, T] - - freqs = [] - for i in range(3): - freq = get_1d_rotary_pos_embed(self.rope_dim[i], grid[i].reshape(-1), self.theta, use_real=True) - freqs.append(freq) - - freqs_cos = torch.cat([f[0] for f in freqs], dim=1) # (W * H * T, D / 2) - freqs_sin = torch.cat([f[1] for f in freqs], dim=1) # (W * H * T, D / 2) - - return freqs_cos, freqs_sin - - -class FramepackClipVisionProjection(nn.Module): - def __init__(self, in_channels: int, out_channels: int): - super().__init__() - self.up = nn.Linear(in_channels, out_channels * 3) - self.down = nn.Linear(out_channels * 3, out_channels) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.up(hidden_states) - hidden_states = F.silu(hidden_states) - hidden_states = self.down(hidden_states) - return hidden_states - - -class HunyuanVideoHistoryPatchEmbed(nn.Module): - def __init__(self, in_channels: int, inner_dim: int): - super().__init__() - self.proj = nn.Conv3d(in_channels, inner_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2)) - self.proj_2x = nn.Conv3d(in_channels, inner_dim, kernel_size=(2, 4, 4), stride=(2, 4, 4)) - self.proj_4x = nn.Conv3d(in_channels, inner_dim, kernel_size=(4, 8, 8), stride=(4, 8, 8)) - - def forward( - self, - latents_clean: torch.Tensor | None = None, - latents_clean_2x: torch.Tensor | None = None, - latents_clean_4x: torch.Tensor | None = None, - ): - if latents_clean is not None: - latents_clean = self.proj(latents_clean) - latents_clean = latents_clean.flatten(2).transpose(1, 2) - if latents_clean_2x is not None: - latents_clean_2x = _pad_for_3d_conv(latents_clean_2x, (2, 4, 4)) - latents_clean_2x = self.proj_2x(latents_clean_2x) - latents_clean_2x = latents_clean_2x.flatten(2).transpose(1, 2) - if latents_clean_4x is not None: - latents_clean_4x = _pad_for_3d_conv(latents_clean_4x, (4, 8, 8)) - latents_clean_4x = self.proj_4x(latents_clean_4x) - latents_clean_4x = latents_clean_4x.flatten(2).transpose(1, 2) - return latents_clean, latents_clean_2x, latents_clean_4x - - -class HunyuanVideoFramepackTransformer3DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin -): - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["x_embedder", "context_embedder", "norm"] - _no_split_modules = [ - "HunyuanVideoTransformerBlock", - "HunyuanVideoSingleTransformerBlock", - "HunyuanVideoHistoryPatchEmbed", - "HunyuanVideoTokenRefiner", - ] - - @register_to_config - def __init__( - self, - in_channels: int = 16, - out_channels: int = 16, - num_attention_heads: int = 24, - attention_head_dim: int = 128, - num_layers: int = 20, - num_single_layers: int = 40, - num_refiner_layers: int = 2, - mlp_ratio: float = 4.0, - patch_size: int = 2, - patch_size_t: int = 1, - qk_norm: str = "rms_norm", - guidance_embeds: bool = True, - text_embed_dim: int = 4096, - pooled_projection_dim: int = 768, - rope_theta: float = 256.0, - rope_axes_dim: tuple[int, ...] = (16, 56, 56), - image_condition_type: str | None = None, - has_image_proj: int = False, - image_proj_dim: int = 1152, - has_clean_x_embedder: int = False, - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - # 1. Latent and condition embedders - self.x_embedder = HunyuanVideoPatchEmbed((patch_size_t, patch_size, patch_size), in_channels, inner_dim) - - # Framepack history projection embedder - self.clean_x_embedder = None - if has_clean_x_embedder: - self.clean_x_embedder = HunyuanVideoHistoryPatchEmbed(in_channels, inner_dim) - - self.context_embedder = HunyuanVideoTokenRefiner( - text_embed_dim, num_attention_heads, attention_head_dim, num_layers=num_refiner_layers - ) - - # Framepack image-conditioning embedder - self.image_projection = FramepackClipVisionProjection(image_proj_dim, inner_dim) if has_image_proj else None - - self.time_text_embed = HunyuanVideoConditionEmbedding( - inner_dim, pooled_projection_dim, guidance_embeds, image_condition_type - ) - - # 2. RoPE - self.rope = HunyuanVideoFramepackRotaryPosEmbed(patch_size, patch_size_t, rope_axes_dim, rope_theta) - - # 3. Dual stream transformer blocks - self.transformer_blocks = nn.ModuleList( - [ - HunyuanVideoTransformerBlock( - num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm - ) - for _ in range(num_layers) - ] - ) - - # 4. Single stream transformer blocks - self.single_transformer_blocks = nn.ModuleList( - [ - HunyuanVideoSingleTransformerBlock( - num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm - ) - for _ in range(num_single_layers) - ] - ) - - # 5. Output projection - self.norm_out = AdaLayerNormContinuous(inner_dim, inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(inner_dim, patch_size_t * patch_size * patch_size * out_channels) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_attention_mask: torch.Tensor, - pooled_projections: torch.Tensor, - image_embeds: torch.Tensor, - indices_latents: torch.Tensor, - guidance: torch.Tensor | None = None, - latents_clean: torch.Tensor | None = None, - indices_latents_clean: torch.Tensor | None = None, - latents_history_2x: torch.Tensor | None = None, - indices_latents_history_2x: torch.Tensor | None = None, - latents_history_4x: torch.Tensor | None = None, - indices_latents_history_4x: torch.Tensor | None = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor] | Transformer2DModelOutput: - """ - The [`HunyuanVideoFramepackTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_attention_mask (`torch.Tensor`): - Mask applied to `encoder_hidden_states` during attention. - pooled_projections (`torch.Tensor` of shape `(batch_size, projection_dim)`): - Embeddings projected from the embeddings of input conditions. - image_embeds (`torch.Tensor`): - Image embeddings for image-conditioned generation. - indices_latents (`torch.Tensor`): - Frame indices for `hidden_states` used to compute the rotary positional embeddings. - guidance (`torch.Tensor`, *optional*): - Guidance scale embedding used for guidance-distilled variants of the model. - latents_clean (`torch.Tensor`, *optional*): - Clean (denoised) history latents conditioning. - indices_latents_clean (`torch.Tensor`, *optional*): - Frame indices for `latents_clean`. - latents_history_2x (`torch.Tensor`, *optional*): - 2x downsampled history latents conditioning. - indices_latents_history_2x (`torch.Tensor`, *optional*): - Frame indices for `latents_history_2x`. - latents_history_4x (`torch.Tensor`, *optional*): - 4x downsampled history latents conditioning. - indices_latents_history_4x (`torch.Tensor`, *optional*): - Frame indices for `latents_history_4x`. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p, p_t = self.config.patch_size, self.config.patch_size_t - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p - post_patch_width = width // p - original_context_length = post_patch_num_frames * post_patch_height * post_patch_width - - if indices_latents is None: - indices_latents = torch.arange(0, num_frames).unsqueeze(0).expand(batch_size, -1) - - hidden_states = self.x_embedder(hidden_states) - image_rotary_emb = self.rope( - frame_indices=indices_latents, height=height, width=width, device=hidden_states.device - ) - - latents_clean, latents_history_2x, latents_history_4x = self.clean_x_embedder( - latents_clean, latents_history_2x, latents_history_4x - ) - - if latents_clean is not None and indices_latents_clean is not None: - image_rotary_emb_clean = self.rope( - frame_indices=indices_latents_clean, height=height, width=width, device=hidden_states.device - ) - if latents_history_2x is not None and indices_latents_history_2x is not None: - image_rotary_emb_history_2x = self.rope( - frame_indices=indices_latents_history_2x, height=height, width=width, device=hidden_states.device - ) - if latents_history_4x is not None and indices_latents_history_4x is not None: - image_rotary_emb_history_4x = self.rope( - frame_indices=indices_latents_history_4x, height=height, width=width, device=hidden_states.device - ) - - hidden_states, image_rotary_emb = self._pack_history_states( - hidden_states, - latents_clean, - latents_history_2x, - latents_history_4x, - image_rotary_emb, - image_rotary_emb_clean, - image_rotary_emb_history_2x, - image_rotary_emb_history_4x, - post_patch_height, - post_patch_width, - ) - - temb, _ = self.time_text_embed(timestep, pooled_projections, guidance) - encoder_hidden_states = self.context_embedder(encoder_hidden_states, timestep, encoder_attention_mask) - - encoder_hidden_states_image = self.image_projection(image_embeds) - attention_mask_image = encoder_attention_mask.new_ones((batch_size, encoder_hidden_states_image.shape[1])) - - # must cat before (not after) encoder_hidden_states, due to attn masking - encoder_hidden_states = torch.cat([encoder_hidden_states_image, encoder_hidden_states], dim=1) - encoder_attention_mask = torch.cat([attention_mask_image, encoder_attention_mask], dim=1) - - latent_sequence_length = hidden_states.shape[1] - condition_sequence_length = encoder_hidden_states.shape[1] - sequence_length = latent_sequence_length + condition_sequence_length - attention_mask = torch.zeros( - batch_size, sequence_length, device=hidden_states.device, dtype=torch.bool - ) # [B, N] - effective_condition_sequence_length = encoder_attention_mask.sum(dim=1, dtype=torch.int) # [B,] - effective_sequence_length = latent_sequence_length + effective_condition_sequence_length - - if batch_size == 1: - encoder_hidden_states = encoder_hidden_states[:, : effective_condition_sequence_length[0]] - attention_mask = None - else: - for i in range(batch_size): - attention_mask[i, : effective_sequence_length[i]] = True - # [B, 1, 1, N], for broadcasting across attention heads - attention_mask = attention_mask.unsqueeze(1).unsqueeze(1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.transformer_blocks: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, hidden_states, encoder_hidden_states, temb, attention_mask, image_rotary_emb - ) - - for block in self.single_transformer_blocks: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, hidden_states, encoder_hidden_states, temb, attention_mask, image_rotary_emb - ) - - else: - for block in self.transformer_blocks: - hidden_states, encoder_hidden_states = block( - hidden_states, encoder_hidden_states, temb, attention_mask, image_rotary_emb - ) - - for block in self.single_transformer_blocks: - hidden_states, encoder_hidden_states = block( - hidden_states, encoder_hidden_states, temb, attention_mask, image_rotary_emb - ) - - hidden_states = hidden_states[:, -original_context_length:] - hidden_states = self.norm_out(hidden_states, temb) - hidden_states = self.proj_out(hidden_states) - - hidden_states = hidden_states.reshape( - batch_size, post_patch_num_frames, post_patch_height, post_patch_width, -1, p_t, p, p - ) - hidden_states = hidden_states.permute(0, 4, 1, 5, 2, 6, 3, 7) - hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (hidden_states,) - return Transformer2DModelOutput(sample=hidden_states) - - def _pack_history_states( - self, - hidden_states: torch.Tensor, - latents_clean: torch.Tensor | None = None, - latents_history_2x: torch.Tensor | None = None, - latents_history_4x: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] = None, - image_rotary_emb_clean: tuple[torch.Tensor, torch.Tensor] | None = None, - image_rotary_emb_history_2x: tuple[torch.Tensor, torch.Tensor] | None = None, - image_rotary_emb_history_4x: tuple[torch.Tensor, torch.Tensor] | None = None, - height: int = None, - width: int = None, - ): - image_rotary_emb = list(image_rotary_emb) # convert tuple to list for in-place modification - - if latents_clean is not None and image_rotary_emb_clean is not None: - hidden_states = torch.cat([latents_clean, hidden_states], dim=1) - image_rotary_emb[0] = torch.cat([image_rotary_emb_clean[0], image_rotary_emb[0]], dim=0) - image_rotary_emb[1] = torch.cat([image_rotary_emb_clean[1], image_rotary_emb[1]], dim=0) - - if latents_history_2x is not None and image_rotary_emb_history_2x is not None: - hidden_states = torch.cat([latents_history_2x, hidden_states], dim=1) - image_rotary_emb_history_2x = self._pad_rotary_emb(image_rotary_emb_history_2x, height, width, (2, 2, 2)) - image_rotary_emb[0] = torch.cat([image_rotary_emb_history_2x[0], image_rotary_emb[0]], dim=0) - image_rotary_emb[1] = torch.cat([image_rotary_emb_history_2x[1], image_rotary_emb[1]], dim=0) - - if latents_history_4x is not None and image_rotary_emb_history_4x is not None: - hidden_states = torch.cat([latents_history_4x, hidden_states], dim=1) - image_rotary_emb_history_4x = self._pad_rotary_emb(image_rotary_emb_history_4x, height, width, (4, 4, 4)) - image_rotary_emb[0] = torch.cat([image_rotary_emb_history_4x[0], image_rotary_emb[0]], dim=0) - image_rotary_emb[1] = torch.cat([image_rotary_emb_history_4x[1], image_rotary_emb[1]], dim=0) - - return hidden_states, tuple(image_rotary_emb) - - def _pad_rotary_emb( - self, - image_rotary_emb: tuple[torch.Tensor], - height: int, - width: int, - kernel_size: tuple[int, int, int], - ): - # freqs_cos, freqs_sin have shape [W * H * T, D / 2], where D is attention head dim - freqs_cos, freqs_sin = image_rotary_emb - freqs_cos = freqs_cos.unsqueeze(0).permute(0, 2, 1).unflatten(2, (-1, height, width)) - freqs_sin = freqs_sin.unsqueeze(0).permute(0, 2, 1).unflatten(2, (-1, height, width)) - freqs_cos = _pad_for_3d_conv(freqs_cos, kernel_size) - freqs_sin = _pad_for_3d_conv(freqs_sin, kernel_size) - freqs_cos = _center_down_sample_3d(freqs_cos, kernel_size) - freqs_sin = _center_down_sample_3d(freqs_sin, kernel_size) - freqs_cos = freqs_cos.flatten(2).permute(0, 2, 1).squeeze(0) - freqs_sin = freqs_sin.flatten(2).permute(0, 2, 1).squeeze(0) - return freqs_cos, freqs_sin - - -def _pad_for_3d_conv(x, kernel_size): - if isinstance(x, (tuple, list)): - return tuple(_pad_for_3d_conv(i, kernel_size) for i in x) - b, c, t, h, w = x.shape - pt, ph, pw = kernel_size - pad_t = (pt - (t % pt)) % pt - pad_h = (ph - (h % ph)) % ph - pad_w = (pw - (w % pw)) % pw - return torch.nn.functional.pad(x, (0, pad_w, 0, pad_h, 0, pad_t), mode="replicate") - - -def _center_down_sample_3d(x, kernel_size): - if isinstance(x, (tuple, list)): - return tuple(_center_down_sample_3d(i, kernel_size) for i in x) - return torch.nn.functional.avg_pool3d(x, kernel_size, stride=kernel_size) diff --git a/diffusers/models/transformers/transformer_hunyuanimage.py b/diffusers/models/transformers/transformer_hunyuanimage.py deleted file mode 100644 index dd2176a4096f903b8cbfedcae50a3052a84654ad..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_hunyuanimage.py +++ /dev/null @@ -1,922 +0,0 @@ -# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from diffusers.loaders import FromOriginalModelMixin - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import AttentionMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..attention_processor import Attention -from ..cache_utils import CacheMixin -from ..embeddings import ( - CombinedTimestepTextProjEmbeddings, - TimestepEmbedding, - Timesteps, - get_1d_rotary_pos_embed, -) -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class HunyuanImageAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "HunyuanImageAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - if attn.add_q_proj is None and encoder_hidden_states is not None: - hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1) - - # 1. QKV projections - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(2, (attn.heads, -1)) # batch_size, seq_len, heads, head_dim - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - # 2. QK normalization - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # 3. Rotational positional embeddings applied to latent stream - if image_rotary_emb is not None: - from ..embeddings import apply_rotary_emb - - if attn.add_q_proj is None and encoder_hidden_states is not None: - query = torch.cat( - [ - apply_rotary_emb( - query[:, : -encoder_hidden_states.shape[1]], image_rotary_emb, sequence_dim=1 - ), - query[:, -encoder_hidden_states.shape[1] :], - ], - dim=1, - ) - key = torch.cat( - [ - apply_rotary_emb(key[:, : -encoder_hidden_states.shape[1]], image_rotary_emb, sequence_dim=1), - key[:, -encoder_hidden_states.shape[1] :], - ], - dim=1, - ) - else: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - # 4. Encoder condition QKV projection and normalization - if attn.add_q_proj is not None and encoder_hidden_states is not None: - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - encoder_query = encoder_query.unflatten(2, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(2, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(2, (attn.heads, -1)) - - if attn.norm_added_q is not None: - encoder_query = attn.norm_added_q(encoder_query) - if attn.norm_added_k is not None: - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([query, encoder_query], dim=1) - key = torch.cat([key, encoder_key], dim=1) - value = torch.cat([value, encoder_value], dim=1) - - # 5. Attention - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - # 6. Output projection - if encoder_hidden_states is not None: - hidden_states, encoder_hidden_states = ( - hidden_states[:, : -encoder_hidden_states.shape[1]], - hidden_states[:, -encoder_hidden_states.shape[1] :], - ) - - if getattr(attn, "to_out", None) is not None: - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - if getattr(attn, "to_add_out", None) is not None: - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - return hidden_states, encoder_hidden_states - - -class HunyuanImagePatchEmbed(nn.Module): - def __init__( - self, - patch_size: tuple[int, int, tuple[int, int, int]] = (16, 16), - in_chans: int = 3, - embed_dim: int = 768, - ) -> None: - super().__init__() - - self.patch_size = patch_size - - if len(patch_size) == 2: - self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) - elif len(patch_size) == 3: - self.proj = nn.Conv3d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) - else: - raise ValueError(f"patch_size must be a tuple of length 2 or 3, got {len(patch_size)}") - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.proj(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) - return hidden_states - - -class HunyuanImageByT5TextProjection(nn.Module): - def __init__(self, in_features: int, hidden_size: int, out_features: int): - super().__init__() - self.norm = nn.LayerNorm(in_features) - self.linear_1 = nn.Linear(in_features, hidden_size) - self.linear_2 = nn.Linear(hidden_size, hidden_size) - self.linear_3 = nn.Linear(hidden_size, out_features) - self.act_fn = nn.GELU() - - def forward(self, encoder_hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.norm(encoder_hidden_states) - hidden_states = self.linear_1(hidden_states) - hidden_states = self.act_fn(hidden_states) - hidden_states = self.linear_2(hidden_states) - hidden_states = self.act_fn(hidden_states) - hidden_states = self.linear_3(hidden_states) - return hidden_states - - -class HunyuanImageAdaNorm(nn.Module): - def __init__(self, in_features: int, out_features: int | None = None) -> None: - super().__init__() - - out_features = out_features or 2 * in_features - self.linear = nn.Linear(in_features, out_features) - self.nonlinearity = nn.SiLU() - - def forward( - self, temb: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - temb = self.linear(self.nonlinearity(temb)) - gate_msa, gate_mlp = temb.chunk(2, dim=1) - gate_msa, gate_mlp = gate_msa.unsqueeze(1), gate_mlp.unsqueeze(1) - return gate_msa, gate_mlp - - -class HunyuanImageCombinedTimeGuidanceEmbedding(nn.Module): - def __init__( - self, - embedding_dim: int, - guidance_embeds: bool = False, - use_meanflow: bool = False, - ): - super().__init__() - - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - self.use_meanflow = use_meanflow - - self.time_proj_r = None - self.timestep_embedder_r = None - if use_meanflow: - self.time_proj_r = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder_r = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - self.guidance_embedder = None - if guidance_embeds: - self.guidance_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - def forward( - self, - timestep: torch.Tensor, - timestep_r: torch.Tensor | None = None, - guidance: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=timestep.dtype)) - - if timestep_r is not None: - timesteps_proj_r = self.time_proj_r(timestep_r) - timesteps_emb_r = self.timestep_embedder_r(timesteps_proj_r.to(dtype=timestep.dtype)) - timesteps_emb = (timesteps_emb + timesteps_emb_r) / 2 - - if self.guidance_embedder is not None: - guidance_proj = self.time_proj(guidance) - guidance_emb = self.guidance_embedder(guidance_proj.to(dtype=timestep.dtype)) - conditioning = timesteps_emb + guidance_emb - else: - conditioning = timesteps_emb - - return conditioning - - -# IndividualTokenRefinerBlock -@maybe_allow_in_graph -class HunyuanImageIndividualTokenRefinerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, # 28 - attention_head_dim: int, # 128 - mlp_width_ratio: str = 4.0, - mlp_drop_rate: float = 0.0, - attention_bias: bool = True, - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6) - self.attn = Attention( - query_dim=hidden_size, - cross_attention_dim=None, - heads=num_attention_heads, - dim_head=attention_head_dim, - bias=attention_bias, - ) - - self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6) - self.ff = FeedForward(hidden_size, mult=mlp_width_ratio, activation_fn="linear-silu", dropout=mlp_drop_rate) - - self.norm_out = HunyuanImageAdaNorm(hidden_size, 2 * hidden_size) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - norm_hidden_states = self.norm1(hidden_states) - - attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=None, - attention_mask=attention_mask, - ) - - gate_msa, gate_mlp = self.norm_out(temb) - hidden_states = hidden_states + attn_output * gate_msa - - ff_output = self.ff(self.norm2(hidden_states)) - hidden_states = hidden_states + ff_output * gate_mlp - - return hidden_states - - -class HunyuanImageIndividualTokenRefiner(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - num_layers: int, - mlp_width_ratio: float = 4.0, - mlp_drop_rate: float = 0.0, - attention_bias: bool = True, - ) -> None: - super().__init__() - - self.refiner_blocks = nn.ModuleList( - [ - HunyuanImageIndividualTokenRefinerBlock( - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - mlp_width_ratio=mlp_width_ratio, - mlp_drop_rate=mlp_drop_rate, - attention_bias=attention_bias, - ) - for _ in range(num_layers) - ] - ) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - ) -> None: - self_attn_mask = None - if attention_mask is not None: - batch_size = attention_mask.shape[0] - seq_len = attention_mask.shape[1] - attention_mask = attention_mask.to(hidden_states.device) - self_attn_mask_1 = attention_mask.view(batch_size, 1, 1, seq_len).repeat(1, 1, seq_len, 1) - self_attn_mask_2 = self_attn_mask_1.transpose(2, 3) - self_attn_mask = (self_attn_mask_1 & self_attn_mask_2).bool() - self_attn_mask[:, :, :, 0] = True - - for block in self.refiner_blocks: - hidden_states = block(hidden_states, temb, self_attn_mask) - - return hidden_states - - -# txt_in -class HunyuanImageTokenRefiner(nn.Module): - def __init__( - self, - in_channels: int, - num_attention_heads: int, - attention_head_dim: int, - num_layers: int, - mlp_ratio: float = 4.0, - mlp_drop_rate: float = 0.0, - attention_bias: bool = True, - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.time_text_embed = CombinedTimestepTextProjEmbeddings( - embedding_dim=hidden_size, pooled_projection_dim=in_channels - ) - self.proj_in = nn.Linear(in_channels, hidden_size, bias=True) - self.token_refiner = HunyuanImageIndividualTokenRefiner( - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - num_layers=num_layers, - mlp_width_ratio=mlp_ratio, - mlp_drop_rate=mlp_drop_rate, - attention_bias=attention_bias, - ) - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - attention_mask: torch.LongTensor | None = None, - ) -> torch.Tensor: - if attention_mask is None: - pooled_hidden_states = hidden_states.mean(dim=1) - else: - original_dtype = hidden_states.dtype - mask_float = attention_mask.float().unsqueeze(-1) - pooled_hidden_states = (hidden_states * mask_float).sum(dim=1) / mask_float.sum(dim=1) - pooled_hidden_states = pooled_hidden_states.to(original_dtype) - - temb = self.time_text_embed(timestep, pooled_hidden_states) - hidden_states = self.proj_in(hidden_states) - hidden_states = self.token_refiner(hidden_states, temb, attention_mask) - - return hidden_states - - -class HunyuanImageRotaryPosEmbed(nn.Module): - def __init__(self, patch_size: tuple | list[int], rope_dim: tuple | list[int], theta: float = 256.0) -> None: - super().__init__() - - if not isinstance(patch_size, (tuple, list)) or len(patch_size) not in [2, 3]: - raise ValueError(f"patch_size must be a tuple or list of length 2 or 3, got {patch_size}") - - if not isinstance(rope_dim, (tuple, list)) or len(rope_dim) not in [2, 3]: - raise ValueError(f"rope_dim must be a tuple or list of length 2 or 3, got {rope_dim}") - - if not len(patch_size) == len(rope_dim): - raise ValueError(f"patch_size and rope_dim must have the same length, got {patch_size} and {rope_dim}") - - self.patch_size = patch_size - self.rope_dim = rope_dim - self.theta = theta - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if hidden_states.ndim == 5: - _, _, frame, height, width = hidden_states.shape - patch_size_frame, patch_size_height, patch_size_width = self.patch_size - rope_sizes = [frame // patch_size_frame, height // patch_size_height, width // patch_size_width] - elif hidden_states.ndim == 4: - _, _, height, width = hidden_states.shape - patch_size_height, patch_size_width = self.patch_size - rope_sizes = [height // patch_size_height, width // patch_size_width] - else: - raise ValueError(f"hidden_states must be a 4D or 5D tensor, got {hidden_states.shape}") - - axes_grids = [] - for i in range(len(rope_sizes)): - grid = torch.arange(0, rope_sizes[i], device=hidden_states.device, dtype=torch.float32) - axes_grids.append(grid) - grid = torch.meshgrid(*axes_grids, indexing="ij") # dim x [H, W] - grid = torch.stack(grid, dim=0) # [2, H, W] - - freqs = [] - for i in range(len(rope_sizes)): - freq = get_1d_rotary_pos_embed(self.rope_dim[i], grid[i].reshape(-1), self.theta, use_real=True) - freqs.append(freq) - - freqs_cos = torch.cat([f[0] for f in freqs], dim=1) # (W * H * T, D / 2) - freqs_sin = torch.cat([f[1] for f in freqs], dim=1) # (W * H * T, D / 2) - return freqs_cos, freqs_sin - - -@maybe_allow_in_graph -class HunyuanImageSingleTransformerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float = 4.0, - qk_norm: str = "rms_norm", - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - mlp_dim = int(hidden_size * mlp_ratio) - - self.attn = Attention( - query_dim=hidden_size, - cross_attention_dim=None, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=hidden_size, - bias=True, - processor=HunyuanImageAttnProcessor(), - qk_norm=qk_norm, - eps=1e-6, - pre_only=True, - ) - - self.norm = AdaLayerNormZeroSingle(hidden_size, norm_type="layer_norm") - self.proj_mlp = nn.Linear(hidden_size, mlp_dim) - self.act_mlp = nn.GELU(approximate="tanh") - self.proj_out = nn.Linear(hidden_size + mlp_dim, hidden_size) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - text_seq_length = encoder_hidden_states.shape[1] - hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1) - - residual = hidden_states - - # 1. Input normalization - norm_hidden_states, gate = self.norm(hidden_states, emb=temb) - mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) - - norm_hidden_states, norm_encoder_hidden_states = ( - norm_hidden_states[:, :-text_seq_length, :], - norm_hidden_states[:, -text_seq_length:, :], - ) - - # 2. Attention - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - attn_output = torch.cat([attn_output, context_attn_output], dim=1) - - # 3. Modulation and residual connection - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) - hidden_states = gate.unsqueeze(1) * self.proj_out(hidden_states) - hidden_states = hidden_states + residual - - hidden_states, encoder_hidden_states = ( - hidden_states[:, :-text_seq_length, :], - hidden_states[:, -text_seq_length:, :], - ) - return hidden_states, encoder_hidden_states - - -@maybe_allow_in_graph -class HunyuanImageTransformerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float, - qk_norm: str = "rms_norm", - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.norm1 = AdaLayerNormZero(hidden_size, norm_type="layer_norm") - self.norm1_context = AdaLayerNormZero(hidden_size, norm_type="layer_norm") - - self.attn = Attention( - query_dim=hidden_size, - cross_attention_dim=None, - added_kv_proj_dim=hidden_size, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=hidden_size, - context_pre_only=False, - bias=True, - processor=HunyuanImageAttnProcessor(), - qk_norm=qk_norm, - eps=1e-6, - ) - - self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate") - - self.norm2_context = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.ff_context = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - *args, - **kwargs, - ) -> tuple[torch.Tensor, torch.Tensor]: - # 1. Input normalization - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - - # 2. Joint attention - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - - # 3. Modulation and residual connection - hidden_states = hidden_states + attn_output * gate_msa.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + context_attn_output * c_gate_msa.unsqueeze(1) - - norm_hidden_states = self.norm2(hidden_states) - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - # 4. Feed-forward - ff_output = self.ff(norm_hidden_states) - context_ff_output = self.ff_context(norm_encoder_hidden_states) - - hidden_states = hidden_states + gate_mlp.unsqueeze(1) * ff_output - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - - return hidden_states, encoder_hidden_states - - -class HunyuanImageTransformer2DModel( - ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin -): - r""" - The Transformer model used in [HunyuanImage-2.1](https://github.com/Tencent-Hunyuan/HunyuanImage-2.1). - - Args: - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - num_attention_heads (`int`, defaults to `24`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - num_layers (`int`, defaults to `20`): - The number of layers of dual-stream blocks to use. - num_single_layers (`int`, defaults to `40`): - The number of layers of single-stream blocks to use. - num_refiner_layers (`int`, defaults to `2`): - The number of layers of refiner blocks to use. - mlp_ratio (`float`, defaults to `4.0`): - The ratio of the hidden layer size to the input size in the feedforward network. - patch_size (`int`, defaults to `2`): - The size of the spatial patches to use in the patch embedding layer. - patch_size_t (`int`, defaults to `1`): - The size of the tmeporal patches to use in the patch embedding layer. - qk_norm (`str`, defaults to `rms_norm`): - The normalization to use for the query and key projections in the attention layers. - guidance_embeds (`bool`, defaults to `True`): - Whether to use guidance embeddings in the model. - text_embed_dim (`int`, defaults to `4096`): - Input dimension of text embeddings from the text encoder. - pooled_projection_dim (`int`, defaults to `768`): - The dimension of the pooled projection of the text embeddings. - rope_theta (`float`, defaults to `256.0`): - The value of theta to use in the RoPE layer. - rope_axes_dim (`tuple[int]`, defaults to `(16, 56, 56)`): - The dimensions of the axes to use in the RoPE layer. - image_condition_type (`str`, *optional*, defaults to `None`): - The type of image conditioning to use. If `None`, no image conditioning is used. If `latent_concat`, the - image is concatenated to the latent stream. If `token_replace`, the image is used to replace first-frame - tokens in the latent stream and apply conditioning. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["x_embedder", "context_embedder", "norm"] - _no_split_modules = [ - "HunyuanImageTransformerBlock", - "HunyuanImageSingleTransformerBlock", - "HunyuanImagePatchEmbed", - "HunyuanImageTokenRefiner", - ] - _repeated_blocks = ["HunyuanImageTransformerBlock", "HunyuanImageSingleTransformerBlock"] - - @register_to_config - def __init__( - self, - in_channels: int = 64, - out_channels: int = 64, - num_attention_heads: int = 28, - attention_head_dim: int = 128, - num_layers: int = 20, - num_single_layers: int = 40, - num_refiner_layers: int = 2, - mlp_ratio: float = 4.0, - patch_size: tuple[int, int] = (1, 1), - qk_norm: str = "rms_norm", - guidance_embeds: bool = False, - text_embed_dim: int = 3584, - text_embed_2_dim: int | None = None, - rope_theta: float = 256.0, - rope_axes_dim: tuple[int, ...] = (64, 64), - use_meanflow: bool = False, - ) -> None: - super().__init__() - - if not (isinstance(patch_size, (tuple, list)) and len(patch_size) in [2, 3]): - raise ValueError(f"patch_size must be a tuple of length 2 or 3, got {patch_size}") - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - # 1. Latent and condition embedders - self.x_embedder = HunyuanImagePatchEmbed(patch_size, in_channels, inner_dim) - self.context_embedder = HunyuanImageTokenRefiner( - text_embed_dim, num_attention_heads, attention_head_dim, num_layers=num_refiner_layers - ) - - if text_embed_2_dim is not None: - self.context_embedder_2 = HunyuanImageByT5TextProjection(text_embed_2_dim, 2048, inner_dim) - else: - self.context_embedder_2 = None - - self.time_guidance_embed = HunyuanImageCombinedTimeGuidanceEmbedding(inner_dim, guidance_embeds, use_meanflow) - - # 2. RoPE - self.rope = HunyuanImageRotaryPosEmbed(patch_size, rope_axes_dim, rope_theta) - - # 3. Dual stream transformer blocks - - self.transformer_blocks = nn.ModuleList( - [ - HunyuanImageTransformerBlock( - num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm - ) - for _ in range(num_layers) - ] - ) - - # 4. Single stream transformer blocks - self.single_transformer_blocks = nn.ModuleList( - [ - HunyuanImageSingleTransformerBlock( - num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm - ) - for _ in range(num_single_layers) - ] - ) - - # 5. Output projection - self.norm_out = AdaLayerNormContinuous(inner_dim, inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(inner_dim, math.prod(patch_size) * out_channels) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_attention_mask: torch.Tensor, - timestep_r: torch.LongTensor | None = None, - encoder_hidden_states_2: torch.Tensor | None = None, - encoder_attention_mask_2: torch.Tensor | None = None, - guidance: torch.Tensor | None = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> torch.Tensor | dict[str, torch.Tensor]: - """ - The [`HunyuanImageTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, height, width)` or `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_attention_mask (`torch.Tensor`): - Mask applied to `encoder_hidden_states` during attention. - timestep_r (`torch.LongTensor`, *optional*): - Refiner timestep conditioning. - encoder_hidden_states_2 (`torch.Tensor`, *optional*): - Additional conditional embeddings computed from a second text encoder. - encoder_attention_mask_2 (`torch.Tensor`, *optional*): - Mask applied to `encoder_hidden_states_2` during attention. - guidance (`torch.Tensor`, *optional*): - Guidance scale embedding used for guidance-distilled variants of the model. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - if hidden_states.ndim == 4: - batch_size, channels, height, width = hidden_states.shape - sizes = (height, width) - elif hidden_states.ndim == 5: - batch_size, channels, frame, height, width = hidden_states.shape - sizes = (frame, height, width) - else: - raise ValueError(f"hidden_states must be a 4D or 5D tensor, got {hidden_states.shape}") - - post_patch_sizes = tuple(d // p for d, p in zip(sizes, self.config.patch_size)) - - # 1. RoPE - image_rotary_emb = self.rope(hidden_states) - - # 2. Conditional embeddings - encoder_attention_mask = encoder_attention_mask.bool() - temb = self.time_guidance_embed(timestep, guidance=guidance, timestep_r=timestep_r) - hidden_states = self.x_embedder(hidden_states) - encoder_hidden_states = self.context_embedder(encoder_hidden_states, timestep, encoder_attention_mask) - - if self.context_embedder_2 is not None and encoder_hidden_states_2 is not None: - encoder_hidden_states_2 = self.context_embedder_2(encoder_hidden_states_2) - - encoder_attention_mask_2 = encoder_attention_mask_2.bool() - - # reorder and combine text tokens: combine valid tokens first, then padding - new_encoder_hidden_states = [] - new_encoder_attention_mask = [] - - for text, text_mask, text_2, text_mask_2 in zip( - encoder_hidden_states, encoder_attention_mask, encoder_hidden_states_2, encoder_attention_mask_2 - ): - # Concatenate: [valid_mllm, valid_byt5, invalid_mllm, invalid_byt5] - new_encoder_hidden_states.append( - torch.cat( - [ - text_2[text_mask_2], # valid byt5 - text[text_mask], # valid mllm - text_2[~text_mask_2], # invalid byt5 - text[~text_mask], # invalid mllm - ], - dim=0, - ) - ) - - # Apply same reordering to attention masks - new_encoder_attention_mask.append( - torch.cat( - [ - text_mask_2[text_mask_2], - text_mask[text_mask], - text_mask_2[~text_mask_2], - text_mask[~text_mask], - ], - dim=0, - ) - ) - - encoder_hidden_states = torch.stack(new_encoder_hidden_states) - encoder_attention_mask = torch.stack(new_encoder_attention_mask) - - attention_mask = torch.nn.functional.pad(encoder_attention_mask, (hidden_states.shape[1], 0), value=True) - attention_mask = attention_mask.unsqueeze(1).unsqueeze(2) - # 3. Transformer blocks - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.transformer_blocks: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - - for block in self.single_transformer_blocks: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - - else: - for block in self.transformer_blocks: - hidden_states, encoder_hidden_states = block( - hidden_states, - encoder_hidden_states, - temb, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - - for block in self.single_transformer_blocks: - hidden_states, encoder_hidden_states = block( - hidden_states, - encoder_hidden_states, - temb, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - - # 4. Output projection - hidden_states = self.norm_out(hidden_states, temb) - hidden_states = self.proj_out(hidden_states) - - # 5. unpatchify - # reshape: [batch_size, *post_patch_dims, channels, *patch_size] - out_channels = self.config.out_channels - reshape_dims = [batch_size] + list(post_patch_sizes) + [out_channels] + list(self.config.patch_size) - hidden_states = hidden_states.reshape(*reshape_dims) - - # create permutation pattern: batch, channels, then interleave post_patch and patch dims - # For 4D: [0, 3, 1, 4, 2, 5] -> batch, channels, post_patch_height, patch_size_height, post_patch_width, patch_size_width - # For 5D: [0, 4, 1, 5, 2, 6, 3, 7] -> batch, channels, post_patch_frame, patch_size_frame, post_patch_height, patch_size_height, post_patch_width, patch_size_width - ndim = len(post_patch_sizes) - permute_pattern = [0, ndim + 1] # batch, channels - for i in range(ndim): - permute_pattern.extend([i + 1, ndim + 2 + i]) # post_patch_sizes[i], patch_sizes[i] - hidden_states = hidden_states.permute(*permute_pattern) - - # flatten patch dimensions: flatten each (post_patch_size, patch_size) pair - # batch_size, channels, post_patch_sizes[0] * patch_sizes[0], post_patch_sizes[1] * patch_sizes[1], ... - final_dims = [batch_size, out_channels] + [ - post_patch * patch for post_patch, patch in zip(post_patch_sizes, self.config.patch_size) - ] - hidden_states = hidden_states.reshape(*final_dims) - - if not return_dict: - return (hidden_states,) - - return Transformer2DModelOutput(sample=hidden_states) diff --git a/diffusers/models/transformers/transformer_ideogram4.py b/diffusers/models/transformers/transformer_ideogram4.py deleted file mode 100644 index 3607c917a7272e95dc0fdf0215ce525931b01738..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_ideogram4.py +++ /dev/null @@ -1,457 +0,0 @@ -# Copyright 2026 Ideogram AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -import math - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import AttentionMixin, AttentionModuleMixin -from ..attention_dispatch import dispatch_attention_fn -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# Per-token role indicators used to label entries of the packed text+image sequence. -SEQUENCE_PADDING_INDICATOR = -1 -OUTPUT_IMAGE_INDICATOR = 2 -LLM_TOKEN_INDICATOR = 3 - -# Image grid coordinates start at this offset so they never collide with text token indices. -IMAGE_POSITION_OFFSET = 65536 - - -def _rotate_half(x: torch.Tensor) -> torch.Tensor: - half = x.shape[-1] // 2 - return torch.cat((-x[..., half:], x[..., :half]), dim=-1) - - -class Ideogram4MRoPE(nn.Module): - """Multi-axis (t, h, w) interleaved rotary position embedding.""" - - inv_freq: torch.Tensor - - def __init__( - self, - head_dim: int, - base: int, - mrope_section: tuple[int, ...], - ) -> None: - super().__init__() - inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)) - self.register_buffer("inv_freq", inv_freq, persistent=False) - self.mrope_section = tuple(mrope_section) - self.head_dim = head_dim - - def forward(self, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - # position_ids: (B, L, 3) of int (axes are t, h, w). - if position_ids.ndim != 3 or position_ids.shape[-1] != 3: - raise ValueError(f"`position_ids` must have shape (B, L, 3), got {tuple(position_ids.shape)}.") - batch_size, seq_len, _ = position_ids.shape - - # Ideogram4's image position ids start at IMAGE_POSITION_OFFSET (65536). If an ambient autocast downcasts the - # matmul to bfloat16, the image positions will collapse to only a few distinct values because bfloat16 cannot - # represent consecutive integers at this value (after pos 65536 each 512-integer block will collapse to the - # same value), which causes the image to become essentially flat. Therefore, we need to disable autocast here. - pos = position_ids.permute(2, 0, 1).to(dtype=torch.float32) - inv_freq = self.inv_freq.to(dtype=torch.float32)[None, None, :, None].expand(3, batch_size, -1, 1) - with torch.autocast(device_type=position_ids.device.type, enabled=False): - freqs = inv_freq @ pos.unsqueeze(2) - freqs = freqs.transpose(2, 3) # (3, B, L, inv_freq_size) - - # Interleaved mrope: pull H freqs into idx 1 mod 3, W freqs into idx 2 mod 3. - freqs_t = freqs[0].clone() - for axis, offset in ((1, 1), (2, 2)): - length = self.mrope_section[axis] * 3 - idx = torch.arange(offset, length, 3, device=freqs_t.device) - freqs_t[..., idx] = freqs[axis][..., idx] - - emb = torch.cat((freqs_t, freqs_t), dim=-1) - return emb.cos().float(), emb.sin().float() - - -class Ideogram4AttnProcessor: - _attention_backend = None - _parallel_config = None - - def __call__( - self, - attn: "Ideogram4Attention", - hidden_states: torch.Tensor, - attention_mask: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor], - ) -> torch.Tensor: - query = attn.to_q(hidden_states).unflatten(-1, (attn.num_heads, attn.head_dim)) - key = attn.to_k(hidden_states).unflatten(-1, (attn.num_heads, attn.head_dim)) - value = attn.to_v(hidden_states).unflatten(-1, (attn.num_heads, attn.head_dim)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - # MRoPE applied in (B, L, num_heads, head_dim) layout; cos/sin broadcast over the head axis. - cos, sin = image_rotary_emb - cos = cos.unsqueeze(2) - sin = sin.unsqueeze(2) - query = (query * cos) + (_rotate_half(query) * sin) - key = (key * cos) + (_rotate_half(key) * sin) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - return attn.to_out[0](hidden_states) - - -class Ideogram4Attention(nn.Module, AttentionModuleMixin): - """Self-attention with split Q/K/V, q/k RMSNorm, MRoPE and a block-diagonal segment mask.""" - - _default_processor_cls = Ideogram4AttnProcessor - _available_processors = [Ideogram4AttnProcessor] - - def __init__(self, hidden_size: int, num_heads: int, eps: float = 1e-5) -> None: - super().__init__() - if hidden_size % num_heads != 0: - raise ValueError(f"hidden_size={hidden_size} must be divisible by num_heads={num_heads}") - self.hidden_size = hidden_size - self.num_heads = num_heads - self.head_dim = hidden_size // num_heads - self.use_bias = False - - self.to_q = nn.Linear(hidden_size, hidden_size, bias=False) - self.to_k = nn.Linear(hidden_size, hidden_size, bias=False) - self.to_v = nn.Linear(hidden_size, hidden_size, bias=False) - self.norm_q = RMSNorm(self.head_dim, eps=eps, elementwise_affine=True) - self.norm_k = RMSNorm(self.head_dim, eps=eps, elementwise_affine=True) - self.to_out = nn.ModuleList([nn.Linear(hidden_size, hidden_size, bias=False), nn.Dropout(0.0)]) - - self.set_processor(self._default_processor_cls()) - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - unused_kwargs = [k for k in kwargs if k not in attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -class Ideogram4MLP(nn.Module): - """SwiGLU feed-forward network.""" - - def __init__(self, dim: int, hidden_dim: int) -> None: - super().__init__() - self.w1 = nn.Linear(dim, hidden_dim, bias=False) - self.w2 = nn.Linear(hidden_dim, dim, bias=False) - self.w3 = nn.Linear(dim, hidden_dim, bias=False) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.w2(F.silu(self.w1(x)) * self.w3(x)) - - -@maybe_allow_in_graph -class Ideogram4TransformerBlock(nn.Module): - def __init__( - self, - hidden_size: int, - intermediate_size: int, - num_heads: int, - norm_eps: float, - adaln_dim: int, - ) -> None: - super().__init__() - self.attention = Ideogram4Attention(hidden_size, num_heads, eps=1e-5) - self.feed_forward = Ideogram4MLP(hidden_size, intermediate_size) - - self.attention_norm1 = RMSNorm(hidden_size, eps=norm_eps, elementwise_affine=True) - self.ffn_norm1 = RMSNorm(hidden_size, eps=norm_eps, elementwise_affine=True) - self.attention_norm2 = RMSNorm(hidden_size, eps=norm_eps, elementwise_affine=True) - self.ffn_norm2 = RMSNorm(hidden_size, eps=norm_eps, elementwise_affine=True) - - self.adaln_modulation = nn.Linear(adaln_dim, 4 * hidden_size, bias=True) - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor], - adaln_input: torch.Tensor, - ) -> torch.Tensor: - mod = self.adaln_modulation(adaln_input) - scale_msa, gate_msa, scale_mlp, gate_mlp = mod.chunk(4, dim=-1) - gate_msa = torch.tanh(gate_msa) - gate_mlp = torch.tanh(gate_mlp) - scale_msa = 1.0 + scale_msa - scale_mlp = 1.0 + scale_mlp - - attn_out = self.attention( - self.attention_norm1(hidden_states) * scale_msa, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - hidden_states = hidden_states + gate_msa * self.attention_norm2(attn_out) - hidden_states = hidden_states + gate_mlp * self.ffn_norm2( - self.feed_forward(self.ffn_norm1(hidden_states) * scale_mlp) - ) - return hidden_states - - -def _sinusoidal_embedding(t: torch.Tensor, dim: int, scale: float = 1e4) -> torch.Tensor: - t = t.to(torch.float32) - half = dim // 2 - freq = math.log(scale) / (half - 1) - freq = torch.exp(torch.arange(half, dtype=torch.float32, device=t.device) * -freq) - emb = t.unsqueeze(-1) * freq - emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) - if dim % 2 == 1: - emb = F.pad(emb, (0, 1)) - return emb - - -class Ideogram4EmbedScalar(nn.Module): - """Sinusoidal scalar embedding followed by a small MLP.""" - - def __init__(self, dim: int, input_range: tuple[float, float]) -> None: - super().__init__() - self.dim = dim - self.range_min, self.range_max = input_range - if self.range_max <= self.range_min: - raise ValueError("input_range[1] must be greater than input_range[0]") - self.mlp_in = nn.Linear(dim, dim, bias=True) - self.mlp_out = nn.Linear(dim, dim, bias=True) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - in_dtype = x.dtype - x = x.to(torch.float32) - scaled = 1e4 * (x - self.range_min) / (self.range_max - self.range_min) - emb = _sinusoidal_embedding(scaled, self.dim) - emb = emb.to(in_dtype) - emb = F.silu(self.mlp_in(emb)) - return self.mlp_out(emb) - - -class Ideogram4FinalLayer(nn.Module): - def __init__(self, hidden_size: int, out_channels: int, adaln_dim: int) -> None: - super().__init__() - self.norm_final = nn.LayerNorm(hidden_size, eps=1e-6, elementwise_affine=False) - self.linear = nn.Linear(hidden_size, out_channels, bias=True) - self.adaln_modulation = nn.Linear(adaln_dim, hidden_size, bias=True) - - def forward(self, hidden_states: torch.Tensor, conditioning: torch.Tensor) -> torch.Tensor: - scale = 1.0 + self.adaln_modulation(F.silu(conditioning)) - return self.linear(self.norm_final(hidden_states) * scale) - - -class Ideogram4Transformer2DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixin, FromOriginalModelMixin): - r""" - The flow-matching transformer backbone used by the Ideogram 4 pipeline. - - The transformer operates on a single packed sequence containing both text-conditioning tokens (produced by a - multimodal text encoder) and the patchified image latents. Per-token indicators distinguish the two roles, and a - block-diagonal attention mask derived from `segment_ids` restricts each sample to attend only to itself within a - packed batch. - - Args: - in_channels (`int`, defaults to 128): - Latent channel count after patchification (`ae_channels * patch_size ** 2`). - num_layers (`int`, defaults to 34): - Number of transformer blocks. - attention_head_dim (`int`, defaults to 256): - Dimension of each attention head; the total hidden size is `attention_head_dim * num_attention_heads`. - num_attention_heads (`int`, defaults to 18): - Number of attention heads. - intermediate_size (`int`, defaults to 12288): - Feed-forward hidden size used by the SwiGLU MLP inside each block. - adaln_dim (`int`, defaults to 512): - Dimensionality of the conditioning vector consumed by the AdaLN modulations. - llm_features_dim (`int`, defaults to 53248): - Dimensionality of the per-token text features fed into the model (typically a concatenation of hidden - states from several layers of the text encoder). - rope_theta (`int`, defaults to 5_000_000): - Base used by the multi-axis rotary position embedding. - mrope_section (`tuple[int, int, int]`, defaults to `(24, 20, 20)`): - Number of frequencies allocated to each of the (t, h, w) axes of MRoPE. - norm_eps (`float`, defaults to 1e-5): - Epsilon used by the RMSNorm modules inside the transformer blocks. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["Ideogram4TransformerBlock"] - _repeated_blocks = ["Ideogram4TransformerBlock"] - _skip_layerwise_casting_patterns = ["t_embedding", "adaln_proj", "embed_image_indicator"] - - @register_to_config - def __init__( - self, - in_channels: int = 128, - num_layers: int = 34, - attention_head_dim: int = 256, - num_attention_heads: int = 18, - intermediate_size: int = 12288, - adaln_dim: int = 512, - llm_features_dim: int = 53248, - rope_theta: int = 5_000_000, - mrope_section: tuple[int, int, int] = (24, 20, 20), - norm_eps: float = 1e-5, - ) -> None: - super().__init__() - - hidden_size = attention_head_dim * num_attention_heads - head_dim = attention_head_dim - - self.in_channels = in_channels - self.out_channels = in_channels - self.hidden_size = hidden_size - self.gradient_checkpointing = False - - self.input_proj = nn.Linear(in_channels, hidden_size, bias=True) - self.llm_cond_norm = RMSNorm(llm_features_dim, eps=1e-6, elementwise_affine=True) - self.llm_cond_proj = nn.Linear(llm_features_dim, hidden_size, bias=True) - self.t_embedding = Ideogram4EmbedScalar(hidden_size, input_range=(0.0, 1.0)) - self.adaln_proj = nn.Linear(hidden_size, adaln_dim, bias=True) - - self.embed_image_indicator = nn.Embedding(2, hidden_size) - - self.rotary_emb = Ideogram4MRoPE( - head_dim=head_dim, - base=rope_theta, - mrope_section=mrope_section, - ) - - self.layers = nn.ModuleList( - [ - Ideogram4TransformerBlock( - hidden_size=hidden_size, - intermediate_size=intermediate_size, - num_heads=num_attention_heads, - norm_eps=norm_eps, - adaln_dim=adaln_dim, - ) - for _ in range(num_layers) - ] - ) - - self.final_layer = Ideogram4FinalLayer( - hidden_size=hidden_size, - out_channels=in_channels, - adaln_dim=adaln_dim, - ) - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - position_ids: torch.Tensor, - segment_ids: torch.Tensor, - indicator: torch.Tensor, - attention_kwargs: dict | None = None, - return_dict: bool = True, - ) -> Transformer2DModelOutput | tuple[torch.Tensor]: - r""" - Predict the flow-matching velocity for the image-token positions of the packed sequence. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, sequence_length, in_channels)`): - Packed sequence of patchified noisy image tokens. Non-image positions are masked out internally. - timestep (`torch.Tensor` of shape `(batch_size,)` or `(batch_size, sequence_length)`): - Flow-matching time in `[0, 1]` (0 is pure noise, 1 is clean data). - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_length, llm_features_dim)`): - Per-token text conditioning features. Non-text positions are masked out internally. - position_ids (`torch.Tensor` of shape `(batch_size, sequence_length, 3)`): - `(t, h, w)` coordinates consumed by the multi-axis RoPE. - segment_ids (`torch.Tensor` of shape `(batch_size, sequence_length)`): - Per-token sample id within a packed batch. Positions sharing a `segment_id` attend to each other. - indicator (`torch.Tensor` of shape `(batch_size, sequence_length)`): - Per-token role: `LLM_TOKEN_INDICATOR` (text) or `OUTPUT_IMAGE_INDICATOR` (image). - attention_kwargs (`dict`, *optional*): - A kwargs dictionary passed along to the attention processor. A `"scale"` entry scales the LoRA weights - (when the PEFT backend is active). - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.modeling_outputs.Transformer2DModelOutput`] instead of a plain tuple. - - Returns: - [`~models.modeling_outputs.Transformer2DModelOutput`] or a `tuple` whose first element is a tensor of shape - `(batch_size, sequence_length, in_channels)` in the model's compute dtype. Only positions tagged with - `OUTPUT_IMAGE_INDICATOR` carry meaningful velocity predictions. - """ - batch_size, seq_len, in_channels = hidden_states.shape - if in_channels != self.in_channels: - raise ValueError(f"Expected last dim {self.in_channels}, got {in_channels}.") - - llm_token_mask = (indicator == LLM_TOKEN_INDICATOR).to(hidden_states.dtype).unsqueeze(-1) - output_image_mask = (indicator == OUTPUT_IMAGE_INDICATOR).to(hidden_states.dtype).unsqueeze(-1) - - encoder_hidden_states = encoder_hidden_states * llm_token_mask - hidden_states = hidden_states * output_image_mask - hidden_states = self.input_proj(hidden_states) * output_image_mask - - # Keep shape (B, 1, ...) when t is per-sample so downstream adaln projections do not pay for L identical copies. - t_cond = self.t_embedding(timestep) - if timestep.dim() == 1: - t_cond = t_cond.unsqueeze(1) - adaln_input = F.silu(self.adaln_proj(t_cond)) - - encoder_hidden_states = self.llm_cond_norm(encoder_hidden_states) - encoder_hidden_states = self.llm_cond_proj(encoder_hidden_states) * llm_token_mask - - hidden_states = hidden_states + encoder_hidden_states - - image_indicator_embedding = self.embed_image_indicator((indicator == OUTPUT_IMAGE_INDICATOR).to(torch.long)) - hidden_states = hidden_states + image_indicator_embedding - - cos, sin = self.rotary_emb(position_ids) - cos = cos.to(hidden_states.dtype) - sin = sin.to(hidden_states.dtype) - image_rotary_emb = (cos, sin) - - # Block-diagonal mask from segment ids: tokens only attend within their segment. Shared by every block. - attention_mask = (segment_ids.unsqueeze(2) == segment_ids.unsqueeze(1)).unsqueeze(1) - - for block in self.layers: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, hidden_states, attention_mask, image_rotary_emb, adaln_input - ) - else: - hidden_states = block(hidden_states, attention_mask, image_rotary_emb, adaln_input) - - output = self.final_layer(hidden_states, conditioning=adaln_input) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_joyimage.py b/diffusers/models/transformers/transformer_joyimage.py deleted file mode 100644 index b17ddb05f799e03da68ccaacc0f5f8d203275c84..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_joyimage.py +++ /dev/null @@ -1,603 +0,0 @@ -# Copyright 2025 The JoyImage Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -import math -from typing import Tuple - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import FP32LayerNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# --------------------------------------------------------------------------- -# Rotary position embedding utilities -# --------------------------------------------------------------------------- - - -def _apply_rotary_emb( - xq: torch.Tensor, - xk: torch.Tensor, - freqs_cis: Tuple[torch.Tensor, torch.Tensor], -) -> Tuple[torch.Tensor, torch.Tensor]: - ndim = xq.ndim - shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(xq.shape)] - cos = freqs_cis[0].view(*shape).to(xq.device) - sin = freqs_cis[1].view(*shape).to(xq.device) - - def _rotate_half(x): - x_real, x_imag = x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1) - return torch.stack([-x_imag, x_real], dim=-1).flatten(3) - - xq_out = (xq.float() * cos + _rotate_half(xq) * sin).type_as(xq) - xk_out = (xk.float() * cos + _rotate_half(xk) * sin).type_as(xk) - return xq_out, xk_out - - -# --------------------------------------------------------------------------- -# Modulation -# --------------------------------------------------------------------------- - - -class JoyImageModulate(nn.Module): - """Wan-style learnable modulation table. - - Produces `factor` modulation vectors by adding the conditioning signal to a learnable parameter table. - """ - - def __init__(self, hidden_size: int, factor: int, dtype=None, device=None): - super().__init__() - self.factor = factor - self.modulate_table = nn.Parameter( - torch.zeros(1, factor, hidden_size, dtype=dtype, device=device) / hidden_size**0.5, - requires_grad=True, - ) - - def forward(self, x: torch.Tensor) -> list[torch.Tensor]: - if x.ndim != 3: - x = x.unsqueeze(1) - return [o.squeeze(1) for o in (self.modulate_table + x).chunk(self.factor, dim=1)] - - -# --------------------------------------------------------------------------- -# Attention processor -# --------------------------------------------------------------------------- - - -class JoyImageAttnProcessor: - """Attention processor for JoyImage double-stream joint attention. - - Implements the joint attention computation where text and image streams are processed together. The - :class:`JoyImageAttention` module stores fused QKV projections (``img_attn_qkv`` / ``txt_attn_qkv``). - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - pass - - def __call__( - self, - attn: "JoyImageAttention", - hidden_states: torch.Tensor, # image stream (B, S_img, D) - encoder_hidden_states: torch.Tensor = None, # text stream (B, S_txt, D) - image_rotary_emb: Tuple[torch.Tensor, torch.Tensor] | None = None, - **kwargs, - ) -> Tuple[torch.Tensor, torch.Tensor]: - if encoder_hidden_states is None: - raise ValueError("JoyImageAttnProcessor requires encoder_hidden_states (text stream)") - - heads = attn.heads - - # image stream: fused QKV -> split - img_qkv = attn.img_attn_qkv(hidden_states) - img_query, img_key, img_value = img_qkv.chunk(3, dim=-1) - - # text stream: fused QKV -> split - txt_qkv = attn.txt_attn_qkv(encoder_hidden_states) - txt_query, txt_key, txt_value = txt_qkv.chunk(3, dim=-1) - - # reshape to multi-head: (B, S, H, D) - img_query = img_query.unflatten(-1, (heads, -1)) - img_key = img_key.unflatten(-1, (heads, -1)) - img_value = img_value.unflatten(-1, (heads, -1)) - - txt_query = txt_query.unflatten(-1, (heads, -1)) - txt_key = txt_key.unflatten(-1, (heads, -1)) - txt_value = txt_value.unflatten(-1, (heads, -1)) - - # QK norm - img_query = attn.img_attn_q_norm(img_query) - img_key = attn.img_attn_k_norm(img_key) - txt_query = attn.txt_attn_q_norm(txt_query) - txt_key = attn.txt_attn_k_norm(txt_key) - - # RoPE (custom implementation) - if image_rotary_emb is not None: - vis_freqs, txt_freqs = image_rotary_emb - if vis_freqs is not None: - img_query, img_key = _apply_rotary_emb(img_query, img_key, vis_freqs) - if txt_freqs is not None: - txt_query, txt_key = _apply_rotary_emb(txt_query, txt_key, txt_freqs) - - # concatenate for joint attention: [img, txt] - joint_query = torch.cat([img_query, txt_query], dim=1) - joint_key = torch.cat([img_key, txt_key], dim=1) - joint_value = torch.cat([img_value, txt_value], dim=1) - - joint_hidden_states = dispatch_attention_fn( - joint_query, - joint_key, - joint_value, - attn_mask=None, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - joint_hidden_states = joint_hidden_states.flatten(2, 3) - joint_hidden_states = joint_hidden_states.to(joint_query.dtype) - - # split back - img_attn_output = joint_hidden_states[:, : hidden_states.shape[1], :] - txt_attn_output = joint_hidden_states[:, hidden_states.shape[1] :, :] - - # output projections - img_attn_output = attn.img_attn_proj(img_attn_output) - txt_attn_output = attn.txt_attn_proj(txt_attn_output) - - return img_attn_output, txt_attn_output - - -# --------------------------------------------------------------------------- -# Attention module -# --------------------------------------------------------------------------- - - -class JoyImageAttention(nn.Module, AttentionModuleMixin): - """Joint attention module for JoyImage double-stream blocks. - - Wraps the fused QKV projections, QK norms, and output projections for both image and text streams. Delegates the - actual attention computation to a pluggable :class:`JoyImageAttnProcessor`. - """ - - _default_processor_cls = JoyImageAttnProcessor - _available_processors = [JoyImageAttnProcessor] - _supports_qkv_fusion = False - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - eps: float = 1e-6, - processor=None, - ): - super().__init__() - - self.heads = num_attention_heads - self.head_dim = attention_head_dim - inner_dim = num_attention_heads * attention_head_dim - - self.img_attn_qkv = nn.Linear(dim, inner_dim * 3, bias=True) - self.img_attn_q_norm = nn.RMSNorm(attention_head_dim, eps=eps) - self.img_attn_k_norm = nn.RMSNorm(attention_head_dim, eps=eps) - self.img_attn_proj = nn.Linear(inner_dim, dim, bias=True) - - self.txt_attn_qkv = nn.Linear(dim, inner_dim * 3, bias=True) - self.txt_attn_q_norm = nn.RMSNorm(attention_head_dim, eps=eps) - self.txt_attn_k_norm = nn.RMSNorm(attention_head_dim, eps=eps) - self.txt_attn_proj = nn.Linear(inner_dim, dim, bias=True) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - image_rotary_emb: Tuple[torch.Tensor, torch.Tensor] | None = None, - **kwargs, - ) -> Tuple[torch.Tensor, torch.Tensor]: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"joint_attention_kwargs {unused_kwargs} are not expected by " - f"{self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, encoder_hidden_states, image_rotary_emb, **kwargs) - - -# --------------------------------------------------------------------------- -# Transformer block -# --------------------------------------------------------------------------- - - -class JoyImageTransformerBlock(nn.Module): - """Double-stream transformer block for JoyImage. - - Each block processes an image stream and a text stream jointly through shared attention, following the SD3 / Flux - double-stream pattern with WAN-style modulation. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - mlp_width_ratio: float = 4.0, - eps: float = 1e-6, - ): - super().__init__() - - self.dim = dim - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - mlp_hidden_dim = int(dim * mlp_width_ratio) - - # image stream - self.img_mod = JoyImageModulate(dim, factor=6) - self.img_norm1 = FP32LayerNorm(dim, elementwise_affine=False, eps=eps) - self.img_norm2 = FP32LayerNorm(dim, elementwise_affine=False, eps=eps) - self.img_mlp = FeedForward(dim, inner_dim=mlp_hidden_dim, activation_fn="gelu-approximate") - - # text stream - self.txt_mod = JoyImageModulate(dim, factor=6) - self.txt_norm1 = FP32LayerNorm(dim, elementwise_affine=False, eps=eps) - self.txt_norm2 = FP32LayerNorm(dim, elementwise_affine=False, eps=eps) - self.txt_mlp = FeedForward(dim, inner_dim=mlp_hidden_dim, activation_fn="gelu-approximate") - - # ---- joint attention ---- - self.attn = JoyImageAttention(dim, num_attention_heads, attention_head_dim, eps=eps) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: Tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> Tuple[torch.Tensor, torch.Tensor]: - # modulation - ( - img_mod1_shift, - img_mod1_scale, - img_mod1_gate, - img_mod2_shift, - img_mod2_scale, - img_mod2_gate, - ) = self.img_mod(temb) - ( - txt_mod1_shift, - txt_mod1_scale, - txt_mod1_gate, - txt_mod2_shift, - txt_mod2_scale, - txt_mod2_gate, - ) = self.txt_mod(temb) - - # --- attention --- - img_normed = self.img_norm1(hidden_states) - txt_normed = self.txt_norm1(encoder_hidden_states) - img_modulated = img_normed * (1 + img_mod1_scale.unsqueeze(1)) + img_mod1_shift.unsqueeze(1) - txt_modulated = txt_normed * (1 + txt_mod1_scale.unsqueeze(1)) + txt_mod1_shift.unsqueeze(1) - - img_attn, txt_attn = self.attn( - hidden_states=img_modulated, - encoder_hidden_states=txt_modulated, - image_rotary_emb=image_rotary_emb, - ) - - hidden_states = hidden_states + img_attn * img_mod1_gate.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + txt_attn * txt_mod1_gate.unsqueeze(1) - - # --- FFN --- - img_ffn_normed = self.img_norm2(hidden_states) - txt_ffn_normed = self.txt_norm2(encoder_hidden_states) - img_ffn_input = img_ffn_normed * (1 + img_mod2_scale.unsqueeze(1)) + img_mod2_shift.unsqueeze(1) - txt_ffn_input = txt_ffn_normed * (1 + txt_mod2_scale.unsqueeze(1)) + txt_mod2_shift.unsqueeze(1) - img_ffn_output = self.img_mlp(img_ffn_input) - txt_ffn_output = self.txt_mlp(txt_ffn_input) - hidden_states = hidden_states + img_ffn_output * img_mod2_gate.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + txt_ffn_output * txt_mod2_gate.unsqueeze(1) - - return hidden_states, encoder_hidden_states - - -class JoyImageTimeTextImageEmbedding(nn.Module): - def __init__( - self, - dim: int, - time_freq_dim: int, - time_proj_dim: int, - text_embed_dim: int, - ): - super().__init__() - - self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim) - self.act_fn = nn.SiLU() - self.time_proj = nn.Linear(dim, time_proj_dim) - self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh") - - def forward( - self, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - ): - timestep = self.timesteps_proj(timestep) - - time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype - if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8: - timestep = timestep.to(time_embedder_dtype) - temb = self.time_embedder(timestep).type_as(encoder_hidden_states) - timestep_proj = self.time_proj(self.act_fn(temb)) - - encoder_hidden_states = self.text_embedder(encoder_hidden_states) - - return temb, timestep_proj, encoder_hidden_states - - -# --------------------------------------------------------------------------- -# Main model -# --------------------------------------------------------------------------- - - -class JoyImageEditTransformer3DModel(ModelMixin, ConfigMixin, AttentionMixin): - """JoyImage Transformer model for image generation / editing. - - Dual-stream DiT architecture with WAN-style conditioning embeddings and custom rotary position embeddings. - """ - - _skip_layerwise_casting_patterns = ["img_in", "condition_embedder", "norm"] - _no_split_modules = ["JoyImageTransformerBlock"] - _supports_gradient_checkpointing = True - _keep_in_fp32_modules = [ - "time_embedder", - "norm1", - "norm2", - "norm_out", - ] - _repeated_blocks = ["JoyImageTransformerBlock"] - - @register_to_config - def __init__( - self, - patch_size: list = [1, 2, 2], - in_channels: int = 16, - out_channels: int | None = None, - hidden_size: int = 3072, - num_attention_heads: int = 24, - text_dim: int = 4096, - mlp_width_ratio: float = 4.0, - num_layers: int = 20, - rope_dim_list: list[int] = [16, 56, 56], - rope_type: str = "rope", - theta: int = 256, - ): - super().__init__() - - self.out_channels = out_channels or in_channels - self.patch_size = patch_size - self.hidden_size = hidden_size - self.num_attention_heads = num_attention_heads - self.rope_dim_list = rope_dim_list - self.rope_type = rope_type - self.theta = theta - - attention_head_dim = hidden_size // num_attention_heads - if hidden_size % num_attention_heads != 0: - raise ValueError( - f"hidden_size ({hidden_size}) must be divisible by num_attention_heads ({num_attention_heads})" - ) - - # image projection - self.img_in = nn.Conv3d(in_channels, hidden_size, kernel_size=patch_size, stride=patch_size) - - # condition embedder - self.condition_embedder = JoyImageTimeTextImageEmbedding( - dim=hidden_size, - time_freq_dim=256, - time_proj_dim=hidden_size * 6, - text_embed_dim=text_dim, - ) - - # double-stream blocks - self.double_blocks = nn.ModuleList( - [ - JoyImageTransformerBlock( - dim=hidden_size, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - mlp_width_ratio=mlp_width_ratio, - ) - for _ in range(num_layers) - ] - ) - - # output head - self.norm_out = FP32LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(hidden_size, self.out_channels * math.prod(patch_size)) - - self.gradient_checkpointing = False - - # ------------------------------------------------------------------ - # RoPE helper - # ------------------------------------------------------------------ - - def get_rotary_pos_embed( - self, - vis_rope_size: list[int], - txt_rope_size: int | None = None, - ): - target_ndim = 3 - if len(vis_rope_size) != target_ndim: - vis_rope_size = [1] * (target_ndim - len(vis_rope_size)) + list(vis_rope_size) - - head_dim = self.hidden_size // self.num_attention_heads - rope_dim_list = self.rope_dim_list - if rope_dim_list is None: - rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)] - if sum(rope_dim_list) != head_dim: - raise ValueError("sum(rope_dim_list) should equal head_dim") - - # Build a 3-D meshgrid [0, size) for each spatial axis - grid = torch.stack( - torch.meshgrid( - *[torch.linspace(0, s, s + 1, dtype=torch.float32)[:s] for s in vis_rope_size], - indexing="ij", - ), - dim=0, - ) - - # Per-axis 1-D rotary embeddings -> concat - vis_cos, vis_sin = [], [] - for i, dim in enumerate(rope_dim_list): - pos = grid[i].reshape(-1) - freqs = 1.0 / (self.theta ** (torch.arange(0, dim, 2, dtype=torch.float32)[: (dim // 2)] / dim)) - freqs = torch.outer(pos.float(), freqs) - vis_cos.append(freqs.cos().repeat_interleave(2, dim=1)) - vis_sin.append(freqs.sin().repeat_interleave(2, dim=1)) - vis_freqs = (torch.cat(vis_cos, dim=1), torch.cat(vis_sin, dim=1)) - - if txt_rope_size is None: - return vis_freqs, None - - # Text positions start right after the largest visual index - grid_txt = torch.arange(txt_rope_size) + grid.view(-1).max().item() + 1 - txt_cos, txt_sin = [], [] - for i, dim in enumerate(rope_dim_list): - freqs = 1.0 / (self.theta ** (torch.arange(0, dim, 2, dtype=torch.float32)[: (dim // 2)] / dim)) - freqs = torch.outer(grid_txt.float(), freqs) - txt_cos.append(freqs.cos().repeat_interleave(2, dim=1)) - txt_sin.append(freqs.sin().repeat_interleave(2, dim=1)) - txt_freqs = (torch.cat(txt_cos, dim=1), torch.cat(txt_sin, dim=1)) - - return vis_freqs, txt_freqs - - # ------------------------------------------------------------------ - # Unpatchify - # ------------------------------------------------------------------ - - def unpatchify(self, x: torch.Tensor, t: int, h: int, w: int) -> torch.Tensor: - c = self.out_channels - pt, ph, pw = self.patch_size - if t * h * w != x.shape[1]: - raise ValueError(f"Expected t*h*w ({t * h * w}) to equal x.shape[1] ({x.shape[1]})") - - x = x.reshape(x.shape[0], t, h, w, pt, ph, pw, c) - x = x.permute(0, 7, 1, 4, 2, 5, 3, 6) # nthwopqc -> nctohpwq - return x.reshape(x.shape[0], c, t * pt, h * ph, w * pw) - - # ------------------------------------------------------------------ - # Forward - # ------------------------------------------------------------------ - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - return_dict: bool = True, - ): - """ - The [`JoyImageEditTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)` or `(batch_size, num_items, num_channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor`, *optional*): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - """ - # handle multi-item input (b, n, c, t, h, w) - is_multi_item = hidden_states.ndim == 6 - num_items = 0 - if is_multi_item: - num_items = hidden_states.shape[1] - if num_items > 1: - if self.patch_size[0] != 1: - raise ValueError("For multi-item input, patch_size[0] must be 1") - hidden_states = torch.cat([hidden_states[:, -1:], hidden_states[:, :-1]], dim=1) - # rearrange: (b, n, c, t, h, w) -> (b, c, n*t, h, w) - b, n, c, t, h, w = hidden_states.shape - hidden_states = hidden_states.permute(0, 2, 1, 3, 4, 5).reshape(b, c, n * t, h, w) - - batch_size, _, ot, oh, ow = hidden_states.shape - tt = ot // self.patch_size[0] - th = oh // self.patch_size[1] - tw = ow // self.patch_size[2] - - # patchify - img = self.img_in(hidden_states).flatten(2).transpose(1, 2) - - # condition embeddings - _, vec, txt = self.condition_embedder(timestep, encoder_hidden_states) - if vec.shape[-1] > self.hidden_size: - vec = vec.unflatten(1, (6, -1)) - - txt_seq_len = txt.shape[1] - - # RoPE - vis_freqs, txt_freqs = self.get_rotary_pos_embed( - vis_rope_size=[tt, th, tw], - txt_rope_size=txt_seq_len if self.rope_type == "mrope" else None, - ) - - # main loop - for block in self.double_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - img, txt = self._gradient_checkpointing_func(block, img, txt, vec, (vis_freqs, txt_freqs)) - else: - img, txt = block( - hidden_states=img, - encoder_hidden_states=txt, - temb=vec, - image_rotary_emb=(vis_freqs, txt_freqs), - ) - - # final layer - img = self.proj_out(self.norm_out(img)) - img = self.unpatchify(img, tt, th, tw) - - # un-multi-item: (b, c, n*t, h, w) -> (b, n, c, t, h, w) - if is_multi_item: - c_out = img.shape[1] - img = img.reshape(batch_size, c_out, num_items, -1, oh, ow) - img = img.permute(0, 2, 1, 3, 4, 5) # (b, n, c, t, h, w) - if num_items > 1: - img = torch.cat([img[:, 1:], img[:, :1]], dim=1) - - if not return_dict: - return (img,) - return Transformer2DModelOutput(sample=img) diff --git a/diffusers/models/transformers/transformer_joyimage_edit_plus.py b/diffusers/models/transformers/transformer_joyimage_edit_plus.py deleted file mode 100644 index 4a13845faad302a9f99c0a86fe15c97efefbb040..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_joyimage_edit_plus.py +++ /dev/null @@ -1,539 +0,0 @@ -# Copyright 2025 The JoyImage Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -import math - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import FP32LayerNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _apply_rotary_emb_batched( - xq: torch.Tensor, - xk: torch.Tensor, - freqs_cis: tuple[torch.Tensor, torch.Tensor], -) -> tuple[torch.Tensor, torch.Tensor]: - """RoPE for batched [B, S, D] freqs.""" - cos, sin = freqs_cis[0].to(xq.device), freqs_cis[1].to(xq.device) - - # batched: [B, S, D] -> [B, S, 1, D] - cos = cos.unsqueeze(2) - sin = sin.unsqueeze(2) - - def _rotate_half(x): - x_real, x_imag = x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1) - return torch.stack([-x_imag, x_real], dim=-1).flatten(3) - - xq_out = (xq.float() * cos + _rotate_half(xq) * sin).type_as(xq) - xk_out = (xk.float() * cos + _rotate_half(xk) * sin).type_as(xk) - return xq_out, xk_out - - -# Copied from diffusers.models.transformers.transformer_joyimage.JoyImageModulate with JoyImage->JoyImageEditPlus -class JoyImageEditPlusModulate(nn.Module): - """Wan-style learnable modulation table. - - Produces `factor` modulation vectors by adding the conditioning signal to a learnable parameter table. - """ - - def __init__(self, hidden_size: int, factor: int, dtype=None, device=None): - super().__init__() - self.factor = factor - self.modulate_table = nn.Parameter( - torch.zeros(1, factor, hidden_size, dtype=dtype, device=device) / hidden_size**0.5, - requires_grad=True, - ) - - def forward(self, x: torch.Tensor) -> list[torch.Tensor]: - if x.ndim != 3: - x = x.unsqueeze(1) - return [o.squeeze(1) for o in (self.modulate_table + x).chunk(self.factor, dim=1)] - - -class JoyImageEditPlusAttnProcessor: - """Attention processor that supports batched RoPE embeddings for edit-plus multi-image input.""" - - _attention_backend = None - _parallel_config = None - - def __call__( - self, - attn: "JoyImageEditPlusAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_mask: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - if encoder_hidden_states is None: - raise ValueError("JoyImageEditPlusAttnProcessor requires encoder_hidden_states") - - heads = attn.heads - - img_qkv = attn.img_attn_qkv(hidden_states) - img_query, img_key, img_value = img_qkv.chunk(3, dim=-1) - - txt_qkv = attn.txt_attn_qkv(encoder_hidden_states) - txt_query, txt_key, txt_value = txt_qkv.chunk(3, dim=-1) - - img_query = img_query.unflatten(-1, (heads, -1)) - img_key = img_key.unflatten(-1, (heads, -1)) - img_value = img_value.unflatten(-1, (heads, -1)) - - txt_query = txt_query.unflatten(-1, (heads, -1)) - txt_key = txt_key.unflatten(-1, (heads, -1)) - txt_value = txt_value.unflatten(-1, (heads, -1)) - - img_query = attn.img_attn_q_norm(img_query) - img_key = attn.img_attn_k_norm(img_key) - txt_query = attn.txt_attn_q_norm(txt_query) - txt_key = attn.txt_attn_k_norm(txt_key) - - if image_rotary_emb is not None: - img_query, img_key = _apply_rotary_emb_batched(img_query, img_key, image_rotary_emb) - - joint_query = torch.cat([img_query, txt_query], dim=1) - joint_key = torch.cat([img_key, txt_key], dim=1) - joint_value = torch.cat([img_value, txt_value], dim=1) - - joint_hidden_states = dispatch_attention_fn( - joint_query, - joint_key, - joint_value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - joint_hidden_states = joint_hidden_states.flatten(2, 3) - joint_hidden_states = joint_hidden_states.to(joint_query.dtype) - - img_attn_output = joint_hidden_states[:, : hidden_states.shape[1], :] - txt_attn_output = joint_hidden_states[:, hidden_states.shape[1] :, :] - - img_attn_output = attn.img_attn_proj(img_attn_output) - txt_attn_output = attn.txt_attn_proj(txt_attn_output) - - return img_attn_output, txt_attn_output - - -class JoyImageEditPlusAttention(nn.Module, AttentionModuleMixin): - """Joint attention module for JoyImage Edit Plus double-stream blocks.""" - - _default_processor_cls = JoyImageEditPlusAttnProcessor - _available_processors = [JoyImageEditPlusAttnProcessor] - _supports_qkv_fusion = False - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - eps: float = 1e-6, - processor=None, - ): - super().__init__() - - self.heads = num_attention_heads - self.head_dim = attention_head_dim - inner_dim = num_attention_heads * attention_head_dim - - self.img_attn_qkv = nn.Linear(dim, inner_dim * 3, bias=True) - self.img_attn_q_norm = nn.RMSNorm(attention_head_dim, eps=eps) - self.img_attn_k_norm = nn.RMSNorm(attention_head_dim, eps=eps) - self.img_attn_proj = nn.Linear(inner_dim, dim, bias=True) - - self.txt_attn_qkv = nn.Linear(dim, inner_dim * 3, bias=True) - self.txt_attn_q_norm = nn.RMSNorm(attention_head_dim, eps=eps) - self.txt_attn_k_norm = nn.RMSNorm(attention_head_dim, eps=eps) - self.txt_attn_proj = nn.Linear(inner_dim, dim, bias=True) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_mask: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - kwargs = {} - if "attention_mask" in attn_parameters: - kwargs["attention_mask"] = attention_mask - return self.processor(self, hidden_states, encoder_hidden_states, image_rotary_emb, **kwargs) - - -class JoyImageEditPlusTransformerBlock(nn.Module): - """Double-stream transformer block for JoyImage Edit Plus.""" - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - mlp_width_ratio: float = 4.0, - eps: float = 1e-6, - ): - super().__init__() - - self.dim = dim - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - mlp_hidden_dim = int(dim * mlp_width_ratio) - - # image stream - self.img_mod = JoyImageEditPlusModulate(dim, factor=6) - self.img_norm1 = FP32LayerNorm(dim, elementwise_affine=False, eps=eps) - self.img_norm2 = FP32LayerNorm(dim, elementwise_affine=False, eps=eps) - self.img_mlp = FeedForward(dim, inner_dim=mlp_hidden_dim, activation_fn="gelu-approximate") - - # text stream - self.txt_mod = JoyImageEditPlusModulate(dim, factor=6) - self.txt_norm1 = FP32LayerNorm(dim, elementwise_affine=False, eps=eps) - self.txt_norm2 = FP32LayerNorm(dim, elementwise_affine=False, eps=eps) - self.txt_mlp = FeedForward(dim, inner_dim=mlp_hidden_dim, activation_fn="gelu-approximate") - - # joint attention - self.attn = JoyImageEditPlusAttention(dim, num_attention_heads, attention_head_dim, eps=eps) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_mask: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - # modulation - ( - img_mod1_shift, - img_mod1_scale, - img_mod1_gate, - img_mod2_shift, - img_mod2_scale, - img_mod2_gate, - ) = self.img_mod(temb) - ( - txt_mod1_shift, - txt_mod1_scale, - txt_mod1_gate, - txt_mod2_shift, - txt_mod2_scale, - txt_mod2_gate, - ) = self.txt_mod(temb) - - # --- attention --- - img_normed = self.img_norm1(hidden_states) - txt_normed = self.txt_norm1(encoder_hidden_states) - img_modulated = img_normed * (1 + img_mod1_scale.unsqueeze(1)) + img_mod1_shift.unsqueeze(1) - txt_modulated = txt_normed * (1 + txt_mod1_scale.unsqueeze(1)) + txt_mod1_shift.unsqueeze(1) - - img_attn, txt_attn = self.attn( - hidden_states=img_modulated, - encoder_hidden_states=txt_modulated, - image_rotary_emb=image_rotary_emb, - attention_mask=attention_mask, - ) - - hidden_states = hidden_states + img_attn * img_mod1_gate.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + txt_attn * txt_mod1_gate.unsqueeze(1) - - # --- FFN --- - img_ffn_normed = self.img_norm2(hidden_states) - txt_ffn_normed = self.txt_norm2(encoder_hidden_states) - img_ffn_input = img_ffn_normed * (1 + img_mod2_scale.unsqueeze(1)) + img_mod2_shift.unsqueeze(1) - txt_ffn_input = txt_ffn_normed * (1 + txt_mod2_scale.unsqueeze(1)) + txt_mod2_shift.unsqueeze(1) - img_ffn_output = self.img_mlp(img_ffn_input) - txt_ffn_output = self.txt_mlp(txt_ffn_input) - hidden_states = hidden_states + img_ffn_output * img_mod2_gate.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + txt_ffn_output * txt_mod2_gate.unsqueeze(1) - - return hidden_states, encoder_hidden_states - - -# Copied from diffusers.models.transformers.transformer_joyimage.JoyImageTimeTextImageEmbedding with JoyImage->JoyImageEditPlus -class JoyImageEditPlusTimeTextImageEmbedding(nn.Module): - def __init__( - self, - dim: int, - time_freq_dim: int, - time_proj_dim: int, - text_embed_dim: int, - ): - super().__init__() - - self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim) - self.act_fn = nn.SiLU() - self.time_proj = nn.Linear(dim, time_proj_dim) - self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh") - - def forward( - self, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - ): - timestep = self.timesteps_proj(timestep) - - time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype - if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8: - timestep = timestep.to(time_embedder_dtype) - temb = self.time_embedder(timestep).type_as(encoder_hidden_states) - timestep_proj = self.time_proj(self.act_fn(temb)) - - encoder_hidden_states = self.text_embedder(encoder_hidden_states) - - return temb, timestep_proj, encoder_hidden_states - - -class JoyImageEditPlusTransformer3DModel(ModelMixin, ConfigMixin, AttentionMixin): - r""" - JoyImage Edit Plus Transformer for multi-image editing. - - Uses a patchify+padding approach where each reference image and the target noise are independently patchified and - concatenated into a flat patch sequence. Supports variable-resolution reference images. - - Input format: `[B, max_patches, C, pt, ph, pw]` (6D padded patches). - - Args: - patch_size (`list`, defaults to `[1, 2, 2]`): - Patch size for patchifying the latent input along `(t, h, w)` dimensions. - in_channels (`int`, defaults to `16`): - The number of channels in the input latent. - out_channels (`int`, *optional*, defaults to `None`): - The number of channels in the output. If not specified, it defaults to `in_channels`. - hidden_size (`int`, defaults to `3072`): - The dimensionality of the hidden representations. - num_attention_heads (`int`, defaults to `24`): - The number of attention heads. - text_dim (`int`, defaults to `4096`): - The dimensionality of the text encoder output. - mlp_width_ratio (`float`, defaults to `4.0`): - The ratio of MLP hidden dimension to `hidden_size`. - num_layers (`int`, defaults to `20`): - The number of double-stream transformer blocks. - rope_dim_list (`list[int]`, defaults to `[16, 56, 56]`): - The dimensions for 3D rotary positional embeddings along `(t, h, w)`. - rope_type (`str`, defaults to `"rope"`): - The type of rotary positional embedding. - theta (`int`, defaults to `256`): - The base frequency for rotary embeddings. - """ - - _skip_layerwise_casting_patterns = ["img_in", "condition_embedder", "norm"] - _no_split_modules = ["JoyImageEditPlusTransformerBlock"] - _supports_gradient_checkpointing = True - _keep_in_fp32_modules = [ - "time_embedder", - "norm1", - "norm2", - "norm_out", - ] - _repeated_blocks = ["JoyImageEditPlusTransformerBlock"] - - @register_to_config - def __init__( - self, - patch_size: list[int] = [1, 2, 2], - in_channels: int = 16, - out_channels: int | None = None, - hidden_size: int = 3072, - num_attention_heads: int = 24, - text_dim: int = 4096, - mlp_width_ratio: float = 4.0, - num_layers: int = 20, - rope_dim_list: list[int] = [16, 56, 56], - rope_type: str = "rope", - theta: int = 256, - ): - super().__init__() - - self.out_channels = out_channels or in_channels - - attention_head_dim = hidden_size // num_attention_heads - if hidden_size % num_attention_heads != 0: - raise ValueError( - f"hidden_size ({hidden_size}) must be divisible by num_attention_heads ({num_attention_heads})" - ) - - self.img_in = nn.Conv3d(in_channels, hidden_size, kernel_size=patch_size, stride=patch_size) - - self.condition_embedder = JoyImageEditPlusTimeTextImageEmbedding( - dim=hidden_size, - time_freq_dim=256, - time_proj_dim=hidden_size * 6, - text_embed_dim=text_dim, - ) - - self.double_blocks = nn.ModuleList( - [ - JoyImageEditPlusTransformerBlock( - dim=hidden_size, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - mlp_width_ratio=mlp_width_ratio, - ) - for _ in range(num_layers) - ] - ) - - self.norm_out = FP32LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(hidden_size, self.out_channels * math.prod(patch_size)) - - self.gradient_checkpointing = False - - # Set batched-RoPE-aware attention processor on all blocks - for block in self.double_blocks: - block.attn.set_processor(JoyImageEditPlusAttnProcessor()) - - def _get_rotary_pos_embed_for_range( - self, - start: tuple[int, int, int], - stop: tuple[int, int, int], - ) -> tuple[torch.Tensor, torch.Tensor]: - """Generate 3D RoPE for a spatial range [start, stop).""" - head_dim = self.config.hidden_size // self.config.num_attention_heads - rope_dim_list = self.config.rope_dim_list - if rope_dim_list is None: - rope_dim_list = [head_dim // 3] * 3 - - grids = [] - for i in range(3): - grids.append(torch.arange(start[i], stop[i], dtype=torch.float32)) - - mesh = torch.stack(torch.meshgrid(*grids, indexing="ij"), dim=0) - - cos_parts, sin_parts = [], [] - for i, dim in enumerate(rope_dim_list): - pos = mesh[i].reshape(-1) - freqs = 1.0 / (self.config.theta ** (torch.arange(0, dim, 2, dtype=torch.float32)[: (dim // 2)] / dim)) - angles = torch.outer(pos, freqs) - cos_parts.append(angles.cos().repeat_interleave(2, dim=1)) - sin_parts.append(angles.sin().repeat_interleave(2, dim=1)) - - return torch.cat(cos_parts, dim=1), torch.cat(sin_parts, dim=1) - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_mask: torch.Tensor | None = None, - shape_list: list[list[tuple[int, int, int]]] = None, - return_dict: bool = True, - ) -> torch.Tensor | tuple: - """ - Args: - hidden_states: [B, max_patches, C, pt, ph, pw] - patchified latent input. - timestep: [B] - diffusion timestep. - encoder_hidden_states: [B, L, D] - text encoder outputs. - encoder_hidden_states_mask: [B, L] - attention mask for text tokens. - shape_list: Per-sample list of (t, h, w) tuples for each component (target + references). - return_dict: Whether to return a dict or tuple. - - Returns: - If `return_dict` is True, an [`~models.modeling_outputs.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, max_num_patches, channels, pt, ph, pw = hidden_states.shape - device = hidden_states.device - - # 1. Condition embeddings - _, vec, txt = self.condition_embedder(timestep, encoder_hidden_states) - vec = vec.unflatten(1, (6, -1)) - - # 2. Patchify via Conv3d: flatten (B, N) -> apply conv -> reshape back - x = hidden_states.reshape(batch_size * max_num_patches, channels, pt, ph, pw) - x = self.img_in(x) # (B*N, D, 1, 1, 1) - img = x.reshape(batch_size, max_num_patches, -1) - - # 3. Build per-component RoPE with temporal offsets - sample_cos_list, sample_sin_list = [], [] - - for i in range(batch_size): - s_cos_parts, s_sin_parts = [], [] - current_t_offset = 0 - - for thw in shape_list[i]: - t, h, w = thw - start = (current_t_offset, 0, 0) - stop = (current_t_offset + t, h, w) - cos_emb, sin_emb = self._get_rotary_pos_embed_for_range(start, stop) - s_cos_parts.append(cos_emb) - s_sin_parts.append(sin_emb) - current_t_offset += t - - s_cos = torch.cat(s_cos_parts, dim=0).to(device) - s_sin = torch.cat(s_sin_parts, dim=0).to(device) - - actual_len = s_cos.shape[0] - pad_len = max_num_patches - actual_len - if pad_len > 0: - s_cos = F.pad(s_cos, (0, 0, 0, pad_len), value=1.0) - s_sin = F.pad(s_sin, (0, 0, 0, pad_len), value=0.0) - - sample_cos_list.append(s_cos) - sample_sin_list.append(s_sin) - - vis_freqs = (torch.stack(sample_cos_list), torch.stack(sample_sin_list)) - - # 4. Build attention mask: [B, 1, 1, img_seq + txt_seq] - attention_mask = None - if encoder_hidden_states_mask is not None: - img_mask = torch.zeros(batch_size, max_num_patches, device=device, dtype=encoder_hidden_states_mask.dtype) - for i in range(batch_size): - actual_len = sum(t * h * w for t, h, w in shape_list[i]) - img_mask[i, :actual_len] = 1.0 - full_mask = torch.cat([img_mask, encoder_hidden_states_mask], dim=1) - attention_mask = full_mask.unsqueeze(1).unsqueeze(1).bool() - - # 5. Run double blocks - for block in self.double_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - img, txt = self._gradient_checkpointing_func(block, img, txt, vec, vis_freqs, attention_mask) - else: - img, txt = block( - hidden_states=img, - encoder_hidden_states=txt, - temb=vec, - image_rotary_emb=vis_freqs, - attention_mask=attention_mask, - ) - - # 6. Output projection + reshape to 6D patches - img = self.proj_out(self.norm_out(img)) - img = img.reshape(batch_size, max_num_patches, pt, ph, pw, self.out_channels).permute( - 0, 1, 5, 2, 3, 4 - ) # -> [B, N, C, pt, ph, pw] - - if not return_dict: - return (img,) - return Transformer2DModelOutput(sample=img) diff --git a/diffusers/models/transformers/transformer_kandinsky.py b/diffusers/models/transformers/transformer_kandinsky.py deleted file mode 100644 index 88ef70d546c8dafc849ca71943248206b61757e8..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_kandinsky.py +++ /dev/null @@ -1,668 +0,0 @@ -# Copyright 2025 The Kandinsky Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -import math -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F -from torch import Tensor - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import ( - logging, -) -from ..attention import AttentionMixin, AttentionModuleMixin -from ..attention_dispatch import _CAN_USE_FLEX_ATTN, dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin - - -logger = logging.get_logger(__name__) - - -def get_freqs(dim, max_period=10000.0): - freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=dim, dtype=torch.float32) / dim) - return freqs - - -def fractal_flatten(x, rope, shape, block_mask=False): - if block_mask: - pixel_size = 8 - x = local_patching(x, shape, (1, pixel_size, pixel_size), dim=1) - rope = local_patching(rope, shape, (1, pixel_size, pixel_size), dim=1) - x = x.flatten(1, 2) - rope = rope.flatten(1, 2) - else: - x = x.flatten(1, 3) - rope = rope.flatten(1, 3) - return x, rope - - -def fractal_unflatten(x, shape, block_mask=False): - if block_mask: - pixel_size = 8 - x = x.reshape(x.shape[0], -1, pixel_size**2, *x.shape[2:]) - x = local_merge(x, shape, (1, pixel_size, pixel_size), dim=1) - else: - x = x.reshape(*shape, *x.shape[2:]) - return x - - -def local_patching(x, shape, group_size, dim=0): - batch_size, duration, height, width = shape - g1, g2, g3 = group_size - x = x.reshape( - *x.shape[:dim], - duration // g1, - g1, - height // g2, - g2, - width // g3, - g3, - *x.shape[dim + 3 :], - ) - x = x.permute( - *range(len(x.shape[:dim])), - dim, - dim + 2, - dim + 4, - dim + 1, - dim + 3, - dim + 5, - *range(dim + 6, len(x.shape)), - ) - x = x.flatten(dim, dim + 2).flatten(dim + 1, dim + 3) - return x - - -def local_merge(x, shape, group_size, dim=0): - batch_size, duration, height, width = shape - g1, g2, g3 = group_size - x = x.reshape( - *x.shape[:dim], - duration // g1, - height // g2, - width // g3, - g1, - g2, - g3, - *x.shape[dim + 2 :], - ) - x = x.permute( - *range(len(x.shape[:dim])), - dim, - dim + 3, - dim + 1, - dim + 4, - dim + 2, - dim + 5, - *range(dim + 6, len(x.shape)), - ) - x = x.flatten(dim, dim + 1).flatten(dim + 1, dim + 2).flatten(dim + 2, dim + 3) - return x - - -def nablaT_v2( - q: Tensor, - k: Tensor, - sta: Tensor, - thr: float = 0.9, -): - if _CAN_USE_FLEX_ATTN: - from torch.nn.attention.flex_attention import BlockMask - else: - raise ValueError("Nabla attention is not supported with this version of PyTorch") - - q = q.transpose(1, 2).contiguous() - k = k.transpose(1, 2).contiguous() - - # Map estimation - B, h, S, D = q.shape - s1 = S // 64 - qa = q.reshape(B, h, s1, 64, D).mean(-2) - ka = k.reshape(B, h, s1, 64, D).mean(-2).transpose(-2, -1) - map = qa @ ka - - map = torch.softmax(map / math.sqrt(D), dim=-1) - # Map binarization - vals, inds = map.sort(-1) - cvals = vals.cumsum_(-1) - mask = (cvals >= 1 - thr).int() - mask = mask.gather(-1, inds.argsort(-1)) - - mask = torch.logical_or(mask, sta) - - # BlockMask creation - kv_nb = mask.sum(-1).to(torch.int32) - kv_inds = mask.argsort(dim=-1, descending=True).to(torch.int32) - return BlockMask.from_kv_blocks(torch.zeros_like(kv_nb), kv_inds, kv_nb, kv_inds, BLOCK_SIZE=64, mask_mod=None) - - -class Kandinsky5TimeEmbeddings(nn.Module): - def __init__(self, model_dim, time_dim, max_period=10000.0): - super().__init__() - assert model_dim % 2 == 0 - self.model_dim = model_dim - self.max_period = max_period - self.freqs = get_freqs(self.model_dim // 2, self.max_period) - self.in_layer = nn.Linear(model_dim, time_dim, bias=True) - self.activation = nn.SiLU() - self.out_layer = nn.Linear(time_dim, time_dim, bias=True) - - def forward(self, time): - args = torch.outer(time.to(torch.float32), self.freqs.to(device=time.device)) - time_embed = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) - time_embed = self.out_layer(self.activation(self.in_layer(time_embed))) - return time_embed - - -class Kandinsky5TextEmbeddings(nn.Module): - def __init__(self, text_dim, model_dim): - super().__init__() - self.in_layer = nn.Linear(text_dim, model_dim, bias=True) - self.norm = nn.LayerNorm(model_dim, elementwise_affine=True) - - def forward(self, text_embed): - text_embed = self.in_layer(text_embed) - return self.norm(text_embed).type_as(text_embed) - - -class Kandinsky5VisualEmbeddings(nn.Module): - def __init__(self, visual_dim, model_dim, patch_size): - super().__init__() - self.patch_size = patch_size - self.in_layer = nn.Linear(math.prod(patch_size) * visual_dim, model_dim) - - def forward(self, x): - batch_size, duration, height, width, dim = x.shape - x = ( - x.view( - batch_size, - duration // self.patch_size[0], - self.patch_size[0], - height // self.patch_size[1], - self.patch_size[1], - width // self.patch_size[2], - self.patch_size[2], - dim, - ) - .permute(0, 1, 3, 5, 2, 4, 6, 7) - .flatten(4, 7) - ) - return self.in_layer(x) - - -class Kandinsky5RoPE1D(nn.Module): - def __init__(self, dim, max_pos=1024, max_period=10000.0): - super().__init__() - self.max_period = max_period - self.dim = dim - self.max_pos = max_pos - freq = get_freqs(dim // 2, max_period) - pos = torch.arange(max_pos, dtype=freq.dtype) - self.register_buffer("args", torch.outer(pos, freq), persistent=False) - - def forward(self, pos): - args = self.args[pos] - cosine = torch.cos(args) - sine = torch.sin(args) - rope = torch.stack([cosine, -sine, sine, cosine], dim=-1) - rope = rope.view(*rope.shape[:-1], 2, 2) - return rope.unsqueeze(-4) - - -class Kandinsky5RoPE3D(nn.Module): - def __init__(self, axes_dims, max_pos=(128, 128, 128), max_period=10000.0): - super().__init__() - self.axes_dims = axes_dims - self.max_pos = max_pos - self.max_period = max_period - - for i, (axes_dim, ax_max_pos) in enumerate(zip(axes_dims, max_pos)): - freq = get_freqs(axes_dim // 2, max_period) - pos = torch.arange(ax_max_pos, dtype=freq.dtype) - self.register_buffer(f"args_{i}", torch.outer(pos, freq), persistent=False) - - def forward(self, shape, pos, scale_factor=(1.0, 1.0, 1.0)): - batch_size, duration, height, width = shape - args_t = self.args_0[pos[0]] / scale_factor[0] - args_h = self.args_1[pos[1]] / scale_factor[1] - args_w = self.args_2[pos[2]] / scale_factor[2] - - args = torch.cat( - [ - args_t.view(1, duration, 1, 1, -1).repeat(batch_size, 1, height, width, 1), - args_h.view(1, 1, height, 1, -1).repeat(batch_size, duration, 1, width, 1), - args_w.view(1, 1, 1, width, -1).repeat(batch_size, duration, height, 1, 1), - ], - dim=-1, - ) - cosine = torch.cos(args) - sine = torch.sin(args) - rope = torch.stack([cosine, -sine, sine, cosine], dim=-1) - rope = rope.view(*rope.shape[:-1], 2, 2) - return rope.unsqueeze(-4) - - -class Kandinsky5Modulation(nn.Module): - def __init__(self, time_dim, model_dim, num_params): - super().__init__() - self.activation = nn.SiLU() - self.out_layer = nn.Linear(time_dim, num_params * model_dim) - self.out_layer.weight.data.zero_() - self.out_layer.bias.data.zero_() - - def forward(self, x): - return self.out_layer(self.activation(x)) - - -class Kandinsky5AttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError(f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version.") - - def __call__(self, attn, hidden_states, encoder_hidden_states=None, rotary_emb=None, sparse_params=None): - # query, key, value = self.get_qkv(x) - query = attn.to_query(hidden_states) - - if encoder_hidden_states is not None: - key = attn.to_key(encoder_hidden_states) - value = attn.to_value(encoder_hidden_states) - - shape, cond_shape = query.shape[:-1], key.shape[:-1] - query = query.reshape(*shape, attn.num_heads, -1) - key = key.reshape(*cond_shape, attn.num_heads, -1) - value = value.reshape(*cond_shape, attn.num_heads, -1) - - else: - key = attn.to_key(hidden_states) - value = attn.to_value(hidden_states) - - shape = query.shape[:-1] - query = query.reshape(*shape, attn.num_heads, -1) - key = key.reshape(*shape, attn.num_heads, -1) - value = value.reshape(*shape, attn.num_heads, -1) - - # query, key = self.norm_qk(query, key) - query = attn.query_norm(query.float()).type_as(query) - key = attn.key_norm(key.float()).type_as(key) - - def apply_rotary(x, rope): - x_ = x.reshape(*x.shape[:-1], -1, 1, 2).to(torch.float32) - x_out = (rope * x_).sum(dim=-1) - return x_out.reshape(*x.shape).to(torch.bfloat16) - - if rotary_emb is not None: - query = apply_rotary(query, rotary_emb).type_as(query) - key = apply_rotary(key, rotary_emb).type_as(key) - - if sparse_params is not None: - attn_mask = nablaT_v2( - query, - key, - sparse_params["sta_mask"], - thr=sparse_params["P"], - ) - - else: - attn_mask = None - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attn_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - hidden_states = hidden_states.flatten(-2, -1) - - attn_out = attn.out_layer(hidden_states) - return attn_out - - -class Kandinsky5Attention(nn.Module, AttentionModuleMixin): - _default_processor_cls = Kandinsky5AttnProcessor - _available_processors = [ - Kandinsky5AttnProcessor, - ] - - def __init__(self, num_channels, head_dim, processor=None): - super().__init__() - assert num_channels % head_dim == 0 - self.num_heads = num_channels // head_dim - - self.to_query = nn.Linear(num_channels, num_channels, bias=True) - self.to_key = nn.Linear(num_channels, num_channels, bias=True) - self.to_value = nn.Linear(num_channels, num_channels, bias=True) - self.query_norm = nn.RMSNorm(head_dim) - self.key_norm = nn.RMSNorm(head_dim) - - self.out_layer = nn.Linear(num_channels, num_channels, bias=True) - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - sparse_params: torch.Tensor | None = None, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - quiet_attn_parameters = {} - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters and k not in quiet_attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"attention_processor_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - - return self.processor( - self, - hidden_states, - encoder_hidden_states=encoder_hidden_states, - sparse_params=sparse_params, - rotary_emb=rotary_emb, - **kwargs, - ) - - -class Kandinsky5FeedForward(nn.Module): - def __init__(self, dim, ff_dim): - super().__init__() - self.in_layer = nn.Linear(dim, ff_dim, bias=False) - self.activation = nn.GELU() - self.out_layer = nn.Linear(ff_dim, dim, bias=False) - - def forward(self, x): - return self.out_layer(self.activation(self.in_layer(x))) - - -class Kandinsky5OutLayer(nn.Module): - def __init__(self, model_dim, time_dim, visual_dim, patch_size): - super().__init__() - self.patch_size = patch_size - self.modulation = Kandinsky5Modulation(time_dim, model_dim, 2) - self.norm = nn.LayerNorm(model_dim, elementwise_affine=False) - self.out_layer = nn.Linear(model_dim, math.prod(patch_size) * visual_dim, bias=True) - - def forward(self, visual_embed, text_embed, time_embed): - shift, scale = torch.chunk(self.modulation(time_embed).unsqueeze(dim=1), 2, dim=-1) - - visual_embed = ( - self.norm(visual_embed.float()) * (scale.float()[:, None, None] + 1.0) + shift.float()[:, None, None] - ).type_as(visual_embed) - - x = self.out_layer(visual_embed) - - batch_size, duration, height, width, _ = x.shape - x = ( - x.view( - batch_size, - duration, - height, - width, - -1, - self.patch_size[0], - self.patch_size[1], - self.patch_size[2], - ) - .permute(0, 1, 5, 2, 6, 3, 7, 4) - .flatten(1, 2) - .flatten(2, 3) - .flatten(3, 4) - ) - return x - - -class Kandinsky5TransformerEncoderBlock(nn.Module): - def __init__(self, model_dim, time_dim, ff_dim, head_dim): - super().__init__() - self.text_modulation = Kandinsky5Modulation(time_dim, model_dim, 6) - - self.self_attention_norm = nn.LayerNorm(model_dim, elementwise_affine=False) - self.self_attention = Kandinsky5Attention(model_dim, head_dim, processor=Kandinsky5AttnProcessor()) - - self.feed_forward_norm = nn.LayerNorm(model_dim, elementwise_affine=False) - self.feed_forward = Kandinsky5FeedForward(model_dim, ff_dim) - - def forward(self, x, time_embed, rope): - self_attn_params, ff_params = torch.chunk(self.text_modulation(time_embed).unsqueeze(dim=1), 2, dim=-1) - shift, scale, gate = torch.chunk(self_attn_params, 3, dim=-1) - out = (self.self_attention_norm(x.float()) * (scale.float() + 1.0) + shift.float()).type_as(x) - out = self.self_attention(out, rotary_emb=rope) - x = (x.float() + gate.float() * out.float()).type_as(x) - - shift, scale, gate = torch.chunk(ff_params, 3, dim=-1) - out = (self.feed_forward_norm(x.float()) * (scale.float() + 1.0) + shift.float()).type_as(x) - out = self.feed_forward(out) - x = (x.float() + gate.float() * out.float()).type_as(x) - - return x - - -class Kandinsky5TransformerDecoderBlock(nn.Module): - def __init__(self, model_dim, time_dim, ff_dim, head_dim): - super().__init__() - self.visual_modulation = Kandinsky5Modulation(time_dim, model_dim, 9) - - self.self_attention_norm = nn.LayerNorm(model_dim, elementwise_affine=False) - self.self_attention = Kandinsky5Attention(model_dim, head_dim, processor=Kandinsky5AttnProcessor()) - - self.cross_attention_norm = nn.LayerNorm(model_dim, elementwise_affine=False) - self.cross_attention = Kandinsky5Attention(model_dim, head_dim, processor=Kandinsky5AttnProcessor()) - - self.feed_forward_norm = nn.LayerNorm(model_dim, elementwise_affine=False) - self.feed_forward = Kandinsky5FeedForward(model_dim, ff_dim) - - def forward(self, visual_embed, text_embed, time_embed, rope, sparse_params): - self_attn_params, cross_attn_params, ff_params = torch.chunk( - self.visual_modulation(time_embed).unsqueeze(dim=1), 3, dim=-1 - ) - - shift, scale, gate = torch.chunk(self_attn_params, 3, dim=-1) - visual_out = (self.self_attention_norm(visual_embed.float()) * (scale.float() + 1.0) + shift.float()).type_as( - visual_embed - ) - visual_out = self.self_attention(visual_out, rotary_emb=rope, sparse_params=sparse_params) - visual_embed = (visual_embed.float() + gate.float() * visual_out.float()).type_as(visual_embed) - - shift, scale, gate = torch.chunk(cross_attn_params, 3, dim=-1) - visual_out = (self.cross_attention_norm(visual_embed.float()) * (scale.float() + 1.0) + shift.float()).type_as( - visual_embed - ) - visual_out = self.cross_attention(visual_out, encoder_hidden_states=text_embed) - visual_embed = (visual_embed.float() + gate.float() * visual_out.float()).type_as(visual_embed) - - shift, scale, gate = torch.chunk(ff_params, 3, dim=-1) - visual_out = (self.feed_forward_norm(visual_embed.float()) * (scale.float() + 1.0) + shift.float()).type_as( - visual_embed - ) - visual_out = self.feed_forward(visual_out) - visual_embed = (visual_embed.float() + gate.float() * visual_out.float()).type_as(visual_embed) - - return visual_embed - - -class Kandinsky5Transformer3DModel( - ModelMixin, - ConfigMixin, - PeftAdapterMixin, - FromOriginalModelMixin, - CacheMixin, - AttentionMixin, -): - """ - A 3D Diffusion Transformer model for video-like data. - """ - - _repeated_blocks = [ - "Kandinsky5TransformerEncoderBlock", - "Kandinsky5TransformerDecoderBlock", - ] - _keep_in_fp32_modules = ["time_embeddings", "modulation", "visual_modulation", "text_modulation"] - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_visual_dim=4, - in_text_dim=3584, - in_text_dim2=768, - time_dim=512, - out_visual_dim=4, - patch_size=(1, 2, 2), - model_dim=2048, - ff_dim=5120, - num_text_blocks=2, - num_visual_blocks=32, - axes_dims=(16, 24, 24), - visual_cond=False, - attention_type: str = "regular", - attention_causal: bool = None, - attention_local: bool = None, - attention_glob: bool = None, - attention_window: int = None, - attention_P: float = None, - attention_wT: int = None, - attention_wW: int = None, - attention_wH: int = None, - attention_add_sta: bool = None, - attention_method: str = None, - ): - super().__init__() - - head_dim = sum(axes_dims) - self.in_visual_dim = in_visual_dim - self.model_dim = model_dim - self.patch_size = patch_size - self.visual_cond = visual_cond - self.attention_type = attention_type - - visual_embed_dim = 2 * in_visual_dim + 1 if visual_cond else in_visual_dim - - # Initialize embeddings - self.time_embeddings = Kandinsky5TimeEmbeddings(model_dim, time_dim) - self.text_embeddings = Kandinsky5TextEmbeddings(in_text_dim, model_dim) - self.pooled_text_embeddings = Kandinsky5TextEmbeddings(in_text_dim2, time_dim) - self.visual_embeddings = Kandinsky5VisualEmbeddings(visual_embed_dim, model_dim, patch_size) - - # Initialize positional embeddings - self.text_rope_embeddings = Kandinsky5RoPE1D(head_dim) - self.visual_rope_embeddings = Kandinsky5RoPE3D(axes_dims) - - # Initialize transformer blocks - self.text_transformer_blocks = nn.ModuleList( - [Kandinsky5TransformerEncoderBlock(model_dim, time_dim, ff_dim, head_dim) for _ in range(num_text_blocks)] - ) - - self.visual_transformer_blocks = nn.ModuleList( - [ - Kandinsky5TransformerDecoderBlock(model_dim, time_dim, ff_dim, head_dim) - for _ in range(num_visual_blocks) - ] - ) - - # Initialize output layer - self.out_layer = Kandinsky5OutLayer(model_dim, time_dim, out_visual_dim, patch_size) - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, # x - encoder_hidden_states: torch.Tensor, # text_embed - timestep: torch.Tensor, # time - pooled_projections: torch.Tensor, # pooled_text_embed - visual_rope_pos: tuple[int, int, int], - text_rope_pos: torch.LongTensor, - scale_factor: tuple[float, float, float] = (1.0, 1.0, 1.0), - sparse_params: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> Transformer2DModelOutput | torch.FloatTensor: - """ - Forward pass of the Kandinsky5 3D Transformer. - - Args: - hidden_states (`torch.FloatTensor`): Input visual states - encoder_hidden_states (`torch.FloatTensor`): Text embeddings - timestep (`torch.Tensor` or `float` or `int`): Current timestep - pooled_projections (`torch.FloatTensor`): Pooled text embeddings - visual_rope_pos (`tuple[int, int, int]`): Position for visual RoPE - text_rope_pos (`torch.LongTensor`): Position for text RoPE - scale_factor (`tuple[float, float, float]`, optional): Scale factor for RoPE - sparse_params (`dict[str, Any]`, optional): Parameters for sparse attention - return_dict (`bool`, optional): Whether to return a dictionary - - Returns: - [`~models.transformer_2d.Transformer2DModelOutput`] or `torch.FloatTensor`: The output of the transformer - """ - x = hidden_states - text_embed = encoder_hidden_states - time = timestep - pooled_text_embed = pooled_projections - - text_embed = self.text_embeddings(text_embed) - time_embed = self.time_embeddings(time) - time_embed = time_embed + self.pooled_text_embeddings(pooled_text_embed) - visual_embed = self.visual_embeddings(x) - text_rope = self.text_rope_embeddings(text_rope_pos) - text_rope = text_rope.unsqueeze(dim=0) - - for text_transformer_block in self.text_transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - text_embed = self._gradient_checkpointing_func( - text_transformer_block, text_embed, time_embed, text_rope - ) - else: - text_embed = text_transformer_block(text_embed, time_embed, text_rope) - - visual_shape = visual_embed.shape[:-1] - visual_rope = self.visual_rope_embeddings(visual_shape, visual_rope_pos, scale_factor) - to_fractal = sparse_params["to_fractal"] if sparse_params is not None else False - visual_embed, visual_rope = fractal_flatten(visual_embed, visual_rope, visual_shape, block_mask=to_fractal) - - for visual_transformer_block in self.visual_transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - visual_embed = self._gradient_checkpointing_func( - visual_transformer_block, - visual_embed, - text_embed, - time_embed, - visual_rope, - sparse_params, - ) - else: - visual_embed = visual_transformer_block( - visual_embed, text_embed, time_embed, visual_rope, sparse_params - ) - - visual_embed = fractal_unflatten(visual_embed, visual_shape, block_mask=to_fractal) - x = self.out_layer(visual_embed, text_embed, time_embed) - - if not return_dict: - return x - - return Transformer2DModelOutput(sample=x) diff --git a/diffusers/models/transformers/transformer_krea2.py b/diffusers/models/transformers/transformer_krea2.py deleted file mode 100644 index d1f6cd0ecdedce8f88d179986bb976a406a3769c..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_krea2.py +++ /dev/null @@ -1,522 +0,0 @@ -# Copyright 2026 Krea AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -import math -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..attention import AttentionMixin, AttentionModuleMixin -from ..attention_dispatch import dispatch_attention_fn -from ..embeddings import apply_rotary_emb, get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class Krea2RMSNorm(nn.Module): - """RMSNorm with a zero-centered scale: the effective multiplier is `1 + weight`, matching the Krea 2 checkpoint - format. The activations are upcast so the normalization runs in float32; the scale weight is kept in float32 by the - model's `_keep_in_fp32_modules`.""" - - def __init__(self, dim: int, eps: float = 1e-5) -> None: - super().__init__() - self.dim = dim - self.eps = eps - self.weight = nn.Parameter(torch.zeros(dim)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - dtype = hidden_states.dtype - hidden_states = F.rms_norm(hidden_states.float(), (self.dim,), weight=self.weight + 1.0, eps=self.eps) - return hidden_states.to(dtype) - - -class Krea2AttnProcessor: - _attention_backend = None - _parallel_config = None - - def __call__( - self, - attn: "Krea2Attention", - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - query = attn.to_q(hidden_states).unflatten(-1, (attn.num_heads, attn.head_dim)) - key = attn.to_k(hidden_states).unflatten(-1, (attn.num_kv_heads, attn.head_dim)) - value = attn.to_v(hidden_states).unflatten(-1, (attn.num_kv_heads, attn.head_dim)) - gate = attn.to_gate(hidden_states) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - enable_gqa=attn.num_heads != attn.num_kv_heads, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states * torch.sigmoid(gate) - return attn.to_out[0](hidden_states) - - -class Krea2Attention(nn.Module, AttentionModuleMixin): - """Self-attention with grouped-query projections, q/k RMSNorm, rotary embeddings and a sigmoid output gate.""" - - _default_processor_cls = Krea2AttnProcessor - _available_processors = [Krea2AttnProcessor] - - def __init__( - self, hidden_size: int, num_heads: int, num_kv_heads: int | None = None, eps: float = 1e-5, processor=None - ) -> None: - super().__init__() - if hidden_size % num_heads != 0: - raise ValueError(f"hidden_size={hidden_size} must be divisible by num_heads={num_heads}") - self.hidden_size = hidden_size - self.num_heads = num_heads - self.num_kv_heads = num_kv_heads if num_kv_heads is not None else num_heads - self.head_dim = hidden_size // num_heads - self.use_bias = False - - self.to_q = nn.Linear(hidden_size, self.head_dim * self.num_heads, bias=False) - self.to_k = nn.Linear(hidden_size, self.head_dim * self.num_kv_heads, bias=False) - self.to_v = nn.Linear(hidden_size, self.head_dim * self.num_kv_heads, bias=False) - self.to_gate = nn.Linear(hidden_size, hidden_size, bias=False) - self.norm_q = Krea2RMSNorm(self.head_dim, eps=eps) - self.norm_k = Krea2RMSNorm(self.head_dim, eps=eps) - self.to_out = nn.ModuleList([nn.Linear(hidden_size, hidden_size, bias=False), nn.Dropout(0.0)]) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - unused_kwargs = [k for k in kwargs if k not in attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -class Krea2SwiGLU(nn.Module): - """SwiGLU feed-forward network.""" - - def __init__(self, dim: int, hidden_dim: int) -> None: - super().__init__() - self.gate = nn.Linear(dim, hidden_dim, bias=False) - self.up = nn.Linear(dim, hidden_dim, bias=False) - self.down = nn.Linear(hidden_dim, dim, bias=False) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return self.down(F.silu(self.gate(hidden_states)) * self.up(hidden_states)) - - -class Krea2TextFusionBlock(nn.Module): - """Pre-norm transformer block (no rotary embeddings, no time modulation) used by the text fusion stage.""" - - def __init__(self, dim: int, num_heads: int, num_kv_heads: int, intermediate_size: int, eps: float) -> None: - super().__init__() - self.norm1 = Krea2RMSNorm(dim, eps=eps) - self.norm2 = Krea2RMSNorm(dim, eps=eps) - self.attn = Krea2Attention(dim, num_heads, num_kv_heads, eps=eps) - self.ff = Krea2SwiGLU(dim, intermediate_size) - - def forward(self, hidden_states: torch.Tensor, attention_mask: torch.Tensor | None = None) -> torch.Tensor: - hidden_states = hidden_states + self.attn(self.norm1(hidden_states), attention_mask=attention_mask) - hidden_states = hidden_states + self.ff(self.norm2(hidden_states)) - return hidden_states - - -class Krea2TextFusion(nn.Module): - """Fuses the stack of tapped text-encoder hidden states into a single sequence of text features. - - Two `layerwise_blocks` attend across the `num_text_layers` axis independently for every token, a linear `projector` - collapses that axis, and two `refiner_blocks` attend across the token sequence. - """ - - def __init__( - self, - num_text_layers: int, - dim: int, - num_heads: int, - num_kv_heads: int, - intermediate_size: int, - num_layerwise_blocks: int, - num_refiner_blocks: int, - eps: float, - ) -> None: - super().__init__() - self.layerwise_blocks = nn.ModuleList( - [ - Krea2TextFusionBlock(dim, num_heads, num_kv_heads, intermediate_size, eps) - for _ in range(num_layerwise_blocks) - ] - ) - self.projector = nn.Linear(num_text_layers, 1, bias=False) - self.refiner_blocks = nn.ModuleList( - [ - Krea2TextFusionBlock(dim, num_heads, num_kv_heads, intermediate_size, eps) - for _ in range(num_refiner_blocks) - ] - ) - - def forward(self, encoder_hidden_states: torch.Tensor, attention_mask: torch.Tensor | None = None) -> torch.Tensor: - batch_size, seq_len, num_text_layers, dim = encoder_hidden_states.shape - - hidden_states = encoder_hidden_states.reshape(batch_size * seq_len, num_text_layers, dim) - for block in self.layerwise_blocks: - hidden_states = block(hidden_states.contiguous()) - - hidden_states = hidden_states.reshape(batch_size, seq_len, num_text_layers, dim).permute(0, 1, 3, 2) - hidden_states = self.projector(hidden_states).squeeze(-1) - - for block in self.refiner_blocks: - hidden_states = block(hidden_states, attention_mask=attention_mask) - - return hidden_states - - -class Krea2TransformerBlock(nn.Module): - def __init__( - self, hidden_size: int, intermediate_size: int, num_heads: int, num_kv_heads: int, norm_eps: float - ) -> None: - super().__init__() - self.scale_shift_table = nn.Parameter(torch.zeros(6, hidden_size)) - self.norm1 = Krea2RMSNorm(hidden_size, eps=norm_eps) - self.norm2 = Krea2RMSNorm(hidden_size, eps=norm_eps) - self.attn = Krea2Attention(hidden_size, num_heads, num_kv_heads, eps=norm_eps) - self.ff = Krea2SwiGLU(hidden_size, intermediate_size) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor], - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - # temb: (B, 1, 6 * hidden_size), shared across all blocks; each block only learns an additive table. - modulation = temb.unflatten(-1, (6, -1)) + self.scale_shift_table - prescale, preshift, pregate, postscale, postshift, postgate = modulation.unbind(-2) - - attn_out = self.attn( - (1.0 + prescale) * self.norm1(hidden_states) + preshift, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - hidden_states = hidden_states + pregate * attn_out - ff_out = self.ff((1.0 + postscale) * self.norm2(hidden_states) + postshift) - hidden_states = hidden_states + postgate * ff_out - return hidden_states - - -class Krea2TimestepEmbedding(nn.Module): - """Sinusoidal flow-time embedding (cos-first, input scaled by 1000) followed by a two-layer MLP. - - Keeps the sequence dimension at size 1 so the per-block modulations broadcast over tokens. - """ - - def __init__(self, embed_dim: int, hidden_size: int) -> None: - super().__init__() - self.embed_dim = embed_dim - self.linear_1 = nn.Linear(embed_dim, hidden_size, bias=True) - self.linear_2 = nn.Linear(hidden_size, hidden_size, bias=True) - - def forward(self, timestep: torch.Tensor, dtype: torch.dtype) -> torch.Tensor: - half = self.embed_dim // 2 - freqs = torch.exp(-math.log(1e4) * torch.arange(half, dtype=torch.float32, device=timestep.device) / half) - args = (timestep.float() * 1e3)[:, None, None] * freqs - emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1).to(dtype) - return self.linear_2(F.gelu(self.linear_1(emb), approximate="tanh")) - - -class Krea2TextProjection(nn.Module): - """Projects the fused text features into the transformer width.""" - - def __init__(self, text_dim: int, hidden_size: int, eps: float) -> None: - super().__init__() - self.norm = Krea2RMSNorm(text_dim, eps=eps) - self.linear_1 = nn.Linear(text_dim, hidden_size, bias=True) - self.linear_2 = nn.Linear(hidden_size, hidden_size, bias=True) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.linear_1(self.norm(hidden_states)) - return self.linear_2(F.gelu(hidden_states, approximate="tanh")) - - -class Krea2FinalLayer(nn.Module): - """Final adaptive RMSNorm and output projection. Kept as one module (and in `_no_split_modules`) so the learned - modulation table, norm and projection stay co-located under device-mapped inference.""" - - def __init__(self, hidden_size: int, out_channels: int, eps: float) -> None: - super().__init__() - self.scale_shift_table = nn.Parameter(torch.zeros(2, hidden_size)) - self.norm = Krea2RMSNorm(hidden_size, eps=eps) - self.linear = nn.Linear(hidden_size, out_channels, bias=True) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor) -> torch.Tensor: - modulation = temb + self.scale_shift_table - scale, shift = modulation.chunk(2, dim=1) - hidden_states = (1.0 + scale) * self.norm(hidden_states) + shift - return self.linear(hidden_states) - - -# Copied from diffusers.models.transformers.transformer_flux.FluxPosEmbed with FluxPosEmbed->Krea2RotaryPosEmbed -class Krea2RotaryPosEmbed(nn.Module): - # modified from https://github.com/black-forest-labs/flux/blob/c00d7c60b085fce8058b9df845e036090873f2ce/src/flux/modules/layers.py#L11 - def __init__(self, theta: int, axes_dim: list[int]): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - - def forward(self, ids: torch.Tensor) -> torch.Tensor: - n_axes = ids.shape[-1] - cos_out = [] - sin_out = [] - pos = ids.float() - freqs_dtype = maybe_adjust_dtype_for_device(torch.float64, ids.device) - for i in range(n_axes): - cos, sin = get_1d_rotary_pos_embed( - self.axes_dim[i], - pos[:, i], - theta=self.theta, - repeat_interleave_real=True, - use_real=True, - freqs_dtype=freqs_dtype, - ) - cos_out.append(cos) - sin_out.append(sin) - freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device) - freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device) - return freqs_cos, freqs_sin - - -class Krea2Transformer2DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixin): - r""" - The single-stream MMDiT flow-matching backbone used by the Krea 2 pipeline. - - Text conditioning enters as a stack of hidden states tapped from several layers of a multimodal text encoder. A - small text-fusion transformer collapses the layer axis and refines the token sequence; the result is concatenated - with the patchified image latents into a single `[text, image]` sequence processed by the transformer blocks. The - timestep conditions every block through one shared modulation vector plus per-block learned tables. - - Args: - in_channels (`int`, defaults to 64): - Latent channel count after patchification (`vae_channels * patch_size ** 2`). - num_layers (`int`, defaults to 28): - Number of transformer blocks. - attention_head_dim (`int`, defaults to 128): - Dimension of each attention head; the total hidden size is `attention_head_dim * num_attention_heads`. - num_attention_heads (`int`, defaults to 48): - Number of query heads. - num_key_value_heads (`int`, defaults to 12): - Number of key/value heads for grouped-query attention. - intermediate_size (`int`, defaults to 16384): - Feed-forward hidden size of the SwiGLU MLP inside each block. - timestep_embed_dim (`int`, defaults to 256): - Width of the sinusoidal timestep embedding before its MLP. - text_hidden_dim (`int`, defaults to 2560): - Hidden size of the text encoder whose hidden states are consumed. - num_text_layers (`int`, defaults to 12): - Number of tapped text-encoder hidden states stacked per token. - text_num_attention_heads (`int`, defaults to 20): - Number of query heads in the text fusion blocks. - text_num_key_value_heads (`int`, defaults to 20): - Number of key/value heads in the text fusion blocks. - text_intermediate_size (`int`, defaults to 6912): - Feed-forward hidden size of the SwiGLU MLP inside the text fusion blocks. - num_layerwise_text_blocks (`int`, defaults to 2): - Number of text fusion blocks applied across the tapped-layer axis (per token). - num_refiner_text_blocks (`int`, defaults to 2): - Number of text fusion blocks applied across the token sequence. - axes_dims_rope (`tuple[int, int, int]`, defaults to `(32, 48, 48)`): - Head-dim split across the (t, h, w) rotary position axes. - rope_theta (`float`, defaults to 1000.0): - Base used by the rotary position embedding. - norm_eps (`float`, defaults to 1e-5): - Epsilon used by all RMSNorm modules. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["Krea2TransformerBlock", "Krea2TextFusionBlock", "Krea2FinalLayer"] - _repeated_blocks = ["Krea2TransformerBlock"] - _keep_in_fp32_modules = ["norm", "norm1", "norm2", "norm_q", "norm_k"] - _skip_layerwise_casting_patterns = ["time_embed", "norm"] - - @register_to_config - def __init__( - self, - in_channels: int = 64, - num_layers: int = 28, - attention_head_dim: int = 128, - num_attention_heads: int = 48, - num_key_value_heads: int = 12, - intermediate_size: int = 16384, - timestep_embed_dim: int = 256, - text_hidden_dim: int = 2560, - num_text_layers: int = 12, - text_num_attention_heads: int = 20, - text_num_key_value_heads: int = 20, - text_intermediate_size: int = 6912, - num_layerwise_text_blocks: int = 2, - num_refiner_text_blocks: int = 2, - axes_dims_rope: tuple[int, int, int] = (32, 48, 48), - rope_theta: float = 1000.0, - norm_eps: float = 1e-5, - ) -> None: - super().__init__() - - hidden_size = attention_head_dim * num_attention_heads - if sum(axes_dims_rope) != attention_head_dim: - raise ValueError( - f"sum(axes_dims_rope)={sum(axes_dims_rope)} must equal attention_head_dim={attention_head_dim}" - ) - - self.in_channels = in_channels - self.out_channels = in_channels - self.hidden_size = hidden_size - self.gradient_checkpointing = False - - self.img_in = nn.Linear(in_channels, hidden_size, bias=True) - self.time_embed = Krea2TimestepEmbedding(timestep_embed_dim, hidden_size) - self.time_mod_proj = nn.Linear(hidden_size, 6 * hidden_size, bias=True) - self.text_fusion = Krea2TextFusion( - num_text_layers=num_text_layers, - dim=text_hidden_dim, - num_heads=text_num_attention_heads, - num_kv_heads=text_num_key_value_heads, - intermediate_size=text_intermediate_size, - num_layerwise_blocks=num_layerwise_text_blocks, - num_refiner_blocks=num_refiner_text_blocks, - eps=norm_eps, - ) - self.txt_in = Krea2TextProjection(text_hidden_dim, hidden_size, eps=norm_eps) - self.rotary_emb = Krea2RotaryPosEmbed(theta=rope_theta, axes_dim=list(axes_dims_rope)) - - self.transformer_blocks = nn.ModuleList( - [ - Krea2TransformerBlock( - hidden_size=hidden_size, - intermediate_size=intermediate_size, - num_heads=num_attention_heads, - num_kv_heads=num_key_value_heads, - norm_eps=norm_eps, - ) - for _ in range(num_layers) - ] - ) - - self.final_layer = Krea2FinalLayer(hidden_size, out_channels=in_channels, eps=norm_eps) - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: torch.Tensor, - position_ids: torch.Tensor, - encoder_attention_mask: torch.Tensor | None = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> Transformer2DModelOutput | tuple[torch.Tensor]: - r""" - Predict the flow-matching velocity for the image tokens. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, image_seq_len, in_channels)`): - Packed (patchified) noisy image latents. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, text_seq_len, num_text_layers, text_hidden_dim)`): - Stack of tapped text-encoder hidden states per token. - timestep (`torch.Tensor` of shape `(batch_size,)`): - Flow-matching time in `[0, 1]` (1 is pure noise, 0 is clean data). - position_ids (`torch.Tensor` of shape `(text_seq_len + image_seq_len, 3)`): - `(t, h, w)` rotary coordinates for the combined sequence. Text rows are all-zero; image rows hold the - latent-grid coordinates. - encoder_attention_mask (`torch.Tensor` of shape `(batch_size, text_seq_len)`, *optional*): - Boolean mask marking valid text tokens. Pass `None` when every text token is valid. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that, when it contains a `scale` entry, sets the LoRA scale applied to this - transformer's adapters for the duration of the forward pass. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.modeling_outputs.Transformer2DModelOutput`] instead of a plain tuple. - - Returns: - [`~models.modeling_outputs.Transformer2DModelOutput`] or a `tuple` whose first element is the velocity - tensor of shape `(batch_size, image_seq_len, in_channels)`. - """ - if position_ids.ndim != 2 or position_ids.shape[-1] != 3: - raise ValueError(f"`position_ids` must have shape (sequence_length, 3), got {tuple(position_ids.shape)}.") - - batch_size, image_seq_len, _ = hidden_states.shape - text_seq_len = encoder_hidden_states.shape[1] - - temb = self.time_embed(timestep, dtype=hidden_states.dtype) - temb_mod = self.time_mod_proj(F.gelu(temb, approximate="tanh")) - - text_attention_mask = None - attention_mask = None - if encoder_attention_mask is not None: - # Key-padding masks of shape (B, 1, 1, L): padded text tokens are excluded as attention keys everywhere; - # their own (garbage) lanes are never read back and are dropped at the output slice. - text_attention_mask = encoder_attention_mask[:, None, None, :] - image_mask = encoder_attention_mask.new_ones((batch_size, image_seq_len)) - attention_mask = torch.cat([encoder_attention_mask, image_mask], dim=1)[:, None, None, :] - - encoder_hidden_states = self.text_fusion(encoder_hidden_states, attention_mask=text_attention_mask) - encoder_hidden_states = self.txt_in(encoder_hidden_states) - - hidden_states = self.img_in(hidden_states) - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - image_rotary_emb = self.rotary_emb(position_ids) - - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, hidden_states, temb_mod, image_rotary_emb, attention_mask - ) - else: - hidden_states = block(hidden_states, temb_mod, image_rotary_emb, attention_mask) - - hidden_states = hidden_states[:, text_seq_len:] - output = self.final_layer(hidden_states, temb) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_longcat_audio_dit.py b/diffusers/models/transformers/transformer_longcat_audio_dit.py deleted file mode 100644 index 9b8c0b4bf147bf96a1abaa6d034a08afeee5020c..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_longcat_audio_dit.py +++ /dev/null @@ -1,630 +0,0 @@ -# Copyright 2026 MeiTuan LongCat-AudioDiT Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -# Adapted from the LongCat-AudioDiT reference implementation: -# https://github.com/meituan-longcat/LongCat-AudioDiT - -import math -from dataclasses import dataclass - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput -from ...utils.torch_utils import lru_cache_unless_export, maybe_allow_in_graph -from ..attention import AttentionModuleMixin -from ..attention_dispatch import dispatch_attention_fn -from ..modeling_utils import ModelMixin -from ..normalization import RMSNorm - - -@dataclass -class LongCatAudioDiTTransformerOutput(BaseOutput): - sample: torch.Tensor - - -class AudioDiTSinusPositionEmbedding(nn.Module): - def __init__(self, dim: int): - super().__init__() - self.dim = dim - - def forward(self, timesteps: torch.Tensor, scale: float = 1000.0) -> torch.Tensor: - device = timesteps.device - half_dim = self.dim // 2 - exponent = math.log(10000) / max(half_dim - 1, 1) - embeddings = torch.exp(torch.arange(half_dim, device=device).float() * -exponent) - embeddings = scale * timesteps.unsqueeze(1) * embeddings.unsqueeze(0) - return torch.cat((embeddings.sin(), embeddings.cos()), dim=-1) - - -class AudioDiTTimestepEmbedding(nn.Module): - def __init__(self, dim: int, freq_embed_dim: int = 256): - super().__init__() - self.time_embed = AudioDiTSinusPositionEmbedding(freq_embed_dim) - self.time_mlp = nn.Sequential(nn.Linear(freq_embed_dim, dim), nn.SiLU(), nn.Linear(dim, dim)) - - def forward(self, timestep: torch.Tensor) -> torch.Tensor: - hidden_states = self.time_embed(timestep) - return self.time_mlp(hidden_states.to(timestep.dtype)) - - -class AudioDiTRotaryEmbedding(nn.Module): - def __init__(self, dim: int, max_position_embeddings: int = 2048, base: float = 100000.0): - super().__init__() - self.dim = dim - self.max_position_embeddings = max_position_embeddings - self.base = base - - @lru_cache_unless_export(maxsize=128) - def _build(self, seq_len: int, device: torch.device | None = None) -> tuple[torch.Tensor, torch.Tensor]: - inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64).float() / self.dim)) - if device is not None: - inv_freq = inv_freq.to(device) - steps = torch.arange(seq_len, dtype=torch.int64, device=inv_freq.device).type_as(inv_freq) - freqs = torch.outer(steps, inv_freq) - embeddings = torch.cat((freqs, freqs), dim=-1) - return embeddings.cos().contiguous(), embeddings.sin().contiguous() - - def forward(self, hidden_states: torch.Tensor, seq_len: int | None = None) -> tuple[torch.Tensor, torch.Tensor]: - seq_len = hidden_states.shape[1] if seq_len is None else seq_len - cos, sin = self._build(max(seq_len, self.max_position_embeddings), hidden_states.device) - return cos[:seq_len].to(dtype=hidden_states.dtype), sin[:seq_len].to(dtype=hidden_states.dtype) - - -def _rotate_half(hidden_states: torch.Tensor) -> torch.Tensor: - first, second = hidden_states.chunk(2, dim=-1) - return torch.cat((-second, first), dim=-1) - - -def _apply_rotary_emb(hidden_states: torch.Tensor, rope: tuple[torch.Tensor, torch.Tensor]) -> torch.Tensor: - cos, sin = rope - cos = cos[None, :, None].to(hidden_states.device) - sin = sin[None, :, None].to(hidden_states.device) - return (hidden_states.float() * cos + _rotate_half(hidden_states).float() * sin).to(hidden_states.dtype) - - -class AudioDiTGRN(nn.Module): - def __init__(self, dim: int): - super().__init__() - self.gamma = nn.Parameter(torch.zeros(1, 1, dim)) - self.beta = nn.Parameter(torch.zeros(1, 1, dim)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - gx = torch.norm(hidden_states, p=2, dim=1, keepdim=True) - nx = gx / (gx.mean(dim=-1, keepdim=True) + 1e-6) - return self.gamma * (hidden_states * nx) + self.beta + hidden_states - - -class AudioDiTConvNeXtV2Block(nn.Module): - def __init__( - self, - dim: int, - intermediate_dim: int, - dilation: int = 1, - kernel_size: int = 7, - bias: bool = True, - eps: float = 1e-6, - ): - super().__init__() - padding = (dilation * (kernel_size - 1)) // 2 - self.dwconv = nn.Conv1d( - dim, dim, kernel_size=kernel_size, padding=padding, groups=dim, dilation=dilation, bias=bias - ) - self.norm = nn.LayerNorm(dim, eps=eps) - self.pwconv1 = nn.Linear(dim, intermediate_dim, bias=bias) - self.act = nn.SiLU() - self.grn = AudioDiTGRN(intermediate_dim) - self.pwconv2 = nn.Linear(intermediate_dim, dim, bias=bias) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - residual = hidden_states - hidden_states = self.dwconv(hidden_states.transpose(1, 2)).transpose(1, 2) - hidden_states = self.norm(hidden_states) - hidden_states = self.pwconv1(hidden_states) - hidden_states = self.act(hidden_states) - hidden_states = self.grn(hidden_states) - hidden_states = self.pwconv2(hidden_states) - return residual + hidden_states - - -class AudioDiTEmbedder(nn.Module): - def __init__(self, in_dim: int, out_dim: int): - super().__init__() - self.proj = nn.Sequential(nn.Linear(in_dim, out_dim), nn.SiLU(), nn.Linear(out_dim, out_dim)) - - def forward(self, hidden_states: torch.Tensor, mask: torch.BoolTensor | None = None) -> torch.Tensor: - if mask is not None: - hidden_states = hidden_states.masked_fill(mask.logical_not().unsqueeze(-1), 0.0) - hidden_states = self.proj(hidden_states) - if mask is not None: - hidden_states = hidden_states.masked_fill(mask.logical_not().unsqueeze(-1), 0.0) - return hidden_states - - -class AudioDiTAdaLNMLP(nn.Module): - def __init__(self, in_dim: int, out_dim: int, bias: bool = True): - super().__init__() - self.mlp = nn.Sequential(nn.SiLU(), nn.Linear(in_dim, out_dim, bias=bias)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return self.mlp(hidden_states) - - -class AudioDiTAdaLayerNormZeroFinal(nn.Module): - def __init__(self, dim: int, bias: bool = True, eps: float = 1e-6): - super().__init__() - self.silu = nn.SiLU() - self.linear = nn.Linear(dim, dim * 2, bias=bias) - self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=eps) - - def forward(self, hidden_states: torch.Tensor, embedding: torch.Tensor) -> torch.Tensor: - embedding = self.linear(self.silu(embedding)) - scale, shift = torch.chunk(embedding, 2, dim=-1) - hidden_states = self.norm(hidden_states.float()).type_as(hidden_states) - if scale.ndim == 2: - hidden_states = hidden_states * (1 + scale)[:, None, :] + shift[:, None, :] - else: - hidden_states = hidden_states * (1 + scale) + shift - return hidden_states - - -class AudioDiTSelfAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __call__( - self, - attn: "AudioDiTAttention", - hidden_states: torch.Tensor, - attention_mask: torch.BoolTensor | None = None, - audio_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - batch_size = hidden_states.shape[0] - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - if attn.qk_norm: - query = attn.q_norm(query) - key = attn.k_norm(key) - - head_dim = attn.inner_dim // attn.heads - query = query.view(batch_size, -1, attn.heads, head_dim) - key = key.view(batch_size, -1, attn.heads, head_dim) - value = value.view(batch_size, -1, attn.heads, head_dim) - - if audio_rotary_emb is not None: - query = _apply_rotary_emb(query, audio_rotary_emb) - key = _apply_rotary_emb(key, audio_rotary_emb) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - if attention_mask is not None: - hidden_states = hidden_states * attention_mask[:, :, None, None].to(hidden_states.dtype) - - hidden_states = hidden_states.flatten(2, 3).to(query.dtype) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class AudioDiTAttention(nn.Module, AttentionModuleMixin): - def __init__( - self, - q_dim: int, - kv_dim: int | None, - heads: int, - dim_head: int, - dropout: float = 0.0, - bias: bool = True, - qk_norm: bool = False, - eps: float = 1e-6, - processor: AttentionModuleMixin | None = None, - ): - super().__init__() - kv_dim = q_dim if kv_dim is None else kv_dim - self.heads = heads - self.inner_dim = dim_head * heads - self.to_q = nn.Linear(q_dim, self.inner_dim, bias=bias) - self.to_k = nn.Linear(kv_dim, self.inner_dim, bias=bias) - self.to_v = nn.Linear(kv_dim, self.inner_dim, bias=bias) - self.qk_norm = qk_norm - if qk_norm: - self.q_norm = RMSNorm(self.inner_dim, eps=eps) - self.k_norm = RMSNorm(self.inner_dim, eps=eps) - self.to_out = nn.ModuleList([nn.Linear(self.inner_dim, q_dim, bias=bias), nn.Dropout(dropout)]) - self.set_processor(processor or AudioDiTSelfAttnProcessor()) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - post_attention_mask: torch.BoolTensor | None = None, - attention_mask: torch.BoolTensor | None = None, - audio_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - prompt_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - if encoder_hidden_states is None: - return self.processor( - self, - hidden_states, - attention_mask=attention_mask, - audio_rotary_emb=audio_rotary_emb, - ) - return self.processor( - self, - hidden_states, - encoder_hidden_states=encoder_hidden_states, - post_attention_mask=post_attention_mask, - attention_mask=attention_mask, - audio_rotary_emb=audio_rotary_emb, - prompt_rotary_emb=prompt_rotary_emb, - ) - - -class AudioDiTCrossAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __call__( - self, - attn: "AudioDiTAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - post_attention_mask: torch.BoolTensor | None = None, - attention_mask: torch.BoolTensor | None = None, - audio_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - prompt_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - batch_size = hidden_states.shape[0] - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if attn.qk_norm: - query = attn.q_norm(query) - key = attn.k_norm(key) - - head_dim = attn.inner_dim // attn.heads - query = query.view(batch_size, -1, attn.heads, head_dim) - key = key.view(batch_size, -1, attn.heads, head_dim) - value = value.view(batch_size, -1, attn.heads, head_dim) - - if audio_rotary_emb is not None: - query = _apply_rotary_emb(query, audio_rotary_emb) - if prompt_rotary_emb is not None: - key = _apply_rotary_emb(key, prompt_rotary_emb) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - if post_attention_mask is not None: - hidden_states = hidden_states * post_attention_mask[:, :, None, None].to(hidden_states.dtype) - - hidden_states = hidden_states.flatten(2, 3).to(query.dtype) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class AudioDiTFeedForward(nn.Module): - def __init__(self, dim: int, mult: float = 4.0, dropout: float = 0.0, bias: bool = True): - super().__init__() - inner_dim = int(dim * mult) - self.ff = nn.Sequential( - nn.Linear(dim, inner_dim, bias=bias), - nn.GELU(approximate="tanh"), - nn.Dropout(dropout), - nn.Linear(inner_dim, dim, bias=bias), - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return self.ff(hidden_states) - - -@maybe_allow_in_graph -class AudioDiTBlock(nn.Module): - def __init__( - self, - dim: int, - cond_dim: int, - heads: int, - dim_head: int, - dropout: float = 0.0, - bias: bool = True, - qk_norm: bool = False, - eps: float = 1e-6, - cross_attn: bool = True, - cross_attn_norm: bool = False, - adaln_type: str = "global", - adaln_use_text_cond: bool = True, - ff_mult: float = 4.0, - ): - super().__init__() - self.adaln_type = adaln_type - self.adaln_use_text_cond = adaln_use_text_cond - if adaln_type == "local": - self.adaln_mlp = AudioDiTAdaLNMLP(dim, dim * 6, bias=True) - elif adaln_type == "global": - self.adaln_scale_shift = nn.Parameter(torch.randn(dim * 6) / dim**0.5) - - self.self_attn = AudioDiTAttention( - dim, None, heads, dim_head, dropout=dropout, bias=bias, qk_norm=qk_norm, eps=eps - ) - - self.use_cross_attn = cross_attn - if cross_attn: - self.cross_attn = AudioDiTAttention( - dim, - cond_dim, - heads, - dim_head, - dropout=dropout, - bias=bias, - qk_norm=qk_norm, - eps=eps, - processor=AudioDiTCrossAttnProcessor(), - ) - self.cross_attn_norm = ( - nn.LayerNorm(dim, elementwise_affine=True, eps=eps) if cross_attn_norm else nn.Identity() - ) - self.cross_attn_norm_c = ( - nn.LayerNorm(cond_dim, elementwise_affine=True, eps=eps) if cross_attn_norm else nn.Identity() - ) - self.ffn = AudioDiTFeedForward(dim=dim, mult=ff_mult, dropout=dropout, bias=bias) - - def forward( - self, - hidden_states: torch.Tensor, - timestep_embed: torch.Tensor, - cond: torch.Tensor, - mask: torch.BoolTensor | None = None, - cond_mask: torch.BoolTensor | None = None, - rope: tuple | None = None, - cond_rope: tuple | None = None, - adaln_global_out: torch.Tensor | None = None, - ) -> torch.Tensor: - if self.adaln_type == "local" and adaln_global_out is None: - if self.adaln_use_text_cond: - denom = cond_mask.sum(1, keepdim=True).clamp(min=1).to(cond.dtype) - cond_mean = cond.sum(1) / denom - norm_cond = timestep_embed + cond_mean - else: - norm_cond = timestep_embed - adaln_out = self.adaln_mlp(norm_cond) - gate_sa, scale_sa, shift_sa, gate_ffn, scale_ffn, shift_ffn = torch.chunk(adaln_out, 6, dim=-1) - else: - adaln_out = adaln_global_out + self.adaln_scale_shift.unsqueeze(0) - gate_sa, scale_sa, shift_sa, gate_ffn, scale_ffn, shift_ffn = torch.chunk(adaln_out, 6, dim=-1) - - norm_hidden_states = F.layer_norm(hidden_states.float(), (hidden_states.shape[-1],), eps=1e-6).type_as( - hidden_states - ) - norm_hidden_states = norm_hidden_states * (1 + scale_sa[:, None]) + shift_sa[:, None] - attn_output = self.self_attn( - norm_hidden_states, - attention_mask=mask, - audio_rotary_emb=rope, - ) - hidden_states = hidden_states + gate_sa.unsqueeze(1) * attn_output - - if self.use_cross_attn: - cross_output = self.cross_attn( - hidden_states=self.cross_attn_norm(hidden_states), - encoder_hidden_states=self.cross_attn_norm_c(cond), - post_attention_mask=mask, - attention_mask=cond_mask, - audio_rotary_emb=rope, - prompt_rotary_emb=cond_rope, - ) - hidden_states = hidden_states + cross_output - - norm_hidden_states = F.layer_norm(hidden_states.float(), (hidden_states.shape[-1],), eps=1e-6).type_as( - hidden_states - ) - norm_hidden_states = norm_hidden_states * (1 + scale_ffn[:, None]) + shift_ffn[:, None] - ff_output = self.ffn(norm_hidden_states) - hidden_states = hidden_states + gate_ffn.unsqueeze(1) * ff_output - return hidden_states - - -class LongCatAudioDiTTransformer(ModelMixin, ConfigMixin): - _supports_gradient_checkpointing = False - _repeated_blocks = ["AudioDiTBlock"] - - @register_to_config - def __init__( - self, - dit_dim: int = 1536, - dit_depth: int = 24, - dit_heads: int = 24, - dit_text_dim: int = 768, - latent_dim: int = 64, - dropout: float = 0.0, - bias: bool = True, - cross_attn: bool = True, - adaln_type: str = "global", - adaln_use_text_cond: bool = True, - long_skip: bool = True, - text_conv: bool = True, - qk_norm: bool = True, - cross_attn_norm: bool = False, - eps: float = 1e-6, - use_latent_condition: bool = True, - ff_mult: float = 4.0, - ): - super().__init__() - dim = dit_dim - dim_head = dim // dit_heads - self.time_embed = AudioDiTTimestepEmbedding(dim) - self.input_embed = AudioDiTEmbedder(latent_dim, dim) - self.text_embed = AudioDiTEmbedder(dit_text_dim, dim) - self.rotary_embed = AudioDiTRotaryEmbedding(dim_head, 2048, base=100000.0) - self.blocks = nn.ModuleList( - [ - AudioDiTBlock( - dim=dim, - cond_dim=dim, - heads=dit_heads, - dim_head=dim_head, - dropout=dropout, - bias=bias, - qk_norm=qk_norm, - eps=eps, - cross_attn=cross_attn, - cross_attn_norm=cross_attn_norm, - adaln_type=adaln_type, - adaln_use_text_cond=adaln_use_text_cond, - ff_mult=ff_mult, - ) - for _ in range(dit_depth) - ] - ) - self.norm_out = AudioDiTAdaLayerNormZeroFinal(dim, bias=bias, eps=eps) - self.proj_out = nn.Linear(dim, latent_dim) - if adaln_type == "global": - self.adaln_global_mlp = AudioDiTAdaLNMLP(dim, dim * 6, bias=True) - self.text_conv = text_conv - if text_conv: - self.text_conv_layer = nn.Sequential( - *[AudioDiTConvNeXtV2Block(dim, dim * 2, bias=bias, eps=eps) for _ in range(4)] - ) - self.use_latent_condition = use_latent_condition - if use_latent_condition: - self.latent_embed = AudioDiTEmbedder(latent_dim, dim) - self.latent_cond_embedder = AudioDiTEmbedder(dim * 2, dim) - self._initialize_weights(bias=bias) - - def _initialize_weights(self, bias: bool = True): - if self.config.adaln_type == "local": - for block in self.blocks: - nn.init.constant_(block.adaln_mlp.mlp[-1].weight, 0) - if bias: - nn.init.constant_(block.adaln_mlp.mlp[-1].bias, 0) - elif self.config.adaln_type == "global": - nn.init.constant_(self.adaln_global_mlp.mlp[-1].weight, 0) - if bias: - nn.init.constant_(self.adaln_global_mlp.mlp[-1].bias, 0) - nn.init.constant_(self.norm_out.linear.weight, 0) - nn.init.constant_(self.proj_out.weight, 0) - if bias: - nn.init.constant_(self.norm_out.linear.bias, 0) - nn.init.constant_(self.proj_out.bias, 0) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_attention_mask: torch.BoolTensor, - timestep: torch.Tensor, - attention_mask: torch.BoolTensor | None = None, - latent_cond: torch.Tensor | None = None, - return_dict: bool = True, - ) -> LongCatAudioDiTTransformerOutput | tuple[torch.Tensor]: - """ - The [`LongCatAudioDiTTransformer`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, sequence_length, in_channels)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_attention_mask (`torch.BoolTensor`): - Mask applied to `encoder_hidden_states` during attention. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - attention_mask (`torch.BoolTensor`, *optional*): - Mask applied to `hidden_states` during self-attention. - latent_cond (`torch.Tensor`, *optional*): - Latent conditioning concatenated to `hidden_states`. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`LongCatAudioDiTTransformerOutput`] instead of a plain tuple. - - Returns: - [`LongCatAudioDiTTransformerOutput`] or `tuple`: - If `return_dict` is True, a [`LongCatAudioDiTTransformerOutput`] is returned, otherwise a plain `tuple` - is returned. - """ - dtype = hidden_states.dtype - encoder_hidden_states = encoder_hidden_states.to(dtype) - timestep = timestep.to(dtype) - batch_size = hidden_states.shape[0] - if timestep.ndim == 0: - timestep = timestep.repeat(batch_size) - timestep_embed = self.time_embed(timestep) - text_mask = encoder_attention_mask.bool() - encoder_hidden_states = self.text_embed(encoder_hidden_states, text_mask) - if self.text_conv: - encoder_hidden_states = self.text_conv_layer(encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states.masked_fill(text_mask.logical_not().unsqueeze(-1), 0.0) - hidden_states = self.input_embed(hidden_states, attention_mask) - if self.use_latent_condition and latent_cond is not None: - latent_cond = self.latent_embed(latent_cond.to(hidden_states.dtype), attention_mask) - hidden_states = self.latent_cond_embedder(torch.cat([hidden_states, latent_cond], dim=-1)) - residual = hidden_states.clone() if self.config.long_skip else None - rope = self.rotary_embed(hidden_states, hidden_states.shape[1]) - cond_rope = self.rotary_embed(encoder_hidden_states, encoder_hidden_states.shape[1]) - if self.config.adaln_type == "global": - if self.config.adaln_use_text_cond: - text_len = text_mask.sum(1).clamp(min=1).to(encoder_hidden_states.dtype) - text_mean = encoder_hidden_states.sum(1) / text_len.unsqueeze(1) - norm_cond = timestep_embed + text_mean - else: - norm_cond = timestep_embed - adaln_global_out = self.adaln_global_mlp(norm_cond) - for block in self.blocks: - hidden_states = block( - hidden_states=hidden_states, - timestep_embed=timestep_embed, - cond=encoder_hidden_states, - mask=attention_mask, - cond_mask=text_mask, - rope=rope, - cond_rope=cond_rope, - adaln_global_out=adaln_global_out, - ) - else: - norm_cond = timestep_embed - for block in self.blocks: - hidden_states = block( - hidden_states=hidden_states, - timestep_embed=timestep_embed, - cond=encoder_hidden_states, - mask=attention_mask, - cond_mask=text_mask, - rope=rope, - cond_rope=cond_rope, - ) - if self.config.long_skip: - hidden_states = hidden_states + residual - hidden_states = self.norm_out(hidden_states, norm_cond) - hidden_states = self.proj_out(hidden_states) - if attention_mask is not None: - hidden_states = hidden_states * attention_mask.unsqueeze(-1).to(hidden_states.dtype) - if not return_dict: - return (hidden_states,) - return LongCatAudioDiTTransformerOutput(sample=hidden_states) diff --git a/diffusers/models/transformers/transformer_longcat_image.py b/diffusers/models/transformers/transformer_longcat_image.py deleted file mode 100644 index 7b842c42132dce83d301326df9951ba415b32e66..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_longcat_image.py +++ /dev/null @@ -1,548 +0,0 @@ -# Copyright 2025 MeiTuan LongCat-Image Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device, maybe_allow_in_graph -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import TimestepEmbedding, Timesteps, apply_rotary_emb, get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _get_projections(attn: "LongCatImageAttention", hidden_states, encoder_hidden_states=None): - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - encoder_query = encoder_key = encoder_value = None - if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None: - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_fused_projections(attn: "LongCatImageAttention", hidden_states, encoder_hidden_states=None): - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - - encoder_query = encoder_key = encoder_value = (None,) - if encoder_hidden_states is not None and hasattr(attn, "to_added_qkv"): - encoder_query, encoder_key, encoder_value = attn.to_added_qkv(encoder_hidden_states).chunk(3, dim=-1) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_qkv_projections(attn: "LongCatImageAttention", hidden_states, encoder_hidden_states=None): - if attn.fused_projections: - return _get_fused_projections(attn, hidden_states, encoder_hidden_states) - return _get_projections(attn, hidden_states, encoder_hidden_states) - - -class LongCatImageAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError(f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version.") - - def __call__( - self, - attn: "LongCatImageAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - query, key, value, encoder_query, encoder_key, encoder_value = _get_qkv_projections( - attn, hidden_states, encoder_hidden_states - ) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if attn.added_kv_proj_dim is not None: - encoder_query = encoder_query.unflatten(-1, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(-1, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(-1, (attn.heads, -1)) - - encoder_query = attn.norm_added_q(encoder_query) - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - if encoder_hidden_states is not None: - encoder_hidden_states, hidden_states = hidden_states.split_with_sizes( - [encoder_hidden_states.shape[1], hidden_states.shape[1] - encoder_hidden_states.shape[1]], dim=1 - ) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -class LongCatImageAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = LongCatImageAttnProcessor - _available_processors = [ - LongCatImageAttnProcessor, - ] - - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - added_kv_proj_dim: int | None = None, - added_proj_bias: bool | None = True, - out_bias: bool = True, - eps: float = 1e-5, - out_dim: int = None, - context_pre_only: bool | None = None, - pre_only: bool = False, - elementwise_affine: bool = True, - processor=None, - ): - super().__init__() - - self.head_dim = dim_head - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.query_dim = query_dim - self.use_bias = bias - self.dropout = dropout - self.out_dim = out_dim if out_dim is not None else query_dim - self.context_pre_only = context_pre_only - self.pre_only = pre_only - self.heads = out_dim // dim_head if out_dim is not None else heads - self.added_kv_proj_dim = added_kv_proj_dim - self.added_proj_bias = added_proj_bias - - self.norm_q = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_v = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - - if not self.pre_only: - self.to_out = torch.nn.ModuleList([]) - self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(torch.nn.Dropout(dropout)) - - if added_kv_proj_dim is not None: - self.norm_added_q = torch.nn.RMSNorm(dim_head, eps=eps) - self.norm_added_k = torch.nn.RMSNorm(dim_head, eps=eps) - self.add_q_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.to_add_out = torch.nn.Linear(self.inner_dim, query_dim, bias=out_bias) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - quiet_attn_parameters = {"ip_adapter_masks", "ip_hidden_states"} - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters and k not in quiet_attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"joint_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -@maybe_allow_in_graph -class LongCatImageSingleTransformerBlock(nn.Module): - def __init__(self, dim: int, num_attention_heads: int, attention_head_dim: int, mlp_ratio: float = 4.0): - super().__init__() - self.mlp_hidden_dim = int(dim * mlp_ratio) - - self.norm = AdaLayerNormZeroSingle(dim) - self.proj_mlp = nn.Linear(dim, self.mlp_hidden_dim) - self.act_mlp = nn.GELU(approximate="tanh") - self.proj_out = nn.Linear(dim + self.mlp_hidden_dim, dim) - - self.attn = LongCatImageAttention( - query_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - bias=True, - processor=LongCatImageAttnProcessor(), - eps=1e-6, - pre_only=True, - ) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - text_seq_len = encoder_hidden_states.shape[1] - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - residual = hidden_states - norm_hidden_states, gate = self.norm(hidden_states, emb=temb) - mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) - joint_attention_kwargs = joint_attention_kwargs or {} - attn_output = self.attn( - hidden_states=norm_hidden_states, - image_rotary_emb=image_rotary_emb, - **joint_attention_kwargs, - ) - - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) - gate = gate.unsqueeze(1) - hidden_states = gate * self.proj_out(hidden_states) - hidden_states = residual + hidden_states - if hidden_states.dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - encoder_hidden_states, hidden_states = hidden_states[:, :text_seq_len], hidden_states[:, text_seq_len:] - return encoder_hidden_states, hidden_states - - -@maybe_allow_in_graph -class LongCatImageTransformerBlock(nn.Module): - def __init__( - self, dim: int, num_attention_heads: int, attention_head_dim: int, qk_norm: str = "rms_norm", eps: float = 1e-6 - ): - super().__init__() - - self.norm1 = AdaLayerNormZero(dim) - self.norm1_context = AdaLayerNormZero(dim) - - self.attn = LongCatImageAttention( - query_dim=dim, - added_kv_proj_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - context_pre_only=False, - bias=True, - processor=LongCatImageAttnProcessor(), - eps=eps, - ) - - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff_context = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - joint_attention_kwargs = joint_attention_kwargs or {} - - # Attention. - attention_outputs = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - **joint_attention_kwargs, - ) - - if len(attention_outputs) == 2: - attn_output, context_attn_output = attention_outputs - elif len(attention_outputs) == 3: - attn_output, context_attn_output, ip_attn_output = attention_outputs - - # Process attention outputs for the `hidden_states`. - attn_output = gate_msa.unsqueeze(1) * attn_output - hidden_states = hidden_states + attn_output - - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - - ff_output = self.ff(norm_hidden_states) - ff_output = gate_mlp.unsqueeze(1) * ff_output - - hidden_states = hidden_states + ff_output - if len(attention_outputs) == 3: - hidden_states = hidden_states + ip_attn_output - - # Process attention outputs for the `encoder_hidden_states`. - context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output - encoder_hidden_states = encoder_hidden_states + context_attn_output - - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - context_ff_output = self.ff_context(norm_encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - if encoder_hidden_states.dtype == torch.float16: - encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) - - return encoder_hidden_states, hidden_states - - -class LongCatImagePosEmbed(nn.Module): - def __init__(self, theta: int, axes_dim: list[int]): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - - def forward(self, ids: torch.Tensor) -> torch.Tensor: - n_axes = ids.shape[-1] - cos_out = [] - sin_out = [] - pos = ids.float() - freqs_dtype = maybe_adjust_dtype_for_device(torch.float64, ids.device) - for i in range(n_axes): - cos, sin = get_1d_rotary_pos_embed( - self.axes_dim[i], - pos[:, i], - theta=self.theta, - repeat_interleave_real=True, - use_real=True, - freqs_dtype=freqs_dtype, - ) - cos_out.append(cos) - sin_out.append(sin) - freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device) - freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device) - return freqs_cos, freqs_sin - - -class LongCatImageTimestepEmbeddings(nn.Module): - def __init__(self, embedding_dim): - super().__init__() - - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - def forward(self, timestep, hidden_dtype): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) - - return timesteps_emb - - -class LongCatImageTransformer2DModel( - ModelMixin, - ConfigMixin, - PeftAdapterMixin, - FromOriginalModelMixin, - CacheMixin, - AttentionMixin, -): - """ - The Transformer model introduced in Longcat-Image. - """ - - _supports_gradient_checkpointing = True - _repeated_blocks = ["LongCatImageTransformerBlock", "LongCatImageSingleTransformerBlock"] - - @register_to_config - def __init__( - self, - patch_size: int = 1, - in_channels: int = 64, - num_layers: int = 19, - num_single_layers: int = 38, - attention_head_dim: int = 128, - num_attention_heads: int = 24, - joint_attention_dim: int = 3584, - pooled_projection_dim: int = 3584, - axes_dims_rope: list[int] = [16, 56, 56], - ): - super().__init__() - self.out_channels = in_channels - self.inner_dim = num_attention_heads * attention_head_dim - self.pooled_projection_dim = pooled_projection_dim - - self.pos_embed = LongCatImagePosEmbed(theta=10000, axes_dim=axes_dims_rope) - - self.time_embed = LongCatImageTimestepEmbeddings(embedding_dim=self.inner_dim) - - self.context_embedder = nn.Linear(joint_attention_dim, self.inner_dim) - self.x_embedder = torch.nn.Linear(in_channels, self.inner_dim) - - self.transformer_blocks = nn.ModuleList( - [ - LongCatImageTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for i in range(num_layers) - ] - ) - - self.single_transformer_blocks = nn.ModuleList( - [ - LongCatImageSingleTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for i in range(num_single_layers) - ] - ) - - self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) - - self.gradient_checkpointing = False - self.use_checkpoint = [True] * num_layers - self.use_single_checkpoint = [True] * num_single_layers - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_ids: torch.Tensor = None, - txt_ids: torch.Tensor = None, - guidance: torch.Tensor = None, - return_dict: bool = True, - ) -> torch.FloatTensor | Transformer2DModelOutput: - """ - The forward method. - - Args: - hidden_states (`torch.FloatTensor` of shape `(batch size, channel, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.FloatTensor` of shape `(batch size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep ( `torch.LongTensor`): - Used to indicate denoising step. - img_ids (`torch.Tensor`): - Image position ids used to compute the rotary positional embeddings. - txt_ids (`torch.Tensor`): - Text position ids used to compute the rotary positional embeddings. - guidance (`torch.Tensor`, *optional*): - Guidance scale embedding used for guidance-distilled variants of the model. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - hidden_states = self.x_embedder(hidden_states) - - timestep = timestep.to(hidden_states.dtype) * 1000 - - temb = self.time_embed(timestep, hidden_states.dtype) - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - - ids = torch.cat((txt_ids, img_ids), dim=0) - image_rotary_emb = self.pos_embed(ids) - - for index_block, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing and self.use_checkpoint[index_block]: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - ) - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - ) - - for index_block, block in enumerate(self.single_transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing and self.use_single_checkpoint[index_block]: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - ) - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - ) - - hidden_states = self.norm_out(hidden_states, temb) - output = self.proj_out(hidden_states) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_ltx.py b/diffusers/models/transformers/transformer_ltx.py deleted file mode 100644 index c33e0f6141fc9e5b29d622f4a6c5072a8c0ee638..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_ltx.py +++ /dev/null @@ -1,601 +0,0 @@ -# Copyright 2025 The Lightricks team and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -import math -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, deprecate, is_torch_version, logging -from ...utils.torch_utils import maybe_allow_in_graph -from .._modeling_parallel import ContextParallelInput, ContextParallelOutput -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import PixArtAlphaTextProjection -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormSingle, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class LTXVideoAttentionProcessor2_0: - def __new__(cls, *args, **kwargs): - deprecation_message = "`LTXVideoAttentionProcessor2_0` is deprecated and this will be removed in a future version. Please use `LTXVideoAttnProcessor`" - deprecate("LTXVideoAttentionProcessor2_0", "1.0.0", deprecation_message) - - return LTXVideoAttnProcessor(*args, **kwargs) - - -class LTXVideoAttnProcessor: - r""" - Processor for implementing attention (SDPA is used by default if you're using PyTorch 2.0). This is used in the LTX - model. It applies a normalization layer and rotary embedding on the query and key vector. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if is_torch_version("<", "2.0"): - raise ValueError( - "LTX attention processors require a minimum PyTorch version of 2.0. Please upgrade your PyTorch installation." - ) - - def __call__( - self, - attn: "LTXAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb) - key = apply_rotary_emb(key, image_rotary_emb) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class LTXAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = LTXVideoAttnProcessor - _available_processors = [LTXVideoAttnProcessor] - - def __init__( - self, - query_dim: int, - heads: int = 8, - kv_heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = True, - cross_attention_dim: int | None = None, - out_bias: bool = True, - qk_norm: str = "rms_norm_across_heads", - processor=None, - ): - super().__init__() - if qk_norm != "rms_norm_across_heads": - raise NotImplementedError("Only 'rms_norm_across_heads' is supported as a valid value for `qk_norm`.") - - self.head_dim = dim_head - self.inner_dim = dim_head * heads - self.inner_kv_dim = self.inner_dim if kv_heads is None else dim_head * kv_heads - self.query_dim = query_dim - self.cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim - self.use_bias = bias - self.dropout = dropout - self.out_dim = query_dim - self.heads = heads - - norm_eps = 1e-5 - norm_elementwise_affine = True - self.norm_q = torch.nn.RMSNorm(dim_head * heads, eps=norm_eps, elementwise_affine=norm_elementwise_affine) - self.norm_k = torch.nn.RMSNorm(dim_head * kv_heads, eps=norm_eps, elementwise_affine=norm_elementwise_affine) - self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = torch.nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias) - self.to_v = torch.nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias) - self.to_out = torch.nn.ModuleList([]) - self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(torch.nn.Dropout(dropout)) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -class LTXVideoRotaryPosEmbed(nn.Module): - def __init__( - self, - dim: int, - base_num_frames: int = 20, - base_height: int = 2048, - base_width: int = 2048, - patch_size: int = 1, - patch_size_t: int = 1, - theta: float = 10000.0, - ) -> None: - super().__init__() - - self.dim = dim - self.base_num_frames = base_num_frames - self.base_height = base_height - self.base_width = base_width - self.patch_size = patch_size - self.patch_size_t = patch_size_t - self.theta = theta - - def _prepare_video_coords( - self, - batch_size: int, - num_frames: int, - height: int, - width: int, - rope_interpolation_scale: tuple[torch.Tensor, float, float], - device: torch.device, - ) -> torch.Tensor: - # Always compute rope in fp32 - grid_h = torch.arange(height, dtype=torch.float32, device=device) - grid_w = torch.arange(width, dtype=torch.float32, device=device) - grid_f = torch.arange(num_frames, dtype=torch.float32, device=device) - grid = torch.meshgrid(grid_f, grid_h, grid_w, indexing="ij") - grid = torch.stack(grid, dim=0) - grid = grid.unsqueeze(0).repeat(batch_size, 1, 1, 1, 1) - - if rope_interpolation_scale is not None: - grid[:, 0:1] = grid[:, 0:1] * rope_interpolation_scale[0] * self.patch_size_t / self.base_num_frames - grid[:, 1:2] = grid[:, 1:2] * rope_interpolation_scale[1] * self.patch_size / self.base_height - grid[:, 2:3] = grid[:, 2:3] * rope_interpolation_scale[2] * self.patch_size / self.base_width - - grid = grid.flatten(2, 4).transpose(1, 2) - - return grid - - def forward( - self, - hidden_states: torch.Tensor, - num_frames: int | None = None, - height: int | None = None, - width: int | None = None, - rope_interpolation_scale: tuple[torch.Tensor, float, float] | None = None, - video_coords: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - batch_size = hidden_states.size(0) - - if video_coords is None: - grid = self._prepare_video_coords( - batch_size, - num_frames, - height, - width, - rope_interpolation_scale=rope_interpolation_scale, - device=hidden_states.device, - ) - else: - grid = torch.stack( - [ - video_coords[:, 0] / self.base_num_frames, - video_coords[:, 1] / self.base_height, - video_coords[:, 2] / self.base_width, - ], - dim=-1, - ) - - start = 1.0 - end = self.theta - freqs = self.theta ** torch.linspace( - math.log(start, self.theta), - math.log(end, self.theta), - self.dim // 6, - device=hidden_states.device, - dtype=torch.float32, - ) - freqs = freqs * math.pi / 2.0 - freqs = freqs * (grid.unsqueeze(-1) * 2 - 1) - freqs = freqs.transpose(-1, -2).flatten(2) - - cos_freqs = freqs.cos().repeat_interleave(2, dim=-1) - sin_freqs = freqs.sin().repeat_interleave(2, dim=-1) - - if self.dim % 6 != 0: - cos_padding = torch.ones_like(cos_freqs[:, :, : self.dim % 6]) - sin_padding = torch.zeros_like(cos_freqs[:, :, : self.dim % 6]) - cos_freqs = torch.cat([cos_padding, cos_freqs], dim=-1) - sin_freqs = torch.cat([sin_padding, sin_freqs], dim=-1) - - return cos_freqs, sin_freqs - - -@maybe_allow_in_graph -class LTXVideoTransformerBlock(nn.Module): - r""" - Transformer block used in [LTX](https://huggingface.co/Lightricks/LTX-Video). - - Args: - dim (`int`): - The number of channels in the input and output. - num_attention_heads (`int`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`): - The number of channels in each head. - qk_norm (`str`, defaults to `"rms_norm"`): - The normalization layer to use. - activation_fn (`str`, defaults to `"gelu-approximate"`): - Activation function to use in feed-forward. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - cross_attention_dim: int, - qk_norm: str = "rms_norm_across_heads", - activation_fn: str = "gelu-approximate", - attention_bias: bool = True, - attention_out_bias: bool = True, - eps: float = 1e-6, - elementwise_affine: bool = False, - ): - super().__init__() - - self.norm1 = RMSNorm(dim, eps=eps, elementwise_affine=elementwise_affine) - self.attn1 = LTXAttention( - query_dim=dim, - heads=num_attention_heads, - kv_heads=num_attention_heads, - dim_head=attention_head_dim, - bias=attention_bias, - cross_attention_dim=None, - out_bias=attention_out_bias, - qk_norm=qk_norm, - ) - - self.norm2 = RMSNorm(dim, eps=eps, elementwise_affine=elementwise_affine) - self.attn2 = LTXAttention( - query_dim=dim, - cross_attention_dim=cross_attention_dim, - heads=num_attention_heads, - kv_heads=num_attention_heads, - dim_head=attention_head_dim, - bias=attention_bias, - out_bias=attention_out_bias, - qk_norm=qk_norm, - ) - - self.ff = FeedForward(dim, activation_fn=activation_fn) - - self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - batch_size = hidden_states.size(0) - norm_hidden_states = self.norm1(hidden_states) - - num_ada_params = self.scale_shift_table.shape[0] - ada_values = self.scale_shift_table[None, None].to(temb.device) + temb.reshape( - batch_size, temb.size(1), num_ada_params, -1 - ) - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ada_values.unbind(dim=2) - norm_hidden_states = norm_hidden_states * (1 + scale_msa) + shift_msa - - attn_hidden_states = self.attn1( - hidden_states=norm_hidden_states, - encoder_hidden_states=None, - image_rotary_emb=image_rotary_emb, - ) - hidden_states = hidden_states + attn_hidden_states * gate_msa - - attn_hidden_states = self.attn2( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - image_rotary_emb=None, - attention_mask=encoder_attention_mask, - ) - hidden_states = hidden_states + attn_hidden_states - norm_hidden_states = self.norm2(hidden_states) * (1 + scale_mlp) + shift_mlp - - ff_output = self.ff(norm_hidden_states) - hidden_states = hidden_states + ff_output * gate_mlp - - return hidden_states - - -@maybe_allow_in_graph -class LTXVideoTransformer3DModel( - ModelMixin, ConfigMixin, AttentionMixin, FromOriginalModelMixin, PeftAdapterMixin, CacheMixin -): - r""" - A Transformer model for video-like data used in [LTX](https://huggingface.co/Lightricks/LTX-Video). - - Args: - in_channels (`int`, defaults to `128`): - The number of channels in the input. - out_channels (`int`, defaults to `128`): - The number of channels in the output. - patch_size (`int`, defaults to `1`): - The size of the spatial patches to use in the patch embedding layer. - patch_size_t (`int`, defaults to `1`): - The size of the tmeporal patches to use in the patch embedding layer. - num_attention_heads (`int`, defaults to `32`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `64`): - The number of channels in each head. - cross_attention_dim (`int`, defaults to `2048 `): - The number of channels for cross attention heads. - num_layers (`int`, defaults to `28`): - The number of layers of Transformer blocks to use. - activation_fn (`str`, defaults to `"gelu-approximate"`): - Activation function to use in feed-forward. - qk_norm (`str`, defaults to `"rms_norm_across_heads"`): - The normalization layer to use. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["norm"] - _repeated_blocks = ["LTXVideoTransformerBlock"] - _cp_plan = { - "": { - "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - "encoder_hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - "encoder_attention_mask": ContextParallelInput(split_dim=1, expected_dims=2, split_output=False), - }, - "rope": { - 0: ContextParallelInput(split_dim=1, expected_dims=3, split_output=True), - 1: ContextParallelInput(split_dim=1, expected_dims=3, split_output=True), - }, - "proj_out": ContextParallelOutput(gather_dim=1, expected_dims=3), - } - - @register_to_config - def __init__( - self, - in_channels: int = 128, - out_channels: int = 128, - patch_size: int = 1, - patch_size_t: int = 1, - num_attention_heads: int = 32, - attention_head_dim: int = 64, - cross_attention_dim: int = 2048, - num_layers: int = 28, - activation_fn: str = "gelu-approximate", - qk_norm: str = "rms_norm_across_heads", - norm_elementwise_affine: bool = False, - norm_eps: float = 1e-6, - caption_channels: int = 4096, - attention_bias: bool = True, - attention_out_bias: bool = True, - ) -> None: - super().__init__() - - out_channels = out_channels or in_channels - inner_dim = num_attention_heads * attention_head_dim - - self.proj_in = nn.Linear(in_channels, inner_dim) - - self.scale_shift_table = nn.Parameter(torch.randn(2, inner_dim) / inner_dim**0.5) - self.time_embed = AdaLayerNormSingle(inner_dim, use_additional_conditions=False) - - self.caption_projection = PixArtAlphaTextProjection(in_features=caption_channels, hidden_size=inner_dim) - - self.rope = LTXVideoRotaryPosEmbed( - dim=inner_dim, - base_num_frames=20, - base_height=2048, - base_width=2048, - patch_size=patch_size, - patch_size_t=patch_size_t, - theta=10000.0, - ) - - self.transformer_blocks = nn.ModuleList( - [ - LTXVideoTransformerBlock( - dim=inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - cross_attention_dim=cross_attention_dim, - qk_norm=qk_norm, - activation_fn=activation_fn, - attention_bias=attention_bias, - attention_out_bias=attention_out_bias, - eps=norm_eps, - elementwise_affine=norm_elementwise_affine, - ) - for _ in range(num_layers) - ] - ) - - self.norm_out = nn.LayerNorm(inner_dim, eps=1e-6, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, out_channels) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_attention_mask: torch.Tensor, - num_frames: int | None = None, - height: int | None = None, - width: int | None = None, - rope_interpolation_scale: tuple[float, float, float] | torch.Tensor | None = None, - video_coords: torch.Tensor | None = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> torch.Tensor: - """ - The [`LTXVideoTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, sequence_length, in_channels)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_attention_mask (`torch.Tensor`): - Mask applied to `encoder_hidden_states` during attention. - num_frames (`int`, *optional*): - Number of frames in the video used to compute the rotary positional embeddings. - height (`int`, *optional*): - Height of the latent used to compute the rotary positional embeddings. - width (`int`, *optional*): - Width of the latent used to compute the rotary positional embeddings. - rope_interpolation_scale (`tuple` of `float` or `torch.Tensor`, *optional*): - Interpolation scale used by the rotary positional embeddings. - video_coords (`torch.Tensor`, *optional*): - Pre-computed video coordinates used by the rotary positional embeddings. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - `torch.Tensor`: - The denoised output tensor of shape `(batch_size, sequence_length, out_channels)`. - """ - image_rotary_emb = self.rope(hidden_states, num_frames, height, width, rope_interpolation_scale, video_coords) - - # convert encoder_attention_mask to a bias the same way we do for attention_mask - if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2: - encoder_attention_mask = (1 - encoder_attention_mask.to(hidden_states.dtype)) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - batch_size = hidden_states.size(0) - hidden_states = self.proj_in(hidden_states) - - temb, embedded_timestep = self.time_embed( - timestep.flatten(), - batch_size=batch_size, - hidden_dtype=hidden_states.dtype, - ) - - temb = temb.view(batch_size, -1, temb.size(-1)) - embedded_timestep = embedded_timestep.view(batch_size, -1, embedded_timestep.size(-1)) - - encoder_hidden_states = self.caption_projection(encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states.view(batch_size, -1, hidden_states.size(-1)) - - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - encoder_attention_mask, - ) - else: - hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - encoder_attention_mask=encoder_attention_mask, - ) - - scale_shift_values = self.scale_shift_table[None, None] + embedded_timestep[:, :, None] - shift, scale = scale_shift_values[:, :, 0], scale_shift_values[:, :, 1] - - hidden_states = self.norm_out(hidden_states) - hidden_states = hidden_states * (1 + scale) + shift - output = self.proj_out(hidden_states) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) - - -def apply_rotary_emb(x, freqs): - cos, sin = freqs - x_real, x_imag = x.unflatten(2, (-1, 2)).unbind(-1) # [B, S, C // 2] - x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(2) - out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype) - return out diff --git a/diffusers/models/transformers/transformer_ltx2.py b/diffusers/models/transformers/transformer_ltx2.py deleted file mode 100644 index 465408d946938b07954f6e7635dbbacf9768990b..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_ltx2.py +++ /dev/null @@ -1,1639 +0,0 @@ -# Copyright 2025 The Lightricks team and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -from dataclasses import dataclass -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import BaseOutput, apply_lora_scale, is_torch_version, logging -from .._modeling_parallel import ContextParallelInput, ContextParallelOutput -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import PixArtAlphaCombinedTimestepSizeEmbeddings, PixArtAlphaTextProjection -from ..modeling_utils import ModelMixin -from ..normalization import RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def apply_interleaved_rotary_emb(x: torch.Tensor, freqs: tuple[torch.Tensor, torch.Tensor]) -> torch.Tensor: - cos, sin = freqs - x_real, x_imag = x.unflatten(2, (-1, 2)).unbind(-1) # [B, S, C // 2] - x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(2) - out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype) - return out - - -def apply_split_rotary_emb(x: torch.Tensor, freqs: tuple[torch.Tensor, torch.Tensor]) -> torch.Tensor: - cos, sin = freqs - - x_dtype = x.dtype - needs_reshape = False - if x.ndim != 4 and cos.ndim == 4: - # cos is (b, h, t, r) -> reshape x to (b, h, t, dim_per_head) - b, h, t, _ = cos.shape - x = x.reshape(b, t, h, -1).swapaxes(1, 2) - needs_reshape = True - - # Split last dim (2*r) into (d=2, r) - last = x.shape[-1] - if last % 2 != 0: - raise ValueError(f"Expected x.shape[-1] to be even for split rotary, got {last}.") - r = last // 2 - - # (..., 2, r) - split_x = x.reshape(*x.shape[:-1], 2, r).float() # Explicitly upcast to float - first_x = split_x[..., :1, :] # (..., 1, r) - second_x = split_x[..., 1:, :] # (..., 1, r) - - cos_u = cos.unsqueeze(-2) # broadcast to (..., 1, r) against (..., 2, r) - sin_u = sin.unsqueeze(-2) - - out = split_x * cos_u - first_out = out[..., :1, :] - second_out = out[..., 1:, :] - - first_out.addcmul_(-sin_u, second_x) - second_out.addcmul_(sin_u, first_x) - - out = out.reshape(*out.shape[:-2], last) - - if needs_reshape: - out = out.swapaxes(1, 2).reshape(b, t, -1) - - out = out.to(dtype=x_dtype) - return out - - -@dataclass -class AudioVisualModelOutput(BaseOutput): - r""" - Holds the output of an audiovisual model which produces both visual (e.g. video) and audio outputs. - - Args: - sample (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - The hidden states output conditioned on the `encoder_hidden_states` input, representing the visual output - of the model. This is typically a video (spatiotemporal) output. - audio_sample (`torch.Tensor` of shape `(batch_size, TODO)`): - The audio output of the audiovisual model. - """ - - sample: "torch.Tensor" # noqa: F821 - audio_sample: "torch.Tensor" # noqa: F821 - - -class LTX2AdaLayerNormSingle(nn.Module): - r""" - Norm layer adaptive layer norm single (adaLN-single). - - As proposed in PixArt-Alpha (see: https://huggingface.co/papers/2310.00426; Section 2.3) and adapted by the LTX-2.0 - model. In particular, the number of modulation parameters to be calculated is now configurable. - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - num_mod_params (`int`, *optional*, defaults to `6`): - The number of modulation parameters which will be calculated in the first return argument. The default of 6 - is standard, but sometimes we may want to have a different (usually smaller) number of modulation - parameters. - use_additional_conditions (`bool`, *optional*, defaults to `False`): - Whether to use additional conditions for normalization or not. - """ - - def __init__(self, embedding_dim: int, num_mod_params: int = 6, use_additional_conditions: bool = False): - super().__init__() - self.num_mod_params = num_mod_params - - self.emb = PixArtAlphaCombinedTimestepSizeEmbeddings( - embedding_dim, size_emb_dim=embedding_dim // 3, use_additional_conditions=use_additional_conditions - ) - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, self.num_mod_params * embedding_dim, bias=True) - - def forward( - self, - timestep: torch.Tensor, - added_cond_kwargs: dict[str, torch.Tensor] | None = None, - batch_size: int | None = None, - hidden_dtype: torch.dtype | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - # No modulation happening here. - added_cond_kwargs = added_cond_kwargs or {"resolution": None, "aspect_ratio": None} - embedded_timestep = self.emb(timestep, **added_cond_kwargs, batch_size=batch_size, hidden_dtype=hidden_dtype) - return self.linear(self.silu(embedded_timestep)), embedded_timestep - - -class LTX2AudioVideoAttnProcessor: - r""" - Processor for implementing attention (SDPA is used by default if you're using PyTorch 2.0) for the LTX-2.0 model. - Compared to the LTX-1.0 model, we allow the RoPE embeddings for the queries and keys to be separate so that we can - support audio-to-video (a2v) and video-to-audio (v2a) cross attention. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if is_torch_version("<", "2.0"): - raise ValueError( - "LTX attention processors require a minimum PyTorch version of 2.0. Please upgrade your PyTorch installation." - ) - - def __call__( - self, - attn: "LTX2Attention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - query_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - key_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - if attn.to_gate_logits is not None: - # Calculate gate logits on original hidden_states - gate_logits = attn.to_gate_logits(hidden_states) - - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if query_rotary_emb is not None: - if attn.rope_type == "interleaved": - query = apply_interleaved_rotary_emb(query, query_rotary_emb) - key = apply_interleaved_rotary_emb( - key, key_rotary_emb if key_rotary_emb is not None else query_rotary_emb - ) - elif attn.rope_type == "split": - query = apply_split_rotary_emb(query, query_rotary_emb) - key = apply_split_rotary_emb(key, key_rotary_emb if key_rotary_emb is not None else query_rotary_emb) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - if attn.to_gate_logits is not None: - hidden_states = hidden_states.unflatten(2, (attn.heads, -1)) # [B, T, H, D] - # The factor of 2.0 is so that if the gates logits are zero-initialized the initial gates are all 1 - gates = 2.0 * torch.sigmoid(gate_logits) # [B, T, H] - hidden_states = hidden_states * gates.unsqueeze(-1) - hidden_states = hidden_states.flatten(2, 3) - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class LTX2PerturbedAttnProcessor: - r""" - Processor which implements attention with perturbation masking and per-head gating for LTX-2.X models. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if is_torch_version("<", "2.0"): - raise ValueError( - "LTX attention processors require a minimum PyTorch version of 2.0. Please upgrade your PyTorch installation." - ) - - def __call__( - self, - attn: "LTX2Attention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - query_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - key_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - perturbation_mask: torch.Tensor | None = None, - all_perturbed: bool | None = None, - ) -> torch.Tensor: - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - if attn.to_gate_logits is not None: - # Calculate gate logits on original hidden_states - gate_logits = attn.to_gate_logits(hidden_states) - - value = attn.to_v(encoder_hidden_states) - if all_perturbed is None: - all_perturbed = torch.all(perturbation_mask == 0) if perturbation_mask is not None else False - - if all_perturbed: - # Skip attention, use the value projection value - hidden_states = value - else: - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if query_rotary_emb is not None: - if attn.rope_type == "interleaved": - query = apply_interleaved_rotary_emb(query, query_rotary_emb) - key = apply_interleaved_rotary_emb( - key, key_rotary_emb if key_rotary_emb is not None else query_rotary_emb - ) - elif attn.rope_type == "split": - query = apply_split_rotary_emb(query, query_rotary_emb) - key = apply_split_rotary_emb( - key, key_rotary_emb if key_rotary_emb is not None else query_rotary_emb - ) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - if perturbation_mask is not None: - value = value.flatten(2, 3) - hidden_states = torch.lerp(value, hidden_states, perturbation_mask) - - if attn.to_gate_logits is not None: - hidden_states = hidden_states.unflatten(2, (attn.heads, -1)) # [B, T, H, D] - # The factor of 2.0 is so that if the gates logits are zero-initialized the initial gates are all 1 - gates = 2.0 * torch.sigmoid(gate_logits) # [B, T, H] - hidden_states = hidden_states * gates.unsqueeze(-1) - hidden_states = hidden_states.flatten(2, 3) - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class LTX2Attention(torch.nn.Module, AttentionModuleMixin): - r""" - Attention class for all LTX-2.0 attention layers. Compared to LTX-1.0, this supports specifying the query and key - RoPE embeddings separately for audio-to-video (a2v) and video-to-audio (v2a) cross-attention. - """ - - _default_processor_cls = LTX2AudioVideoAttnProcessor - _available_processors = [LTX2AudioVideoAttnProcessor, LTX2PerturbedAttnProcessor] - - def __init__( - self, - query_dim: int, - heads: int = 8, - kv_heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = True, - cross_attention_dim: int | None = None, - out_bias: bool = True, - qk_norm: str = "rms_norm_across_heads", - norm_eps: float = 1e-6, - norm_elementwise_affine: bool = True, - rope_type: str = "interleaved", - apply_gated_attention: bool = False, - processor=None, - ): - super().__init__() - if qk_norm != "rms_norm_across_heads": - raise NotImplementedError("Only 'rms_norm_across_heads' is supported as a valid value for `qk_norm`.") - - self.head_dim = dim_head - self.inner_dim = dim_head * heads - self.inner_kv_dim = self.inner_dim if kv_heads is None else dim_head * kv_heads - self.query_dim = query_dim - self.cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim - self.use_bias = bias - self.dropout = dropout - self.out_dim = query_dim - self.heads = heads - self.rope_type = rope_type - - self.norm_q = torch.nn.RMSNorm(dim_head * heads, eps=norm_eps, elementwise_affine=norm_elementwise_affine) - self.norm_k = torch.nn.RMSNorm(dim_head * kv_heads, eps=norm_eps, elementwise_affine=norm_elementwise_affine) - self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = torch.nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias) - self.to_v = torch.nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias) - self.to_out = torch.nn.ModuleList([]) - self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(torch.nn.Dropout(dropout)) - - if apply_gated_attention: - # Per head gate values - self.to_gate_logits = torch.nn.Linear(query_dim, heads, bias=True) - else: - self.to_gate_logits = None - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - query_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - key_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - hidden_states = self.processor( - self, hidden_states, encoder_hidden_states, attention_mask, query_rotary_emb, key_rotary_emb, **kwargs - ) - return hidden_states - - -class LTX2VideoTransformerBlock(nn.Module): - r""" - Transformer block used in [LTX-2.0](https://huggingface.co/Lightricks/LTX-Video). - - Args: - dim (`int`): - The number of channels in the input and output. - num_attention_heads (`int`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`): - The number of channels in each head. - qk_norm (`str`, defaults to `"rms_norm"`): - The normalization layer to use. - activation_fn (`str`, defaults to `"gelu-approximate"`): - Activation function to use in feed-forward. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - cross_attention_dim: int, - audio_dim: int, - audio_num_attention_heads: int, - audio_attention_head_dim, - audio_cross_attention_dim: int, - video_gated_attn: bool = False, - video_cross_attn_adaln: bool = False, - audio_gated_attn: bool = False, - audio_cross_attn_adaln: bool = False, - qk_norm: str = "rms_norm_across_heads", - activation_fn: str = "gelu-approximate", - attention_bias: bool = True, - attention_out_bias: bool = True, - eps: float = 1e-6, - elementwise_affine: bool = False, - rope_type: str = "interleaved", - perturbed_attn: bool = False, - ): - super().__init__() - - self.perturbed_attn = perturbed_attn - if perturbed_attn: - attn_processor_cls = LTX2PerturbedAttnProcessor - else: - attn_processor_cls = LTX2AudioVideoAttnProcessor - - # 1. Self-Attention (video and audio) - self.norm1 = RMSNorm(dim, eps=eps, elementwise_affine=elementwise_affine) - self.attn1 = LTX2Attention( - query_dim=dim, - heads=num_attention_heads, - kv_heads=num_attention_heads, - dim_head=attention_head_dim, - bias=attention_bias, - cross_attention_dim=None, - out_bias=attention_out_bias, - qk_norm=qk_norm, - rope_type=rope_type, - apply_gated_attention=video_gated_attn, - processor=attn_processor_cls(), - ) - - self.audio_norm1 = RMSNorm(audio_dim, eps=eps, elementwise_affine=elementwise_affine) - self.audio_attn1 = LTX2Attention( - query_dim=audio_dim, - heads=audio_num_attention_heads, - kv_heads=audio_num_attention_heads, - dim_head=audio_attention_head_dim, - bias=attention_bias, - cross_attention_dim=None, - out_bias=attention_out_bias, - qk_norm=qk_norm, - rope_type=rope_type, - apply_gated_attention=audio_gated_attn, - processor=attn_processor_cls(), - ) - - # 2. Prompt Cross-Attention - self.norm2 = RMSNorm(dim, eps=eps, elementwise_affine=elementwise_affine) - self.attn2 = LTX2Attention( - query_dim=dim, - cross_attention_dim=cross_attention_dim, - heads=num_attention_heads, - kv_heads=num_attention_heads, - dim_head=attention_head_dim, - bias=attention_bias, - out_bias=attention_out_bias, - qk_norm=qk_norm, - rope_type=rope_type, - apply_gated_attention=video_gated_attn, - processor=attn_processor_cls(), - ) - - self.audio_norm2 = RMSNorm(audio_dim, eps=eps, elementwise_affine=elementwise_affine) - self.audio_attn2 = LTX2Attention( - query_dim=audio_dim, - cross_attention_dim=audio_cross_attention_dim, - heads=audio_num_attention_heads, - kv_heads=audio_num_attention_heads, - dim_head=audio_attention_head_dim, - bias=attention_bias, - out_bias=attention_out_bias, - qk_norm=qk_norm, - rope_type=rope_type, - apply_gated_attention=audio_gated_attn, - processor=attn_processor_cls(), - ) - - # 3. Audio-to-Video (a2v) and Video-to-Audio (v2a) Cross-Attention - # Audio-to-Video (a2v) Attention --> Q: Video; K,V: Audio - self.audio_to_video_norm = RMSNorm(dim, eps=eps, elementwise_affine=elementwise_affine) - self.audio_to_video_attn = LTX2Attention( - query_dim=dim, - cross_attention_dim=audio_dim, - heads=audio_num_attention_heads, - kv_heads=audio_num_attention_heads, - dim_head=audio_attention_head_dim, - bias=attention_bias, - out_bias=attention_out_bias, - qk_norm=qk_norm, - rope_type=rope_type, - apply_gated_attention=video_gated_attn, - processor=attn_processor_cls(), - ) - - # Video-to-Audio (v2a) Attention --> Q: Audio; K,V: Video - self.video_to_audio_norm = RMSNorm(audio_dim, eps=eps, elementwise_affine=elementwise_affine) - self.video_to_audio_attn = LTX2Attention( - query_dim=audio_dim, - cross_attention_dim=dim, - heads=audio_num_attention_heads, - kv_heads=audio_num_attention_heads, - dim_head=audio_attention_head_dim, - bias=attention_bias, - out_bias=attention_out_bias, - qk_norm=qk_norm, - rope_type=rope_type, - apply_gated_attention=audio_gated_attn, - processor=attn_processor_cls(), - ) - - # 4. Feedforward layers - self.norm3 = RMSNorm(dim, eps=eps, elementwise_affine=elementwise_affine) - self.ff = FeedForward(dim, activation_fn=activation_fn) - - self.audio_norm3 = RMSNorm(audio_dim, eps=eps, elementwise_affine=elementwise_affine) - self.audio_ff = FeedForward(audio_dim, activation_fn=activation_fn) - - # 5. Per-Layer Modulation Parameters - # Self-Attention (attn1) / Feedforward AdaLayerNorm-Zero mod params - # 6 base mod params for text cross-attn K,V; if cross_attn_adaln, also has mod params for Q - self.video_cross_attn_adaln = video_cross_attn_adaln - self.audio_cross_attn_adaln = audio_cross_attn_adaln - video_mod_param_num = 9 if self.video_cross_attn_adaln else 6 - audio_mod_param_num = 9 if self.audio_cross_attn_adaln else 6 - self.scale_shift_table = nn.Parameter(torch.randn(video_mod_param_num, dim) / dim**0.5) - self.audio_scale_shift_table = nn.Parameter(torch.randn(audio_mod_param_num, audio_dim) / audio_dim**0.5) - - # Prompt cross-attn (attn2) additional modulation params - self.cross_attn_adaln = video_cross_attn_adaln or audio_cross_attn_adaln - if self.cross_attn_adaln: - self.prompt_scale_shift_table = nn.Parameter(torch.randn(2, dim)) - self.audio_prompt_scale_shift_table = nn.Parameter(torch.randn(2, audio_dim)) - - # Per-layer a2v, v2a Cross-Attention mod params - self.video_a2v_cross_attn_scale_shift_table = nn.Parameter(torch.randn(5, dim)) - self.audio_a2v_cross_attn_scale_shift_table = nn.Parameter(torch.randn(5, audio_dim)) - - @staticmethod - def get_mod_params( - scale_shift_table: torch.Tensor, temb: torch.Tensor, batch_size: int - ) -> tuple[torch.Tensor, ...]: - num_ada_params = scale_shift_table.shape[0] - ada_values = scale_shift_table[None, None].to(temb.device) + temb.reshape( - batch_size, temb.shape[1], num_ada_params, -1 - ) - ada_params = ada_values.unbind(dim=2) - return ada_params - - def forward( - self, - hidden_states: torch.Tensor, - audio_hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - audio_encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - temb_audio: torch.Tensor, - temb_ca_scale_shift: torch.Tensor, - temb_ca_audio_scale_shift: torch.Tensor, - temb_ca_gate: torch.Tensor, - temb_ca_audio_gate: torch.Tensor, - temb_prompt: torch.Tensor | None = None, - temb_prompt_audio: torch.Tensor | None = None, - video_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - audio_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ca_video_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ca_audio_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - audio_encoder_attention_mask: torch.Tensor | None = None, - self_attention_mask: torch.Tensor | None = None, - audio_self_attention_mask: torch.Tensor | None = None, - a2v_cross_attention_mask: torch.Tensor | None = None, - v2a_cross_attention_mask: torch.Tensor | None = None, - use_a2v_cross_attention: bool = True, - use_v2a_cross_attention: bool = True, - perturbation_mask: torch.Tensor | None = None, - all_perturbed: bool | None = None, - ) -> torch.Tensor: - batch_size = hidden_states.size(0) - - # 1. Video and Audio Self-Attention - # 1.1. Video Self-Attention - video_ada_params = self.get_mod_params(self.scale_shift_table, temb, batch_size) - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = video_ada_params[:6] - if self.video_cross_attn_adaln: - shift_text_q, scale_text_q, gate_text_q = video_ada_params[6:9] - - norm_hidden_states = self.norm1(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_msa) + shift_msa - - video_self_attn_args = { - "hidden_states": norm_hidden_states, - "encoder_hidden_states": None, - "query_rotary_emb": video_rotary_emb, - "attention_mask": self_attention_mask, - } - if self.perturbed_attn: - video_self_attn_args["perturbation_mask"] = perturbation_mask - video_self_attn_args["all_perturbed"] = all_perturbed - - attn_hidden_states = self.attn1(**video_self_attn_args) - hidden_states = hidden_states + attn_hidden_states * gate_msa - - # 1.2. Audio Self-Attention - audio_ada_params = self.get_mod_params(self.audio_scale_shift_table, temb_audio, batch_size) - audio_shift_msa, audio_scale_msa, audio_gate_msa, audio_shift_mlp, audio_scale_mlp, audio_gate_mlp = ( - audio_ada_params[:6] - ) - if self.audio_cross_attn_adaln: - audio_shift_text_q, audio_scale_text_q, audio_gate_text_q = audio_ada_params[6:9] - - norm_audio_hidden_states = self.audio_norm1(audio_hidden_states) - norm_audio_hidden_states = norm_audio_hidden_states * (1 + audio_scale_msa) + audio_shift_msa - - audio_self_attn_args = { - "hidden_states": norm_audio_hidden_states, - "encoder_hidden_states": None, - "query_rotary_emb": audio_rotary_emb, - "attention_mask": audio_self_attention_mask, - } - if self.perturbed_attn: - audio_self_attn_args["perturbation_mask"] = perturbation_mask - audio_self_attn_args["all_perturbed"] = all_perturbed - - attn_audio_hidden_states = self.audio_attn1(**audio_self_attn_args) - audio_hidden_states = audio_hidden_states + attn_audio_hidden_states * audio_gate_msa - - # 2. Video and Audio Cross-Attention with the text embeddings (Q: Video or Audio; K,V: Text) - if self.cross_attn_adaln: - video_prompt_ada_params = self.get_mod_params(self.prompt_scale_shift_table, temb_prompt, batch_size) - shift_text_kv, scale_text_kv = video_prompt_ada_params - - audio_prompt_ada_params = self.get_mod_params( - self.audio_prompt_scale_shift_table, temb_prompt_audio, batch_size - ) - audio_shift_text_kv, audio_scale_text_kv = audio_prompt_ada_params - - # 2.1. Video-Text Cross-Attention (Q: Video; K,V: Text) - norm_hidden_states = self.norm2(hidden_states) - if self.video_cross_attn_adaln: - norm_hidden_states = norm_hidden_states * (1 + scale_text_q) + shift_text_q - if self.cross_attn_adaln: - encoder_hidden_states = encoder_hidden_states * (1 + scale_text_kv) + shift_text_kv - - attn_hidden_states = self.attn2( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - query_rotary_emb=None, - attention_mask=encoder_attention_mask, - ) - if self.video_cross_attn_adaln: - attn_hidden_states = attn_hidden_states * gate_text_q - hidden_states = hidden_states + attn_hidden_states - - # 2.2. Audio-Text Cross-Attention - norm_audio_hidden_states = self.audio_norm2(audio_hidden_states) - if self.audio_cross_attn_adaln: - norm_audio_hidden_states = norm_audio_hidden_states * (1 + audio_scale_text_q) + audio_shift_text_q - if self.cross_attn_adaln: - audio_encoder_hidden_states = audio_encoder_hidden_states * (1 + audio_scale_text_kv) + audio_shift_text_kv - - attn_audio_hidden_states = self.audio_attn2( - norm_audio_hidden_states, - encoder_hidden_states=audio_encoder_hidden_states, - query_rotary_emb=None, - attention_mask=audio_encoder_attention_mask, - ) - if self.audio_cross_attn_adaln: - attn_audio_hidden_states = attn_audio_hidden_states * audio_gate_text_q - audio_hidden_states = audio_hidden_states + attn_audio_hidden_states - - # 3. Audio-to-Video (a2v) and Video-to-Audio (v2a) Cross-Attention - if use_a2v_cross_attention or use_v2a_cross_attention: - norm_hidden_states = self.audio_to_video_norm(hidden_states) - norm_audio_hidden_states = self.video_to_audio_norm(audio_hidden_states) - - # 3.1. Combine global and per-layer cross attention modulation parameters - # Video - video_per_layer_ca_scale_shift = self.video_a2v_cross_attn_scale_shift_table[:4, :] - video_per_layer_ca_gate = self.video_a2v_cross_attn_scale_shift_table[4:, :] - - video_ca_ada_params = self.get_mod_params(video_per_layer_ca_scale_shift, temb_ca_scale_shift, batch_size) - video_ca_gate_param = self.get_mod_params(video_per_layer_ca_gate, temb_ca_gate, batch_size) - - video_a2v_ca_scale, video_a2v_ca_shift, video_v2a_ca_scale, video_v2a_ca_shift = video_ca_ada_params - a2v_gate = video_ca_gate_param[0].squeeze(2) - - # Audio - audio_per_layer_ca_scale_shift = self.audio_a2v_cross_attn_scale_shift_table[:4, :] - audio_per_layer_ca_gate = self.audio_a2v_cross_attn_scale_shift_table[4:, :] - - audio_ca_ada_params = self.get_mod_params( - audio_per_layer_ca_scale_shift, temb_ca_audio_scale_shift, batch_size - ) - audio_ca_gate_param = self.get_mod_params(audio_per_layer_ca_gate, temb_ca_audio_gate, batch_size) - - audio_a2v_ca_scale, audio_a2v_ca_shift, audio_v2a_ca_scale, audio_v2a_ca_shift = audio_ca_ada_params - v2a_gate = audio_ca_gate_param[0].squeeze(2) - - # 3.2. Audio-to-Video Cross Attention: Q: Video; K,V: Audio - if use_a2v_cross_attention: - mod_norm_hidden_states = norm_hidden_states * ( - 1 + video_a2v_ca_scale.squeeze(2) - ) + video_a2v_ca_shift.squeeze(2) - mod_norm_audio_hidden_states = norm_audio_hidden_states * ( - 1 + audio_a2v_ca_scale.squeeze(2) - ) + audio_a2v_ca_shift.squeeze(2) - - a2v_attn_hidden_states = self.audio_to_video_attn( - mod_norm_hidden_states, - encoder_hidden_states=mod_norm_audio_hidden_states, - query_rotary_emb=ca_video_rotary_emb, - key_rotary_emb=ca_audio_rotary_emb, - attention_mask=a2v_cross_attention_mask, - ) - - hidden_states = hidden_states + a2v_gate * a2v_attn_hidden_states - - # 3.3. Video-to-Audio Cross Attention: Q: Audio; K,V: Video - if use_v2a_cross_attention: - mod_norm_hidden_states = norm_hidden_states * ( - 1 + video_v2a_ca_scale.squeeze(2) - ) + video_v2a_ca_shift.squeeze(2) - mod_norm_audio_hidden_states = norm_audio_hidden_states * ( - 1 + audio_v2a_ca_scale.squeeze(2) - ) + audio_v2a_ca_shift.squeeze(2) - - v2a_attn_hidden_states = self.video_to_audio_attn( - mod_norm_audio_hidden_states, - encoder_hidden_states=mod_norm_hidden_states, - query_rotary_emb=ca_audio_rotary_emb, - key_rotary_emb=ca_video_rotary_emb, - attention_mask=v2a_cross_attention_mask, - ) - - audio_hidden_states = audio_hidden_states + v2a_gate * v2a_attn_hidden_states - - # 4. Feedforward - norm_hidden_states = self.norm3(hidden_states) * (1 + scale_mlp) + shift_mlp - ff_output = self.ff(norm_hidden_states) - hidden_states = hidden_states + ff_output * gate_mlp - - norm_audio_hidden_states = self.audio_norm3(audio_hidden_states) * (1 + audio_scale_mlp) + audio_shift_mlp - audio_ff_output = self.audio_ff(norm_audio_hidden_states) - audio_hidden_states = audio_hidden_states + audio_ff_output * audio_gate_mlp - - return hidden_states, audio_hidden_states - - -class LTX2AudioVideoRotaryPosEmbed(nn.Module): - """ - Video and audio rotary positional embeddings (RoPE) for the LTX-2.0 model. - - Args: - causal_offset (`int`, *optional*, defaults to `1`): - Offset in the temporal axis for causal VAE modeling. This is typically 1 (for causal modeling where the VAE - treats the very first frame differently), but could also be 0 (for non-causal modeling). - """ - - def __init__( - self, - dim: int, - patch_size: int = 1, - patch_size_t: int = 1, - base_num_frames: int = 20, - base_height: int = 2048, - base_width: int = 2048, - sampling_rate: int = 16000, - hop_length: int = 160, - scale_factors: tuple[int, ...] = (8, 32, 32), - theta: float = 10000.0, - causal_offset: int = 1, - modality: str = "video", - double_precision: bool = True, - rope_type: str = "interleaved", - num_attention_heads: int = 32, - ) -> None: - super().__init__() - - self.dim = dim - self.patch_size = patch_size - self.patch_size_t = patch_size_t - - if rope_type not in ["interleaved", "split"]: - raise ValueError(f"{rope_type=} not supported. Choose between 'interleaved' and 'split'.") - self.rope_type = rope_type - - self.base_num_frames = base_num_frames - self.num_attention_heads = num_attention_heads - - # Video-specific - self.base_height = base_height - self.base_width = base_width - - # Audio-specific - self.sampling_rate = sampling_rate - self.hop_length = hop_length - self.audio_latents_per_second = float(sampling_rate) / float(hop_length) / float(scale_factors[0]) - - self.scale_factors = scale_factors - self.theta = theta - self.causal_offset = causal_offset - - self.modality = modality - if self.modality not in ["video", "audio"]: - raise ValueError(f"Modality {modality} is not supported. Supported modalities are `video` and `audio`.") - self.double_precision = double_precision - - def prepare_video_coords( - self, - batch_size: int, - num_frames: int, - height: int, - width: int, - device: torch.device, - fps: float = 24.0, - ) -> torch.Tensor: - """ - Create per-dimension bounds [inclusive start, exclusive end) for each patch with respect to the original pixel - space video grid (num_frames, height, width). This will ultimately have shape (batch_size, 3, num_patches, 2) - where - - axis 1 (size 3) enumerates (frame, height, width) dimensions (e.g. idx 0 corresponds to frames) - - axis 3 (size 2) stores `[start, end)` indices within each dimension - - Args: - batch_size (`int`): - Batch size of the video latents. - num_frames (`int`): - Number of latent frames in the video latents. - height (`int`): - Latent height of the video latents. - width (`int`): - Latent width of the video latents. - device (`torch.device`): - Device on which to create the video grid. - - Returns: - `torch.Tensor`: - Per-dimension patch boundaries tensor of shape [batch_size, 3, num_patches, 2]. - """ - - # 1. Generate grid coordinates for each spatiotemporal dimension (frames, height, width) - # Always compute rope in fp32 - grid_f = torch.arange(start=0, end=num_frames, step=self.patch_size_t, dtype=torch.float32, device=device) - grid_h = torch.arange(start=0, end=height, step=self.patch_size, dtype=torch.float32, device=device) - grid_w = torch.arange(start=0, end=width, step=self.patch_size, dtype=torch.float32, device=device) - # indexing='ij' ensures that the dimensions are kept in order as (frames, height, width) - grid = torch.meshgrid(grid_f, grid_h, grid_w, indexing="ij") - grid = torch.stack(grid, dim=0) # [3, N_F, N_H, N_W], where e.g. N_F is the number of temporal patches - - # 2. Get the patch boundaries with respect to the latent video grid - patch_size = (self.patch_size_t, self.patch_size, self.patch_size) - patch_size_delta = torch.tensor(patch_size, dtype=grid.dtype, device=grid.device) - patch_ends = grid + patch_size_delta.view(3, 1, 1, 1) - - # Combine the start (grid) and end (patch_ends) coordinates along new trailing dimension - latent_coords = torch.stack([grid, patch_ends], dim=-1) # [3, N_F, N_H, N_W, 2] - # Reshape to (batch_size, 3, num_patches, 2) - latent_coords = latent_coords.flatten(1, 3) - latent_coords = latent_coords.unsqueeze(0).repeat(batch_size, 1, 1, 1) - - # 3. Calculate the pixel space patch boundaries from the latent boundaries. - scale_tensor = torch.tensor(self.scale_factors, device=latent_coords.device) - # Broadcast the VAE scale factors such that they are compatible with latent_coords's shape - broadcast_shape = [1] * latent_coords.ndim - broadcast_shape[1] = -1 # This is the (frame, height, width) dim - # Apply per-axis scaling to convert latent coordinates to pixel space coordinates - pixel_coords = latent_coords * scale_tensor.view(*broadcast_shape) - - # As the VAE temporal stride for the first frame is 1 instead of self.vae_scale_factors[0], we need to shift - # and clamp to keep the first-frame timestamps causal and non-negative. - pixel_coords[:, 0, ...] = (pixel_coords[:, 0, ...] + self.causal_offset - self.scale_factors[0]).clamp(min=0) - - # Scale the temporal coordinates by the video FPS - pixel_coords[:, 0, ...] = pixel_coords[:, 0, ...] / fps - - return pixel_coords - - def prepare_audio_coords( - self, - batch_size: int, - num_frames: int, - device: torch.device, - shift: int = 0, - ) -> torch.Tensor: - """ - Create per-dimension bounds [inclusive start, exclusive end) of start and end timestamps for each latent frame. - This will ultimately have shape (batch_size, 3, num_patches, 2) where - - axis 1 (size 1) represents the temporal dimension - - axis 3 (size 2) stores `[start, end)` indices within each dimension - - Args: - batch_size (`int`): - Batch size of the audio latents. - num_frames (`int`): - Number of latent frames in the audio latents. - device (`torch.device`): - Device on which to create the audio grid. - shift (`int`, *optional*, defaults to `0`): - Offset on the latent indices. Different shift values correspond to different overlapping windows with - respect to the same underlying latent grid. - - Returns: - `torch.Tensor`: - Per-dimension patch boundaries tensor of shape [batch_size, 1, num_patches, 2]. - """ - - # 1. Generate coordinates in the frame (time) dimension. - # Always compute rope in fp32 - grid_f = torch.arange( - start=shift, end=num_frames + shift, step=self.patch_size_t, dtype=torch.float32, device=device - ) - - # 2. Calculate start timstamps in seconds with respect to the original spectrogram grid - audio_scale_factor = self.scale_factors[0] - # Scale back to mel spectrogram space - grid_start_mel = grid_f * audio_scale_factor - # Handle first frame causal offset, ensuring non-negative timestamps - grid_start_mel = (grid_start_mel + self.causal_offset - audio_scale_factor).clip(min=0) - # Convert mel bins back into seconds - grid_start_s = grid_start_mel * self.hop_length / self.sampling_rate - - # 3. Calculate start timstamps in seconds with respect to the original spectrogram grid - grid_end_mel = (grid_f + self.patch_size_t) * audio_scale_factor - grid_end_mel = (grid_end_mel + self.causal_offset - audio_scale_factor).clip(min=0) - grid_end_s = grid_end_mel * self.hop_length / self.sampling_rate - - audio_coords = torch.stack([grid_start_s, grid_end_s], dim=-1) # [num_patches, 2] - audio_coords = audio_coords.unsqueeze(0).expand(batch_size, -1, -1) # [batch_size, num_patches, 2] - audio_coords = audio_coords.unsqueeze(1) # [batch_size, 1, num_patches, 2] - return audio_coords - - def prepare_coords(self, *args, **kwargs): - if self.modality == "video": - return self.prepare_video_coords(*args, **kwargs) - elif self.modality == "audio": - return self.prepare_audio_coords(*args, **kwargs) - - def forward( - self, coords: torch.Tensor, device: str | torch.device | None = None - ) -> tuple[torch.Tensor, torch.Tensor]: - device = device or coords.device - - # Number of spatiotemporal dimensions (3 for video, 1 (temporal) for audio and cross attn) - num_pos_dims = coords.shape[1] - - # 1. If the coords are patch boundaries [start, end), use the midpoint of these boundaries as the patch - # position index - if coords.ndim == 4: - coords_start, coords_end = coords.chunk(2, dim=-1) - coords = (coords_start + coords_end) / 2.0 - coords = coords.squeeze(-1) # [B, num_pos_dims, num_patches] - - # 2. Get coordinates as a fraction of the base data shape - if self.modality == "video": - max_positions = (self.base_num_frames, self.base_height, self.base_width) - elif self.modality == "audio": - max_positions = (self.base_num_frames,) - # [B, num_pos_dims, num_patches] --> [B, num_patches, num_pos_dims] - grid = torch.stack([coords[:, i] / max_positions[i] for i in range(num_pos_dims)], dim=-1).to(device) - # Number of spatiotemporal dimensions (3 for video, 1 for audio and cross attn) times 2 for cos, sin - num_rope_elems = num_pos_dims * 2 - - # 3. Create a 1D grid of frequencies for RoPE - freqs_dtype = torch.float64 if self.double_precision else torch.float32 - pow_indices = torch.pow( - self.theta, - torch.linspace(start=0.0, end=1.0, steps=self.dim // num_rope_elems, dtype=freqs_dtype, device=device), - ) - freqs = (pow_indices * torch.pi / 2.0).to(dtype=torch.float32) - - # 4. Tensor-vector outer product between pos ids tensor of shape (B, 3, num_patches) and freqs vector of shape - # (self.dim // num_elems,) - freqs = (grid.unsqueeze(-1) * 2 - 1) * freqs # [B, num_patches, num_pos_dims, self.dim // num_elems] - freqs = freqs.transpose(-1, -2).flatten(2) # [B, num_patches, self.dim // 2] - - # 5. Get real, interleaved (cos, sin) frequencies, padded to self.dim - # TODO: consider implementing this as a utility and reuse in `connectors.py`. - # src/diffusers/pipelines/ltx2/connectors.py - if self.rope_type == "interleaved": - cos_freqs = freqs.cos().repeat_interleave(2, dim=-1) - sin_freqs = freqs.sin().repeat_interleave(2, dim=-1) - - if self.dim % num_rope_elems != 0: - cos_padding = torch.ones_like(cos_freqs[:, :, : self.dim % num_rope_elems]) - sin_padding = torch.zeros_like(cos_freqs[:, :, : self.dim % num_rope_elems]) - cos_freqs = torch.cat([cos_padding, cos_freqs], dim=-1) - sin_freqs = torch.cat([sin_padding, sin_freqs], dim=-1) - - elif self.rope_type == "split": - expected_freqs = self.dim // 2 - current_freqs = freqs.shape[-1] - pad_size = expected_freqs - current_freqs - cos_freq = freqs.cos() - sin_freq = freqs.sin() - - if pad_size != 0: - cos_padding = torch.ones_like(cos_freq[:, :, :pad_size]) - sin_padding = torch.zeros_like(sin_freq[:, :, :pad_size]) - - cos_freq = torch.concatenate([cos_padding, cos_freq], axis=-1) - sin_freq = torch.concatenate([sin_padding, sin_freq], axis=-1) - - # Reshape freqs to be compatible with multi-head attention - b = cos_freq.shape[0] - t = cos_freq.shape[1] - - cos_freq = cos_freq.reshape(b, t, self.num_attention_heads, -1) - sin_freq = sin_freq.reshape(b, t, self.num_attention_heads, -1) - - cos_freqs = torch.swapaxes(cos_freq, 1, 2) # (B,H,T,D//2) - sin_freqs = torch.swapaxes(sin_freq, 1, 2) # (B,H,T,D//2) - - return cos_freqs, sin_freqs - - -class LTX2VideoTransformer3DModel( - ModelMixin, ConfigMixin, AttentionMixin, FromOriginalModelMixin, PeftAdapterMixin, CacheMixin -): - r""" - A Transformer model for video-like data used in [LTX](https://huggingface.co/Lightricks/LTX-Video). - - Args: - in_channels (`int`, defaults to `128`): - The number of channels in the input. - out_channels (`int`, defaults to `128`): - The number of channels in the output. - patch_size (`int`, defaults to `1`): - The size of the spatial patches to use in the patch embedding layer. - patch_size_t (`int`, defaults to `1`): - The size of the tmeporal patches to use in the patch embedding layer. - num_attention_heads (`int`, defaults to `32`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `64`): - The number of channels in each head. - cross_attention_dim (`int`, defaults to `2048 `): - The number of channels for cross attention heads. - num_layers (`int`, defaults to `28`): - The number of layers of Transformer blocks to use. - activation_fn (`str`, defaults to `"gelu-approximate"`): - Activation function to use in feed-forward. - qk_norm (`str`, defaults to `"rms_norm_across_heads"`): - The normalization layer to use. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["norm"] - _repeated_blocks = ["LTX2VideoTransformerBlock"] - _cp_plan = { - "": { - "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - "encoder_hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - "encoder_attention_mask": ContextParallelInput(split_dim=1, expected_dims=2, split_output=False), - }, - "rope": { - 0: ContextParallelInput(split_dim=1, expected_dims=3, split_output=True), - 1: ContextParallelInput(split_dim=1, expected_dims=3, split_output=True), - }, - "proj_out": ContextParallelOutput(gather_dim=1, expected_dims=3), - } - - @register_to_config - def __init__( - self, - in_channels: int = 128, # Video Arguments - out_channels: int | None = 128, - patch_size: int = 1, - patch_size_t: int = 1, - num_attention_heads: int = 32, - attention_head_dim: int = 128, - cross_attention_dim: int = 4096, - vae_scale_factors: tuple[int, int, int] = (8, 32, 32), - pos_embed_max_pos: int = 20, - base_height: int = 2048, - base_width: int = 2048, - gated_attn: bool = False, - cross_attn_mod: bool = False, - audio_in_channels: int = 128, # Audio Arguments - audio_out_channels: int | None = 128, - audio_patch_size: int = 1, - audio_patch_size_t: int = 1, - audio_num_attention_heads: int = 32, - audio_attention_head_dim: int = 64, - audio_cross_attention_dim: int = 2048, - audio_scale_factor: int = 4, - audio_pos_embed_max_pos: int = 20, - audio_sampling_rate: int = 16000, - audio_hop_length: int = 160, - audio_gated_attn: bool = False, - audio_cross_attn_mod: bool = False, - num_layers: int = 48, # Shared arguments - activation_fn: str = "gelu-approximate", - qk_norm: str = "rms_norm_across_heads", - norm_elementwise_affine: bool = False, - norm_eps: float = 1e-6, - caption_channels: int = 3840, - attention_bias: bool = True, - attention_out_bias: bool = True, - rope_theta: float = 10000.0, - rope_double_precision: bool = True, - causal_offset: int = 1, - timestep_scale_multiplier: int = 1000, - cross_attn_timestep_scale_multiplier: int = 1000, - rope_type: str = "interleaved", - use_prompt_embeddings=True, - perturbed_attn: bool = False, - ) -> None: - super().__init__() - - out_channels = out_channels or in_channels - audio_out_channels = audio_out_channels or audio_in_channels - inner_dim = num_attention_heads * attention_head_dim - audio_inner_dim = audio_num_attention_heads * audio_attention_head_dim - - # 1. Patchification input projections - self.proj_in = nn.Linear(in_channels, inner_dim) - self.audio_proj_in = nn.Linear(audio_in_channels, audio_inner_dim) - - # 2. Prompt embeddings - if use_prompt_embeddings: - # LTX-2.0; LTX-2.3 uses per-modality feature projections in the connector instead - self.caption_projection = PixArtAlphaTextProjection(in_features=caption_channels, hidden_size=inner_dim) - self.audio_caption_projection = PixArtAlphaTextProjection( - in_features=caption_channels, hidden_size=audio_inner_dim - ) - - # 3. Timestep Modulation Params and Embedding - self.prompt_modulation = cross_attn_mod or audio_cross_attn_mod # used by LTX-2.3 - - # 3.1. Global Timestep Modulation Parameters (except for cross-attention) and timestep + size embedding - # time_embed and audio_time_embed calculate both the timestep embedding and (global) modulation parameters - video_time_emb_mod_params = 9 if cross_attn_mod else 6 - audio_time_emb_mod_params = 9 if audio_cross_attn_mod else 6 - self.time_embed = LTX2AdaLayerNormSingle( - inner_dim, num_mod_params=video_time_emb_mod_params, use_additional_conditions=False - ) - self.audio_time_embed = LTX2AdaLayerNormSingle( - audio_inner_dim, num_mod_params=audio_time_emb_mod_params, use_additional_conditions=False - ) - - # 3.2. Global Cross Attention Modulation Parameters - # Used in the audio-to-video and video-to-audio cross attention layers as a global set of modulation params, - # which are then further modified by per-block modulaton params in each transformer block. - # There are 2 sets of scale/shift parameters for each modality, 1 each for audio-to-video (a2v) and - # video-to-audio (v2a) cross attention - self.av_cross_attn_video_scale_shift = LTX2AdaLayerNormSingle( - inner_dim, num_mod_params=4, use_additional_conditions=False - ) - self.av_cross_attn_audio_scale_shift = LTX2AdaLayerNormSingle( - audio_inner_dim, num_mod_params=4, use_additional_conditions=False - ) - # Gate param for audio-to-video (a2v) cross attn (where the video is the queries (Q) and the audio is the keys - # and values (KV)) - self.av_cross_attn_video_a2v_gate = LTX2AdaLayerNormSingle( - inner_dim, num_mod_params=1, use_additional_conditions=False - ) - # Gate param for video-to-audio (v2a) cross attn (where the audio is the queries (Q) and the video is the keys - # and values (KV)) - self.av_cross_attn_audio_v2a_gate = LTX2AdaLayerNormSingle( - audio_inner_dim, num_mod_params=1, use_additional_conditions=False - ) - - # 3.3. Output Layer Scale/Shift Modulation parameters - self.scale_shift_table = nn.Parameter(torch.randn(2, inner_dim) / inner_dim**0.5) - self.audio_scale_shift_table = nn.Parameter(torch.randn(2, audio_inner_dim) / audio_inner_dim**0.5) - - # 3.4. Prompt Scale/Shift Modulation parameters (LTX-2.3) - if self.prompt_modulation: - self.prompt_adaln = LTX2AdaLayerNormSingle(inner_dim, num_mod_params=2, use_additional_conditions=False) - self.audio_prompt_adaln = LTX2AdaLayerNormSingle( - audio_inner_dim, num_mod_params=2, use_additional_conditions=False - ) - - # 4. Rotary Positional Embeddings (RoPE) - # Self-Attention - self.rope = LTX2AudioVideoRotaryPosEmbed( - dim=inner_dim, - patch_size=patch_size, - patch_size_t=patch_size_t, - base_num_frames=pos_embed_max_pos, - base_height=base_height, - base_width=base_width, - scale_factors=vae_scale_factors, - theta=rope_theta, - causal_offset=causal_offset, - modality="video", - double_precision=rope_double_precision, - rope_type=rope_type, - num_attention_heads=num_attention_heads, - ) - self.audio_rope = LTX2AudioVideoRotaryPosEmbed( - dim=audio_inner_dim, - patch_size=audio_patch_size, - patch_size_t=audio_patch_size_t, - base_num_frames=audio_pos_embed_max_pos, - sampling_rate=audio_sampling_rate, - hop_length=audio_hop_length, - scale_factors=[audio_scale_factor], - theta=rope_theta, - causal_offset=causal_offset, - modality="audio", - double_precision=rope_double_precision, - rope_type=rope_type, - num_attention_heads=audio_num_attention_heads, - ) - - # Audio-to-Video, Video-to-Audio Cross-Attention - cross_attn_pos_embed_max_pos = max(pos_embed_max_pos, audio_pos_embed_max_pos) - self.cross_attn_rope = LTX2AudioVideoRotaryPosEmbed( - dim=audio_cross_attention_dim, - patch_size=patch_size, - patch_size_t=patch_size_t, - base_num_frames=cross_attn_pos_embed_max_pos, - base_height=base_height, - base_width=base_width, - theta=rope_theta, - causal_offset=causal_offset, - modality="video", - double_precision=rope_double_precision, - rope_type=rope_type, - num_attention_heads=num_attention_heads, - ) - self.cross_attn_audio_rope = LTX2AudioVideoRotaryPosEmbed( - dim=audio_cross_attention_dim, - patch_size=audio_patch_size, - patch_size_t=audio_patch_size_t, - base_num_frames=cross_attn_pos_embed_max_pos, - sampling_rate=audio_sampling_rate, - hop_length=audio_hop_length, - theta=rope_theta, - causal_offset=causal_offset, - modality="audio", - double_precision=rope_double_precision, - rope_type=rope_type, - num_attention_heads=audio_num_attention_heads, - ) - - # 5. Transformer Blocks - self.transformer_blocks = nn.ModuleList( - [ - LTX2VideoTransformerBlock( - dim=inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - cross_attention_dim=cross_attention_dim, - audio_dim=audio_inner_dim, - audio_num_attention_heads=audio_num_attention_heads, - audio_attention_head_dim=audio_attention_head_dim, - audio_cross_attention_dim=audio_cross_attention_dim, - video_gated_attn=gated_attn, - video_cross_attn_adaln=cross_attn_mod, - audio_gated_attn=audio_gated_attn, - audio_cross_attn_adaln=audio_cross_attn_mod, - qk_norm=qk_norm, - activation_fn=activation_fn, - attention_bias=attention_bias, - attention_out_bias=attention_out_bias, - eps=norm_eps, - elementwise_affine=norm_elementwise_affine, - rope_type=rope_type, - perturbed_attn=perturbed_attn, - ) - for _ in range(num_layers) - ] - ) - - # 6. Output layers - self.norm_out = nn.LayerNorm(inner_dim, eps=1e-6, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, out_channels) - - self.audio_norm_out = nn.LayerNorm(audio_inner_dim, eps=1e-6, elementwise_affine=False) - self.audio_proj_out = nn.Linear(audio_inner_dim, audio_out_channels) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - audio_hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - audio_encoder_hidden_states: torch.Tensor, - timestep: torch.LongTensor, - audio_timestep: torch.LongTensor | None = None, - sigma: torch.Tensor | None = None, - audio_sigma: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - audio_encoder_attention_mask: torch.Tensor | None = None, - num_frames: int | None = None, - height: int | None = None, - width: int | None = None, - fps: float = 24.0, - audio_num_frames: int | None = None, - video_coords: torch.Tensor | None = None, - audio_coords: torch.Tensor | None = None, - isolate_modalities: bool = False, - spatio_temporal_guidance_blocks: list[int] | None = None, - perturbation_mask: torch.Tensor | None = None, - use_cross_timestep: bool = False, - attention_kwargs: dict[str, Any] | None = None, - video_self_attention_mask: torch.Tensor | None = None, - return_dict: bool = True, - ) -> torch.Tensor: - """ - Forward pass for LTX-2.0 audiovisual video transformer. - - Args: - hidden_states (`torch.Tensor`): - Input patchified video latents of shape `(batch_size, num_video_tokens, in_channels)`. - audio_hidden_states (`torch.Tensor`): - Input patchified audio latents of shape `(batch_size, num_audio_tokens, audio_in_channels)`. - encoder_hidden_states (`torch.Tensor`): - Input video text embeddings of shape `(batch_size, text_seq_len, self.config.caption_channels)`. - audio_encoder_hidden_states (`torch.Tensor`): - Input audio text embeddings of shape `(batch_size, text_seq_len, self.config.caption_channels)`. - timestep (`torch.Tensor`): - Input timestep of shape `(batch_size, num_video_tokens)`. These should already be scaled by - `self.config.timestep_scale_multiplier`. - audio_timestep (`torch.Tensor`, *optional*): - Input timestep of shape `(batch_size,)` or `(batch_size, num_audio_tokens)` for audio modulation - params. This is only used by certain pipelines such as the I2V pipeline. - sigma (`torch.Tensor`, *optional*): - Input scaled timestep of shape (batch_size,). Used for video prompt cross attention modulation in - models such as LTX-2.3. - audio_sigma (`torch.Tensor`, *optional*): - Input scaled timestep of shape (batch_size,). Used for audio prompt cross attention modulation in - models such as LTX-2.3. If `sigma` is supplied but `audio_sigma` is not, `audio_sigma` will be set to - the provided `sigma` value. - encoder_attention_mask (`torch.Tensor`, *optional*): - Optional multiplicative text attention mask of shape `(batch_size, text_seq_len)`. - audio_encoder_attention_mask (`torch.Tensor`, *optional*): - Optional multiplicative text attention mask of shape `(batch_size, text_seq_len)` for audio modeling. - num_frames (`int`, *optional*): - The number of latent video frames. Used if calculating the video coordinates for RoPE. - height (`int`, *optional*): - The latent video height. Used if calculating the video coordinates for RoPE. - width (`int`, *optional*): - The latent video width. Used if calculating the video coordinates for RoPE. - fps: (`float`, *optional*, defaults to `24.0`): - The desired frames per second of the generated video. Used if calculating the video coordinates for - RoPE. - audio_num_frames: (`int`, *optional*): - The number of latent audio frames. Used if calculating the audio coordinates for RoPE. - video_coords (`torch.Tensor`, *optional*): - The video coordinates to be used when calculating the rotary positional embeddings (RoPE) of shape - `(batch_size, 3, num_video_tokens, 2)`. If not supplied, this will be calculated inside `forward`. - audio_coords (`torch.Tensor`, *optional*): - The audio coordinates to be used when calculating the rotary positional embeddings (RoPE) of shape - `(batch_size, 1, num_audio_tokens, 2)`. If not supplied, this will be calculated inside `forward`. - isolate_modalities (`bool`, *optional*, defaults to `False`): - Whether to isolate each modality by turning off cross-modality (audio-to-video and video-to-audio) - cross attention (for all blocks). Use for modality guidance in LTX-2.3. - spatio_temporal_guidance_blocks (`list[int]`, *optional*, defaults to `None`): - The transformer block indices at which to apply spatio-temporal guidance (STG), which shortcuts the - self-attention operations by simply using the values rather than the full scaled dot-product attention - (SDPA) operation. If `None` or empty, STG will not be applied to any block. - perturbation_mask (`torch.Tensor`, *optional*): - Perturbation mask for STG of shape `(batch_size,)` or `(batch_size, 1, 1)`. Should be 0 at batch - elements where STG should be applied and 1 elsewhere. If STG is being used but `peturbation_mask` is - not supplied, will default to applying STG (perturbing) all batch elements. - use_cross_timestep (`bool` *optional*, defaults to `False`): - Whether to use the cross modality (audio is the cross modality of video, and vice versa) sigma when - calculating the cross attention modulation parameters. `True` is the newer (e.g. LTX-2.3) behavior; - `False` is the legacy LTX-2.0 behavior. - attention_kwargs (`dict[str, Any]`, *optional*): - Optional dict of keyword args to be passed to the attention processor. - video_self_attention_mask (`torch.Tensor`, *optional*): - Optional multiplicative self-attention mask of shape `(batch_size, num_video_tokens, num_video_tokens)` - applied to the video self-attention in each transformer block. Values in `[0, 1]` where `1` means full - attention and `0` means masked. Used e.g. by the IC-LoRA pipeline to control attention strength between - noisy tokens and appended reference tokens. Audio self-attention is not affected. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a dict-like structured output of type `AudioVisualModelOutput` or a tuple. - - Returns: - `AudioVisualModelOutput` or `tuple`: - If `return_dict` is `True`, returns a structured output of type `AudioVisualModelOutput`, otherwise a - `tuple` is returned where the first element is the denoised video latent patch sequence and the second - element is the denoised audio latent patch sequence. - """ - # Determine timestep for audio. - audio_timestep = audio_timestep if audio_timestep is not None else timestep - audio_sigma = audio_sigma if audio_sigma is not None else sigma - - # convert encoder_attention_mask to a bias the same way we do for attention_mask - if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2: - encoder_attention_mask = (1 - encoder_attention_mask.to(hidden_states.dtype)) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - if audio_encoder_attention_mask is not None and audio_encoder_attention_mask.ndim == 2: - audio_encoder_attention_mask = (1 - audio_encoder_attention_mask.to(audio_hidden_states.dtype)) * -10000.0 - audio_encoder_attention_mask = audio_encoder_attention_mask.unsqueeze(1) - - # Convert video_self_attention_mask from multiplicative mask ([0, 1]) to additive bias form (0 / -10000) - # matching the encoder_attention_mask convention above. Shape is preserved: (B, T_v, T_v). - if video_self_attention_mask is not None: - video_self_attention_mask = (1 - video_self_attention_mask.to(hidden_states.dtype)) * -10000.0 - - batch_size = hidden_states.size(0) - - # 1. Prepare RoPE positional embeddings - if video_coords is None: - video_coords = self.rope.prepare_video_coords( - batch_size, num_frames, height, width, hidden_states.device, fps=fps - ) - if audio_coords is None: - audio_coords = self.audio_rope.prepare_audio_coords( - batch_size, audio_num_frames, audio_hidden_states.device - ) - - video_rotary_emb = self.rope(video_coords, device=hidden_states.device) - audio_rotary_emb = self.audio_rope(audio_coords, device=audio_hidden_states.device) - - video_cross_attn_rotary_emb = self.cross_attn_rope(video_coords[:, 0:1, :], device=hidden_states.device) - audio_cross_attn_rotary_emb = self.cross_attn_audio_rope( - audio_coords[:, 0:1, :], device=audio_hidden_states.device - ) - - # 2. Patchify input projections - hidden_states = self.proj_in(hidden_states) - audio_hidden_states = self.audio_proj_in(audio_hidden_states) - - # 3. Prepare timestep embeddings and modulation parameters - timestep_cross_attn_gate_scale_factor = ( - self.config.cross_attn_timestep_scale_multiplier / self.config.timestep_scale_multiplier - ) - - # 3.1. Prepare global modality (video and audio) timestep embedding and modulation parameters - # temb is used in the transformer blocks (as expected), while embedded_timestep is used for the output layer - # modulation with scale_shift_table (and similarly for audio) - temb, embedded_timestep = self.time_embed( - timestep.flatten(), - batch_size=batch_size, - hidden_dtype=hidden_states.dtype, - ) - temb = temb.view(batch_size, -1, temb.size(-1)) - embedded_timestep = embedded_timestep.view(batch_size, -1, embedded_timestep.size(-1)) - - temb_audio, audio_embedded_timestep = self.audio_time_embed( - audio_timestep.flatten(), - batch_size=batch_size, - hidden_dtype=audio_hidden_states.dtype, - ) - temb_audio = temb_audio.view(batch_size, -1, temb_audio.size(-1)) - audio_embedded_timestep = audio_embedded_timestep.view(batch_size, -1, audio_embedded_timestep.size(-1)) - - if self.prompt_modulation: - # LTX-2.3 - temb_prompt, _ = self.prompt_adaln( - sigma.flatten(), batch_size=batch_size, hidden_dtype=hidden_states.dtype - ) - temb_prompt_audio, _ = self.audio_prompt_adaln( - audio_sigma.flatten(), batch_size=batch_size, hidden_dtype=audio_hidden_states.dtype - ) - temb_prompt = temb_prompt.view(batch_size, -1, temb_prompt.size(-1)) - temb_prompt_audio = temb_prompt_audio.view(batch_size, -1, temb_prompt_audio.size(-1)) - else: - temb_prompt = temb_prompt_audio = None - - # 3.2. Prepare global modality cross attention modulation parameters - video_ca_timestep = audio_sigma.flatten() if use_cross_timestep else timestep.flatten() - video_cross_attn_scale_shift, _ = self.av_cross_attn_video_scale_shift( - video_ca_timestep, - batch_size=batch_size, - hidden_dtype=hidden_states.dtype, - ) - video_cross_attn_a2v_gate, _ = self.av_cross_attn_video_a2v_gate( - video_ca_timestep * timestep_cross_attn_gate_scale_factor, - batch_size=batch_size, - hidden_dtype=hidden_states.dtype, - ) - video_cross_attn_scale_shift = video_cross_attn_scale_shift.view( - batch_size, -1, video_cross_attn_scale_shift.shape[-1] - ) - video_cross_attn_a2v_gate = video_cross_attn_a2v_gate.view(batch_size, -1, video_cross_attn_a2v_gate.shape[-1]) - - audio_ca_timestep = sigma.flatten() if use_cross_timestep else audio_timestep.flatten() - audio_cross_attn_scale_shift, _ = self.av_cross_attn_audio_scale_shift( - audio_ca_timestep, - batch_size=batch_size, - hidden_dtype=audio_hidden_states.dtype, - ) - audio_cross_attn_v2a_gate, _ = self.av_cross_attn_audio_v2a_gate( - audio_ca_timestep * timestep_cross_attn_gate_scale_factor, - batch_size=batch_size, - hidden_dtype=audio_hidden_states.dtype, - ) - audio_cross_attn_scale_shift = audio_cross_attn_scale_shift.view( - batch_size, -1, audio_cross_attn_scale_shift.shape[-1] - ) - audio_cross_attn_v2a_gate = audio_cross_attn_v2a_gate.view(batch_size, -1, audio_cross_attn_v2a_gate.shape[-1]) - - # 4. Prepare prompt embeddings (LTX-2.0) - if self.config.use_prompt_embeddings: - encoder_hidden_states = self.caption_projection(encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states.view(batch_size, -1, hidden_states.size(-1)) - - audio_encoder_hidden_states = self.audio_caption_projection(audio_encoder_hidden_states) - audio_encoder_hidden_states = audio_encoder_hidden_states.view( - batch_size, -1, audio_hidden_states.size(-1) - ) - - # 5. Run transformer blocks - spatio_temporal_guidance_blocks = spatio_temporal_guidance_blocks or [] - if len(spatio_temporal_guidance_blocks) > 0 and perturbation_mask is None: - # If STG is being used and perturbation_mask is not supplied, default to perturbing all batch elements. - perturbation_mask = torch.zeros((batch_size,)) - if perturbation_mask is not None and perturbation_mask.ndim == 1: - perturbation_mask = perturbation_mask[:, None, None] # unsqueeze to 3D to broadcast with hidden_states - all_perturbed = torch.all(perturbation_mask == 0) if perturbation_mask is not None else False - stg_blocks = set(spatio_temporal_guidance_blocks) - - for block_idx, block in enumerate(self.transformer_blocks): - block_perturbation_mask = perturbation_mask if block_idx in stg_blocks else None - block_all_perturbed = all_perturbed if block_idx in stg_blocks else False - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, audio_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - audio_hidden_states, - encoder_hidden_states, - audio_encoder_hidden_states, - temb, - temb_audio, - video_cross_attn_scale_shift, - audio_cross_attn_scale_shift, - video_cross_attn_a2v_gate, - audio_cross_attn_v2a_gate, - temb_prompt, - temb_prompt_audio, - video_rotary_emb, - audio_rotary_emb, - video_cross_attn_rotary_emb, - audio_cross_attn_rotary_emb, - encoder_attention_mask, - audio_encoder_attention_mask, - video_self_attention_mask, # self_attention_mask (video-only) - None, # audio_self_attention_mask - None, # a2v_cross_attention_mask - None, # v2a_cross_attention_mask - not isolate_modalities, # use_a2v_cross_attention - not isolate_modalities, # use_v2a_cross_attention - block_perturbation_mask, - block_all_perturbed, - ) - else: - hidden_states, audio_hidden_states = block( - hidden_states=hidden_states, - audio_hidden_states=audio_hidden_states, - encoder_hidden_states=encoder_hidden_states, - audio_encoder_hidden_states=audio_encoder_hidden_states, - temb=temb, - temb_audio=temb_audio, - temb_ca_scale_shift=video_cross_attn_scale_shift, - temb_ca_audio_scale_shift=audio_cross_attn_scale_shift, - temb_ca_gate=video_cross_attn_a2v_gate, - temb_ca_audio_gate=audio_cross_attn_v2a_gate, - temb_prompt=temb_prompt, - temb_prompt_audio=temb_prompt_audio, - video_rotary_emb=video_rotary_emb, - audio_rotary_emb=audio_rotary_emb, - ca_video_rotary_emb=video_cross_attn_rotary_emb, - ca_audio_rotary_emb=audio_cross_attn_rotary_emb, - encoder_attention_mask=encoder_attention_mask, - audio_encoder_attention_mask=audio_encoder_attention_mask, - self_attention_mask=video_self_attention_mask, - audio_self_attention_mask=None, - a2v_cross_attention_mask=None, - v2a_cross_attention_mask=None, - use_a2v_cross_attention=not isolate_modalities, - use_v2a_cross_attention=not isolate_modalities, - perturbation_mask=block_perturbation_mask, - all_perturbed=block_all_perturbed, - ) - - # 6. Output layers (including unpatchification) - scale_shift_values = self.scale_shift_table[None, None] + embedded_timestep[:, :, None] - shift, scale = scale_shift_values[:, :, 0], scale_shift_values[:, :, 1] - - hidden_states = self.norm_out(hidden_states) - hidden_states = hidden_states * (1 + scale) + shift - output = self.proj_out(hidden_states) - - audio_scale_shift_values = self.audio_scale_shift_table[None, None] + audio_embedded_timestep[:, :, None] - audio_shift, audio_scale = audio_scale_shift_values[:, :, 0], audio_scale_shift_values[:, :, 1] - - audio_hidden_states = self.audio_norm_out(audio_hidden_states) - audio_hidden_states = audio_hidden_states * (1 + audio_scale) + audio_shift - audio_output = self.audio_proj_out(audio_hidden_states) - - if not return_dict: - return (output, audio_output) - return AudioVisualModelOutput(sample=output, audio_sample=audio_output) diff --git a/diffusers/models/transformers/transformer_lumina2.py b/diffusers/models/transformers/transformer_lumina2.py deleted file mode 100644 index ba822730cb32cf46280f9017a09ed024be8fc991..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_lumina2.py +++ /dev/null @@ -1,554 +0,0 @@ -# Copyright 2025 Alpha-VLLM Authors and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...loaders.single_file_model import FromOriginalModelMixin -from ...utils import apply_lora_scale, logging -from ..attention import LuminaFeedForward -from ..attention_processor import Attention -from ..embeddings import TimestepEmbedding, Timesteps, apply_rotary_emb, get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import LuminaLayerNormContinuous, LuminaRMSNormZero, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class Lumina2CombinedTimestepCaptionEmbedding(nn.Module): - def __init__( - self, - hidden_size: int = 4096, - cap_feat_dim: int = 2048, - frequency_embedding_size: int = 256, - norm_eps: float = 1e-5, - ) -> None: - super().__init__() - - self.time_proj = Timesteps( - num_channels=frequency_embedding_size, flip_sin_to_cos=True, downscale_freq_shift=0.0 - ) - - self.timestep_embedder = TimestepEmbedding( - in_channels=frequency_embedding_size, time_embed_dim=min(hidden_size, 1024) - ) - - self.caption_embedder = nn.Sequential( - RMSNorm(cap_feat_dim, eps=norm_eps), nn.Linear(cap_feat_dim, hidden_size, bias=True) - ) - - def forward( - self, hidden_states: torch.Tensor, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor]: - timestep_proj = self.time_proj(timestep).type_as(hidden_states) - time_embed = self.timestep_embedder(timestep_proj) - caption_embed = self.caption_embedder(encoder_hidden_states) - return time_embed, caption_embed - - -class Lumina2AttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is - used in the Lumina2Transformer2DModel model. It applies normalization and RoPE on query and key vectors. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - base_sequence_length: int | None = None, - ) -> torch.Tensor: - batch_size, sequence_length, _ = hidden_states.shape - - # Get Query-Key-Value Pair - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - query_dim = query.shape[-1] - inner_dim = key.shape[-1] - head_dim = query_dim // attn.heads - dtype = query.dtype - - # Get key-value heads - kv_heads = inner_dim // head_dim - - query = query.view(batch_size, -1, attn.heads, head_dim) - key = key.view(batch_size, -1, kv_heads, head_dim) - value = value.view(batch_size, -1, kv_heads, head_dim) - - # Apply Query-Key Norm if needed - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE if needed - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, use_real=False) - key = apply_rotary_emb(key, image_rotary_emb, use_real=False) - - query, key = query.to(dtype), key.to(dtype) - - # Apply proportional attention if true - if base_sequence_length is not None: - softmax_scale = math.sqrt(math.log(sequence_length, base_sequence_length)) * attn.scale - else: - softmax_scale = attn.scale - - # perform Grouped-qurey Attention (GQA) - n_rep = attn.heads // kv_heads - if n_rep >= 1: - key = key.unsqueeze(3).repeat(1, 1, 1, n_rep, 1).flatten(2, 3) - value = value.unsqueeze(3).repeat(1, 1, 1, n_rep, 1).flatten(2, 3) - - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - if attention_mask is not None: - attention_mask = attention_mask.bool().view(batch_size, 1, 1, -1) - - query = query.transpose(1, 2) - key = key.transpose(1, 2) - value = value.transpose(1, 2) - - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, scale=softmax_scale - ) - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.type_as(query) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class Lumina2TransformerBlock(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - num_kv_heads: int, - multiple_of: int, - ffn_dim_multiplier: float, - norm_eps: float, - modulation: bool = True, - ) -> None: - super().__init__() - self.head_dim = dim // num_attention_heads - self.modulation = modulation - - self.attn = Attention( - query_dim=dim, - cross_attention_dim=None, - dim_head=dim // num_attention_heads, - qk_norm="rms_norm", - heads=num_attention_heads, - kv_heads=num_kv_heads, - eps=1e-5, - bias=False, - out_bias=False, - processor=Lumina2AttnProcessor2_0(), - ) - - self.feed_forward = LuminaFeedForward( - dim=dim, - inner_dim=4 * dim, - multiple_of=multiple_of, - ffn_dim_multiplier=ffn_dim_multiplier, - ) - - if modulation: - self.norm1 = LuminaRMSNormZero( - embedding_dim=dim, - norm_eps=norm_eps, - norm_elementwise_affine=True, - ) - else: - self.norm1 = RMSNorm(dim, eps=norm_eps) - self.ffn_norm1 = RMSNorm(dim, eps=norm_eps) - - self.norm2 = RMSNorm(dim, eps=norm_eps) - self.ffn_norm2 = RMSNorm(dim, eps=norm_eps) - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor, - image_rotary_emb: torch.Tensor, - temb: torch.Tensor | None = None, - ) -> torch.Tensor: - if self.modulation: - norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(hidden_states, temb) - attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - hidden_states = hidden_states + gate_msa.unsqueeze(1).tanh() * self.norm2(attn_output) - mlp_output = self.feed_forward(self.ffn_norm1(hidden_states) * (1 + scale_mlp.unsqueeze(1))) - hidden_states = hidden_states + gate_mlp.unsqueeze(1).tanh() * self.ffn_norm2(mlp_output) - else: - norm_hidden_states = self.norm1(hidden_states) - attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - hidden_states = hidden_states + self.norm2(attn_output) - mlp_output = self.feed_forward(self.ffn_norm1(hidden_states)) - hidden_states = hidden_states + self.ffn_norm2(mlp_output) - - return hidden_states - - -class Lumina2RotaryPosEmbed(nn.Module): - def __init__(self, theta: int, axes_dim: list[int], axes_lens: list[int] = (300, 512, 512), patch_size: int = 2): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - self.axes_lens = axes_lens - self.patch_size = patch_size - - self.freqs_cis = self._precompute_freqs_cis(axes_dim, axes_lens, theta) - - def _precompute_freqs_cis(self, axes_dim: list[int], axes_lens: list[int], theta: int) -> list[torch.Tensor]: - freqs_cis = [] - freqs_dtype = torch.float32 if torch.backends.mps.is_available() else torch.float64 - for i, (d, e) in enumerate(zip(axes_dim, axes_lens)): - emb = get_1d_rotary_pos_embed(d, e, theta=self.theta, freqs_dtype=freqs_dtype) - freqs_cis.append(emb) - return freqs_cis - - def _get_freqs_cis(self, ids: torch.Tensor) -> torch.Tensor: - device = ids.device - if ids.device.type == "mps": - ids = ids.to("cpu") - - result = [] - for i in range(len(self.axes_dim)): - freqs = self.freqs_cis[i].to(ids.device) - index = ids[:, :, i : i + 1].repeat(1, 1, freqs.shape[-1]).to(torch.int64) - result.append(torch.gather(freqs.unsqueeze(0).repeat(index.shape[0], 1, 1), dim=1, index=index)) - return torch.cat(result, dim=-1).to(device) - - def forward(self, hidden_states: torch.Tensor, attention_mask: torch.Tensor): - batch_size, channels, height, width = hidden_states.shape - p = self.patch_size - post_patch_height, post_patch_width = height // p, width // p - image_seq_len = post_patch_height * post_patch_width - device = hidden_states.device - - encoder_seq_len = attention_mask.shape[1] - l_effective_cap_len = attention_mask.sum(dim=1).tolist() - seq_lengths = [cap_seq_len + image_seq_len for cap_seq_len in l_effective_cap_len] - max_seq_len = max(seq_lengths) - - # Create position IDs - position_ids = torch.zeros(batch_size, max_seq_len, 3, dtype=torch.int32, device=device) - - for i, (cap_seq_len, seq_len) in enumerate(zip(l_effective_cap_len, seq_lengths)): - # add caption position ids - position_ids[i, :cap_seq_len, 0] = torch.arange(cap_seq_len, dtype=torch.int32, device=device) - position_ids[i, cap_seq_len:seq_len, 0] = cap_seq_len - - # add image position ids - row_ids = ( - torch.arange(post_patch_height, dtype=torch.int32, device=device) - .view(-1, 1) - .repeat(1, post_patch_width) - .flatten() - ) - col_ids = ( - torch.arange(post_patch_width, dtype=torch.int32, device=device) - .view(1, -1) - .repeat(post_patch_height, 1) - .flatten() - ) - position_ids[i, cap_seq_len:seq_len, 1] = row_ids - position_ids[i, cap_seq_len:seq_len, 2] = col_ids - - # Get combined rotary embeddings - freqs_cis = self._get_freqs_cis(position_ids) - - # create separate rotary embeddings for captions and images - cap_freqs_cis = torch.zeros( - batch_size, encoder_seq_len, freqs_cis.shape[-1], device=device, dtype=freqs_cis.dtype - ) - img_freqs_cis = torch.zeros( - batch_size, image_seq_len, freqs_cis.shape[-1], device=device, dtype=freqs_cis.dtype - ) - - for i, (cap_seq_len, seq_len) in enumerate(zip(l_effective_cap_len, seq_lengths)): - cap_freqs_cis[i, :cap_seq_len] = freqs_cis[i, :cap_seq_len] - img_freqs_cis[i, :image_seq_len] = freqs_cis[i, cap_seq_len:seq_len] - - # image patch embeddings - hidden_states = ( - hidden_states.view(batch_size, channels, post_patch_height, p, post_patch_width, p) - .permute(0, 2, 4, 3, 5, 1) - .flatten(3) - .flatten(1, 2) - ) - - return hidden_states, cap_freqs_cis, img_freqs_cis, freqs_cis, l_effective_cap_len, seq_lengths - - -class Lumina2Transformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): - r""" - Lumina2NextDiT: Diffusion model with a Transformer backbone. - - Parameters: - sample_size (`int`): The width of the latent images. This is fixed during training since - it is used to learn a number of position embeddings. - patch_size (`int`, *optional*, (`int`, *optional*, defaults to 2): - The size of each patch in the image. This parameter defines the resolution of patches fed into the model. - in_channels (`int`, *optional*, defaults to 4): - The number of input channels for the model. Typically, this matches the number of channels in the input - images. - hidden_size (`int`, *optional*, defaults to 4096): - The dimensionality of the hidden layers in the model. This parameter determines the width of the model's - hidden representations. - num_layers (`int`, *optional*, default to 32): - The number of layers in the model. This defines the depth of the neural network. - num_attention_heads (`int`, *optional*, defaults to 32): - The number of attention heads in each attention layer. This parameter specifies how many separate attention - mechanisms are used. - num_kv_heads (`int`, *optional*, defaults to 8): - The number of key-value heads in the attention mechanism, if different from the number of attention heads. - If None, it defaults to num_attention_heads. - multiple_of (`int`, *optional*, defaults to 256): - A factor that the hidden size should be a multiple of. This can help optimize certain hardware - configurations. - ffn_dim_multiplier (`float`, *optional*): - A multiplier for the dimensionality of the feed-forward network. If None, it uses a default value based on - the model configuration. - norm_eps (`float`, *optional*, defaults to 1e-5): - A small value added to the denominator for numerical stability in normalization layers. - scaling_factor (`float`, *optional*, defaults to 1.0): - A scaling factor applied to certain parameters or layers in the model. This can be used for adjusting the - overall scale of the model's operations. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["Lumina2TransformerBlock"] - _skip_layerwise_casting_patterns = ["x_embedder", "norm"] - - @register_to_config - def __init__( - self, - sample_size: int = 128, - patch_size: int = 2, - in_channels: int = 16, - out_channels: int | None = None, - hidden_size: int = 2304, - num_layers: int = 26, - num_refiner_layers: int = 2, - num_attention_heads: int = 24, - num_kv_heads: int = 8, - multiple_of: int = 256, - ffn_dim_multiplier: float | None = None, - norm_eps: float = 1e-5, - scaling_factor: float = 1.0, - axes_dim_rope: tuple[int, int, int] = (32, 32, 32), - axes_lens: tuple[int, int, int] = (300, 512, 512), - cap_feat_dim: int = 1024, - ) -> None: - super().__init__() - self.out_channels = out_channels or in_channels - - # 1. Positional, patch & conditional embeddings - self.rope_embedder = Lumina2RotaryPosEmbed( - theta=10000, axes_dim=axes_dim_rope, axes_lens=axes_lens, patch_size=patch_size - ) - - self.x_embedder = nn.Linear(in_features=patch_size * patch_size * in_channels, out_features=hidden_size) - - self.time_caption_embed = Lumina2CombinedTimestepCaptionEmbedding( - hidden_size=hidden_size, cap_feat_dim=cap_feat_dim, norm_eps=norm_eps - ) - - # 2. Noise and context refinement blocks - self.noise_refiner = nn.ModuleList( - [ - Lumina2TransformerBlock( - hidden_size, - num_attention_heads, - num_kv_heads, - multiple_of, - ffn_dim_multiplier, - norm_eps, - modulation=True, - ) - for _ in range(num_refiner_layers) - ] - ) - - self.context_refiner = nn.ModuleList( - [ - Lumina2TransformerBlock( - hidden_size, - num_attention_heads, - num_kv_heads, - multiple_of, - ffn_dim_multiplier, - norm_eps, - modulation=False, - ) - for _ in range(num_refiner_layers) - ] - ) - - # 3. Transformer blocks - self.layers = nn.ModuleList( - [ - Lumina2TransformerBlock( - hidden_size, - num_attention_heads, - num_kv_heads, - multiple_of, - ffn_dim_multiplier, - norm_eps, - modulation=True, - ) - for _ in range(num_layers) - ] - ) - - # 4. Output norm & projection - self.norm_out = LuminaLayerNormContinuous( - embedding_dim=hidden_size, - conditioning_embedding_dim=min(hidden_size, 1024), - elementwise_affine=False, - eps=1e-6, - bias=True, - out_dim=patch_size * patch_size * self.out_channels, - ) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_attention_mask: torch.Tensor, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> torch.Tensor | Transformer2DModelOutput: - """ - The [`Lumina2Transformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, in_channels, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_attention_mask (`torch.Tensor`): - Mask applied to `encoder_hidden_states` during attention. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - # 1. Condition, positional & patch embedding - batch_size, _, height, width = hidden_states.shape - - temb, encoder_hidden_states = self.time_caption_embed(hidden_states, timestep, encoder_hidden_states) - - ( - hidden_states, - context_rotary_emb, - noise_rotary_emb, - rotary_emb, - encoder_seq_lengths, - seq_lengths, - ) = self.rope_embedder(hidden_states, encoder_attention_mask) - - hidden_states = self.x_embedder(hidden_states) - - # 2. Context & noise refinement - for layer in self.context_refiner: - encoder_hidden_states = layer(encoder_hidden_states, encoder_attention_mask, context_rotary_emb) - - for layer in self.noise_refiner: - hidden_states = layer(hidden_states, None, noise_rotary_emb, temb) - - # 3. Joint Transformer blocks - max_seq_len = max(seq_lengths) - use_mask = len(set(seq_lengths)) > 1 - - attention_mask = hidden_states.new_zeros(batch_size, max_seq_len, dtype=torch.bool) - joint_hidden_states = hidden_states.new_zeros(batch_size, max_seq_len, self.config.hidden_size) - for i, (encoder_seq_len, seq_len) in enumerate(zip(encoder_seq_lengths, seq_lengths)): - attention_mask[i, :seq_len] = True - joint_hidden_states[i, :encoder_seq_len] = encoder_hidden_states[i, :encoder_seq_len] - joint_hidden_states[i, encoder_seq_len:seq_len] = hidden_states[i] - - hidden_states = joint_hidden_states - - for layer in self.layers: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - layer, hidden_states, attention_mask if use_mask else None, rotary_emb, temb - ) - else: - hidden_states = layer(hidden_states, attention_mask if use_mask else None, rotary_emb, temb) - - # 4. Output norm & projection - hidden_states = self.norm_out(hidden_states, temb) - - # 5. Unpatchify - p = self.config.patch_size - output = [] - for i, (encoder_seq_len, seq_len) in enumerate(zip(encoder_seq_lengths, seq_lengths)): - output.append( - hidden_states[i][encoder_seq_len:seq_len] - .view(height // p, width // p, p, p, self.out_channels) - .permute(4, 0, 2, 1, 3) - .flatten(3, 4) - .flatten(1, 2) - ) - output = torch.stack(output, dim=0) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_minimax_h3.py b/diffusers/models/transformers/transformer_minimax_h3.py deleted file mode 100644 index 5170f149a8eef1acc01f6e7b24f74e2073f09d1c..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_minimax_h3.py +++ /dev/null @@ -1,644 +0,0 @@ -# Copyright 2025 The MiniMax Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import BaseOutput, apply_lora_scale, logging -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# MiniMax-H3 tags every row of the packed sequence with the modality it belongs to and keeps one set of AdaLN -# modulation parameters per (timestep, modality) pair: 0 = video, 1 = text, 2 = audio. -MINIMAX_H3_MODALITY_NUM = 3 - - -@dataclass -class MiniMaxH3TransformerOutput(BaseOutput): - r""" - The output of [`MiniMaxH3Transformer3DModel`]. - - Args: - sample (`torch.Tensor` of shape `(batch_size, num_video_tokens, in_channels * prod(patch_size))`): - The video velocity prediction for the rows addressed by `video_indices`, in the same order. Conditioning - rows are returned unmasked — masking them out before the scheduler step is the caller's job. - audio_sample (`torch.Tensor` of shape `(batch_size, num_audio_tokens, audio_in_channels)`): - The audio velocity prediction for the rows addressed by `audio_indices`, in the same order. - """ - - sample: torch.Tensor - audio_sample: torch.Tensor - - -def _apply_rotary_emb(hidden_states: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: - r""" - Rotate the leading `rotary_dim` channels of every head and pass the remaining channels through unchanged. - `hidden_states` is `(batch_size, seq_len, num_heads, head_dim)` and `cos`/`sin` are `(seq_len, rotary_dim)`. - """ - rotary_dim = cos.shape[-1] - hidden_states_rotary = hidden_states[..., :rotary_dim] - hidden_states_pass = hidden_states[..., rotary_dim:] - - cos = cos.to(hidden_states.dtype)[None, :, None, :] - sin = sin.to(hidden_states.dtype)[None, :, None, :] - x1, x2 = hidden_states_rotary.chunk(2, dim=-1) - hidden_states_rotated = torch.cat((-x2, x1), dim=-1) - hidden_states_rotary = hidden_states_rotary * cos + hidden_states_rotated * sin - return torch.cat((hidden_states_rotary, hidden_states_pass), dim=-1).contiguous() - - -class MiniMaxH3RotaryPosEmbed(nn.Module): - r""" - 3-axis rotary embedding over the `(t, h, w)` coordinates of the packed sequence. - - A single `inv_freq` buffer of `rope_freq_dim` frequencies is shared by the three axes. Each axis contributes - `rope_freq_dim` angles, the three blocks are concatenated to `3 * rope_freq_dim` and then concatenated with - themselves so that the `rotate_half` convention rotates `2 * 3 * rope_freq_dim` of the `head_dim` channels. - """ - - def __init__(self, rope_freq_dim: int = 16, rope_theta: float = 10000.0): - super().__init__() - self.rope_freq_dim = rope_freq_dim - inv_freq = 1.0 / ( - rope_theta ** (torch.arange(0, 2 * rope_freq_dim, 2, dtype=torch.float32) / (2 * rope_freq_dim)) - ) - self.register_buffer("inv_freq", inv_freq, persistent=False) - - def forward(self, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - # position_ids: (seq_len, 3) -> cos/sin: (seq_len, 2 * 3 * rope_freq_dim) - position_ids = position_ids.to(torch.float32) - freqs = position_ids.unsqueeze(-1) * self.inv_freq.view(1, 1, -1) # (seq_len, 3, rope_freq_dim) - freqs_t, freqs_h, freqs_w = freqs.unbind(dim=1) - freqs = torch.cat((freqs_t, freqs_h, freqs_w), dim=-1) - freqs = torch.cat((freqs, freqs), dim=-1) - return freqs.cos(), freqs.sin() - - -class MiniMaxH3AdaLayerNormModulation(nn.Module): - r""" - Projects the shared timestep embedding into the six per-(timestep, modality) modulation parameters of one - transformer block. - - `(num_timesteps, time_embed_dim)` -> six tensors of shape `(num_timesteps * MINIMAX_H3_MODALITY_NUM, - hidden_size)`, in the diffusers `shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp` order. The row - layout of the returned tensors is `[t0_mod0, t0_mod1, t0_mod2, t1_mod0, ...]`, which is what `timestep_indices * - MINIMAX_H3_MODALITY_NUM + token_tags` addresses. - - A single projection is shared by `norm1` and `norm2` and by the three modalities, so it cannot be folded into - either norm the way [`~models.normalization.AdaLayerNormZero`] does. It is therefore a block-level module of its - own, named after the checkpoint's `adaln_proj`, with the modulation projection under the `linear` name diffusers - uses inside every AdaLN module. - """ - - def __init__(self, time_embed_dim: int, hidden_size: int): - super().__init__() - self.hidden_size = hidden_size - self.linear = nn.Linear(time_embed_dim, 6 * hidden_size * MINIMAX_H3_MODALITY_NUM, bias=True) - - def forward(self, temb: torch.Tensor) -> tuple[torch.Tensor, ...]: - # The activation runs at `temb`'s own precision — float32, since `time_embedder` is a float32 module in this - # mixed-precision checkpoint — and only its result is cast down to the bfloat16 projection. Every block reads - # the same `temb`, so a rounding applied before the activation biases every block's modulation parameters - # identically at every sampling step, which accumulates coherently over the denoising trajectory. - temb = self.linear(nn.functional.silu(temb).to(self.linear.weight.dtype)) - temb = temb.view(-1, 6 * self.hidden_size) - return temb.chunk(6, dim=-1) - - -class MiniMaxH3AdaLayerNormOut(nn.Module): - r""" - Final norm of the packed sequence, shift/scale modulated per row. - - Same module layout and checkpoint keys as [`~models.normalization.AdaLayerNormContinuous`] (`norm` plus a `linear` - projecting the conditioning embedding to `2 * hidden_size`), with two MiniMax-H3 specifics: the modulation table - holds one row per *timestep* and is addressed per row of the packed sequence rather than per batch item, and the - two halves of the projection are `shift` then `scale`, the order `LTX2Transformer3DModel` and - `WanTransformer3DModel` also use in their output layers. - """ - - def __init__(self, hidden_size: int, time_embed_dim: int, eps: float): - super().__init__() - self.norm = nn.RMSNorm(hidden_size, eps=eps) - self.linear = nn.Linear(time_embed_dim, 2 * hidden_size, bias=True) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor, timestep_indices: torch.Tensor) -> torch.Tensor: - # As in `MiniMaxH3AdaLayerNormModulation`: activate at `temb`'s precision, cast to the projection's dtype after. - shift, scale = self.linear(nn.functional.silu(temb).to(self.linear.weight.dtype)).chunk(2, dim=-1) - # The modulation itself stays at the block stack's precision; `forward` casts to the output heads' dtype. - hidden_states = self.norm(hidden_states) - return hidden_states * (1.0 + scale.index_select(0, timestep_indices)) + shift.index_select( - 0, timestep_indices - ) - - -class MiniMaxH3AttnProcessor: - r""" - Full self-attention over one packed sequence. There is no cross-attention anywhere in MiniMax-H3. - """ - - _attention_backend = None - _parallel_config = None - - def __call__( - self, - attn: "MiniMaxH3Attention", - hidden_states: torch.Tensor, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - if attn.fused_projections: - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - else: - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if rotary_emb is not None: - query = _apply_rotary_emb(query, *rotary_emb) - key = _apply_rotary_emb(key, *rotary_emb) - - # Without padding rows the packed sequence is a single attention document and no mask is needed (passing an - # all-zero float mask here would hard-fail the flash / sage backends). When padding rows are present, the - # caller supplies a boolean mask that keeps them in their own attention document, mirroring the reference's - # `cu_seqlens = [0, used, S]` split; masked backends (SDPA & co.) are required in that case. - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3).type_as(query) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class MiniMaxH3Attention(nn.Module, AttentionModuleMixin): - _default_processor_cls = MiniMaxH3AttnProcessor - _available_processors = [MiniMaxH3AttnProcessor] - - def __init__( - self, - hidden_size: int, - heads: int, - dim_head: int, - qk_norm_eps: float = 1e-5, - processor=None, - ): - super().__init__() - self.heads = heads - self.head_dim = dim_head - self.inner_dim = heads * dim_head - self.use_bias = False - - self.to_q = nn.Linear(hidden_size, self.inner_dim, bias=False) - self.to_k = nn.Linear(hidden_size, self.inner_dim, bias=False) - self.to_v = nn.Linear(hidden_size, self.inner_dim, bias=False) - self.norm_q = nn.RMSNorm(dim_head, eps=qk_norm_eps) - self.norm_k = nn.RMSNorm(dim_head, eps=qk_norm_eps) - self.to_out = nn.ModuleList([nn.Linear(self.inner_dim, hidden_size, bias=False), nn.Dropout(0.0)]) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - return self.processor(self, hidden_states, rotary_emb, attention_mask) - - -class MiniMaxH3TokenRefinerBlock(nn.Module): - r""" - Plain pre-norm transformer block used to refine the projected text stream. No AdaLN and no rotary embedding. - """ - - def __init__( - self, - hidden_size: int, - num_attention_heads: int, - attention_head_dim: int, - ffn_dim: int, - norm_eps: float, - qk_norm_eps: float, - ): - super().__init__() - self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps) - self.attn = MiniMaxH3Attention( - hidden_size=hidden_size, - heads=num_attention_heads, - dim_head=attention_head_dim, - qk_norm_eps=qk_norm_eps, - ) - self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps) - self.ff = FeedForward(hidden_size, inner_dim=ffn_dim, activation_fn="swiglu", bias=False) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = hidden_states + self.attn(self.norm1(hidden_states)) - hidden_states = hidden_states + self.ff(self.norm2(hidden_states)) - return hidden_states - - -class MiniMaxH3TokenRefiner(nn.Module): - def __init__( - self, - hidden_size: int, - num_attention_heads: int, - attention_head_dim: int, - ffn_dim: int, - num_layers: int, - norm_eps: float, - qk_norm_eps: float, - final_norm_eps: float, - ): - super().__init__() - self.refiner_blocks = nn.ModuleList( - [ - MiniMaxH3TokenRefinerBlock( - hidden_size=hidden_size, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ffn_dim=ffn_dim, - norm_eps=norm_eps, - qk_norm_eps=qk_norm_eps, - ) - for _ in range(num_layers) - ] - ) - self.final_norm = nn.RMSNorm(hidden_size, eps=final_norm_eps) - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - for block in self.refiner_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(block, hidden_states) - else: - hidden_states = block(hidden_states) - return self.final_norm(hidden_states) - - -class MiniMaxH3TransformerBlock(nn.Module): - r""" - MiniMax-H3 block: pre-norm self-attention and feed-forward, each modulated by AdaLN parameters selected per row of - the packed sequence from the `(timestep, modality)` table. - """ - - def __init__( - self, - hidden_size: int, - num_attention_heads: int, - attention_head_dim: int, - ffn_dim: int, - time_embed_dim: int, - norm_eps: float, - qk_norm_eps: float, - ): - super().__init__() - self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps) - self.attn = MiniMaxH3Attention( - hidden_size=hidden_size, - heads=num_attention_heads, - dim_head=attention_head_dim, - qk_norm_eps=qk_norm_eps, - ) - self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps) - self.ff = FeedForward(hidden_size, inner_dim=ffn_dim, activation_fn="swiglu", bias=False) - self.adaln_proj = MiniMaxH3AdaLayerNormModulation(time_embed_dim=time_embed_dim, hidden_size=hidden_size) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor, - adaln_indices: torch.Tensor, - rotary_emb: tuple[torch.Tensor, torch.Tensor], - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaln_proj(temb) - - residual = hidden_states - norm_hidden_states = self.norm1(hidden_states) - norm_hidden_states = norm_hidden_states * ( - 1.0 + scale_msa.index_select(0, adaln_indices) - ) + shift_msa.index_select(0, adaln_indices) - attn_output = self.attn(norm_hidden_states, rotary_emb, attention_mask) - hidden_states = residual + gate_msa.index_select(0, adaln_indices) * attn_output - - residual = hidden_states - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * ( - 1.0 + scale_mlp.index_select(0, adaln_indices) - ) + shift_mlp.index_select(0, adaln_indices) - ff_output = self.ff(norm_hidden_states) - hidden_states = residual + gate_mlp.index_select(0, adaln_indices) * ff_output - - return hidden_states - - -class MiniMaxH3Transformer3DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixin, CacheMixin): - r""" - A Transformer model for joint video + audio generation, introduced in MiniMax-H3. - - MiniMax-H3 runs a single stack of blocks over **one packed 1-D sequence** that holds the text condition, the - conditioning image / video rows, the audio rows and the target video rows. Attention is full self-attention over - that sequence; there is no cross-attention and no per-modality block weights. Modality-specific behaviour comes - only from the two input patch projections, the per-row AdaLN modality tag, and the two output heads. - - The caller is responsible for building the packed layout: patchifying the video latents, ordering the rows, and - producing the `(t, h, w)` position grid, the per-row modality tags and the per-row timestep indices. Padding rows - (tag `-1`) are kept in a separate attention document, matching the reference implementation, which pads to a - multiple of 64 for FlashAttention with `cu_seqlens = [0, used, S]`. Prefer dropping them — a padless sequence - needs no attention mask, keeping the unmasked attention backends available. - - The batch axis is a pure replication axis: the structural arguments (`timestep`, `timestep_indices`, `token_tags`, - `position_ids` and the three index tensors) describe one packed layout that every batch item shares, and each item - is a single attention document. - - Args: - num_attention_heads (`int`, defaults to `56`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each attention head. Note that `num_attention_heads * attention_head_dim` is - *larger* than `hidden_size` in MiniMax-H3. - hidden_size (`int`, defaults to `5376`): - The number of channels of the packed sequence (the residual stream). - num_layers (`int`, defaults to `50`): - The number of transformer blocks. - num_refiner_layers (`int`, defaults to `2`): - The number of token refiner blocks applied to the projected text stream. - ffn_dim (`int`, defaults to `14336`): - The inner dimension of the SwiGLU feed-forward layers. - in_channels (`int`, defaults to `24`): - The number of channels of the video latents. - audio_in_channels (`int`, defaults to `32`): - The number of channels of the audio latents. - patch_size (`tuple[int, int, int]`, defaults to `(1, 2, 2)`): - The `(t, h, w)` patch used to pack the video latents into rows. - text_dim (`int`, defaults to `5120`): - The number of channels of the text conditioning produced by the text encoder. - freq_dim (`int`, defaults to `256`): - The dimension of the sinusoidal timestep embedding. Timesteps are consumed unscaled in `[0, 1]`. - time_embed_hidden_dim (`int`, defaults to `5376`): - The inner dimension of the timestep MLP. - time_embed_dim (`int`, defaults to `2688`): - The output dimension of the timestep MLP, i.e. the input of every AdaLN projection. - rope_freq_dim (`int`, defaults to `16`): - The number of rotary frequencies per axis. The `(t, h, w)` axes share one `inv_freq` buffer of this length - and `2 * 3 * rope_freq_dim` of the `attention_head_dim` channels are rotated. - rope_theta (`float`, defaults to `10000.0`): - The base of the rotary frequency schedule the `rope.inv_freq` buffer is computed from. - norm_eps (`float`, defaults to `1e-5`): - Epsilon of the pre-attention and pre-feed-forward norms. - qk_norm_eps (`float`, defaults to `1e-5`): - Epsilon of the per-head query/key norms. - final_norm_eps (`float`, defaults to `1e-5`): - Epsilon of the token refiner output norm and of `norm_out`. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["MiniMaxH3TransformerBlock", "MiniMaxH3TokenRefinerBlock", "MiniMaxH3AdaLayerNormOut"] - _repeated_blocks = ["MiniMaxH3TransformerBlock", "MiniMaxH3TokenRefinerBlock"] - _skip_layerwise_casting_patterns = ["norm"] - # MiniMax-H3 ships a mixed-precision checkpoint: the two input patch projections, the timestep MLP and the two - # output heads are float32 while everything else (including the AdaLN projections) is bfloat16. The `rope.inv_freq` - # buffer is computed rather than loaded and is kept float32 for the same reason the reference ships it float32. - # Entries are matched as substrings of the parameter name, so `proj_in` / `proj_out` also cover the audio heads. - _keep_in_fp32_modules = [ - "proj_in", - "audio_proj_in", - "time_embedder", - "proj_out", - "audio_proj_out", - "rope", - ] - - @register_to_config - def __init__( - self, - num_attention_heads: int = 56, - attention_head_dim: int = 128, - hidden_size: int = 5376, - num_layers: int = 50, - num_refiner_layers: int = 2, - ffn_dim: int = 14336, - in_channels: int = 24, - audio_in_channels: int = 32, - patch_size: tuple[int, int, int] = (1, 2, 2), - text_dim: int = 5120, - freq_dim: int = 256, - time_embed_hidden_dim: int = 5376, - time_embed_dim: int = 2688, - rope_freq_dim: int = 16, - rope_theta: float = 10000.0, - norm_eps: float = 1e-5, - qk_norm_eps: float = 1e-5, - final_norm_eps: float = 1e-5, - ) -> None: - super().__init__() - - video_patch_dim = in_channels * patch_size[0] * patch_size[1] * patch_size[2] - - # 1. Per-modality input projections - self.proj_in = nn.Linear(video_patch_dim, hidden_size, bias=True) - self.audio_proj_in = nn.Linear(audio_in_channels, hidden_size, bias=True) - self.context_embedder = nn.Linear(text_dim, hidden_size, bias=True) - - # 2. Timestep embedding, shared by every AdaLN projection - self.time_proj = Timesteps(num_channels=freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.time_embedder = TimestepEmbedding( - in_channels=freq_dim, time_embed_dim=time_embed_hidden_dim, out_dim=time_embed_dim - ) - - # 3. Rotary embedding over the packed (t, h, w) grid - self.rope = MiniMaxH3RotaryPosEmbed(rope_freq_dim=rope_freq_dim, rope_theta=rope_theta) - - # 4. Text stream refiner - self.token_refiner = MiniMaxH3TokenRefiner( - hidden_size=hidden_size, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ffn_dim=ffn_dim, - num_layers=num_refiner_layers, - norm_eps=norm_eps, - qk_norm_eps=qk_norm_eps, - final_norm_eps=final_norm_eps, - ) - - # 5. The block stack - self.transformer_blocks = nn.ModuleList( - [ - MiniMaxH3TransformerBlock( - hidden_size=hidden_size, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ffn_dim=ffn_dim, - time_embed_dim=time_embed_dim, - norm_eps=norm_eps, - qk_norm_eps=qk_norm_eps, - ) - for _ in range(num_layers) - ] - ) - - # 6. Shared output norm and the two per-modality output heads. Both heads run over every row of the packed - # sequence; the rows of each modality are selected afterwards. - self.norm_out = MiniMaxH3AdaLayerNormOut( - hidden_size=hidden_size, time_embed_dim=time_embed_dim, eps=final_norm_eps - ) - self.proj_out = nn.Linear(hidden_size, video_patch_dim, bias=True) - self.audio_proj_out = nn.Linear(hidden_size, audio_in_channels, bias=True) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - audio_hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: torch.Tensor, - timestep_indices: torch.Tensor, - token_tags: torch.Tensor, - position_ids: torch.Tensor, - video_indices: torch.Tensor, - audio_indices: torch.Tensor, - text_indices: torch.Tensor, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> MiniMaxH3TransformerOutput | tuple[torch.Tensor, torch.Tensor]: - r""" - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_video_tokens, in_channels * prod(patch_size))`): - Patchified video latent rows — conditioning rows and target rows — ordered as they appear in the packed - sequence, i.e. matching `video_indices`. - audio_hidden_states (`torch.Tensor` of shape `(batch_size, num_audio_tokens, audio_in_channels)`): - Audio latent rows, ordered to match `audio_indices`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, num_text_tokens, text_dim)`): - Text conditioning, ordered to match `text_indices`. - timestep (`torch.Tensor` of shape `(num_timesteps,)`): - The *distinct* timestep values present in the packed sequence, in `[0, 1]` and unscaled. One forward - serves rows at different noise levels (target video, target audio, conditioning rows). - timestep_indices (`torch.Tensor` of shape `(seq_len,)`): - For every row of the packed sequence, the index of its timestep in `timestep`. - token_tags (`torch.Tensor` of shape `(seq_len,)`): - For every row of the packed sequence, its modality: `0` video, `1` text, `2` audio, `-1` padding. - Padding rows form their own attention document and never reach the outputs. - position_ids (`torch.Tensor` of shape `(seq_len, 3)`): - The `(t, h, w)` rotary coordinates of every row of the packed sequence. - video_indices (`torch.Tensor` of shape `(num_video_tokens,)`): - Positions of the video rows in the packed sequence. - audio_indices (`torch.Tensor` of shape `(num_audio_tokens,)`): - Positions of the audio rows in the packed sequence. - text_indices (`torch.Tensor` of shape `(num_text_tokens,)`): - Positions of the text rows in the packed sequence. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that, if specified, may carry a `scale` entry which is applied to the LoRA layers. - return_dict (`bool`, defaults to `True`): - Whether to return a [`MiniMaxH3TransformerOutput`] instead of a plain tuple. - - Returns: - [`MiniMaxH3TransformerOutput`] or `tuple`: - The video velocity of shape `(batch_size, num_video_tokens, in_channels * prod(patch_size))` and the - audio velocity of shape `(batch_size, num_audio_tokens, audio_in_channels)`, in the row order of - `video_indices` and `audio_indices`. - """ - # `attention_kwargs` is consumed by the `@apply_lora_scale` decorator on this method. - if position_ids.ndim != 2 or position_ids.shape[-1] != 3: - raise ValueError(f"`position_ids` must be a `(seq_len, 3)` tensor, got {list(position_ids.shape)}.") - sequence_length = position_ids.shape[0] - if token_tags.shape != (sequence_length,) or timestep_indices.shape != (sequence_length,): - raise ValueError( - "`token_tags` and `timestep_indices` must both be `(seq_len,)` tensors matching `position_ids`, got " - f"{list(token_tags.shape)} and {list(timestep_indices.shape)} for seq_len={sequence_length}." - ) - - rotary_emb = self.rope(position_ids) - - # 1. Project each modality and scatter the rows into the packed sequence buffer. The checkpoint is - # mixed-precision (the two patch projections are float32 while `context_embedder` and the block stack are - # bfloat16 — see `_keep_in_fp32_modules`), so every input is aligned with its projection's parameter dtype, - # mirroring the reference's explicit casts. The text stream sets the dtype of the packed sequence. - video_embeds = self.proj_in(hidden_states.to(self.proj_in.weight.dtype)) - audio_embeds = self.audio_proj_in(audio_hidden_states.to(self.audio_proj_in.weight.dtype)) - text_embeds = self.context_embedder(encoder_hidden_states.to(self.context_embedder.weight.dtype)) - text_embeds = self.token_refiner(text_embeds) - - hidden_states = text_embeds.new_zeros((text_embeds.shape[0], sequence_length, text_embeds.shape[-1])) - hidden_states = hidden_states.index_copy(1, text_indices, text_embeds) - hidden_states = hidden_states.index_copy(1, video_indices, video_embeds.to(text_embeds.dtype)) - hidden_states = hidden_states.index_copy(1, audio_indices, audio_embeds.to(text_embeds.dtype)) - - # 2. One timestep embedding per distinct noise level. `temb` is shared by all AdaLN projections, which are - # bfloat16 in the checkpoint while `time_embedder` is float32, so it stays at the time embedder's precision: - # each AdaLN module applies its own activation to it and casts to its projection's dtype afterwards. - temb = self.time_proj(timestep) - temb = self.time_embedder(temb.to(self.time_embedder.linear_1.weight.dtype)) - - # 3. Row -> AdaLN table row. `clamp(min=0)` mirrors the reference, where padding rows carry the tag `-1`; the - # clamp keeps the `-1` from indexing backwards (padding rows never reach the outputs, which are selected by - # `video_indices` / `audio_indices`). - adaln_indices = timestep_indices * MINIMAX_H3_MODALITY_NUM + token_tags.clamp(min=0) - - # 4. Padding rows (tag `-1`) must not exchange attention with live rows: the reference keeps the padding tail - # as a separate attention document (`cu_seqlens = [0, used, S]`). A boolean mask that pairs live rows with live - # rows and padding rows with padding rows reproduces that split exactly. Padless sequences keep `None` so the - # unmasked fast paths (flash & co.) stay available. - attention_mask = None - is_pad = token_tags < 0 - if bool(is_pad.any()): - attention_mask = is_pad[None, :] == is_pad[:, None] - - for block in self.transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, hidden_states, temb, adaln_indices, rotary_emb, attention_mask - ) - else: - hidden_states = block(hidden_states, temb, adaln_indices, rotary_emb, attention_mask) - - # 5. Both heads run over every row, then the rows of each modality are selected. The heads are listed in - # `_keep_in_fp32_modules`, so they stay float32 while the block stack runs in the requested `torch_dtype`; - # align the activation with their parameter dtype. - hidden_states = self.norm_out(hidden_states, temb, timestep_indices).to(self.proj_out.weight.dtype) - video_output = self.proj_out(hidden_states).index_select(1, video_indices) - audio_output = self.audio_proj_out(hidden_states).index_select(1, audio_indices) - - if not return_dict: - return (video_output, audio_output) - return MiniMaxH3TransformerOutput(sample=video_output, audio_sample=audio_output) diff --git a/diffusers/models/transformers/transformer_mochi.py b/diffusers/models/transformers/transformer_mochi.py deleted file mode 100644 index a1a1f5e9c9002a5c549c1c3641b11c9cbece21d8..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_mochi.py +++ /dev/null @@ -1,494 +0,0 @@ -# Copyright 2025 The Genmo team and The HuggingFace Team. -# All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...loaders.single_file_model import FromOriginalModelMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import FeedForward -from ..attention_processor import MochiAttention, MochiAttnProcessor2_0 -from ..cache_utils import CacheMixin -from ..embeddings import MochiCombinedTimestepCaptionEmbedding, PatchEmbed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class MochiModulatedRMSNorm(nn.Module): - def __init__(self, eps: float): - super().__init__() - - self.eps = eps - self.norm = RMSNorm(0, eps, False) - - def forward(self, hidden_states, scale=None): - hidden_states_dtype = hidden_states.dtype - hidden_states = hidden_states.to(torch.float32) - - hidden_states = self.norm(hidden_states) - - if scale is not None: - hidden_states = hidden_states * scale - - hidden_states = hidden_states.to(hidden_states_dtype) - - return hidden_states - - -class MochiLayerNormContinuous(nn.Module): - def __init__( - self, - embedding_dim: int, - conditioning_embedding_dim: int, - eps=1e-5, - bias=True, - ): - super().__init__() - - # AdaLN - self.silu = nn.SiLU() - self.linear_1 = nn.Linear(conditioning_embedding_dim, embedding_dim, bias=bias) - self.norm = MochiModulatedRMSNorm(eps=eps) - - def forward( - self, - x: torch.Tensor, - conditioning_embedding: torch.Tensor, - ) -> torch.Tensor: - input_dtype = x.dtype - - # convert back to the original dtype in case `conditioning_embedding`` is upcasted to float32 (needed for hunyuanDiT) - scale = self.linear_1(self.silu(conditioning_embedding).to(x.dtype)) - x = self.norm(x, (1 + scale.unsqueeze(1).to(torch.float32))) - - return x.to(input_dtype) - - -class MochiRMSNormZero(nn.Module): - r""" - Adaptive RMS Norm used in Mochi. - - Parameters: - embedding_dim (`int`): The size of each embedding vector. - """ - - def __init__( - self, embedding_dim: int, hidden_dim: int, eps: float = 1e-5, elementwise_affine: bool = False - ) -> None: - super().__init__() - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, hidden_dim) - self.norm = RMSNorm(0, eps, False) - - def forward( - self, hidden_states: torch.Tensor, emb: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - hidden_states_dtype = hidden_states.dtype - - emb = self.linear(self.silu(emb)) - scale_msa, gate_msa, scale_mlp, gate_mlp = emb.chunk(4, dim=1) - hidden_states = self.norm(hidden_states.to(torch.float32)) * (1 + scale_msa[:, None].to(torch.float32)) - hidden_states = hidden_states.to(hidden_states_dtype) - - return hidden_states, gate_msa, scale_mlp, gate_mlp - - -@maybe_allow_in_graph -class MochiTransformerBlock(nn.Module): - r""" - Transformer block used in [Mochi](https://huggingface.co/genmo/mochi-1-preview). - - Args: - dim (`int`): - The number of channels in the input and output. - num_attention_heads (`int`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`): - The number of channels in each head. - qk_norm (`str`, defaults to `"rms_norm"`): - The normalization layer to use. - activation_fn (`str`, defaults to `"swiglu"`): - Activation function to use in feed-forward. - context_pre_only (`bool`, defaults to `False`): - Whether or not to process context-related conditions with additional layers. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - pooled_projection_dim: int, - qk_norm: str = "rms_norm", - activation_fn: str = "swiglu", - context_pre_only: bool = False, - eps: float = 1e-6, - ) -> None: - super().__init__() - - self.context_pre_only = context_pre_only - self.ff_inner_dim = (4 * dim * 2) // 3 - self.ff_context_inner_dim = (4 * pooled_projection_dim * 2) // 3 - - self.norm1 = MochiRMSNormZero(dim, 4 * dim, eps=eps, elementwise_affine=False) - - if not context_pre_only: - self.norm1_context = MochiRMSNormZero(dim, 4 * pooled_projection_dim, eps=eps, elementwise_affine=False) - else: - self.norm1_context = MochiLayerNormContinuous( - embedding_dim=pooled_projection_dim, - conditioning_embedding_dim=dim, - eps=eps, - ) - - self.attn1 = MochiAttention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - bias=False, - added_kv_proj_dim=pooled_projection_dim, - added_proj_bias=False, - out_dim=dim, - out_context_dim=pooled_projection_dim, - context_pre_only=context_pre_only, - processor=MochiAttnProcessor2_0(), - eps=1e-5, - ) - - # TODO(aryan): norm_context layers are not needed when `context_pre_only` is True - self.norm2 = MochiModulatedRMSNorm(eps=eps) - self.norm2_context = MochiModulatedRMSNorm(eps=eps) if not self.context_pre_only else None - - self.norm3 = MochiModulatedRMSNorm(eps) - self.norm3_context = MochiModulatedRMSNorm(eps=eps) if not self.context_pre_only else None - - self.ff = FeedForward(dim, inner_dim=self.ff_inner_dim, activation_fn=activation_fn, bias=False) - self.ff_context = None - if not context_pre_only: - self.ff_context = FeedForward( - pooled_projection_dim, - inner_dim=self.ff_context_inner_dim, - activation_fn=activation_fn, - bias=False, - ) - - self.norm4 = MochiModulatedRMSNorm(eps=eps) - self.norm4_context = MochiModulatedRMSNorm(eps=eps) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - encoder_attention_mask: torch.Tensor, - image_rotary_emb: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(hidden_states, temb) - - if not self.context_pre_only: - norm_encoder_hidden_states, enc_gate_msa, enc_scale_mlp, enc_gate_mlp = self.norm1_context( - encoder_hidden_states, temb - ) - else: - norm_encoder_hidden_states = self.norm1_context(encoder_hidden_states, temb) - - attn_hidden_states, context_attn_hidden_states = self.attn1( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - attention_mask=encoder_attention_mask, - ) - - hidden_states = hidden_states + self.norm2(attn_hidden_states, torch.tanh(gate_msa).unsqueeze(1)) - norm_hidden_states = self.norm3(hidden_states, (1 + scale_mlp.unsqueeze(1).to(torch.float32))) - ff_output = self.ff(norm_hidden_states) - hidden_states = hidden_states + self.norm4(ff_output, torch.tanh(gate_mlp).unsqueeze(1)) - - if not self.context_pre_only: - encoder_hidden_states = encoder_hidden_states + self.norm2_context( - context_attn_hidden_states, torch.tanh(enc_gate_msa).unsqueeze(1) - ) - norm_encoder_hidden_states = self.norm3_context( - encoder_hidden_states, (1 + enc_scale_mlp.unsqueeze(1).to(torch.float32)) - ) - context_ff_output = self.ff_context(norm_encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states + self.norm4_context( - context_ff_output, torch.tanh(enc_gate_mlp).unsqueeze(1) - ) - - return hidden_states, encoder_hidden_states - - -class MochiRoPE(nn.Module): - r""" - RoPE implementation used in [Mochi](https://huggingface.co/genmo/mochi-1-preview). - - Args: - base_height (`int`, defaults to `192`): - Base height used to compute interpolation scale for rotary positional embeddings. - base_width (`int`, defaults to `192`): - Base width used to compute interpolation scale for rotary positional embeddings. - """ - - def __init__(self, base_height: int = 192, base_width: int = 192) -> None: - super().__init__() - - self.target_area = base_height * base_width - - def _centers(self, start, stop, num, device, dtype) -> torch.Tensor: - edges = torch.linspace(start, stop, num + 1, device=device, dtype=dtype) - return (edges[:-1] + edges[1:]) / 2 - - def _get_positions( - self, - num_frames: int, - height: int, - width: int, - device: torch.device | None = None, - dtype: torch.dtype | None = None, - ) -> torch.Tensor: - scale = (self.target_area / (height * width)) ** 0.5 - - t = torch.arange(num_frames, device=device, dtype=dtype) - h = self._centers(-height * scale / 2, height * scale / 2, height, device, dtype) - w = self._centers(-width * scale / 2, width * scale / 2, width, device, dtype) - - grid_t, grid_h, grid_w = torch.meshgrid(t, h, w, indexing="ij") - - positions = torch.stack([grid_t, grid_h, grid_w], dim=-1).view(-1, 3) - return positions - - def _create_rope(self, freqs: torch.Tensor, pos: torch.Tensor) -> torch.Tensor: - with torch.autocast(freqs.device.type, torch.float32): - # Always run ROPE freqs computation in FP32 - freqs = torch.einsum("nd,dhf->nhf", pos.to(torch.float32), freqs.to(torch.float32)) - - freqs_cos = torch.cos(freqs) - freqs_sin = torch.sin(freqs) - return freqs_cos, freqs_sin - - def forward( - self, - pos_frequencies: torch.Tensor, - num_frames: int, - height: int, - width: int, - device: torch.device | None = None, - dtype: torch.dtype | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - pos = self._get_positions(num_frames, height, width, device, dtype) - rope_cos, rope_sin = self._create_rope(pos_frequencies, pos) - return rope_cos, rope_sin - - -@maybe_allow_in_graph -class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin): - r""" - A Transformer model for video-like data introduced in [Mochi](https://huggingface.co/genmo/mochi-1-preview). - - Args: - patch_size (`int`, defaults to `2`): - The size of the patches to use in the patch embedding layer. - num_attention_heads (`int`, defaults to `24`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - num_layers (`int`, defaults to `48`): - The number of layers of Transformer blocks to use. - in_channels (`int`, defaults to `12`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `None`): - The number of channels in the output. - qk_norm (`str`, defaults to `"rms_norm"`): - The normalization layer to use. - text_embed_dim (`int`, defaults to `4096`): - Input dimension of text embeddings from the text encoder. - time_embed_dim (`int`, defaults to `256`): - Output dimension of timestep embeddings. - activation_fn (`str`, defaults to `"swiglu"`): - Activation function to use in feed-forward. - max_sequence_length (`int`, defaults to `256`): - The maximum sequence length of text embeddings supported. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["MochiTransformerBlock"] - _skip_layerwise_casting_patterns = ["patch_embed", "norm"] - - @register_to_config - def __init__( - self, - patch_size: int = 2, - num_attention_heads: int = 24, - attention_head_dim: int = 128, - num_layers: int = 48, - pooled_projection_dim: int = 1536, - in_channels: int = 12, - out_channels: int | None = None, - qk_norm: str = "rms_norm", - text_embed_dim: int = 4096, - time_embed_dim: int = 256, - activation_fn: str = "swiglu", - max_sequence_length: int = 256, - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - self.patch_embed = PatchEmbed( - patch_size=patch_size, - in_channels=in_channels, - embed_dim=inner_dim, - pos_embed_type=None, - ) - - self.time_embed = MochiCombinedTimestepCaptionEmbedding( - embedding_dim=inner_dim, - pooled_projection_dim=pooled_projection_dim, - text_embed_dim=text_embed_dim, - time_embed_dim=time_embed_dim, - num_attention_heads=8, - ) - - self.pos_frequencies = nn.Parameter(torch.full((3, num_attention_heads, attention_head_dim // 2), 0.0)) - self.rope = MochiRoPE() - - self.transformer_blocks = nn.ModuleList( - [ - MochiTransformerBlock( - dim=inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - pooled_projection_dim=pooled_projection_dim, - qk_norm=qk_norm, - activation_fn=activation_fn, - context_pre_only=i == num_layers - 1, - ) - for i in range(num_layers) - ] - ) - - self.norm_out = AdaLayerNormContinuous( - inner_dim, - inner_dim, - elementwise_affine=False, - eps=1e-6, - norm_type="layer_norm", - ) - self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_attention_mask: torch.Tensor, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> torch.Tensor: - """ - The [`MochiTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_attention_mask (`torch.Tensor`): - Mask applied to `encoder_hidden_states` during attention. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - `torch.Tensor`: - The denoised output tensor of shape `(batch_size, out_channels, num_frames, height, width)`. - """ - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p = self.config.patch_size - - post_patch_height = height // p - post_patch_width = width // p - - temb, encoder_hidden_states = self.time_embed( - timestep, - encoder_hidden_states, - encoder_attention_mask, - hidden_dtype=hidden_states.dtype, - ) - - hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) - hidden_states = self.patch_embed(hidden_states) - hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(1, 2) - - image_rotary_emb = self.rope( - self.pos_frequencies, - num_frames, - post_patch_height, - post_patch_width, - device=hidden_states.device, - dtype=torch.float32, - ) - - for i, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states, encoder_hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - encoder_attention_mask, - image_rotary_emb, - ) - else: - hidden_states, encoder_hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - encoder_attention_mask=encoder_attention_mask, - image_rotary_emb=image_rotary_emb, - ) - hidden_states = self.norm_out(hidden_states, temb) - hidden_states = self.proj_out(hidden_states) - - hidden_states = hidden_states.reshape(batch_size, num_frames, post_patch_height, post_patch_width, p, p, -1) - hidden_states = hidden_states.permute(0, 6, 1, 2, 4, 3, 5) - output = hidden_states.reshape(batch_size, -1, num_frames, height, width) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_motif_video.py b/diffusers/models/transformers/transformer_motif_video.py deleted file mode 100644 index fb3ff0666f9561d4353e907eabc6b4b4d0f37667..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_motif_video.py +++ /dev/null @@ -1,1057 +0,0 @@ -# Copyright 2026 Motif Technologies and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -from typing import Any, Dict, List, Optional, Tuple, Union - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import ( - PixArtAlphaTextProjection, - TimestepEmbedding, - Timesteps, - apply_rotary_emb, - get_1d_rotary_pos_embed, -) -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin, get_parameter_dtype -from ..normalization import ( - AdaLayerNormContinuous, - AdaLayerNormZero, - AdaLayerNormZeroSingle, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class MotifVideoCrossAttnProcessor2_0: - """Attention processor for Motif-Video text cross-attention.""" - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "MotifVideoCrossAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: "MotifVideoCrossAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: Optional[torch.Tensor] = None, - image_rotary_emb: Optional[torch.Tensor] = None, - image_embed_seq_len: int = 0, - ) -> torch.Tensor: - txt_kv = encoder_hidden_states[:, image_embed_seq_len:, :] - - text_mask = None - if attention_mask is not None: - text_mask = attention_mask[:, :, :, image_embed_seq_len - encoder_hidden_states.shape[1] :] - - query = attn.to_q(hidden_states) - key = attn.to_k(txt_kv) - value = attn.to_v(txt_kv) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=text_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class MotifVideoAttnProcessor2_0: - """Attention processor for Motif-Video self-attention.""" - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "MotifVideoAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: "MotifVideoAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - image_rotary_emb: Optional[torch.Tensor] = None, - ) -> torch.Tensor: - # Concatenate hidden states with encoder hidden states for joint attention if needed - if attn.add_q_proj is None and encoder_hidden_states is not None: - hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1) - - # Project QKV - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - # Normalize QK - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE - if image_rotary_emb is not None: - if attn.add_q_proj is None and encoder_hidden_states is not None: - split_idx = -encoder_hidden_states.shape[1] - query = torch.cat( - [ - apply_rotary_emb(query[:, :split_idx, :, :], image_rotary_emb, sequence_dim=1), - query[:, split_idx:, :, :], - ], - dim=1, - ) - key = torch.cat( - [ - apply_rotary_emb(key[:, :split_idx, :, :], image_rotary_emb, sequence_dim=1), - key[:, split_idx:, :, :], - ], - dim=1, - ) - else: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - # Add encoder conditioning QKV projections and normalization - if attn.add_q_proj is not None and encoder_hidden_states is not None: - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - encoder_query = encoder_query.unflatten(2, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(2, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(2, (attn.heads, -1)) - - if attn.norm_added_q is not None: - encoder_query = attn.norm_added_q(encoder_query) - if attn.norm_added_k is not None: - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([query, encoder_query], dim=1) - key = torch.cat([key, encoder_key], dim=1) - value = torch.cat([value, encoder_value], dim=1) - - # Compute attention with backend dispatch - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - # Apply output projections and split encoder states - if encoder_hidden_states is not None: - hidden_states, encoder_hidden_states = ( - hidden_states[:, : -encoder_hidden_states.shape[1]], - hidden_states[:, -encoder_hidden_states.shape[1] :], - ) - - if attn.to_out is not None: - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - if attn.to_add_out is not None: - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - return hidden_states, encoder_hidden_states - - if attn.to_out is not None: - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class MotifVideoCrossAttention(nn.Module, AttentionModuleMixin): - """Dedicated cross-attention module for Motif-Video text cross-attention.""" - - _default_processor_cls = MotifVideoCrossAttnProcessor2_0 - _available_processors = [MotifVideoCrossAttnProcessor2_0] - - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - out_bias: bool = True, - eps: float = 1e-5, - qk_norm: str = "rms_norm", - elementwise_affine: bool = True, - processor=None, - ): - super().__init__() - - self.head_dim = dim_head - self.inner_dim = dim_head * heads - self.heads = heads - - self.to_q = nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_v = nn.Linear(query_dim, self.inner_dim, bias=bias) - - if qk_norm == "rms_norm": - self.norm_q = nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - elif qk_norm == "layer_norm": - self.norm_q = nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - else: - self.norm_q = None - self.norm_k = None - - self.to_out = nn.ModuleList( - [ - nn.Linear(self.inner_dim, query_dim, bias=out_bias), - nn.Dropout(dropout), - ] - ) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: Optional[torch.Tensor] = None, - image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, - image_embed_seq_len: int = 0, - ) -> torch.Tensor: - return self.processor( - self, - hidden_states, - encoder_hidden_states, - attention_mask, - image_rotary_emb, - image_embed_seq_len, - ) - - -class MotifVideoAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = MotifVideoAttnProcessor2_0 - _available_processors = [MotifVideoAttnProcessor2_0] - - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - added_kv_proj_dim: int | None = None, - added_proj_bias: bool | None = True, - out_bias: bool = True, - eps: float = 1e-5, - out_dim: int = None, - elementwise_affine: bool = True, - pre_only: bool = False, - context_pre_only: bool = False, - qk_norm: str = "rms_norm", - processor=None, - ): - super().__init__() - - self.head_dim = dim_head - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.query_dim = query_dim - self.out_dim = out_dim if out_dim is not None else query_dim - self.heads = out_dim // dim_head if out_dim is not None else heads - self.pre_only = pre_only - - self.use_bias = bias - self.dropout = dropout - - self.added_kv_proj_dim = added_kv_proj_dim - self.added_proj_bias = added_proj_bias - self.context_pre_only = context_pre_only - - self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_v = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - - # QK Norm - if qk_norm == "rms_norm": - self.norm_q = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - elif qk_norm == "layer_norm": - self.norm_q = torch.nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = torch.nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - else: - self.norm_q = None - self.norm_k = None - - if not pre_only: - self.to_out = torch.nn.ModuleList([]) - self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(torch.nn.Dropout(dropout)) - else: - self.to_out = None - - if added_kv_proj_dim is not None: - self.norm_added_q = torch.nn.RMSNorm(dim_head, eps=eps) - self.norm_added_k = torch.nn.RMSNorm(dim_head, eps=eps) - self.add_q_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - if not context_pre_only: - self.to_add_out = torch.nn.Linear(self.inner_dim, query_dim, bias=out_bias) - else: - self.to_add_out = None - else: - self.norm_added_q = None - self.norm_added_k = None - self.add_q_proj = None - self.add_k_proj = None - self.add_v_proj = None - self.to_add_out = None - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"joint_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -class MotifVideoPatchEmbed(nn.Module): - def __init__( - self, - patch_size: Union[int, Tuple[int, int, int]] = 16, - in_chans: int = 3, - embed_dim: int = 768, - ) -> None: - super().__init__() - - patch_size = (patch_size, patch_size, patch_size) if isinstance(patch_size, int) else patch_size - self.proj = nn.Conv3d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.proj(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) # BCFHW -> BNC - return hidden_states - - -class MotifVideoAdaNorm(nn.Module): - def __init__(self, in_features: int, out_features: Optional[int] = None) -> None: - super().__init__() - - out_features = out_features or 2 * in_features - self.linear = nn.Linear(in_features, out_features) - self.nonlinearity = nn.SiLU() - - def forward(self, temb: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: - temb = self.linear(self.nonlinearity(temb)) - gate_msa, gate_mlp = temb.chunk(2, dim=1) - gate_msa, gate_mlp = gate_msa.unsqueeze(1), gate_mlp.unsqueeze(1) - return gate_msa, gate_mlp - - -class MotifVideoConditionEmbedding(nn.Module): - def __init__( - self, - embedding_dim: int, - ): - super().__init__() - - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - def forward( - self, - timestep: torch.Tensor, - ) -> torch.Tensor: - timesteps_proj = self.time_proj(timestep) - param_dtype = get_parameter_dtype(self.timestep_embedder) - # Timesteps always returns FP32 output, so cast to the weight dtype of timestep_embedder if we're operating in - # FP16 or BF16 (and no quantization) - if param_dtype in (torch.float16, torch.bfloat16): - timesteps_proj = timesteps_proj.to(param_dtype) - conditioning = self.timestep_embedder(timesteps_proj) # (N, D) - - return conditioning - - -class MotifVideoRotaryPosEmbed(nn.Module): - def __init__( - self, - patch_size: int, - patch_size_t: int, - rope_dim: List[int], - theta: float = 256.0, - ): - """ - Rotary Positional Embedding (RoPE) for video latents. - - Args: - patch_size (`int`): Spatial patch size. - patch_size_t (`int`): Temporal patch size. - rope_dim (`List[int]`): Dimensions for RoPE across [Time, Height, Width] axes. - theta (`float`, *optional*, defaults to 256.0): Base frequency for rotary embeddings. - """ - super().__init__() - - self.patch_size = patch_size - self.patch_size_t = patch_size_t - self.rope_dim = rope_dim - self.theta = theta - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - rope_sizes = [ - num_frames // self.patch_size_t, - height // self.patch_size, - width // self.patch_size, - ] - - axes_grids = [] - for i in range(3): - grid = torch.arange(0, rope_sizes[i], device=hidden_states.device, dtype=torch.float32) - axes_grids.append(grid) - grid = torch.meshgrid(*axes_grids, indexing="ij") - grid = torch.stack(grid, dim=0) - - freqs = [] - freqs_dtype = maybe_adjust_dtype_for_device(torch.float64, hidden_states.device) - for i in range(3): - freq = get_1d_rotary_pos_embed( - dim=self.rope_dim[i], - pos=grid[i].reshape(-1), - theta=self.theta, - use_real=True, - freqs_dtype=freqs_dtype, - ) - freqs.append(freq) - - freqs_cos = torch.cat([f[0] for f in freqs], dim=1) - freqs_sin = torch.cat([f[1] for f in freqs], dim=1) - return freqs_cos, freqs_sin - - -class MotifVideoImageProjection(nn.Module): - def __init__(self, in_features: int, hidden_size: int): - super().__init__() - self.norm_in = nn.LayerNorm(in_features) - self.linear_1 = nn.Linear(in_features, in_features) - self.act_fn = nn.GELU() - self.linear_2 = nn.Linear(in_features, hidden_size) - self.norm_out = nn.LayerNorm(hidden_size) - - def forward(self, image_embeds: torch.Tensor) -> torch.Tensor: - hidden_states = self.norm_in(image_embeds) - hidden_states = self.linear_1(hidden_states) - hidden_states = self.act_fn(hidden_states) - hidden_states = self.linear_2(hidden_states) - hidden_states = self.norm_out(hidden_states) - return hidden_states - - -class MotifVideoSingleTransformerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float = 4.0, - qk_norm: str = "rms_norm", - norm_type: str = "layer_norm", - enable_text_cross_attention: bool = False, - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - mlp_dim = int(hidden_size * mlp_ratio) - - self.attn = MotifVideoAttention( - query_dim=hidden_size, - heads=num_attention_heads, - dim_head=attention_head_dim, - out_dim=hidden_size, - bias=True, - pre_only=True, - qk_norm=qk_norm, - eps=1e-6, - processor=MotifVideoAttnProcessor2_0(), - ) - - self.cross_attn = ( - MotifVideoCrossAttention( - query_dim=hidden_size, - heads=num_attention_heads, - dim_head=attention_head_dim, - bias=True, - qk_norm=qk_norm, - eps=1e-6, - ) - if enable_text_cross_attention - else None - ) - - self.enable_text_cross_attention = enable_text_cross_attention - - self.norm = AdaLayerNormZeroSingle(hidden_size, norm_type=norm_type) - self.proj_mlp = nn.Linear(hidden_size, mlp_dim) - self.act_mlp = nn.GELU(approximate="tanh") - self.proj_out = nn.Linear(hidden_size + mlp_dim, hidden_size) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: Optional[torch.Tensor] = None, - image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, - image_embed_seq_len: int = 0, - ) -> torch.Tensor: - encoder_seq_length = encoder_hidden_states.shape[1] - hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1) - - residual = hidden_states - - # 1. Input normalization - norm_hidden_states, gate = self.norm(hidden_states, emb=temb) - mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) - - norm_hidden_states, norm_encoder_hidden_states = ( - norm_hidden_states[:, :-encoder_seq_length, :], - norm_hidden_states[:, -encoder_seq_length:, :], - ) - - # 2. Attention - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - - # 3. Text cross-attention - if self.cross_attn is not None: - cross_output = self.cross_attn( - hidden_states=attn_output, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - image_embed_seq_len=image_embed_seq_len, - ) - attn_output = attn_output + cross_output - - attn_output = torch.cat([attn_output, context_attn_output], dim=1) - - # 4. Modulation and residual connection - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) - hidden_states = gate.unsqueeze(1) * self.proj_out(hidden_states) - hidden_states = hidden_states + residual - - hidden_states, encoder_hidden_states = ( - hidden_states[:, :-encoder_seq_length, :], - hidden_states[:, -encoder_seq_length:, :], - ) - return hidden_states, encoder_hidden_states - - -class MotifVideoTransformerBlock(nn.Module): - def __init__( - self, - num_attention_heads: int, - attention_head_dim: int, - mlp_ratio: float, - qk_norm: str = "rms_norm", - norm_type: str = "layer_norm", - enable_text_cross_attention: bool = False, - ) -> None: - super().__init__() - - hidden_size = num_attention_heads * attention_head_dim - - self.norm1 = AdaLayerNormZero(hidden_size, norm_type=norm_type) - self.norm1_context = AdaLayerNormZero(hidden_size, norm_type=norm_type) - - self.attn = MotifVideoAttention( - query_dim=hidden_size, - added_kv_proj_dim=hidden_size, - heads=num_attention_heads, - dim_head=attention_head_dim, - out_dim=hidden_size, - bias=True, - context_pre_only=False, - qk_norm=qk_norm, - eps=1e-6, - processor=MotifVideoAttnProcessor2_0(), - ) - - self.cross_attn = ( - MotifVideoCrossAttention( - query_dim=hidden_size, - heads=num_attention_heads, - dim_head=attention_head_dim, - bias=True, - qk_norm=qk_norm, - eps=1e-6, - ) - if enable_text_cross_attention - else None - ) - - self.enable_text_cross_attention = enable_text_cross_attention - - self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.norm2_context = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - - self.ff = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate") - self.ff_context = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - attention_mask: Optional[torch.Tensor] = None, - image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, - image_embed_seq_len: int = 0, - ) -> Tuple[torch.Tensor, torch.Tensor]: - # 1. Input normalization - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - - # 2. Joint attention - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - - # 3. Modulation and residual connection - hidden_states = hidden_states + attn_output * gate_msa.unsqueeze(1) - - # 4. Text cross-attention - if self.cross_attn is not None: - cross_output = self.cross_attn( - hidden_states=attn_output, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - image_embed_seq_len=image_embed_seq_len, - ) - hidden_states = hidden_states + cross_output - - encoder_hidden_states = encoder_hidden_states + context_attn_output * c_gate_msa.unsqueeze(1) - - norm_hidden_states = self.norm2(hidden_states) - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - # 5. Feed-forward - ff_output = self.ff(norm_hidden_states) - context_ff_output = self.ff_context(norm_encoder_hidden_states) - - hidden_states = hidden_states + gate_mlp.unsqueeze(1) * ff_output - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - - return hidden_states, encoder_hidden_states - - -class MotifVideoTransformer3DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin, AttentionMixin -): - r""" - A Transformer model for video-like data used in the Motif-Video model. - - Args: - in_channels (`int`, defaults to `33`): - The number of channels in the input. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - num_attention_heads (`int`, defaults to `24`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - num_layers (`int`, defaults to `20`): - The number of layers of dual-stream blocks to use. - num_single_layers (`int`, defaults to `40`): - The number of layers of single-stream blocks to use. - num_decoder_layers (`int`, defaults to `0`): - The number of decoder layers in single-stream blocks. - mlp_ratio (`float`, defaults to `4.0`): - The ratio of the hidden layer size to the input size in the feedforward network. - patch_size (`int`, defaults to `2`): - The size of the spatial patches to use in the patch embedding layer. - patch_size_t (`int`, defaults to `1`): - The size of the temporal patches to use in the patch embedding layer. - qk_norm (`str`, defaults to `rms_norm`): - The normalization to use for the query and key projections in the attention layers. - text_embed_dim (`int`, defaults to `4096`): - Input dimension of text embeddings from the text encoder. - image_embed_dim (`int`, *optional*): - Input dimension of image embeddings from a vision encoder. If provided, enables image conditioning. - rope_theta (`float`, defaults to `256.0`): - The value of theta to use in the RoPE layer. - rope_axes_dim (`Tuple[int]`, defaults to `(16, 56, 56)`): - The dimensions of the axes to use in the RoPE layer. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["x_embedder", "context_embedder", "norm"] - _repeated_blocks = ["MotifVideoSingleTransformerBlock", "MotifVideoTransformerBlock"] - _no_split_modules = [ - "MotifVideoTransformerBlock", - "MotifVideoSingleTransformerBlock", - "MotifVideoPatchEmbed", - ] - - @register_to_config - def __init__( - self, - in_channels: int = 33, - out_channels: int = 16, - num_attention_heads: int = 24, - attention_head_dim: int = 128, - num_layers: int = 20, - num_single_layers: int = 40, - num_decoder_layers: int = 0, - mlp_ratio: float = 4.0, - patch_size: int = 2, - patch_size_t: int = 1, - qk_norm: str = "rms_norm", - norm_type: str = "layer_norm", - text_embed_dim: int = 4096, - image_embed_dim: int | None = None, - rope_theta: float = 256.0, - rope_axes_dim: Tuple[int, ...] = (16, 56, 56), - enable_text_cross_attention_dual: bool = False, - enable_text_cross_attention_single: bool = False, - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - # 1. Latent and condition embedders - self.x_embedder = MotifVideoPatchEmbed((patch_size_t, patch_size, patch_size), in_channels, inner_dim) - self.context_embedder = PixArtAlphaTextProjection(in_features=text_embed_dim, hidden_size=inner_dim) - - # First frame conditioning: Image conditioning embedders - self.image_embed_dim = image_embed_dim - if image_embed_dim is not None: - self.image_embedder = MotifVideoImageProjection(in_features=image_embed_dim, hidden_size=inner_dim) - - self.time_text_embed = MotifVideoConditionEmbedding(inner_dim) - - # 2. RoPE - self.rope = MotifVideoRotaryPosEmbed(patch_size, patch_size_t, rope_axes_dim, rope_theta) - - # Cross-attention config - self.enable_text_cross_attention_dual = enable_text_cross_attention_dual - self.enable_text_cross_attention_single = enable_text_cross_attention_single - - # 3. Dual stream transformer blocks - self.transformer_blocks = nn.ModuleList( - [ - MotifVideoTransformerBlock( - num_attention_heads, - attention_head_dim, - mlp_ratio=mlp_ratio, - qk_norm=qk_norm, - norm_type=norm_type, - enable_text_cross_attention=enable_text_cross_attention_dual, - ) - for _ in range(num_layers) - ] - ) - - # 4. Single stream transformer blocks - # Encoder blocks get cross-attention; decoder blocks do not (no text stream in decoder) - num_encoder_single = num_single_layers - num_decoder_layers - self.single_transformer_blocks = nn.ModuleList( - [ - MotifVideoSingleTransformerBlock( - num_attention_heads, - attention_head_dim, - mlp_ratio=mlp_ratio, - qk_norm=qk_norm, - norm_type=norm_type, - enable_text_cross_attention=enable_text_cross_attention_single - if i < num_encoder_single - else False, - ) - for i in range(num_single_layers) - ] - ) - - # 5. Output projection - self.norm_out = AdaLayerNormContinuous( - inner_dim, - inner_dim, - elementwise_affine=False, - eps=1e-6, - norm_type=norm_type, - ) - self.proj_out = nn.Linear(inner_dim, patch_size_t * patch_size * patch_size * out_channels) - - # Verify cross-attention config matches actual block state. - # Catches silent misconfiguration (e.g. checkpoint config with renamed keys). - for i, block in enumerate(self.transformer_blocks): - if block.enable_text_cross_attention != enable_text_cross_attention_dual: - raise ValueError( - f"transformer_blocks[{i}].enable_text_cross_attention=" - f"{block.enable_text_cross_attention}, expected {enable_text_cross_attention_dual}. " - f"Check checkpoint config.json key names match __init__ parameters." - ) - for i, block in enumerate(self.single_transformer_blocks): - expected = enable_text_cross_attention_single if i < num_encoder_single else False - if block.enable_text_cross_attention != expected: - raise ValueError( - f"single_transformer_blocks[{i}].enable_text_cross_attention=" - f"{block.enable_text_cross_attention}, expected {expected}. " - f"Check checkpoint config.json key names match __init__ parameters." - ) - - self.gradient_checkpointing = False - self.num_decoder_layers = num_decoder_layers - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_attention_mask: torch.Tensor | None = None, - image_embeds: torch.Tensor | None = None, - attention_kwargs: Optional[Dict[str, Any]] = None, - return_dict: bool = True, - ) -> Union[torch.Tensor, Dict[str, torch.Tensor]]: - """ - Forward pass of the MotifVideoTransformer3DModel. - - Args: - hidden_states (`torch.Tensor`): - Input latent tensor of shape `(batch_size, channels, num_frames, height, width)`. - timestep (`torch.LongTensor`): - Diffusion timesteps of shape `(batch_size,)`. - encoder_hidden_states (`torch.Tensor`): - Text conditioning of shape `(batch_size, sequence_length, embed_dim)`. - encoder_attention_mask (`torch.Tensor`): - Mask for text conditioning of shape `(batch_size, sequence_length)`. - image_embeds (`torch.Tensor`, *optional*): - Image embeddings from vision encoder of shape `(batch_size, num_tokens, embed_dim)`. - attention_kwargs (`dict`, *optional*): - Additional arguments for attention processors. - return_dict (`bool`, defaults to `True`): - Whether to return a [`~models.modeling_outputs.Transformer2DModelOutput`]. - - Returns: - [`~models.modeling_outputs.Transformer2DModelOutput`] or `tuple`: - The predicted samples. - """ - if attention_kwargs is not None: - attention_kwargs = attention_kwargs.copy() - lora_scale = attention_kwargs.pop("scale", 1.0) - else: - lora_scale = 1.0 - - if USE_PEFT_BACKEND: - scale_lora_layers(self, lora_scale) - else: - if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None: - logger.warning( - "Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective." - ) - - batch_size, _, num_frames, height, width = hidden_states.shape - p, p_t = self.config.patch_size, self.config.patch_size_t - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p - post_patch_width = width // p - - # 1. RoPE - image_rotary_emb = self.rope(hidden_states) - - # 2. Conditional embeddings - temb = self.time_text_embed(timestep) - hidden_states = self.x_embedder(hidden_states) - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - - # First frame conditioning: Image embeddings from vision encoder - if image_embeds is not None: - image_embeds = self.image_embedder(image_embeds) - encoder_hidden_states = torch.cat([image_embeds, encoder_hidden_states], dim=1) - if encoder_attention_mask is not None: - image_mask = torch.ones( - image_embeds.shape[0], - image_embeds.shape[1], - device=encoder_attention_mask.device, - dtype=encoder_attention_mask.dtype, - ) - encoder_attention_mask = torch.cat([image_mask, encoder_attention_mask], dim=1) - - # image_embed_seq_len: used by cross-attention blocks to slice text from encoder_hidden_states - image_embed_seq_len = image_embeds.shape[1] if image_embeds is not None else 0 - - if self.num_decoder_layers > 0: - decoder_hidden_states = hidden_states.clone() - - if encoder_attention_mask is not None: - attention_mask = F.pad( - encoder_attention_mask.to(torch.bool), - (hidden_states.shape[1], 0), - value=True, - ) - attention_mask = attention_mask.unsqueeze(1).unsqueeze(1) - else: - attention_mask = None - - # 3. Dual stream transformer blocks - for block in self.transformer_blocks: - hidden_states, encoder_hidden_states = ( - self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - attention_mask, - image_rotary_emb, - image_embed_seq_len, - ) - if torch.is_grad_enabled() and self.gradient_checkpointing - else block( - hidden_states, encoder_hidden_states, temb, attention_mask, image_rotary_emb, image_embed_seq_len - ) - ) - - # 4. Single stream transformer blocks (Encoder) - single_transformer_blocks = self.single_transformer_blocks - - for block in single_transformer_blocks[: len(single_transformer_blocks) - self.num_decoder_layers]: - hidden_states, encoder_hidden_states = ( - self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - attention_mask, - image_rotary_emb, - image_embed_seq_len, - ) - if torch.is_grad_enabled() and self.gradient_checkpointing - else block( - hidden_states, encoder_hidden_states, temb, attention_mask, image_rotary_emb, image_embed_seq_len - ) - ) - - # 5. Single stream transformer blocks (Decoder) - if self.num_decoder_layers > 0: - encoder_hidden_states = hidden_states - attention_mask = None - - for block in single_transformer_blocks[-self.num_decoder_layers :]: - decoder_hidden_states, encoder_hidden_states = ( - self._gradient_checkpointing_func( - block, decoder_hidden_states, encoder_hidden_states, temb, attention_mask, image_rotary_emb - ) - if torch.is_grad_enabled() and self.gradient_checkpointing - else block(decoder_hidden_states, encoder_hidden_states, temb, attention_mask, image_rotary_emb) - ) - - hidden_states = decoder_hidden_states - - # 6. Output projection - hidden_states = self.norm_out(hidden_states, temb) - hidden_states = self.proj_out(hidden_states) - - hidden_states = hidden_states.reshape( - batch_size, - post_patch_num_frames, - post_patch_height, - post_patch_width, - -1, - p_t, - p, - p, - ) - hidden_states = hidden_states.permute(0, 4, 1, 5, 2, 6, 3, 7) - hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if USE_PEFT_BACKEND: - unscale_lora_layers(self, lora_scale) - - if not return_dict: - return (hidden_states,) - - return Transformer2DModelOutput( - sample=hidden_states, - ) diff --git a/diffusers/models/transformers/transformer_nucleusmoe_image.py b/diffusers/models/transformers/transformer_nucleusmoe_image.py deleted file mode 100644 index f1c0eee949f797600805e2870198ea3296e1298f..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_nucleusmoe_image.py +++ /dev/null @@ -1,925 +0,0 @@ -# Copyright 2025 Nucleus-Image Team, The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import functools -import math -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers -from ..attention import AttentionMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..attention_processor import Attention -from ..cache_utils import CacheMixin -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, RMSNorm - - -logger = logging.get_logger(__name__) - - -# Copied from diffusers.models.transformers.transformer_qwenimage.apply_rotary_emb_qwen with qwen->nucleus -def _apply_rotary_emb_nucleus( - x: torch.Tensor, - freqs_cis: torch.Tensor | tuple[torch.Tensor], - use_real: bool = True, - use_real_unbind_dim: int = -1, -) -> tuple[torch.Tensor, torch.Tensor]: - """ - Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings - to the given query or key 'x' tensors using the provided frequency tensor 'freqs_cis'. The input tensors are - reshaped as complex numbers, and the frequency tensor is reshaped for broadcasting compatibility. The resulting - tensors contain rotary embeddings and are returned as real tensors. - - Args: - x (`torch.Tensor`): - Query or key tensor to apply rotary embeddings. [B, S, H, D] xk (torch.Tensor): Key tensor to apply - freqs_cis (`tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],) - - Returns: - tuple[torch.Tensor, torch.Tensor]: tuple of modified query tensor and key tensor with rotary embeddings. - """ - if use_real: - cos, sin = freqs_cis # [S, D] - cos = cos[None, None] - sin = sin[None, None] - cos, sin = cos.to(x.device), sin.to(x.device) - - if use_real_unbind_dim == -1: - # Used for flux, cogvideox, hunyuan-dit - x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2] - x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3) - elif use_real_unbind_dim == -2: - # Used for Stable Audio, OmniGen, CogView4 and Cosmos - x_real, x_imag = x.reshape(*x.shape[:-1], 2, -1).unbind(-2) # [B, S, H, D//2] - x_rotated = torch.cat([-x_imag, x_real], dim=-1) - else: - raise ValueError(f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2.") - - out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype) - - return out - else: - x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2)) - freqs_cis = freqs_cis.unsqueeze(1) - x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3) - - return x_out.type_as(x) - - -def _compute_text_seq_len_from_mask( - encoder_hidden_states: torch.Tensor, encoder_hidden_states_mask: torch.Tensor | None -) -> tuple[int, torch.Tensor | None, torch.Tensor | None]: - batch_size, text_seq_len = encoder_hidden_states.shape[:2] - if encoder_hidden_states_mask is None: - return text_seq_len, None, None - - if encoder_hidden_states_mask.shape[:2] != (batch_size, text_seq_len): - raise ValueError( - f"`encoder_hidden_states_mask` shape {encoder_hidden_states_mask.shape} must match " - f"(batch_size, text_seq_len)=({batch_size}, {text_seq_len})." - ) - - if encoder_hidden_states_mask.dtype != torch.bool: - encoder_hidden_states_mask = encoder_hidden_states_mask.to(torch.bool) - - position_ids = torch.arange(text_seq_len, device=encoder_hidden_states.device, dtype=torch.long) - active_positions = torch.where(encoder_hidden_states_mask, position_ids, position_ids.new_zeros(())) - has_active = encoder_hidden_states_mask.any(dim=1) - per_sample_len = torch.where( - has_active, - active_positions.max(dim=1).values + 1, - torch.as_tensor(text_seq_len, device=encoder_hidden_states.device), - ) - return text_seq_len, per_sample_len, encoder_hidden_states_mask - - -class NucleusMoETimestepProjEmbeddings(nn.Module): - def __init__(self, embedding_dim, use_additional_t_cond=False): - super().__init__() - - self.time_proj = Timesteps( - num_channels=embedding_dim, flip_sin_to_cos=True, downscale_freq_shift=0, scale=1000 - ) - self.timestep_embedder = TimestepEmbedding( - in_channels=embedding_dim, time_embed_dim=4 * embedding_dim, out_dim=embedding_dim - ) - self.norm = RMSNorm(embedding_dim, eps=1e-6) - self.use_additional_t_cond = use_additional_t_cond - if use_additional_t_cond: - self.addition_t_embedding = nn.Embedding(2, embedding_dim) - - def forward(self, timestep, hidden_states, addition_t_cond=None): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_states.dtype)) - - conditioning = timesteps_emb - if self.use_additional_t_cond: - if addition_t_cond is None: - raise ValueError("When additional_t_cond is True, addition_t_cond must be provided.") - addition_t_emb = self.addition_t_embedding(addition_t_cond) - addition_t_emb = addition_t_emb.to(dtype=hidden_states.dtype) - conditioning = conditioning + addition_t_emb - - return self.norm(conditioning) - - -class NucleusMoEEmbedRope(nn.Module): - def __init__(self, theta: int, axes_dim: list[int], scale_rope=False): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - pos_index = torch.arange(4096) - neg_index = torch.arange(4096).flip(0) * -1 - 1 - self.pos_freqs = torch.cat( - [ - self._rope_params(pos_index, self.axes_dim[0], self.theta), - self._rope_params(pos_index, self.axes_dim[1], self.theta), - self._rope_params(pos_index, self.axes_dim[2], self.theta), - ], - dim=1, - ) - self.neg_freqs = torch.cat( - [ - self._rope_params(neg_index, self.axes_dim[0], self.theta), - self._rope_params(neg_index, self.axes_dim[1], self.theta), - self._rope_params(neg_index, self.axes_dim[2], self.theta), - ], - dim=1, - ) - - self.scale_rope = scale_rope - - @staticmethod - def _rope_params(index, dim, theta=10000): - assert dim % 2 == 0 - freqs = torch.outer(index, 1.0 / torch.pow(theta, torch.arange(0, dim, 2).to(torch.float32).div(dim))) - freqs = torch.polar(torch.ones_like(freqs), freqs) - return freqs - - def forward( - self, - video_fhw: tuple[int, int, int] | list[tuple[int, int, int]], - device: torch.device = None, - max_txt_seq_len: int | torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - """ - Args: - video_fhw (`tuple[int, int, int]` or `list[tuple[int, int, int]]`): - A list of 3 integers [frame, height, width] representing the shape of the video. - device: (`torch.device`, *optional*): - The device on which to perform the RoPE computation. - max_txt_seq_len (`int` or `torch.Tensor`, *optional*): - The maximum text sequence length for RoPE computation. - """ - if max_txt_seq_len is None: - raise ValueError("Either `max_txt_seq_len` must be provided.") - - if isinstance(video_fhw, list) and len(video_fhw) > 1: - first_fhw = video_fhw[0] - if not all(fhw == first_fhw for fhw in video_fhw): - logger.warning( - "Batch inference with variable-sized images is not currently supported in NucleusMoEEmbedRope. " - "All images in the batch should have the same dimensions (frame, height, width). " - f"Detected sizes: {video_fhw}. Using the first image's dimensions {first_fhw} " - "for RoPE computation, which may lead to incorrect results for other images in the batch." - ) - - if isinstance(video_fhw, list): - video_fhw = video_fhw[0] - if not isinstance(video_fhw, list): - video_fhw = [video_fhw] - - vid_freqs = [] - for idx, fhw in enumerate(video_fhw): - frame, height, width = fhw - video_freq = self._compute_video_freqs(frame, height, width, idx, device) - vid_freqs.append(video_freq) - - max_txt_seq_len_int = int(max_txt_seq_len) - if self.scale_rope: - max_vid_index = torch.maximum( - torch.tensor(height // 2, device=device, dtype=torch.long), - torch.tensor(width // 2, device=device, dtype=torch.long), - ) - else: - max_vid_index = torch.maximum( - torch.tensor(height, device=device, dtype=torch.long), - torch.tensor(width, device=device, dtype=torch.long), - ) - - txt_freqs = self.pos_freqs.to(device)[max_vid_index + torch.arange(max_txt_seq_len_int, device=device)] - vid_freqs = torch.cat(vid_freqs, dim=0) - - return vid_freqs, txt_freqs - - @functools.lru_cache(maxsize=128) - def _compute_video_freqs( - self, frame: int, height: int, width: int, idx: int = 0, device: torch.device = None - ) -> torch.Tensor: - seq_lens = frame * height * width - pos_freqs = self.pos_freqs.to(device) if device is not None else self.pos_freqs - neg_freqs = self.neg_freqs.to(device) if device is not None else self.neg_freqs - - freqs_pos = pos_freqs.split([x // 2 for x in self.axes_dim], dim=1) - freqs_neg = neg_freqs.split([x // 2 for x in self.axes_dim], dim=1) - - freqs_frame = freqs_pos[0][idx : idx + frame].view(frame, 1, 1, -1).expand(frame, height, width, -1) - if self.scale_rope: - freqs_height = torch.cat([freqs_neg[1][-(height - height // 2) :], freqs_pos[1][: height // 2]], dim=0) - freqs_height = freqs_height.view(1, height, 1, -1).expand(frame, height, width, -1) - freqs_width = torch.cat([freqs_neg[2][-(width - width // 2) :], freqs_pos[2][: width // 2]], dim=0) - freqs_width = freqs_width.view(1, 1, width, -1).expand(frame, height, width, -1) - else: - freqs_height = freqs_pos[1][:height].view(1, height, 1, -1).expand(frame, height, width, -1) - freqs_width = freqs_pos[2][:width].view(1, 1, width, -1).expand(frame, height, width, -1) - - freqs = torch.cat([freqs_frame, freqs_height, freqs_width], dim=-1).reshape(seq_lens, -1) - return freqs.clone().contiguous() - - -class NucleusMoEAttnProcessor2_0: - """ - Attention processor for the NucleusMoE architecture. Image queries attend to concatenated image+text keys/values - (cross-attention style, no text query). Supports grouped-query attention (GQA) when num_key_value_heads is set on - the Attention module. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "NucleusMoEAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states: torch.FloatTensor = None, - attention_mask: torch.FloatTensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - cached_txt_key: torch.FloatTensor | None = None, - cached_txt_value: torch.FloatTensor | None = None, - ) -> torch.FloatTensor: - head_dim = attn.inner_dim // attn.heads - num_kv_heads = attn.inner_kv_dim // head_dim - num_kv_groups = attn.heads // num_kv_heads - - img_query = attn.to_q(hidden_states).unflatten(-1, (attn.heads, -1)) - img_key = attn.to_k(hidden_states).unflatten(-1, (num_kv_heads, -1)) - img_value = attn.to_v(hidden_states).unflatten(-1, (num_kv_heads, -1)) - - if attn.norm_q is not None: - img_query = attn.norm_q(img_query) - if attn.norm_k is not None: - img_key = attn.norm_k(img_key) - - if image_rotary_emb is not None: - img_freqs, txt_freqs = image_rotary_emb - img_query = _apply_rotary_emb_nucleus(img_query, img_freqs, use_real=False) - img_key = _apply_rotary_emb_nucleus(img_key, img_freqs, use_real=False) - - if cached_txt_key is not None and cached_txt_value is not None: - txt_key, txt_value = cached_txt_key, cached_txt_value - joint_key = torch.cat([img_key, txt_key], dim=1) - joint_value = torch.cat([img_value, txt_value], dim=1) - elif encoder_hidden_states is not None: - txt_key = attn.add_k_proj(encoder_hidden_states).unflatten(-1, (num_kv_heads, -1)) - txt_value = attn.add_v_proj(encoder_hidden_states).unflatten(-1, (num_kv_heads, -1)) - - if attn.norm_added_k is not None: - txt_key = attn.norm_added_k(txt_key) - - if image_rotary_emb is not None: - txt_key = _apply_rotary_emb_nucleus(txt_key, txt_freqs, use_real=False) - - joint_key = torch.cat([img_key, txt_key], dim=1) - joint_value = torch.cat([img_value, txt_value], dim=1) - else: - joint_key = img_key - joint_value = img_value - - if num_kv_groups > 1: - joint_key = joint_key.repeat_interleave(num_kv_groups, dim=2) - joint_value = joint_value.repeat_interleave(num_kv_groups, dim=2) - - hidden_states = dispatch_attention_fn( - img_query, - joint_key, - joint_value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(img_query.dtype) - - hidden_states = attn.to_out[0](hidden_states) - if len(attn.to_out) > 1: - hidden_states = attn.to_out[1](hidden_states) - - return hidden_states - - -def _is_moe_layer(strategy: str, layer_idx: int, num_layers: int) -> bool: - if strategy == "leave_first_three_and_last_block_dense": - return layer_idx >= 3 and layer_idx < num_layers - 1 - elif strategy == "leave_first_three_blocks_dense": - return layer_idx >= 3 - elif strategy == "leave_first_block_dense": - return layer_idx >= 1 - elif strategy == "all_moe": - return True - elif strategy == "all_dense": - return False - return True - - -class SwiGLUExperts(nn.Module): - """ - Packed SwiGLU feed-forward experts for MoE: ``gate, up = (x @ gate_up_proj).chunk(2); out = (silu(gate) * up) @ - down_proj``. - - Gate and up projections are fused into a single weight ``gate_up_proj`` so that only two grouped matmuls are needed - at runtime (gate+up combined, then down). - - Weights are stored pre-transposed relative to the standard linear-layer convention so that matmuls can be issued - without a transpose at runtime. - - Weight shapes: - gate_up_proj: (num_experts, hidden_size, 2 * moe_intermediate_dim) -- fused gate + up projection down_proj: - (num_experts, moe_intermediate_dim, hidden_size) -- down projection - """ - - def __init__( - self, - hidden_size: int, - moe_intermediate_dim: int, - num_experts: int, - use_grouped_mm: bool = False, - ): - super().__init__() - self.num_experts = num_experts - self.moe_intermediate_dim = moe_intermediate_dim - self.hidden_size = hidden_size - self.use_grouped_mm = use_grouped_mm - - self.gate_up_proj = nn.Parameter(torch.empty(num_experts, hidden_size, 2 * moe_intermediate_dim)) - self.down_proj = nn.Parameter(torch.empty(num_experts, moe_intermediate_dim, hidden_size)) - - def _run_experts_for_loop( - self, - x: torch.Tensor, - num_tokens_per_expert: torch.Tensor, - ) -> torch.Tensor: - """ - Compute SwiGLU MoE expert outputs using a sequential per-expert for loop. - - Tokens in ``x`` must be pre-sorted so that all tokens assigned to expert 0 come first, followed by expert 1, - and so on — i.e. the layout produced by a standard token-permutation step (e.g. ``generate_permute_indices``). - - ``x`` may contain trailing padding rows appended by the permutation utility to reach a length that is a - multiple of some alignment requirement. The padding rows are stripped before expert computation and re-appended - as zeros so that the output shape matches ``x.shape``, keeping downstream scatter/gather indices valid. - - .. note:: - ``num_tokens_per_expert.tolist()`` synchronises the device with the host. This is acceptable for the loop - path but means the method introduces a pipeline bubble. Use :meth:`forward` with ``use_grouped_mm=True`` - when a fully device-resident kernel is required (e.g. inside ``torch.compile``). - - SwiGLU formula:: - - gate, up = (x @ gate_up_proj).chunk(2) out = (silu(gate) * up) @ down_proj - - Args: - x (Tensor): Pre-permuted input tokens of shape - ``(total_tokens_including_padding, hidden_dim)``. - num_tokens_per_expert (Tensor): 1-D integer tensor of length - ``num_experts`` giving the number of real (non-padding) tokens assigned to each expert. Values may - differ across experts to support load-imbalanced routing. - - Returns: - Tensor of shape ``(total_tokens_including_padding, hidden_dim)``. Positions corresponding to padding rows - contain zeros. - """ - # .tolist() triggers a host-device sync; see docstring note above. - num_tokens_per_expert_list = num_tokens_per_expert.tolist() - - # x may be padded to a larger buffer size by the permutation utility. - # Track the padding count so we can restore the original buffer shape. - num_real_tokens = sum(num_tokens_per_expert_list) - num_padding = x.shape[0] - num_real_tokens - - # Split the real-token prefix of x into per-expert slices (variable length). - x_per_expert = torch.split( - x[:num_real_tokens], - split_size_or_sections=num_tokens_per_expert_list, - dim=0, - ) - - expert_outputs = [] - for expert_idx, x_expert in enumerate(x_per_expert): - gate_up = torch.matmul(x_expert, self.gate_up_proj[expert_idx]) - gate, up = gate_up.chunk(2, dim=-1) - out_expert = torch.matmul(F.silu(gate) * up, self.down_proj[expert_idx]) - expert_outputs.append(out_expert) - - # Concatenate real-token outputs, then re-append zero rows for the padding. - out = torch.cat(expert_outputs, dim=0) - out = torch.vstack((out, out.new_zeros((num_padding, out.shape[-1])))) - return out - - def _run_experts_grouped_mm( - self, - x: torch.Tensor, - num_tokens_per_expert: torch.Tensor, - ) -> torch.Tensor: - """ - Compute SwiGLU MoE expert outputs using fused grouped GEMM kernels. - - Tokens in ``x`` must be pre-sorted so that all tokens assigned to expert 0 come first, followed by expert 1, - and so on — the same layout required by :meth:`_run_experts_for_loop`. - - This method is fully device-resident (no host-device sync) and is compatible with ``torch.compile``. - - ``F.grouped_mm`` is called with *exclusive end* offsets: ``offsets[k]`` is the exclusive end index of expert - ``k``'s token range in ``x`` (equivalently the inclusive start of expert ``k+1``'s range). This is the - cumulative sum of ``num_tokens_per_expert``. - - SwiGLU formula:: - - gate, up = (x @ gate_up_proj).chunk(2) out = (silu(gate) * up) @ down_proj - - Args: - x (Tensor): Pre-permuted input tokens of shape - ``(total_tokens, hidden_dim)``. No padding rows expected; ``total_tokens`` must equal - ``num_tokens_per_expert.sum()``. - num_tokens_per_expert (Tensor): 1-D integer tensor of length - ``num_experts`` giving the number of tokens assigned to each expert. - - Returns: - Tensor of shape ``(total_tokens, hidden_dim)`` with dtype matching ``x``. - """ - offsets = torch.cumsum(num_tokens_per_expert, dim=0, dtype=torch.int32) - - gate_up = F.grouped_mm(x, self.gate_up_proj, offs=offsets) - gate, up = gate_up.chunk(2, dim=-1) - out = F.grouped_mm(F.silu(gate) * up, self.down_proj, offs=offsets) - - return out.type_as(x) - - def forward(self, x: torch.Tensor, num_tokens_per_expert: torch.Tensor) -> torch.Tensor: - if self.use_grouped_mm: - return self._run_experts_grouped_mm(x, num_tokens_per_expert) - return self._run_experts_for_loop(x, num_tokens_per_expert) - - -class NucleusMoELayer(nn.Module): - """ - Mixture-of-Experts layer with expert-choice routing and a shared expert. - - Routed expert weights live in :class:`SwiGLUExperts`. The router concatenates a timestep embedding with the - (unmodulated) hidden state to produce per-token affinity scores, then selects the top-C tokens per expert - (expert-choice routing). A shared expert processes all tokens in parallel and its output is combined with the - routed expert outputs via scatter-add. - - SwiGLU expert computation is implemented by :class:`SwiGLUExperts`. - """ - - def __init__( - self, - hidden_size: int, - moe_intermediate_dim: int, - num_experts: int, - capacity_factor: float, - use_sigmoid: bool, - route_scale: float, - use_grouped_mm: bool = False, - ): - super().__init__() - self.num_experts = num_experts - self.moe_intermediate_dim = moe_intermediate_dim - self.hidden_size = hidden_size - self.capacity_factor = capacity_factor - self.use_sigmoid = use_sigmoid - self.route_scale = route_scale - - self.gate = nn.Linear(hidden_size * 2, num_experts, bias=False) - - self.experts = SwiGLUExperts( - hidden_size=hidden_size, - moe_intermediate_dim=moe_intermediate_dim, - num_experts=num_experts, - use_grouped_mm=use_grouped_mm, - ) - - self.shared_expert = FeedForward( - dim=hidden_size, - dim_out=hidden_size, - inner_dim=moe_intermediate_dim, - activation_fn="swiglu", - bias=False, - ) - - def forward( - self, - hidden_states: torch.Tensor, - hidden_states_unmodulated: torch.Tensor, - timestep: torch.Tensor | None = None, - ) -> torch.Tensor: - bs, slen, dim = hidden_states.shape - - if timestep is not None: - timestep_expanded = timestep.unsqueeze(1).expand(-1, slen, -1) - router_input = torch.cat([timestep_expanded, hidden_states_unmodulated], dim=-1) - else: - router_input = hidden_states_unmodulated - - logits = self.gate(router_input) - - if self.use_sigmoid: - scores = torch.sigmoid(logits.float()).to(logits.dtype) - else: - scores = F.softmax(logits.float(), dim=-1).to(logits.dtype) - - affinity = scores.transpose(1, 2) # (B, E, S) - capacity = max(1, math.ceil(self.capacity_factor * slen / self.num_experts)) - - topk = torch.topk(affinity, k=capacity, dim=-1) - top_indices = topk.indices # (B, E, C) - gating = affinity.gather(dim=-1, index=top_indices) # (B, E, C) - - batch_offsets = torch.arange(bs, device=hidden_states.device, dtype=torch.long).view(bs, 1, 1) * slen - global_token_indices = (batch_offsets + top_indices).transpose(0, 1).reshape(self.num_experts, -1).reshape(-1) - gating_flat = gating.transpose(0, 1).reshape(self.num_experts, -1).reshape(-1) - - token_score_sums = torch.zeros(bs * slen, device=hidden_states.device, dtype=gating_flat.dtype) - token_score_sums.scatter_add_(0, global_token_indices, gating_flat) - gating_flat = gating_flat / (token_score_sums[global_token_indices] + 1e-12) - gating_flat = gating_flat * self.route_scale - - x_flat = hidden_states.reshape(bs * slen, dim) - routed_input = x_flat[global_token_indices] - - tokens_per_expert = bs * capacity - num_tokens_per_expert = torch.full( - (self.num_experts,), - tokens_per_expert, - device=hidden_states.device, - dtype=torch.long, - ) - routed_output = self.experts(routed_input, num_tokens_per_expert) - routed_output = (routed_output.float() * gating_flat.unsqueeze(-1)).to(hidden_states.dtype) - - out = self.shared_expert(hidden_states).reshape(bs * slen, dim) - - scatter_idx = global_token_indices.reshape(-1, 1).expand(-1, dim) - out = out.scatter_add(dim=0, index=scatter_idx, src=routed_output) - out = out.reshape(bs, slen, dim) - - return out - - -class NucleusMoEImageTransformerBlock(nn.Module): - """ - Single-stream DiT block with optional Mixture-of-Experts MLP. Only the image stream receives adaptive modulation; - the text context is projected per-block and used as cross-attention keys/values. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - num_key_value_heads: int | None = None, - joint_attention_dim: int = 3584, - qk_norm: str = "rms_norm", - eps: float = 1e-6, - mlp_ratio: float = 4.0, - moe_enabled: bool = False, - num_experts: int = 128, - moe_intermediate_dim: int = 1344, - capacity_factor: float = 8.0, - use_sigmoid: bool = False, - route_scale: float = 2.5, - use_grouped_mm: bool = False, - ): - super().__init__() - self.dim = dim - self.moe_enabled = moe_enabled - - self.img_mod = nn.Sequential( - nn.SiLU(), - nn.Linear(dim, 4 * dim, bias=True), - ) - - self.encoder_proj = nn.Linear(joint_attention_dim, dim) - - self.pre_attn_norm = nn.LayerNorm(dim, eps=eps, elementwise_affine=False, bias=False) - self.attn = Attention( - query_dim=dim, - heads=num_attention_heads, - kv_heads=num_key_value_heads, - dim_head=attention_head_dim, - added_kv_proj_dim=dim, - added_proj_bias=False, - out_dim=dim, - out_bias=False, - bias=False, - processor=NucleusMoEAttnProcessor2_0(), - qk_norm=qk_norm, - eps=eps, - context_pre_only=None, - ) - - self.pre_mlp_norm = nn.LayerNorm(dim, eps=eps, elementwise_affine=False, bias=False) - - if moe_enabled: - self.img_mlp = NucleusMoELayer( - hidden_size=dim, - moe_intermediate_dim=moe_intermediate_dim, - num_experts=num_experts, - capacity_factor=capacity_factor, - use_sigmoid=use_sigmoid, - route_scale=route_scale, - use_grouped_mm=use_grouped_mm, - ) - else: - mlp_inner_dim = int(dim * mlp_ratio * 2 / 3) // 128 * 128 - self.img_mlp = FeedForward( - dim=dim, - dim_out=dim, - inner_dim=mlp_inner_dim, - activation_fn="swiglu", - bias=False, - ) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor: - scale1, gate1, scale2, gate2 = self.img_mod(temb).unsqueeze(1).chunk(4, dim=-1) - - gate1 = gate1.clamp(min=-2.0, max=2.0) - gate2 = gate2.clamp(min=-2.0, max=2.0) - - attn_kwargs = attention_kwargs or {} - context = None if attn_kwargs.get("cached_txt_key") is not None else self.encoder_proj(encoder_hidden_states) - - img_normed = self.pre_attn_norm(hidden_states) - img_modulated = img_normed * (1 + scale1) - - img_attn_output = self.attn( - hidden_states=img_modulated, - encoder_hidden_states=context, - image_rotary_emb=image_rotary_emb, - **attn_kwargs, - ) - - hidden_states = hidden_states + gate1.tanh() * img_attn_output - - img_normed2 = self.pre_mlp_norm(hidden_states) - img_modulated2 = img_normed2 * (1 + scale2) - - if self.moe_enabled: - img_mlp_output = self.img_mlp(img_modulated2, img_normed2, timestep=temb) - else: - img_mlp_output = self.img_mlp(img_modulated2) - - hidden_states = hidden_states + gate2.tanh() * img_mlp_output - - if hidden_states.dtype == torch.float16: - fp16_finfo = torch.finfo(torch.float16) - hidden_states = hidden_states.clip(fp16_finfo.min, fp16_finfo.max) - - return hidden_states - - -class NucleusMoEImageTransformer2DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin, AttentionMixin -): - """ - Nucleus MoE Transformer for image generation. Single-stream DiT with cross-attention to text and optional - Mixture-of-Experts feed-forward layers. - - Args: - patch_size (`int`, defaults to `2`): - Patch size to turn the input data into small patches. - in_channels (`int`, defaults to `64`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `None`): - The number of channels in the output. If not specified, it defaults to `in_channels`. - num_layers (`int`, defaults to `24`): - The number of transformer blocks. - attention_head_dim (`int`, defaults to `128`): - The number of dimensions to use for each attention head. - num_attention_heads (`int`, defaults to `16`): - The number of attention heads to use. - num_key_value_heads (`int`, *optional*): - The number of key/value heads for grouped-query attention. Defaults to `num_attention_heads`. - joint_attention_dim (`int`, defaults to `3584`): - The embedding dimension of the encoder hidden states (text). - axes_dims_rope (`tuple[int]`, defaults to `(16, 56, 56)`): - The dimensions to use for the rotary positional embeddings. - mlp_ratio (`float`, defaults to `4.0`): - Multiplier for the MLP hidden dimension in dense (non-MoE) blocks. - moe_enabled (`bool`, defaults to `True`): - Whether to use Mixture-of-Experts layers. - dense_moe_strategy (`str`, defaults to ``"leave_first_three_and_last_block_dense"``): - Strategy for choosing which layers are MoE vs dense. - num_experts (`int`, defaults to `128`): - Number of experts per MoE layer. - moe_intermediate_dim (`int`, defaults to `1344`): - Hidden dimension inside each expert. - capacity_factors (`float | list[float]`, defaults to `8.0`): - Expert-choice capacity factor per layer. - use_sigmoid (`bool`, defaults to `False`): - Use sigmoid instead of softmax for routing scores. - route_scale (`float`, defaults to `2.5`): - Scaling factor applied to routing weights. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["NucleusMoEImageTransformerBlock"] - _skip_layerwise_casting_patterns = ["pos_embed", "norm"] - _repeated_blocks = ["NucleusMoEImageTransformerBlock"] - - @register_to_config - def __init__( - self, - patch_size: int = 2, - in_channels: int = 64, - out_channels: int | None = None, - num_layers: int = 24, - attention_head_dim: int = 128, - num_attention_heads: int = 16, - num_key_value_heads: int | None = None, - joint_attention_dim: int = 3584, - axes_dims_rope: tuple[int, int, int] = (16, 56, 56), - mlp_ratio: float = 4.0, - moe_enabled: bool = True, - dense_moe_strategy: str = "leave_first_three_and_last_block_dense", - num_experts: int = 128, - moe_intermediate_dim: int = 1344, - capacity_factors: float | list[float] = 8.0, - use_sigmoid: bool = False, - route_scale: float = 2.5, - use_grouped_mm: bool = False, - ): - super().__init__() - self.out_channels = out_channels or in_channels - self.inner_dim = num_attention_heads * attention_head_dim - capacity_factors = capacity_factors if isinstance(capacity_factors, list) else [capacity_factors] * num_layers - - self.pos_embed = NucleusMoEEmbedRope(theta=10000, axes_dim=list(axes_dims_rope), scale_rope=True) - - self.time_text_embed = NucleusMoETimestepProjEmbeddings(embedding_dim=self.inner_dim) - - self.txt_norm = RMSNorm(joint_attention_dim, eps=1e-6) - self.img_in = nn.Linear(in_channels, self.inner_dim) - - self.transformer_blocks = nn.ModuleList( - [ - NucleusMoEImageTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - num_key_value_heads=num_key_value_heads, - joint_attention_dim=joint_attention_dim, - mlp_ratio=mlp_ratio, - moe_enabled=moe_enabled and _is_moe_layer(dense_moe_strategy, idx, num_layers), - num_experts=num_experts, - moe_intermediate_dim=moe_intermediate_dim, - capacity_factor=capacity_factors[idx], - use_sigmoid=use_sigmoid, - route_scale=route_scale, - use_grouped_mm=use_grouped_mm, - ) - for idx in range(num_layers) - ] - ) - - self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=False) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - img_shapes: tuple[int, int, int] | list[tuple[int, int, int]], - encoder_hidden_states: torch.Tensor = None, - encoder_hidden_states_mask: torch.Tensor = None, - timestep: torch.LongTensor = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> torch.Tensor | Transformer2DModelOutput: - """ - The [`NucleusMoEImageTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, image_sequence_length, in_channels)`): - Input `hidden_states`. - img_shapes (`list[tuple[int, int, int]]`, *optional*): - Image shapes ``(frame, height, width)`` for RoPE computation. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, text_sequence_length, joint_attention_dim)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_hidden_states_mask (`torch.Tensor` of shape `(batch_size, text_sequence_length)`, *optional*): - Boolean mask for the encoder hidden states. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - attention_kwargs (`dict`, *optional*): - Extra kwargs forwarded to the attention processor. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a [`~models.transformer_2d.Transformer2DModelOutput`]. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - if attention_kwargs is not None: - attention_kwargs = attention_kwargs.copy() - lora_scale = attention_kwargs.pop("scale", 1.0) - else: - lora_scale = 1.0 - - if USE_PEFT_BACKEND: - scale_lora_layers(self, lora_scale) - - hidden_states = self.img_in(hidden_states) - timestep = timestep.to(hidden_states.dtype) - - encoder_hidden_states = self.txt_norm(encoder_hidden_states) - - text_seq_len, _, encoder_hidden_states_mask = _compute_text_seq_len_from_mask( - encoder_hidden_states, encoder_hidden_states_mask - ) - - temb = self.time_text_embed(timestep, hidden_states) - - image_rotary_emb = self.pos_embed(img_shapes, max_txt_seq_len=text_seq_len, device=hidden_states.device) - - block_attention_kwargs = attention_kwargs.copy() if attention_kwargs is not None else {} - if encoder_hidden_states_mask is not None: - batch_size, image_seq_len = hidden_states.shape[:2] - image_mask = torch.ones((batch_size, image_seq_len), dtype=torch.bool, device=hidden_states.device) - joint_attention_mask = torch.cat([image_mask, encoder_hidden_states_mask], dim=1) - block_attention_kwargs["attention_mask"] = joint_attention_mask - - for index_block, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - block_attention_kwargs, - ) - else: - hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - attention_kwargs=block_attention_kwargs, - ) - - hidden_states = self.norm_out(hidden_states, temb) - output = self.proj_out(hidden_states) - - if USE_PEFT_BACKEND: - unscale_lora_layers(self, lora_scale) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_omnigen.py b/diffusers/models/transformers/transformer_omnigen.py deleted file mode 100644 index f860f5d5ab3e4cf4e3c04c1ba84ed17a377c7db0..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_omnigen.py +++ /dev/null @@ -1,496 +0,0 @@ -# Copyright 2025 OmniGen team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ..attention_processor import Attention -from ..embeddings import TimestepEmbedding, Timesteps, get_2d_sincos_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNorm, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class OmniGenFeedForward(nn.Module): - def __init__(self, hidden_size: int, intermediate_size: int): - super().__init__() - - self.gate_up_proj = nn.Linear(hidden_size, 2 * intermediate_size, bias=False) - self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False) - self.activation_fn = nn.SiLU() - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - up_states = self.gate_up_proj(hidden_states) - gate, up_states = up_states.chunk(2, dim=-1) - up_states = up_states * self.activation_fn(gate) - return self.down_proj(up_states) - - -class OmniGenPatchEmbed(nn.Module): - def __init__( - self, - patch_size: int = 2, - in_channels: int = 4, - embed_dim: int = 768, - bias: bool = True, - interpolation_scale: float = 1, - pos_embed_max_size: int = 192, - base_size: int = 64, - ): - super().__init__() - - self.output_image_proj = nn.Conv2d( - in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias - ) - self.input_image_proj = nn.Conv2d( - in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias - ) - - self.patch_size = patch_size - self.interpolation_scale = interpolation_scale - self.pos_embed_max_size = pos_embed_max_size - - pos_embed = get_2d_sincos_pos_embed( - embed_dim, - self.pos_embed_max_size, - base_size=base_size, - interpolation_scale=self.interpolation_scale, - output_type="pt", - ) - self.register_buffer("pos_embed", pos_embed.float().unsqueeze(0), persistent=True) - - def _cropped_pos_embed(self, height, width): - """Crops positional embeddings for SD3 compatibility.""" - if self.pos_embed_max_size is None: - raise ValueError("`pos_embed_max_size` must be set for cropping.") - - height = height // self.patch_size - width = width // self.patch_size - if height > self.pos_embed_max_size: - raise ValueError( - f"Height ({height}) cannot be greater than `pos_embed_max_size`: {self.pos_embed_max_size}." - ) - if width > self.pos_embed_max_size: - raise ValueError( - f"Width ({width}) cannot be greater than `pos_embed_max_size`: {self.pos_embed_max_size}." - ) - - top = (self.pos_embed_max_size - height) // 2 - left = (self.pos_embed_max_size - width) // 2 - spatial_pos_embed = self.pos_embed.reshape(1, self.pos_embed_max_size, self.pos_embed_max_size, -1) - spatial_pos_embed = spatial_pos_embed[:, top : top + height, left : left + width, :] - spatial_pos_embed = spatial_pos_embed.reshape(1, -1, spatial_pos_embed.shape[-1]) - return spatial_pos_embed - - def _patch_embeddings(self, hidden_states: torch.Tensor, is_input_image: bool) -> torch.Tensor: - if is_input_image: - hidden_states = self.input_image_proj(hidden_states) - else: - hidden_states = self.output_image_proj(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) - return hidden_states - - def forward( - self, hidden_states: torch.Tensor, is_input_image: bool, padding_latent: torch.Tensor = None - ) -> torch.Tensor: - if isinstance(hidden_states, list): - if padding_latent is None: - padding_latent = [None] * len(hidden_states) - patched_latents = [] - for sub_latent, padding in zip(hidden_states, padding_latent): - height, width = sub_latent.shape[-2:] - sub_latent = self._patch_embeddings(sub_latent, is_input_image) - pos_embed = self._cropped_pos_embed(height, width) - sub_latent = sub_latent + pos_embed - if padding is not None: - sub_latent = torch.cat([sub_latent, padding.to(sub_latent.device)], dim=-2) - patched_latents.append(sub_latent) - else: - height, width = hidden_states.shape[-2:] - pos_embed = self._cropped_pos_embed(height, width) - hidden_states = self._patch_embeddings(hidden_states, is_input_image) - patched_latents = hidden_states + pos_embed - - return patched_latents - - -class OmniGenSuScaledRotaryEmbedding(nn.Module): - def __init__( - self, dim, max_position_embeddings=131072, original_max_position_embeddings=4096, base=10000, rope_scaling=None - ): - super().__init__() - - self.dim = dim - self.max_position_embeddings = max_position_embeddings - self.base = base - - inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64).float() / self.dim)) - self.register_buffer("inv_freq", tensor=inv_freq, persistent=False) - - self.short_factor = rope_scaling["short_factor"] - self.long_factor = rope_scaling["long_factor"] - self.original_max_position_embeddings = original_max_position_embeddings - - def forward(self, hidden_states, position_ids): - seq_len = torch.max(position_ids) + 1 - if seq_len > self.original_max_position_embeddings: - ext_factors = torch.tensor(self.long_factor, dtype=torch.float32, device=hidden_states.device) - else: - ext_factors = torch.tensor(self.short_factor, dtype=torch.float32, device=hidden_states.device) - - inv_freq_shape = ( - torch.arange(0, self.dim, 2, dtype=torch.int64, device=hidden_states.device).float() / self.dim - ) - self.inv_freq = 1.0 / (ext_factors * self.base**inv_freq_shape) - - inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) - position_ids_expanded = position_ids[:, None, :].float() - - # Force float32 since bfloat16 loses precision on long contexts - # See https://github.com/huggingface/transformers/pull/29285 - device_type = hidden_states.device.type - device_type = device_type if isinstance(device_type, str) and device_type != "mps" else "cpu" - with torch.autocast(device_type=device_type, enabled=False): - freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2) - emb = torch.cat((freqs, freqs), dim=-1)[0] - - scale = self.max_position_embeddings / self.original_max_position_embeddings - if scale <= 1.0: - scaling_factor = 1.0 - else: - scaling_factor = math.sqrt(1 + math.log(scale) / math.log(self.original_max_position_embeddings)) - - cos = emb.cos() * scaling_factor - sin = emb.sin() * scaling_factor - return cos, sin - - -class OmniGenAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is - used in the OmniGen model. - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - batch_size, sequence_length, _ = hidden_states.shape - - # Get Query-Key-Value Pair - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - bsz, q_len, query_dim = query.size() - inner_dim = key.shape[-1] - head_dim = query_dim // attn.heads - - # Get key-value heads - kv_heads = inner_dim // head_dim - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - key = key.view(batch_size, -1, kv_heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, kv_heads, head_dim).transpose(1, 2) - - # Apply RoPE if needed - if image_rotary_emb is not None: - from ..embeddings import apply_rotary_emb - - query = apply_rotary_emb(query, image_rotary_emb, use_real_unbind_dim=-2) - key = apply_rotary_emb(key, image_rotary_emb, use_real_unbind_dim=-2) - - hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask) - hidden_states = hidden_states.transpose(1, 2).type_as(query) - hidden_states = hidden_states.reshape(bsz, q_len, attn.out_dim) - hidden_states = attn.to_out[0](hidden_states) - return hidden_states - - -class OmniGenBlock(nn.Module): - def __init__( - self, - hidden_size: int, - num_attention_heads: int, - num_key_value_heads: int, - intermediate_size: int, - rms_norm_eps: float, - ) -> None: - super().__init__() - - self.input_layernorm = RMSNorm(hidden_size, eps=rms_norm_eps) - self.self_attn = Attention( - query_dim=hidden_size, - cross_attention_dim=hidden_size, - dim_head=hidden_size // num_attention_heads, - heads=num_attention_heads, - kv_heads=num_key_value_heads, - bias=False, - out_dim=hidden_size, - out_bias=False, - processor=OmniGenAttnProcessor2_0(), - ) - self.post_attention_layernorm = RMSNorm(hidden_size, eps=rms_norm_eps) - self.mlp = OmniGenFeedForward(hidden_size, intermediate_size) - - def forward( - self, hidden_states: torch.Tensor, attention_mask: torch.Tensor, image_rotary_emb: torch.Tensor - ) -> torch.Tensor: - # 1. Attention - norm_hidden_states = self.input_layernorm(hidden_states) - attn_output = self.self_attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - hidden_states = hidden_states + attn_output - - # 2. Feed Forward - norm_hidden_states = self.post_attention_layernorm(hidden_states) - ff_output = self.mlp(norm_hidden_states) - hidden_states = hidden_states + ff_output - return hidden_states - - -class OmniGenTransformer2DModel(ModelMixin, ConfigMixin): - """ - The Transformer model introduced in OmniGen (https://huggingface.co/papers/2409.11340). - - Parameters: - in_channels (`int`, defaults to `4`): - The number of channels in the input. - patch_size (`int`, defaults to `2`): - The size of the spatial patches to use in the patch embedding layer. - hidden_size (`int`, defaults to `3072`): - The dimensionality of the hidden layers in the model. - rms_norm_eps (`float`, defaults to `1e-5`): - Eps for RMSNorm layer. - num_attention_heads (`int`, defaults to `32`): - The number of heads to use for multi-head attention. - num_key_value_heads (`int`, defaults to `32`): - The number of heads to use for keys and values in multi-head attention. - intermediate_size (`int`, defaults to `8192`): - Dimension of the hidden layer in FeedForward layers. - num_layers (`int`, default to `32`): - The number of layers of transformer blocks to use. - pad_token_id (`int`, default to `32000`): - The id of the padding token. - vocab_size (`int`, default to `32064`): - The size of the vocabulary of the embedding vocabulary. - rope_base (`int`, default to `10000`): - The default theta value to use when creating RoPE. - rope_scaling (`dict`, optional): - The scaling factors for the RoPE. Must contain `short_factor` and `long_factor`. - pos_embed_max_size (`int`, default to `192`): - The maximum size of the positional embeddings. - time_step_dim (`int`, default to `256`): - Output dimension of timestep embeddings. - flip_sin_to_cos (`bool`, default to `True`): - Whether to flip the sin and cos in the positional embeddings when preparing timestep embeddings. - downscale_freq_shift (`int`, default to `0`): - The frequency shift to use when downscaling the timestep embeddings. - timestep_activation_fn (`str`, default to `silu`): - The activation function to use for the timestep embeddings. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["OmniGenBlock"] - _skip_layerwise_casting_patterns = ["patch_embedding", "embed_tokens", "norm"] - - @register_to_config - def __init__( - self, - in_channels: int = 4, - patch_size: int = 2, - hidden_size: int = 3072, - rms_norm_eps: float = 1e-5, - num_attention_heads: int = 32, - num_key_value_heads: int = 32, - intermediate_size: int = 8192, - num_layers: int = 32, - pad_token_id: int = 32000, - vocab_size: int = 32064, - max_position_embeddings: int = 131072, - original_max_position_embeddings: int = 4096, - rope_base: int = 10000, - rope_scaling: dict = None, - pos_embed_max_size: int = 192, - time_step_dim: int = 256, - flip_sin_to_cos: bool = True, - downscale_freq_shift: int = 0, - timestep_activation_fn: str = "silu", - ): - super().__init__() - self.in_channels = in_channels - self.out_channels = in_channels - - self.patch_embedding = OmniGenPatchEmbed( - patch_size=patch_size, - in_channels=in_channels, - embed_dim=hidden_size, - pos_embed_max_size=pos_embed_max_size, - ) - - self.time_proj = Timesteps(time_step_dim, flip_sin_to_cos, downscale_freq_shift) - self.time_token = TimestepEmbedding(time_step_dim, hidden_size, timestep_activation_fn) - self.t_embedder = TimestepEmbedding(time_step_dim, hidden_size, timestep_activation_fn) - - self.embed_tokens = nn.Embedding(vocab_size, hidden_size, pad_token_id) - self.rope = OmniGenSuScaledRotaryEmbedding( - hidden_size // num_attention_heads, - max_position_embeddings=max_position_embeddings, - original_max_position_embeddings=original_max_position_embeddings, - base=rope_base, - rope_scaling=rope_scaling, - ) - - self.layers = nn.ModuleList( - [ - OmniGenBlock(hidden_size, num_attention_heads, num_key_value_heads, intermediate_size, rms_norm_eps) - for _ in range(num_layers) - ] - ) - - self.norm = RMSNorm(hidden_size, eps=rms_norm_eps) - self.norm_out = AdaLayerNorm(hidden_size, norm_elementwise_affine=False, norm_eps=1e-6, chunk_dim=1) - self.proj_out = nn.Linear(hidden_size, patch_size * patch_size * self.out_channels, bias=True) - - self.gradient_checkpointing = False - - def _get_multimodal_embeddings( - self, input_ids: torch.Tensor, input_img_latents: list[torch.Tensor], input_image_sizes: dict - ) -> torch.Tensor | None: - if input_ids is None: - return None - - input_img_latents = [x.to(self.dtype) for x in input_img_latents] - condition_tokens = self.embed_tokens(input_ids) - input_img_inx = 0 - input_image_tokens = self.patch_embedding(input_img_latents, is_input_image=True) - for b_inx in input_image_sizes.keys(): - for start_inx, end_inx in input_image_sizes[b_inx]: - # replace the placeholder in text tokens with the image embedding. - condition_tokens[b_inx, start_inx:end_inx] = input_image_tokens[input_img_inx].to( - condition_tokens.dtype - ) - input_img_inx += 1 - return condition_tokens - - def forward( - self, - hidden_states: torch.Tensor, - timestep: int | float | torch.FloatTensor, - input_ids: torch.Tensor, - input_img_latents: list[torch.Tensor], - input_image_sizes: dict[int, list[int]], - attention_mask: torch.Tensor, - position_ids: torch.Tensor, - return_dict: bool = True, - ) -> Transformer2DModelOutput | tuple[torch.Tensor]: - """ - The [`OmniGenTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, in_channels, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - input_ids (`torch.Tensor`): - Multimodal text token ids used as conditioning. - input_img_latents (`list` of `torch.Tensor`): - List of latents for input images used as conditioning. - input_image_sizes (`dict` of `int` to `list` of `int`): - Mapping from sample index to the positions where input image embeddings should be placed in the - conditioning sequence. - attention_mask (`torch.Tensor`): - Attention mask for the joint multimodal sequence. - position_ids (`torch.Tensor`): - Position ids used to compute the positional embeddings. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - [`~models.transformer_2d.Transformer2DModelOutput`] or `tuple`: - If `return_dict` is True, a [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise - a plain `tuple` is returned. - """ - batch_size, num_channels, height, width = hidden_states.shape - p = self.config.patch_size - post_patch_height, post_patch_width = height // p, width // p - - # 1. Patch & Timestep & Conditional Embedding - hidden_states = self.patch_embedding(hidden_states, is_input_image=False) - num_tokens_for_output_image = hidden_states.size(1) - - timestep_proj = self.time_proj(timestep).type_as(hidden_states) - time_token = self.time_token(timestep_proj).unsqueeze(1) - temb = self.t_embedder(timestep_proj) - - condition_tokens = self._get_multimodal_embeddings(input_ids, input_img_latents, input_image_sizes) - if condition_tokens is not None: - hidden_states = torch.cat([condition_tokens, time_token, hidden_states], dim=1) - else: - hidden_states = torch.cat([time_token, hidden_states], dim=1) - - seq_length = hidden_states.size(1) - position_ids = position_ids.view(-1, seq_length).long() - - # 2. Attention mask preprocessing - if attention_mask is not None and attention_mask.dim() == 3: - dtype = hidden_states.dtype - min_dtype = torch.finfo(dtype).min - attention_mask = (1 - attention_mask) * min_dtype - attention_mask = attention_mask.unsqueeze(1).type_as(hidden_states) - - # 3. Rotary position embedding - image_rotary_emb = self.rope(hidden_states, position_ids) - - # 4. Transformer blocks - for block in self.layers: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, hidden_states, attention_mask, image_rotary_emb - ) - else: - hidden_states = block(hidden_states, attention_mask=attention_mask, image_rotary_emb=image_rotary_emb) - - # 5. Output norm & projection - hidden_states = self.norm(hidden_states) - hidden_states = hidden_states[:, -num_tokens_for_output_image:] - hidden_states = self.norm_out(hidden_states, temb=temb) - hidden_states = self.proj_out(hidden_states) - hidden_states = hidden_states.reshape(batch_size, post_patch_height, post_patch_width, p, p, -1) - output = hidden_states.permute(0, 5, 1, 3, 2, 4).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_ovis_image.py b/diffusers/models/transformers/transformer_ovis_image.py deleted file mode 100644 index 44723bc44fd07ffcb5cdc0460ebcfaca4def7599..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_ovis_image.py +++ /dev/null @@ -1,585 +0,0 @@ -# Copyright 2025 Alibaba Ovis-Image Team and The HuggingFace. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device, maybe_allow_in_graph -from ..attention import AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import TimestepEmbedding, Timesteps, apply_rotary_emb, get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _get_projections(attn: "OvisImageAttention", hidden_states, encoder_hidden_states=None): - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - encoder_query = encoder_key = encoder_value = None - if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None: - encoder_query = attn.add_q_proj(encoder_hidden_states) - encoder_key = attn.add_k_proj(encoder_hidden_states) - encoder_value = attn.add_v_proj(encoder_hidden_states) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_fused_projections(attn: "OvisImageAttention", hidden_states, encoder_hidden_states=None): - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - - encoder_query = encoder_key = encoder_value = (None,) - if encoder_hidden_states is not None and hasattr(attn, "to_added_qkv"): - encoder_query, encoder_key, encoder_value = attn.to_added_qkv(encoder_hidden_states).chunk(3, dim=-1) - - return query, key, value, encoder_query, encoder_key, encoder_value - - -def _get_qkv_projections(attn: "OvisImageAttention", hidden_states, encoder_hidden_states=None): - if attn.fused_projections: - return _get_fused_projections(attn, hidden_states, encoder_hidden_states) - return _get_projections(attn, hidden_states, encoder_hidden_states) - - -class OvisImageAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError(f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version.") - - def __call__( - self, - attn: "OvisImageAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - query, key, value, encoder_query, encoder_key, encoder_value = _get_qkv_projections( - attn, hidden_states, encoder_hidden_states - ) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - if attn.added_kv_proj_dim is not None: - encoder_query = encoder_query.unflatten(-1, (attn.heads, -1)) - encoder_key = encoder_key.unflatten(-1, (attn.heads, -1)) - encoder_value = encoder_value.unflatten(-1, (attn.heads, -1)) - - encoder_query = attn.norm_added_q(encoder_query) - encoder_key = attn.norm_added_k(encoder_key) - - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) - - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - if encoder_hidden_states is not None: - encoder_hidden_states, hidden_states = hidden_states.split_with_sizes( - [encoder_hidden_states.shape[1], hidden_states.shape[1] - encoder_hidden_states.shape[1]], dim=1 - ) - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - encoder_hidden_states = attn.to_add_out(encoder_hidden_states) - - return hidden_states, encoder_hidden_states - else: - return hidden_states - - -class OvisImageAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = OvisImageAttnProcessor - _available_processors = [ - OvisImageAttnProcessor, - ] - - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - added_kv_proj_dim: int | None = None, - added_proj_bias: bool | None = True, - out_bias: bool = True, - eps: float = 1e-5, - out_dim: int = None, - context_pre_only: bool | None = None, - pre_only: bool = False, - elementwise_affine: bool = True, - processor=None, - ): - super().__init__() - - self.head_dim = dim_head - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.query_dim = query_dim - self.use_bias = bias - self.dropout = dropout - self.out_dim = out_dim if out_dim is not None else query_dim - self.context_pre_only = context_pre_only - self.pre_only = pre_only - self.heads = out_dim // dim_head if out_dim is not None else heads - self.added_kv_proj_dim = added_kv_proj_dim - self.added_proj_bias = added_proj_bias - - self.norm_q = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.norm_k = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine) - self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_v = torch.nn.Linear(query_dim, self.inner_dim, bias=bias) - - if not self.pre_only: - self.to_out = torch.nn.ModuleList([]) - self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(torch.nn.Dropout(dropout)) - - if added_kv_proj_dim is not None: - self.norm_added_q = torch.nn.RMSNorm(dim_head, eps=eps) - self.norm_added_k = torch.nn.RMSNorm(dim_head, eps=eps) - self.add_q_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) - self.to_add_out = torch.nn.Linear(self.inner_dim, query_dim, bias=out_bias) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) - quiet_attn_parameters = {"ip_adapter_masks", "ip_hidden_states"} - unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters and k not in quiet_attn_parameters] - if len(unused_kwargs) > 0: - logger.warning( - f"joint_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." - ) - kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters} - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, image_rotary_emb, **kwargs) - - -@maybe_allow_in_graph -class OvisImageSingleTransformerBlock(nn.Module): - def __init__(self, dim: int, num_attention_heads: int, attention_head_dim: int, mlp_ratio: float = 4.0): - super().__init__() - self.mlp_hidden_dim = int(dim * mlp_ratio) - - self.norm = AdaLayerNormZeroSingle(dim) - self.proj_mlp = nn.Linear(dim, self.mlp_hidden_dim * 2) - self.act_mlp = nn.SiLU() - self.proj_out = nn.Linear(dim + self.mlp_hidden_dim, dim) - - self.attn = OvisImageAttention( - query_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - bias=True, - processor=OvisImageAttnProcessor(), - eps=1e-6, - pre_only=True, - ) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - text_seq_len = encoder_hidden_states.shape[1] - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) - - residual = hidden_states - norm_hidden_states, gate = self.norm(hidden_states, emb=temb) - mlp_hidden_states, mlp_hidden_gate = torch.split( - self.proj_mlp(norm_hidden_states), [self.mlp_hidden_dim, self.mlp_hidden_dim], dim=-1 - ) - mlp_hidden_states = self.act_mlp(mlp_hidden_gate) * mlp_hidden_states - joint_attention_kwargs = joint_attention_kwargs or {} - attn_output = self.attn( - hidden_states=norm_hidden_states, - image_rotary_emb=image_rotary_emb, - **joint_attention_kwargs, - ) - - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) - gate = gate.unsqueeze(1) - hidden_states = gate * self.proj_out(hidden_states) - hidden_states = residual + hidden_states - if hidden_states.dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - encoder_hidden_states, hidden_states = hidden_states[:, :text_seq_len], hidden_states[:, text_seq_len:] - return encoder_hidden_states, hidden_states - - -@maybe_allow_in_graph -class OvisImageTransformerBlock(nn.Module): - def __init__( - self, dim: int, num_attention_heads: int, attention_head_dim: int, qk_norm: str = "rms_norm", eps: float = 1e-6 - ): - super().__init__() - - self.norm1 = AdaLayerNormZero(dim) - self.norm1_context = AdaLayerNormZero(dim) - - self.attn = OvisImageAttention( - query_dim=dim, - added_kv_proj_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - context_pre_only=False, - bias=True, - processor=OvisImageAttnProcessor(), - eps=eps, - ) - - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="swiglu") - - self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff_context = FeedForward(dim=dim, dim_out=dim, activation_fn="swiglu") - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - - norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( - encoder_hidden_states, emb=temb - ) - joint_attention_kwargs = joint_attention_kwargs or {} - - # Attention. - attention_outputs = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - **joint_attention_kwargs, - ) - - if len(attention_outputs) == 2: - attn_output, context_attn_output = attention_outputs - elif len(attention_outputs) == 3: - attn_output, context_attn_output, ip_attn_output = attention_outputs - - # Process attention outputs for the `hidden_states`. - attn_output = gate_msa.unsqueeze(1) * attn_output - hidden_states = hidden_states + attn_output - - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - - ff_output = self.ff(norm_hidden_states) - ff_output = gate_mlp.unsqueeze(1) * ff_output - - hidden_states = hidden_states + ff_output - if len(attention_outputs) == 3: - hidden_states = hidden_states + ip_attn_output - - # Process attention outputs for the `encoder_hidden_states`. - context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output - encoder_hidden_states = encoder_hidden_states + context_attn_output - - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - context_ff_output = self.ff_context(norm_encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output - if encoder_hidden_states.dtype == torch.float16: - encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) - - return encoder_hidden_states, hidden_states - - -class OvisImagePosEmbed(nn.Module): - def __init__(self, theta: int, axes_dim: list[int]): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - - def forward(self, ids: torch.Tensor) -> torch.Tensor: - n_axes = ids.shape[-1] - cos_out = [] - sin_out = [] - pos = ids.float() - freqs_dtype = maybe_adjust_dtype_for_device(torch.float64, ids.device) - for i in range(n_axes): - cos, sin = get_1d_rotary_pos_embed( - self.axes_dim[i], - pos[:, i], - theta=self.theta, - repeat_interleave_real=True, - use_real=True, - freqs_dtype=freqs_dtype, - ) - cos_out.append(cos) - sin_out.append(sin) - freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device) - freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device) - return freqs_cos, freqs_sin - - -class OvisImageTransformer2DModel( - ModelMixin, - ConfigMixin, - PeftAdapterMixin, - FromOriginalModelMixin, - CacheMixin, -): - """ - The Transformer model introduced in Ovis-Image. - - Reference: https://github.com/AIDC-AI/Ovis-Image - - Args: - patch_size (`int`, defaults to `1`): - Patch size to turn the input data into small patches. - in_channels (`int`, defaults to `64`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `None`): - The number of channels in the output. If not specified, it defaults to `in_channels`. - num_layers (`int`, defaults to `6`): - The number of layers of dual stream DiT blocks to use. - num_single_layers (`int`, defaults to `27`): - The number of layers of single stream DiT blocks to use. - attention_head_dim (`int`, defaults to `128`): - The number of dimensions to use for each attention head. - num_attention_heads (`int`, defaults to `24`): - The number of attention heads to use. - joint_attention_dim (`int`, defaults to `2048`): - The number of dimensions to use for the joint attention (embedding/channel dimension of - `encoder_hidden_states`). - axes_dims_rope (`tuple[int]`, defaults to `(16, 56, 56)`): - The dimensions to use for the rotary positional embeddings. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["OvisImageTransformerBlock", "OvisImageSingleTransformerBlock"] - _skip_layerwise_casting_patterns = ["pos_embed", "norm"] - _repeated_blocks = ["OvisImageTransformerBlock", "OvisImageSingleTransformerBlock"] - - @register_to_config - def __init__( - self, - patch_size: int = 1, - in_channels: int = 64, - out_channels: int | None = 64, - num_layers: int = 6, - num_single_layers: int = 27, - attention_head_dim: int = 128, - num_attention_heads: int = 24, - joint_attention_dim: int = 2048, - axes_dims_rope: tuple[int, int, int] = (16, 56, 56), - ): - super().__init__() - self.out_channels = out_channels or in_channels - self.inner_dim = num_attention_heads * attention_head_dim - - self.pos_embed = OvisImagePosEmbed(theta=10000, axes_dim=axes_dims_rope) - - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=self.inner_dim) - - self.context_embedder_norm = nn.RMSNorm(joint_attention_dim, eps=1e-6) - self.context_embedder = nn.Linear(joint_attention_dim, self.inner_dim) - self.x_embedder = nn.Linear(in_channels, self.inner_dim) - - self.transformer_blocks = nn.ModuleList( - [ - OvisImageTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for _ in range(num_layers) - ] - ) - - self.single_transformer_blocks = nn.ModuleList( - [ - OvisImageSingleTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - ) - for _ in range(num_single_layers) - ] - ) - - self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) - - self.gradient_checkpointing = False - - @apply_lora_scale("joint_attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_ids: torch.Tensor = None, - txt_ids: torch.Tensor = None, - joint_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> torch.Tensor | Transformer2DModelOutput: - """ - The [`OvisImageTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, image_sequence_length, in_channels)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, text_sequence_length, joint_attention_dim)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - img_ids: (`torch.Tensor`): - The position ids for image tokens. - txt_ids (`torch.Tensor`): - The position ids for text tokens. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - hidden_states = self.x_embedder(hidden_states) - - timestep = timestep.to(hidden_states.dtype) * 1000 - - timesteps_proj = self.time_proj(timestep) - temb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_states.dtype)) - - encoder_hidden_states = self.context_embedder_norm(encoder_hidden_states) - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - - if txt_ids.ndim == 3: - logger.warning( - "Passing `txt_ids` 3d torch.Tensor is deprecated." - "Please remove the batch dimension and pass it as a 2d torch Tensor" - ) - txt_ids = txt_ids[0] - if img_ids.ndim == 3: - logger.warning( - "Passing `img_ids` 3d torch.Tensor is deprecated." - "Please remove the batch dimension and pass it as a 2d torch Tensor" - ) - img_ids = img_ids[0] - - ids = torch.cat((txt_ids, img_ids), dim=0) - image_rotary_emb = self.pos_embed(ids) - - for index_block, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - joint_attention_kwargs, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - joint_attention_kwargs=joint_attention_kwargs, - ) - - for index_block, block in enumerate(self.single_transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - joint_attention_kwargs, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - joint_attention_kwargs=joint_attention_kwargs, - ) - - hidden_states = self.norm_out(hidden_states, temb) - output = self.proj_out(hidden_states) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_prx.py b/diffusers/models/transformers/transformer_prx.py deleted file mode 100644 index 2676db2e715822532e3e9ff88d95af06bf7142a3..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_prx.py +++ /dev/null @@ -1,870 +0,0 @@ -# Copyright 2025 The Photoroom and The HuggingFace Teams. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..attention import AttentionMixin, AttentionModuleMixin -from ..attention_dispatch import dispatch_attention_fn -from ..embeddings import get_timestep_embedding -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import RMSNorm - - -logger = logging.get_logger(__name__) - - -def get_image_ids(batch_size: int, height: int, width: int, patch_size: int, device: torch.device) -> torch.Tensor: - r""" - Generates 2D patch coordinate indices for a batch of images. - - Args: - batch_size (`int`): - Number of images in the batch. - height (`int`): - Height of the input images (in pixels). - width (`int`): - Width of the input images (in pixels). - patch_size (`int`): - Size of the square patches that the image is divided into. - device (`torch.device`): - The device on which to create the tensor. - - Returns: - `torch.Tensor`: - Tensor of shape `(batch_size, num_patches, 2)` containing the (row, col) coordinates of each patch in the - image grid. - """ - - img_ids = torch.zeros(height // patch_size, width // patch_size, 2, device=device) - img_ids[..., 0] = torch.arange(height // patch_size, device=device)[:, None] - img_ids[..., 1] = torch.arange(width // patch_size, device=device)[None, :] - return img_ids.reshape((height // patch_size) * (width // patch_size), 2).unsqueeze(0).repeat(batch_size, 1, 1) - - -def apply_rope(xq: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: - r""" - Applies rotary positional embeddings (RoPE) to a query tensor. - - Args: - xq (`torch.Tensor`): - Input tensor of shape `(..., dim)` representing the queries. - freqs_cis (`torch.Tensor`): - Precomputed rotary frequency components of shape `(..., dim/2, 2)` containing cosine and sine pairs. - - Returns: - `torch.Tensor`: - Tensor of the same shape as `xq` with rotary embeddings applied. - """ - xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2) - # Ensure freqs_cis is on the same device as queries to avoid device mismatches with offloading - freqs_cis = freqs_cis.to(device=xq.device, dtype=xq_.dtype) - xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1] - return xq_out.reshape(*xq.shape).type_as(xq) - - -class PRXAttnProcessor2_0: - r""" - Processor for implementing PRX-style attention with multi-source tokens and RoPE. Supports multiple attention - backends (Flash Attention, Sage Attention, etc.) via dispatch_attention_fn. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(torch.nn.functional, "scaled_dot_product_attention"): - raise ImportError("PRXAttnProcessor2_0 requires PyTorch 2.0, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: "PRXAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - """ - Apply PRX attention using PRXAttention module. - - Args: - attn: PRXAttention module containing projection layers - hidden_states: Image tokens [B, L_img, D] - encoder_hidden_states: Text tokens [B, L_txt, D] - attention_mask: Boolean mask for text tokens [B, L_txt] - image_rotary_emb: Rotary positional embeddings [B, 1, L_img, head_dim//2, 2, 2] - """ - - if encoder_hidden_states is None: - raise ValueError("PRXAttnProcessor2_0 requires 'encoder_hidden_states' containing text tokens.") - - # Project image tokens to Q, K, V - img_qkv = attn.img_qkv_proj(hidden_states) - B, L_img, _ = img_qkv.shape - img_qkv = img_qkv.reshape(B, L_img, 3, attn.heads, attn.head_dim) - img_qkv = img_qkv.permute(2, 0, 3, 1, 4) # [3, B, H, L_img, D] - img_q, img_k, img_v = img_qkv[0], img_qkv[1], img_qkv[2] - - # Apply QK normalization to image tokens - img_q = attn.norm_q(img_q) - img_k = attn.norm_k(img_k) - - # Project text tokens to K, V - txt_kv = attn.txt_kv_proj(encoder_hidden_states) - B, L_txt, _ = txt_kv.shape - txt_kv = txt_kv.reshape(B, L_txt, 2, attn.heads, attn.head_dim) - txt_kv = txt_kv.permute(2, 0, 3, 1, 4) # [2, B, H, L_txt, D] - txt_k, txt_v = txt_kv[0], txt_kv[1] - - # Apply K normalization to text tokens - txt_k = attn.norm_added_k(txt_k) - - # Apply RoPE to image queries and keys - if image_rotary_emb is not None: - img_q = apply_rope(img_q, image_rotary_emb) - img_k = apply_rope(img_k, image_rotary_emb) - - # Concatenate text and image keys/values - k = torch.cat((txt_k, img_k), dim=2) # [B, H, L_txt + L_img, D] - v = torch.cat((txt_v, img_v), dim=2) # [B, H, L_txt + L_img, D] - - # Build attention mask if provided - attn_mask_tensor = None - if attention_mask is not None: - bs, _, l_img, _ = img_q.shape - l_txt = txt_k.shape[2] - - if attention_mask.dim() != 2: - raise ValueError(f"Unsupported attention_mask shape: {attention_mask.shape}") - if attention_mask.shape[-1] != l_txt: - raise ValueError(f"attention_mask last dim {attention_mask.shape[-1]} must equal text length {l_txt}") - - device = img_q.device - ones_img = torch.ones((bs, l_img), dtype=torch.bool, device=device) - attention_mask = attention_mask.to(device=device, dtype=torch.bool) - joint_mask = torch.cat([attention_mask, ones_img], dim=-1) - attn_mask_tensor = joint_mask[:, None, None, :].expand(-1, attn.heads, l_img, -1) - - # Apply attention using dispatch_attention_fn for backend support - # Reshape to match dispatch_attention_fn expectations: [B, L, H, D] - query = img_q.transpose(1, 2) # [B, L_img, H, D] - key = k.transpose(1, 2) # [B, L_txt + L_img, H, D] - value = v.transpose(1, 2) # [B, L_txt + L_img, H, D] - - attn_output = dispatch_attention_fn( - query, - key, - value, - attn_mask=attn_mask_tensor, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - # Reshape from [B, L_img, H, D] to [B, L_img, H*D] - batch_size, seq_len, num_heads, head_dim = attn_output.shape - attn_output = attn_output.reshape(batch_size, seq_len, num_heads * head_dim) - - # Apply output projection - attn_output = attn.to_out[0](attn_output) - if len(attn.to_out) > 1: - attn_output = attn.to_out[1](attn_output) # dropout if present - - return attn_output - - -class PRXAttention(nn.Module, AttentionModuleMixin): - r""" - PRX-style attention module that handles multi-source tokens and RoPE. Similar to FluxAttention but adapted for - PRX's architecture. - """ - - _default_processor_cls = PRXAttnProcessor2_0 - _available_processors = [PRXAttnProcessor2_0] - - def __init__( - self, - query_dim: int, - heads: int = 8, - dim_head: int = 64, - bias: bool = False, - out_bias: bool = False, - eps: float = 1e-6, - processor=None, - ): - super().__init__() - - self.heads = heads - self.head_dim = dim_head - self.inner_dim = dim_head * heads - self.query_dim = query_dim - - self.img_qkv_proj = nn.Linear(query_dim, query_dim * 3, bias=bias) - - self.norm_q = RMSNorm(self.head_dim, eps=eps, elementwise_affine=True) - self.norm_k = RMSNorm(self.head_dim, eps=eps, elementwise_affine=True) - - self.txt_kv_proj = nn.Linear(query_dim, query_dim * 2, bias=bias) - self.norm_added_k = RMSNorm(self.head_dim, eps=eps, elementwise_affine=True) - - self.to_out = nn.ModuleList([]) - self.to_out.append(nn.Linear(self.inner_dim, query_dim, bias=out_bias)) - self.to_out.append(nn.Dropout(0.0)) - - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - return self.processor( - self, - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - **kwargs, - ) - - -# inspired from https://github.com/black-forest-labs/flux/blob/main/src/flux/modules/layers.py -class PRXEmbedND(nn.Module): - r""" - N-dimensional rotary positional embedding. - - This module creates rotary embeddings (RoPE) across multiple axes, where each axis can have its own embedding - dimension. The embeddings are combined and returned as a single tensor - - Args: - dim (int): - Base embedding dimension (must be even). - theta (int): - Scaling factor that controls the frequency spectrum of the rotary embeddings. - axes_dim (list[int]): - list of embedding dimensions for each axis (each must be even). - """ - - def __init__(self, dim: int, theta: int, axes_dim: list[int]): - super().__init__() - self.dim = dim - self.theta = theta - self.axes_dim = axes_dim - - def rope(self, pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor: - assert dim % 2 == 0 - - dtype = maybe_adjust_dtype_for_device(torch.float64, pos.device) - - scale = torch.arange(0, dim, 2, dtype=dtype, device=pos.device) / dim - omega = 1.0 / (theta**scale) - out = pos.unsqueeze(-1) * omega.unsqueeze(0) - out = torch.stack([torch.cos(out), -torch.sin(out), torch.sin(out), torch.cos(out)], dim=-1) - # Native PyTorch equivalent of: Rearrange("b n d (i j) -> b n d i j", i=2, j=2) - # out shape: (b, n, d, 4) -> reshape to (b, n, d, 2, 2) - out = out.reshape(*out.shape[:-1], 2, 2) - return out.float() - - def forward(self, ids: torch.Tensor) -> torch.Tensor: - n_axes = ids.shape[-1] - emb = torch.cat( - [self.rope(ids[:, :, i], self.axes_dim[i], self.theta) for i in range(n_axes)], - dim=-3, - ) - return emb.unsqueeze(1) - - -class MLPEmbedder(nn.Module): - r""" - A simple 2-layer MLP used for embedding inputs. - - Args: - in_dim (`int`): - Dimensionality of the input features. - hidden_dim (`int`): - Dimensionality of the hidden and output embedding space. - - Returns: - `torch.Tensor`: - Tensor of shape `(..., hidden_dim)` containing the embedded representations. - """ - - def __init__(self, in_dim: int, hidden_dim: int): - super().__init__() - self.in_layer = nn.Linear(in_dim, hidden_dim, bias=True) - self.silu = nn.SiLU() - self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.out_layer(self.silu(self.in_layer(x))) - - -class PRXResolutionEmbedder(nn.Module): - r""" - Embeds the spatial resolution `(height, width)` of the latent into a vector that is added to the timestep - embedding, so the model can condition its modulation on the generation resolution. - - A sinusoidal embedding of dimension 128 is built for the height and the width separately and concatenated into a - 256-dim vector, which is then projected to `hidden_size` by a 2-layer MLP. This matches the `"vec"` mode of the - resolution-aware conditioning used during PRX-7B training. - - Args: - hidden_size (`int`): - Dimension of the output embedding (must match the timestep embedding dimension). - max_period (`int`, *optional*, defaults to 10000): - Maximum frequency period for the sinusoidal resolution embedding. - """ - - def __init__(self, hidden_size: int, max_period: int = 10000): - super().__init__() - self.max_period = max_period - self.mlp = MLPEmbedder(in_dim=256, hidden_dim=hidden_size) - - def forward(self, height: torch.Tensor, width: torch.Tensor, dtype: torch.dtype) -> torch.Tensor: - h_emb = get_timestep_embedding( - timesteps=height, - embedding_dim=128, - max_period=self.max_period, - scale=1.0, - flip_sin_to_cos=True, - downscale_freq_shift=0.0, - ) - w_emb = get_timestep_embedding( - timesteps=width, - embedding_dim=128, - max_period=self.max_period, - scale=1.0, - flip_sin_to_cos=True, - downscale_freq_shift=0.0, - ) - hw_emb = torch.cat([h_emb, w_emb], dim=-1).to(dtype) - return self.mlp(hw_emb) - - -class Modulation(nn.Module): - r""" - Modulation network that generates scale, shift, and gating parameters. - - Given an input vector, the module projects it through a linear layer to produce six chunks, which are grouped into - two tuples `(shift, scale, gate)`. - - Args: - dim (`int`): - Dimensionality of the input vector. The output will have `6 * dim` features internally. - - Returns: - ((`torch.Tensor`, `torch.Tensor`, `torch.Tensor`), (`torch.Tensor`, `torch.Tensor`, `torch.Tensor`)): - Two tuples `(shift, scale, gate)`. - """ - - def __init__(self, dim: int): - super().__init__() - self.lin = nn.Linear(dim, 6 * dim, bias=True) - nn.init.constant_(self.lin.weight, 0) - nn.init.constant_(self.lin.bias, 0) - - def forward( - self, vec: torch.Tensor - ) -> tuple[tuple[torch.Tensor, torch.Tensor, torch.Tensor], tuple[torch.Tensor, torch.Tensor, torch.Tensor]]: - out = self.lin(nn.functional.silu(vec))[:, None, :].chunk(6, dim=-1) - return tuple(out[:3]), tuple(out[3:]) - - -class PRXBlock(nn.Module): - r""" - Multimodal transformer block with text–image cross-attention, modulation, and MLP. - - Args: - hidden_size (`int`): - Dimension of the hidden representations. - num_heads (`int`): - Number of attention heads. - mlp_ratio (`float`, *optional*, defaults to 4.0): - Expansion ratio for the hidden dimension inside the MLP. - qk_scale (`float`, *optional*): - Scale factor for queries and keys. If not provided, defaults to ``head_dim**-0.5``. - - Attributes: - img_pre_norm (`nn.LayerNorm`): - Pre-normalization applied to image tokens before attention. - attention (`PRXAttention`): - Multi-head attention module with built-in QKV projections and normalizations for cross-attention between - image and text tokens. - post_attention_layernorm (`nn.LayerNorm`): - Normalization applied after attention. - gate_proj / up_proj / down_proj (`nn.Linear`): - Feedforward layers forming the gated MLP. - mlp_act (`nn.GELU`): - Nonlinear activation used in the MLP. - modulation (`Modulation`): - Produces scale/shift/gating parameters for modulated layers. - - Methods: - The forward method performs cross-attention and the MLP with modulation. - """ - - def __init__( - self, - hidden_size: int, - num_heads: int, - mlp_ratio: float = 4.0, - qk_scale: float | None = None, - ): - super().__init__() - - self.hidden_dim = hidden_size - self.num_heads = num_heads - self.head_dim = hidden_size // num_heads - self.scale = qk_scale or self.head_dim**-0.5 - - self.mlp_hidden_dim = int(hidden_size * mlp_ratio) - self.hidden_size = hidden_size - - # Pre-attention normalization for image tokens - self.img_pre_norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - - # PRXAttention module with built-in projections and norms - self.attention = PRXAttention( - query_dim=hidden_size, - heads=num_heads, - dim_head=self.head_dim, - bias=False, - out_bias=False, - eps=1e-6, - processor=PRXAttnProcessor2_0(), - ) - - # mlp - self.post_attention_layernorm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.gate_proj = nn.Linear(hidden_size, self.mlp_hidden_dim, bias=False) - self.up_proj = nn.Linear(hidden_size, self.mlp_hidden_dim, bias=False) - self.down_proj = nn.Linear(self.mlp_hidden_dim, hidden_size, bias=False) - self.mlp_act = nn.GELU(approximate="tanh") - - self.modulation = Modulation(hidden_size) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: torch.Tensor, - attention_mask: torch.Tensor | None = None, - **kwargs: dict[str, Any], - ) -> torch.Tensor: - r""" - Runs modulation-gated cross-attention and MLP, with residual connections. - - Args: - hidden_states (`torch.Tensor`): - Image tokens of shape `(B, L_img, hidden_size)`. - encoder_hidden_states (`torch.Tensor`): - Text tokens of shape `(B, L_txt, hidden_size)`. - temb (`torch.Tensor`): - Conditioning vector used by `Modulation` to produce scale/shift/gates, shape `(B, hidden_size)` (or - broadcastable). - image_rotary_emb (`torch.Tensor`): - Rotary positional embeddings applied inside attention. - attention_mask (`torch.Tensor`, *optional*): - Boolean mask for text tokens of shape `(B, L_txt)`, where `0` marks padding. - **kwargs: - Additional keyword arguments for API compatibility. - - Returns: - `torch.Tensor`: - Updated image tokens of shape `(B, L_img, hidden_size)`. - """ - - mod_attn, mod_mlp = self.modulation(temb) - attn_shift, attn_scale, attn_gate = mod_attn - mlp_shift, mlp_scale, mlp_gate = mod_mlp - - hidden_states_mod = (1 + attn_scale) * self.img_pre_norm(hidden_states) + attn_shift - - attn_out = self.attention( - hidden_states=hidden_states_mod, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) - - hidden_states = hidden_states + attn_gate * attn_out - - x = (1 + mlp_scale) * self.post_attention_layernorm(hidden_states) + mlp_shift - hidden_states = hidden_states + mlp_gate * (self.down_proj(self.mlp_act(self.gate_proj(x)) * self.up_proj(x))) - return hidden_states - - -class FinalLayer(nn.Module): - r""" - Final projection layer with adaptive LayerNorm modulation. - - This layer applies a normalized and modulated transformation to input tokens and projects them into patch-level - outputs. - - Args: - hidden_size (`int`): - Dimensionality of the input tokens. - patch_size (`int`): - Size of the square image patches. - out_channels (`int`): - Number of output channels per pixel (e.g. RGB = 3). - - Forward Inputs: - x (`torch.Tensor`): - Input tokens of shape `(B, L, hidden_size)`, where `L` is the number of patches. - vec (`torch.Tensor`): - Conditioning vector of shape `(B, hidden_size)` used to generate shift and scale parameters for adaptive - LayerNorm. - - Returns: - `torch.Tensor`: - Projected patch outputs of shape `(B, L, patch_size * patch_size * out_channels)`. - """ - - def __init__(self, hidden_size: int, patch_size: int, out_channels: int): - super().__init__() - self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) - self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True)) - - def forward(self, x: torch.Tensor, vec: torch.Tensor) -> torch.Tensor: - shift, scale = self.adaLN_modulation(vec).chunk(2, dim=1) - x = (1 + scale[:, None, :]) * self.norm_final(x) + shift[:, None, :] - x = self.linear(x) - return x - - -def img2seq(img: torch.Tensor, patch_size: int) -> torch.Tensor: - r""" - Flattens an image tensor into a sequence of non-overlapping patches. - - Args: - img (`torch.Tensor`): - Input image tensor of shape `(B, C, H, W)`. - patch_size (`int`): - Size of each square patch. Must evenly divide both `H` and `W`. - - Returns: - `torch.Tensor`: - Flattened patch sequence of shape `(B, L, C * patch_size * patch_size)`, where `L = (H // patch_size) * (W - // patch_size)` is the number of patches. - """ - b, c, h, w = img.shape - p = patch_size - - # Reshape to (B, C, H//p, p, W//p, p) separating grid and patch dimensions - img = img.reshape(b, c, h // p, p, w // p, p) - - # Permute to (B, H//p, W//p, C, p, p) using einsum - # n=batch, c=channels, h=grid_height, p=patch_height, w=grid_width, q=patch_width - img = torch.einsum("nchpwq->nhwcpq", img) - - # Flatten to (B, L, C * p * p) - img = img.reshape(b, -1, c * p * p) - return img - - -def seq2img(seq: torch.Tensor, patch_size: int, shape: torch.Tensor) -> torch.Tensor: - r""" - Reconstructs an image tensor from a sequence of patches (inverse of `img2seq`). - - Args: - seq (`torch.Tensor`): - Patch sequence of shape `(B, L, C * patch_size * patch_size)`, where `L = (H // patch_size) * (W // - patch_size)`. - patch_size (`int`): - Size of each square patch. - shape (`tuple` or `torch.Tensor`): - The original image spatial shape `(H, W)`. If a tensor is provided, the first two values are interpreted as - height and width. - - Returns: - `torch.Tensor`: - Reconstructed image tensor of shape `(B, C, H, W)`. - """ - if isinstance(shape, tuple): - h, w = shape[-2:] - elif isinstance(shape, torch.Tensor): - h, w = (int(shape[0]), int(shape[1])) - else: - raise NotImplementedError(f"shape type {type(shape)} not supported") - - b, l, d = seq.shape - p = patch_size - c = d // (p * p) - - # Reshape back to grid structure: (B, H//p, W//p, C, p, p) - seq = seq.reshape(b, h // p, w // p, c, p, p) - - # Permute back to image layout: (B, C, H//p, p, W//p, p) - # n=batch, h=grid_height, w=grid_width, c=channels, p=patch_height, q=patch_width - seq = torch.einsum("nhwcpq->nchpwq", seq) - - # Final reshape to (B, C, H, W) - seq = seq.reshape(b, c, h, w) - return seq - - -class PRXTransformer2DModel(ModelMixin, ConfigMixin, AttentionMixin): - r""" - Transformer-based 2D model for text to image generation. - - Args: - in_channels (`int`, *optional*, defaults to 16): - Number of input channels in the latent image. - patch_size (`int`, *optional*, defaults to 2): - Size of the square patches used to flatten the input image. - context_in_dim (`int`, *optional*, defaults to 2304): - Dimensionality of the text conditioning input. - hidden_size (`int`, *optional*, defaults to 1792): - Dimension of the hidden representation. - mlp_ratio (`float`, *optional*, defaults to 3.5): - Expansion ratio for the hidden dimension inside MLP blocks. - num_heads (`int`, *optional*, defaults to 28): - Number of attention heads. - depth (`int`, *optional*, defaults to 16): - Number of transformer blocks. - axes_dim (`list[int]`, *optional*): - list of dimensions for each positional embedding axis. Defaults to `[32, 32]`. - theta (`int`, *optional*, defaults to 10000): - Frequency scaling factor for rotary embeddings. - time_factor (`float`, *optional*, defaults to 1000.0): - Scaling factor applied in timestep embeddings. - time_max_period (`int`, *optional*, defaults to 10000): - Maximum frequency period for timestep embeddings. - bottleneck_size (`int`, *optional*): - If set, the image patch projection (`img_in`) uses a two-layer bottleneck (`patch_dim -> bottleneck_size -> - hidden_size`) instead of a single linear layer. Used by the pixel-space PRX-7B variant where the patch - dimension is large. - resolution_embeds (`bool`, *optional*, defaults to `False`): - Whether to condition the timestep modulation on the latent resolution `(H, W)` via a - `PRXResolutionEmbedder`. Used by the PRX-7B variant. - - Attributes: - pe_embedder (`EmbedND`): - Multi-axis rotary embedding generator for positional encodings. - img_in (`nn.Linear` or `nn.Sequential`): - Projection layer for image patch tokens (a two-layer bottleneck when `bottleneck_size` is set). - time_in (`MLPEmbedder`): - Embedding layer for timestep embeddings. - txt_in (`nn.Linear`): - Projection layer for text conditioning. - blocks (`nn.ModuleList`): - Stack of transformer blocks (`PRXBlock`). - final_layer (`LastLayer`): - Projection layer mapping hidden tokens back to patch outputs. - - Methods: - attn_processors: - Returns a dictionary of all attention processors in the model. - set_attn_processor(processor): - Replaces attention processors across all attention layers. - process_inputs(image_latent, txt): - Converts inputs into patch tokens, encodes text, and produces positional encodings. - compute_timestep_embedding(timestep, dtype): - Creates a timestep embedding of dimension 256, scaled and projected. - forward_transformers(image_latent, cross_attn_conditioning, timestep, time_embedding, attention_mask, - **block_kwargs): - Runs the sequence of transformer blocks over image and text tokens. - forward(image_latent, timestep, cross_attn_conditioning, micro_conditioning, cross_attn_mask=None, - attention_kwargs=None, return_dict=True): - Full forward pass from latent input to reconstructed output image. - - Returns: - `Transformer2DModelOutput` if `return_dict=True` (default), otherwise a tuple containing: - - `sample` (`torch.Tensor`): Reconstructed image of shape `(B, C, H, W)`. - """ - - config_name = "config.json" - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 16, - patch_size: int = 2, - context_in_dim: int = 2304, - hidden_size: int = 1792, - mlp_ratio: float = 3.5, - num_heads: int = 28, - depth: int = 16, - axes_dim: list = None, - theta: int = 10000, - time_factor: float = 1000.0, - time_max_period: int = 10000, - bottleneck_size: int | None = None, - resolution_embeds: bool = False, - ): - super().__init__() - - if axes_dim is None: - axes_dim = [32, 32] - - # Store parameters directly - self.in_channels = in_channels - self.patch_size = patch_size - self.out_channels = self.in_channels * self.patch_size**2 - - self.time_factor = time_factor - self.time_max_period = time_max_period - - if hidden_size % num_heads != 0: - raise ValueError(f"Hidden size {hidden_size} must be divisible by num_heads {num_heads}") - - pe_dim = hidden_size // num_heads - - if sum(axes_dim) != pe_dim: - raise ValueError(f"Got {axes_dim} but expected positional dim {pe_dim}") - - self.hidden_size = hidden_size - self.num_heads = num_heads - self.pe_embedder = PRXEmbedND(dim=pe_dim, theta=theta, axes_dim=axes_dim) - patch_dim = self.in_channels * self.patch_size**2 - if bottleneck_size is not None: - # Two-layer bottleneck projection (used by pixel-space PRX where the patch dimension is large). - self.img_in = nn.Sequential( - nn.Linear(patch_dim, bottleneck_size, bias=True), - nn.Linear(bottleneck_size, self.hidden_size, bias=True), - ) - else: - self.img_in = nn.Linear(patch_dim, self.hidden_size, bias=True) - self.time_in = MLPEmbedder(in_dim=256, hidden_dim=self.hidden_size) - self.txt_in = nn.Linear(context_in_dim, self.hidden_size) - - self.resolution_embedder = ( - PRXResolutionEmbedder(self.hidden_size, max_period=time_max_period) if resolution_embeds else None - ) - - self.blocks = nn.ModuleList( - [ - PRXBlock( - self.hidden_size, - self.num_heads, - mlp_ratio=mlp_ratio, - ) - for i in range(depth) - ] - ) - - self.final_layer = FinalLayer(self.hidden_size, 1, self.out_channels) - - self.gradient_checkpointing = False - - def _compute_timestep_embedding(self, timestep: torch.Tensor, dtype: torch.dtype) -> torch.Tensor: - return self.time_in( - get_timestep_embedding( - timesteps=timestep, - embedding_dim=256, - max_period=self.time_max_period, - scale=self.time_factor, - flip_sin_to_cos=True, # Match original cos, sin order - downscale_freq_shift=0.0, - ).to(dtype) - ) - - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor, ...] | Transformer2DModelOutput: - r""" - Forward pass of the PRXTransformer2DModel. - - The latent image is split into patch tokens, combined with text conditioning, and processed through a stack of - transformer blocks modulated by the timestep. The output is reconstructed into the latent image space. - - Args: - hidden_states (`torch.Tensor`): - Input latent image tensor of shape `(B, C, H, W)`. - timestep (`torch.Tensor`): - Timestep tensor of shape `(B,)` or `(1,)`, used for temporal conditioning. - encoder_hidden_states (`torch.Tensor`): - Text conditioning tensor of shape `(B, L_txt, context_in_dim)`. - attention_mask (`torch.Tensor`, *optional*): - Boolean mask of shape `(B, L_txt)`, where `0` marks padding in the text sequence. - attention_kwargs (`dict`, *optional*): - Additional arguments passed to attention layers. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return a `Transformer2DModelOutput` or a tuple. - - Returns: - `Transformer2DModelOutput` if `return_dict=True`, otherwise a tuple: - - - `sample` (`torch.Tensor`): Output latent image of shape `(B, C, H, W)`. - """ - # Process text conditioning - txt = self.txt_in(encoder_hidden_states) - - # Convert image to sequence and embed - img = img2seq(hidden_states, self.patch_size) - img = self.img_in(img) - - # Generate positional embeddings - bs, _, h, w = hidden_states.shape - img_ids = get_image_ids(bs, h, w, patch_size=self.patch_size, device=hidden_states.device) - pe = self.pe_embedder(img_ids) - - # Compute time embedding - vec = self._compute_timestep_embedding(timestep, dtype=img.dtype) - - # Add resolution conditioning (PRX-7B "vec" mode): embed the latent (H, W) and add it to the timestep vector - # so every block's modulation is resolution-aware. - if self.resolution_embedder is not None: - height = torch.full((bs,), h, device=hidden_states.device, dtype=torch.float32) - width = torch.full((bs,), w, device=hidden_states.device, dtype=torch.float32) - vec = vec + self.resolution_embedder(height, width, dtype=vec.dtype) - - # Apply transformer blocks - for block in self.blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - img = self._gradient_checkpointing_func( - block.__call__, - img, - txt, - vec, - pe, - attention_mask, - ) - else: - img = block( - hidden_states=img, - encoder_hidden_states=txt, - temb=vec, - image_rotary_emb=pe, - attention_mask=attention_mask, - ) - - # Final layer and convert back to image - img = self.final_layer(img, vec) - output = seq2img(img, self.patch_size, hidden_states.shape) - - if not return_dict: - return (output,) - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_qwenimage.py b/diffusers/models/transformers/transformer_qwenimage.py deleted file mode 100644 index 464712bd94fdc095246a3f0e0d0191c2e4502817..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_qwenimage.py +++ /dev/null @@ -1,966 +0,0 @@ -# Copyright 2025 Qwen-Image Team, The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from math import prod -from typing import Any - -import numpy as np -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import lru_cache_unless_export, maybe_allow_in_graph -from .._modeling_parallel import ContextParallelInput, ContextParallelOutput -from ..attention import AttentionMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..attention_processor import Attention -from ..cache_utils import CacheMixin -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def get_timestep_embedding( - timesteps: torch.Tensor, - embedding_dim: int, - flip_sin_to_cos: bool = False, - downscale_freq_shift: float = 1, - scale: float = 1, - max_period: int = 10000, -) -> torch.Tensor: - """ - This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings. - - Args - timesteps (torch.Tensor): - a 1-D Tensor of N indices, one per batch element. These may be fractional. - embedding_dim (int): - the dimension of the output. - flip_sin_to_cos (bool): - Whether the embedding order should be `cos, sin` (if True) or `sin, cos` (if False) - downscale_freq_shift (float): - Controls the delta between frequencies between dimensions - scale (float): - Scaling factor applied to the embeddings. - max_period (int): - Controls the maximum frequency of the embeddings - Returns - torch.Tensor: an [N x dim] Tensor of positional embeddings. - """ - assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array" - - half_dim = embedding_dim // 2 - exponent = -math.log(max_period) * torch.arange( - start=0, end=half_dim, dtype=torch.float32, device=timesteps.device - ) - exponent = exponent / (half_dim - downscale_freq_shift) - - emb = torch.exp(exponent).to(timesteps.dtype) - emb = timesteps[:, None].float() * emb[None, :] - - # scale embeddings - emb = scale * emb - - # concat sine and cosine embeddings - emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) - - # flip sine and cosine embeddings - if flip_sin_to_cos: - emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) - - # zero pad - if embedding_dim % 2 == 1: - emb = torch.nn.functional.pad(emb, (0, 1, 0, 0)) - return emb - - -def apply_rotary_emb_qwen( - x: torch.Tensor, - freqs_cis: torch.Tensor | tuple[torch.Tensor], - use_real: bool = True, - use_real_unbind_dim: int = -1, -) -> tuple[torch.Tensor, torch.Tensor]: - """ - Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings - to the given query or key 'x' tensors using the provided frequency tensor 'freqs_cis'. The input tensors are - reshaped as complex numbers, and the frequency tensor is reshaped for broadcasting compatibility. The resulting - tensors contain rotary embeddings and are returned as real tensors. - - Args: - x (`torch.Tensor`): - Query or key tensor to apply rotary embeddings. [B, S, H, D] xk (torch.Tensor): Key tensor to apply - freqs_cis (`tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],) - - Returns: - tuple[torch.Tensor, torch.Tensor]: tuple of modified query tensor and key tensor with rotary embeddings. - """ - if use_real: - cos, sin = freqs_cis # [S, D] - cos = cos[None, None] - sin = sin[None, None] - cos, sin = cos.to(x.device), sin.to(x.device) - - if use_real_unbind_dim == -1: - # Used for flux, cogvideox, hunyuan-dit - x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2] - x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3) - elif use_real_unbind_dim == -2: - # Used for Stable Audio, OmniGen, CogView4 and Cosmos - x_real, x_imag = x.reshape(*x.shape[:-1], 2, -1).unbind(-2) # [B, S, H, D//2] - x_rotated = torch.cat([-x_imag, x_real], dim=-1) - else: - raise ValueError(f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2.") - - out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype) - - return out - else: - x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2)) - freqs_cis = freqs_cis.unsqueeze(1) - x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3) - - return x_out.type_as(x) - - -def compute_text_seq_len_from_mask( - encoder_hidden_states: torch.Tensor, encoder_hidden_states_mask: torch.Tensor | None -) -> tuple[int, torch.Tensor | None, torch.Tensor | None]: - """ - Compute text sequence length without assuming contiguous masks. Returns length for RoPE and a normalized bool mask. - """ - batch_size, text_seq_len = encoder_hidden_states.shape[:2] - if encoder_hidden_states_mask is None: - return text_seq_len, None, None - - if encoder_hidden_states_mask.shape[:2] != (batch_size, text_seq_len): - raise ValueError( - f"`encoder_hidden_states_mask` shape {encoder_hidden_states_mask.shape} must match " - f"(batch_size, text_seq_len)=({batch_size}, {text_seq_len})." - ) - - if encoder_hidden_states_mask.dtype != torch.bool: - encoder_hidden_states_mask = encoder_hidden_states_mask.to(torch.bool) - - position_ids = torch.arange(text_seq_len, device=encoder_hidden_states.device, dtype=torch.long) - active_positions = torch.where(encoder_hidden_states_mask, position_ids, position_ids.new_zeros(())) - has_active = encoder_hidden_states_mask.any(dim=1) - per_sample_len = torch.where( - has_active, - active_positions.max(dim=1).values + 1, - torch.as_tensor(text_seq_len, device=encoder_hidden_states.device), - ) - return text_seq_len, per_sample_len, encoder_hidden_states_mask - - -class QwenTimestepProjEmbeddings(nn.Module): - def __init__(self, embedding_dim, use_additional_t_cond=False): - super().__init__() - - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0, scale=1000) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - self.use_additional_t_cond = use_additional_t_cond - if use_additional_t_cond: - self.addition_t_embedding = nn.Embedding(2, embedding_dim) - - def forward(self, timestep, hidden_states, addition_t_cond=None): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_states.dtype)) # (N, D) - - conditioning = timesteps_emb - if self.use_additional_t_cond: - if addition_t_cond is None: - raise ValueError("When additional_t_cond is True, addition_t_cond must be provided.") - addition_t_emb = self.addition_t_embedding(addition_t_cond) - addition_t_emb = addition_t_emb.to(dtype=hidden_states.dtype) - conditioning = conditioning + addition_t_emb - - return conditioning - - -class QwenEmbedRope(nn.Module): - def __init__(self, theta: int, axes_dim: list[int], scale_rope=False): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - pos_index = torch.arange(4096) - neg_index = torch.arange(4096).flip(0) * -1 - 1 - self.pos_freqs = torch.cat( - [ - self.rope_params(pos_index, self.axes_dim[0], self.theta), - self.rope_params(pos_index, self.axes_dim[1], self.theta), - self.rope_params(pos_index, self.axes_dim[2], self.theta), - ], - dim=1, - ) - self.neg_freqs = torch.cat( - [ - self.rope_params(neg_index, self.axes_dim[0], self.theta), - self.rope_params(neg_index, self.axes_dim[1], self.theta), - self.rope_params(neg_index, self.axes_dim[2], self.theta), - ], - dim=1, - ) - - # DO NOT USING REGISTER BUFFER HERE, IT WILL CAUSE COMPLEX NUMBERS LOSE ITS IMAGINARY PART - self.scale_rope = scale_rope - - def rope_params(self, index, dim, theta=10000): - """ - Args: - index: [0, 1, 2, 3] 1D Tensor representing the position index of the token - """ - assert dim % 2 == 0 - freqs = torch.outer(index, 1.0 / torch.pow(theta, torch.arange(0, dim, 2).to(torch.float32).div(dim))) - freqs = torch.polar(torch.ones_like(freqs), freqs) - return freqs - - @lru_cache_unless_export(maxsize=None) - def _get_device_freqs(self, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]: - """Return pos_freqs and neg_freqs on the given device.""" - return self.pos_freqs.to(device), self.neg_freqs.to(device) - - def forward( - self, - video_fhw: tuple[int, int, int, list[tuple[int, int, int]]], - device: torch.device = None, - max_txt_seq_len: int | torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - """ - Args: - video_fhw (`tuple[int, int, int]` or `list[tuple[int, int, int]]`): - A list of 3 integers [frame, height, width] representing the shape of the video. - device: (`torch.device`, *optional*): - The device on which to perform the RoPE computation. - max_txt_seq_len (`int` or `torch.Tensor`, *optional*): - The maximum text sequence length for RoPE computation. This should match the encoder hidden states - sequence length. Can be either an int or a scalar tensor (for torch.compile compatibility). - """ - if max_txt_seq_len is None: - raise ValueError("`max_txt_seq_len` must be provided.") - - # Validate batch inference with variable-sized images - if isinstance(video_fhw, list) and len(video_fhw) > 1: - # Check if all instances have the same size - first_fhw = video_fhw[0] - if not all(fhw == first_fhw for fhw in video_fhw): - logger.warning( - "Batch inference with variable-sized images is not currently supported in QwenEmbedRope. " - "All images in the batch should have the same dimensions (frame, height, width). " - f"Detected sizes: {video_fhw}. Using the first image's dimensions {first_fhw} " - "for RoPE computation, which may lead to incorrect results for other images in the batch." - ) - - if isinstance(video_fhw, list): - video_fhw = video_fhw[0] - if not isinstance(video_fhw, list): - video_fhw = [video_fhw] - - vid_freqs = [] - max_vid_index = 0 - for idx, fhw in enumerate(video_fhw): - frame, height, width = fhw - # RoPE frequencies are cached via a lru_cache decorator on _compute_video_freqs - video_freq = self._compute_video_freqs(frame, height, width, idx, device) - vid_freqs.append(video_freq) - - if self.scale_rope: - max_vid_index = max(height // 2, width // 2, max_vid_index) - else: - max_vid_index = max(height, width, max_vid_index) - - max_txt_seq_len_int = int(max_txt_seq_len) - # Use cached device-transferred freqs to avoid CPU→GPU sync every forward call - pos_freqs_device, _ = self._get_device_freqs(device) - txt_freqs = pos_freqs_device[max_vid_index : max_vid_index + max_txt_seq_len_int, ...] - vid_freqs = torch.cat(vid_freqs, dim=0) - - return vid_freqs, txt_freqs - - @lru_cache_unless_export(maxsize=128) - def _compute_video_freqs( - self, frame: int, height: int, width: int, idx: int = 0, device: torch.device = None - ) -> torch.Tensor: - seq_lens = frame * height * width - pos_freqs, neg_freqs = ( - self._get_device_freqs(device) if device is not None else (self.pos_freqs, self.neg_freqs) - ) - - freqs_pos = pos_freqs.split([x // 2 for x in self.axes_dim], dim=1) - freqs_neg = neg_freqs.split([x // 2 for x in self.axes_dim], dim=1) - - freqs_frame = freqs_pos[0][idx : idx + frame].view(frame, 1, 1, -1).expand(frame, height, width, -1) - if self.scale_rope: - freqs_height = torch.cat([freqs_neg[1][-(height - height // 2) :], freqs_pos[1][: height // 2]], dim=0) - freqs_height = freqs_height.view(1, height, 1, -1).expand(frame, height, width, -1) - freqs_width = torch.cat([freqs_neg[2][-(width - width // 2) :], freqs_pos[2][: width // 2]], dim=0) - freqs_width = freqs_width.view(1, 1, width, -1).expand(frame, height, width, -1) - else: - freqs_height = freqs_pos[1][:height].view(1, height, 1, -1).expand(frame, height, width, -1) - freqs_width = freqs_pos[2][:width].view(1, 1, width, -1).expand(frame, height, width, -1) - - freqs = torch.cat([freqs_frame, freqs_height, freqs_width], dim=-1).reshape(seq_lens, -1) - return freqs.clone().contiguous() - - -class QwenEmbedLayer3DRope(nn.Module): - def __init__(self, theta: int, axes_dim: list[int], scale_rope=False): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - pos_index = torch.arange(4096) - neg_index = torch.arange(4096).flip(0) * -1 - 1 - self.pos_freqs = torch.cat( - [ - self.rope_params(pos_index, self.axes_dim[0], self.theta), - self.rope_params(pos_index, self.axes_dim[1], self.theta), - self.rope_params(pos_index, self.axes_dim[2], self.theta), - ], - dim=1, - ) - self.neg_freqs = torch.cat( - [ - self.rope_params(neg_index, self.axes_dim[0], self.theta), - self.rope_params(neg_index, self.axes_dim[1], self.theta), - self.rope_params(neg_index, self.axes_dim[2], self.theta), - ], - dim=1, - ) - - self.scale_rope = scale_rope - - def rope_params(self, index, dim, theta=10000): - """ - Args: - index: [0, 1, 2, 3] 1D Tensor representing the position index of the token - """ - assert dim % 2 == 0 - freqs = torch.outer(index, 1.0 / torch.pow(theta, torch.arange(0, dim, 2).to(torch.float32).div(dim))) - freqs = torch.polar(torch.ones_like(freqs), freqs) - return freqs - - @lru_cache_unless_export(maxsize=None) - def _get_device_freqs(self, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]: - """Return pos_freqs and neg_freqs on the given device.""" - return self.pos_freqs.to(device), self.neg_freqs.to(device) - - def forward( - self, - video_fhw: tuple[int, int, int, list[tuple[int, int, int]]], - max_txt_seq_len: int | torch.Tensor, - device: torch.device = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - """ - Args: - video_fhw (`tuple[int, int, int]` or `list[tuple[int, int, int]]`): - A list of 3 integers [frame, height, width] representing the shape of the video, or a list of layer - structures. - max_txt_seq_len (`int` or `torch.Tensor`): - The maximum text sequence length for RoPE computation. This should match the encoder hidden states - sequence length. Can be either an int or a scalar tensor (for torch.compile compatibility). - device: (`torch.device`, *optional*): - The device on which to perform the RoPE computation. - """ - # Validate batch inference with variable-sized images - # In Layer3DRope, the outer list represents batch, inner list/tuple represents layers - if isinstance(video_fhw, list) and len(video_fhw) > 1: - # Check if this is batch inference (list of layer lists/tuples) - first_entry = video_fhw[0] - if not all(entry == first_entry for entry in video_fhw): - logger.warning( - "Batch inference with variable-sized images is not currently supported in QwenEmbedLayer3DRope. " - "All images in the batch should have the same layer structure. " - f"Detected sizes: {video_fhw}. Using the first image's layer structure {first_entry} " - "for RoPE computation, which may lead to incorrect results for other images in the batch." - ) - - if isinstance(video_fhw, list): - video_fhw = video_fhw[0] - if not isinstance(video_fhw, list): - video_fhw = [video_fhw] - - vid_freqs = [] - max_vid_index = 0 - layer_num = len(video_fhw) - 1 - for idx, fhw in enumerate(video_fhw): - frame, height, width = fhw - if idx != layer_num: - video_freq = self._compute_video_freqs(frame, height, width, idx, device) - else: - ### For the condition image, we set the layer index to -1 - video_freq = self._compute_condition_freqs(frame, height, width, device) - vid_freqs.append(video_freq) - - if self.scale_rope: - max_vid_index = max(height // 2, width // 2, max_vid_index) - else: - max_vid_index = max(height, width, max_vid_index) - - max_vid_index = max(max_vid_index, layer_num) - max_txt_seq_len_int = int(max_txt_seq_len) - # Use cached device-transferred freqs to avoid CPU→GPU sync every forward call - pos_freqs_device, _ = self._get_device_freqs(device) - txt_freqs = pos_freqs_device[max_vid_index : max_vid_index + max_txt_seq_len_int, ...] - vid_freqs = torch.cat(vid_freqs, dim=0) - - return vid_freqs, txt_freqs - - @lru_cache_unless_export(maxsize=None) - def _compute_video_freqs(self, frame, height, width, idx=0, device: torch.device = None): - seq_lens = frame * height * width - pos_freqs, neg_freqs = ( - self._get_device_freqs(device) if device is not None else (self.pos_freqs, self.neg_freqs) - ) - - freqs_pos = pos_freqs.split([x // 2 for x in self.axes_dim], dim=1) - freqs_neg = neg_freqs.split([x // 2 for x in self.axes_dim], dim=1) - - freqs_frame = freqs_pos[0][idx : idx + frame].view(frame, 1, 1, -1).expand(frame, height, width, -1) - if self.scale_rope: - freqs_height = torch.cat([freqs_neg[1][-(height - height // 2) :], freqs_pos[1][: height // 2]], dim=0) - freqs_height = freqs_height.view(1, height, 1, -1).expand(frame, height, width, -1) - freqs_width = torch.cat([freqs_neg[2][-(width - width // 2) :], freqs_pos[2][: width // 2]], dim=0) - freqs_width = freqs_width.view(1, 1, width, -1).expand(frame, height, width, -1) - else: - freqs_height = freqs_pos[1][:height].view(1, height, 1, -1).expand(frame, height, width, -1) - freqs_width = freqs_pos[2][:width].view(1, 1, width, -1).expand(frame, height, width, -1) - - freqs = torch.cat([freqs_frame, freqs_height, freqs_width], dim=-1).reshape(seq_lens, -1) - return freqs.clone().contiguous() - - @lru_cache_unless_export(maxsize=None) - def _compute_condition_freqs(self, frame, height, width, device: torch.device = None): - seq_lens = frame * height * width - pos_freqs, neg_freqs = ( - self._get_device_freqs(device) if device is not None else (self.pos_freqs, self.neg_freqs) - ) - - freqs_pos = pos_freqs.split([x // 2 for x in self.axes_dim], dim=1) - freqs_neg = neg_freqs.split([x // 2 for x in self.axes_dim], dim=1) - - freqs_frame = freqs_neg[0][-1:].view(frame, 1, 1, -1).expand(frame, height, width, -1) - if self.scale_rope: - freqs_height = torch.cat([freqs_neg[1][-(height - height // 2) :], freqs_pos[1][: height // 2]], dim=0) - freqs_height = freqs_height.view(1, height, 1, -1).expand(frame, height, width, -1) - freqs_width = torch.cat([freqs_neg[2][-(width - width // 2) :], freqs_pos[2][: width // 2]], dim=0) - freqs_width = freqs_width.view(1, 1, width, -1).expand(frame, height, width, -1) - else: - freqs_height = freqs_pos[1][:height].view(1, height, 1, -1).expand(frame, height, width, -1) - freqs_width = freqs_pos[2][:width].view(1, 1, width, -1).expand(frame, height, width, -1) - - freqs = torch.cat([freqs_frame, freqs_height, freqs_width], dim=-1).reshape(seq_lens, -1) - return freqs.clone().contiguous() - - -class QwenDoubleStreamAttnProcessor2_0: - """ - Attention processor for Qwen double-stream architecture, matching DoubleStreamLayerMegatron logic. This processor - implements joint attention computation where text and image streams are processed together. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "QwenDoubleStreamAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, # Image stream - encoder_hidden_states: torch.FloatTensor = None, # Text stream - encoder_hidden_states_mask: torch.FloatTensor = None, - attention_mask: torch.FloatTensor | None = None, - image_rotary_emb: torch.Tensor | None = None, - ) -> torch.FloatTensor: - if encoder_hidden_states is None: - raise ValueError("QwenDoubleStreamAttnProcessor2_0 requires encoder_hidden_states (text stream)") - - if attention_mask is not None: - raise ValueError( - "QwenDoubleStreamAttnProcessor2_0 does not accept an external attention_mask. " - "Pass encoder_hidden_states_mask to let the processor build the joint mask." - ) - - if encoder_hidden_states_mask is not None: - seq_img = hidden_states.shape[1] - image_mask = torch.ones((hidden_states.shape[0], seq_img), dtype=torch.bool, device=hidden_states.device) - attention_mask = torch.cat([encoder_hidden_states_mask, image_mask], dim=1) - attention_mask = attention_mask[:, None, None, :] - - seq_txt = encoder_hidden_states.shape[1] - - # Compute QKV for image stream (sample projections) - img_query = attn.to_q(hidden_states) - img_key = attn.to_k(hidden_states) - img_value = attn.to_v(hidden_states) - - # Compute QKV for text stream (context projections) - txt_query = attn.add_q_proj(encoder_hidden_states) - txt_key = attn.add_k_proj(encoder_hidden_states) - txt_value = attn.add_v_proj(encoder_hidden_states) - - # Reshape for multi-head attention - img_query = img_query.unflatten(-1, (attn.heads, -1)) - img_key = img_key.unflatten(-1, (attn.heads, -1)) - img_value = img_value.unflatten(-1, (attn.heads, -1)) - - txt_query = txt_query.unflatten(-1, (attn.heads, -1)) - txt_key = txt_key.unflatten(-1, (attn.heads, -1)) - txt_value = txt_value.unflatten(-1, (attn.heads, -1)) - - # Apply QK normalization - if attn.norm_q is not None: - img_query = attn.norm_q(img_query) - if attn.norm_k is not None: - img_key = attn.norm_k(img_key) - if attn.norm_added_q is not None: - txt_query = attn.norm_added_q(txt_query) - if attn.norm_added_k is not None: - txt_key = attn.norm_added_k(txt_key) - - # Apply RoPE - if image_rotary_emb is not None: - img_freqs, txt_freqs = image_rotary_emb - img_query = apply_rotary_emb_qwen(img_query, img_freqs, use_real=False) - img_key = apply_rotary_emb_qwen(img_key, img_freqs, use_real=False) - txt_query = apply_rotary_emb_qwen(txt_query, txt_freqs, use_real=False) - txt_key = apply_rotary_emb_qwen(txt_key, txt_freqs, use_real=False) - - # Concatenate for joint attention - # Order: [text, image] - joint_query = torch.cat([txt_query, img_query], dim=1) - joint_key = torch.cat([txt_key, img_key], dim=1) - joint_value = torch.cat([txt_value, img_value], dim=1) - - joint_hidden_states = dispatch_attention_fn( - joint_query, - joint_key, - joint_value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - # Reshape back - joint_hidden_states = joint_hidden_states.flatten(2, 3) - joint_hidden_states = joint_hidden_states.to(joint_query.dtype) - - # Split attention outputs back - txt_attn_output = joint_hidden_states[:, :seq_txt, :] # Text part - img_attn_output = joint_hidden_states[:, seq_txt:, :] # Image part - - # Apply output projections - img_attn_output = attn.to_out[0](img_attn_output.contiguous()) - if len(attn.to_out) > 1: - img_attn_output = attn.to_out[1](img_attn_output) # dropout - - txt_attn_output = attn.to_add_out(txt_attn_output.contiguous()) - - return img_attn_output, txt_attn_output - - -@maybe_allow_in_graph -class QwenImageTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - qk_norm: str = "rms_norm", - eps: float = 1e-6, - zero_cond_t: bool = False, - ): - super().__init__() - - self.dim = dim - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - - # Image processing modules - self.img_mod = nn.Sequential( - nn.SiLU(), - nn.Linear(dim, 6 * dim, bias=True), # For scale, shift, gate for norm1 and norm2 - ) - self.img_norm1 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps) - self.attn = Attention( - query_dim=dim, - cross_attention_dim=None, # Enable cross attention for joint computation - added_kv_proj_dim=dim, # Enable added KV projections for text stream - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - context_pre_only=False, - bias=True, - processor=QwenDoubleStreamAttnProcessor2_0(), - qk_norm=qk_norm, - eps=eps, - ) - self.img_norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps) - self.img_mlp = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - # Text processing modules - self.txt_mod = nn.Sequential( - nn.SiLU(), - nn.Linear(dim, 6 * dim, bias=True), # For scale, shift, gate for norm1 and norm2 - ) - self.txt_norm1 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps) - # Text doesn't need separate attention - it's handled by img_attn joint computation - self.txt_norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps) - self.txt_mlp = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - self.zero_cond_t = zero_cond_t - - def _modulate(self, x, mod_params, index=None): - """Apply modulation to input tensor""" - # x: b l d, shift: b d, scale: b d, gate: b d - shift, scale, gate = mod_params.chunk(3, dim=-1) - - if index is not None: - # Assuming mod_params batch dim is 2*actual_batch (chunked into 2 parts) - # So shift, scale, gate have shape [2*actual_batch, d] - actual_batch = shift.size(0) // 2 - shift_0, shift_1 = shift[:actual_batch], shift[actual_batch:] # each: [actual_batch, d] - scale_0, scale_1 = scale[:actual_batch], scale[actual_batch:] - gate_0, gate_1 = gate[:actual_batch], gate[actual_batch:] - - # index: [b, l] where b is actual batch size - # Expand to [b, l, 1] to match feature dimension - index_expanded = index.unsqueeze(-1) # [b, l, 1] - - # Expand chunks to [b, 1, d] then broadcast to [b, l, d] - shift_0_exp = shift_0.unsqueeze(1) # [b, 1, d] - shift_1_exp = shift_1.unsqueeze(1) # [b, 1, d] - scale_0_exp = scale_0.unsqueeze(1) - scale_1_exp = scale_1.unsqueeze(1) - gate_0_exp = gate_0.unsqueeze(1) - gate_1_exp = gate_1.unsqueeze(1) - - # Use torch.where to select based on index - shift_result = torch.where(index_expanded == 0, shift_0_exp, shift_1_exp) - scale_result = torch.where(index_expanded == 0, scale_0_exp, scale_1_exp) - gate_result = torch.where(index_expanded == 0, gate_0_exp, gate_1_exp) - else: - shift_result = shift.unsqueeze(1) - scale_result = scale.unsqueeze(1) - gate_result = gate.unsqueeze(1) - - return x * (1 + scale_result) + shift_result, gate_result - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_mask: torch.Tensor, - temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - joint_attention_kwargs: dict[str, Any] | None = None, - modulate_index: list[int] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - # Get modulation parameters for both streams - img_mod_params = self.img_mod(temb) # [B, 6*dim] - - if self.zero_cond_t: - temb = torch.chunk(temb, 2, dim=0)[0] - txt_mod_params = self.txt_mod(temb) # [B, 6*dim] - - # Split modulation parameters for norm1 and norm2 - img_mod1, img_mod2 = img_mod_params.chunk(2, dim=-1) # Each [B, 3*dim] - txt_mod1, txt_mod2 = txt_mod_params.chunk(2, dim=-1) # Each [B, 3*dim] - - # Process image stream - norm1 + modulation - img_normed = self.img_norm1(hidden_states) - img_modulated, img_gate1 = self._modulate(img_normed, img_mod1, modulate_index) - - # Process text stream - norm1 + modulation - txt_normed = self.txt_norm1(encoder_hidden_states) - txt_modulated, txt_gate1 = self._modulate(txt_normed, txt_mod1) - - # Use QwenAttnProcessor2_0 for joint attention computation - # This directly implements the DoubleStreamLayerMegatron logic: - # 1. Computes QKV for both streams - # 2. Applies QK normalization and RoPE - # 3. Concatenates and runs joint attention - # 4. Splits results back to separate streams - joint_attention_kwargs = joint_attention_kwargs or {} - attn_output = self.attn( - hidden_states=img_modulated, # Image stream (will be processed as "sample") - encoder_hidden_states=txt_modulated, # Text stream (will be processed as "context") - encoder_hidden_states_mask=encoder_hidden_states_mask, - image_rotary_emb=image_rotary_emb, - **joint_attention_kwargs, - ) - - # QwenAttnProcessor2_0 returns (img_output, txt_output) when encoder_hidden_states is provided - img_attn_output, txt_attn_output = attn_output - - # Apply attention gates and add residual (like in Megatron) - hidden_states = hidden_states + img_gate1 * img_attn_output - encoder_hidden_states = encoder_hidden_states + txt_gate1 * txt_attn_output - - # Process image stream - norm2 + MLP - img_normed2 = self.img_norm2(hidden_states) - img_modulated2, img_gate2 = self._modulate(img_normed2, img_mod2, modulate_index) - img_mlp_output = self.img_mlp(img_modulated2) - hidden_states = hidden_states + img_gate2 * img_mlp_output - - # Process text stream - norm2 + MLP - txt_normed2 = self.txt_norm2(encoder_hidden_states) - txt_modulated2, txt_gate2 = self._modulate(txt_normed2, txt_mod2) - txt_mlp_output = self.txt_mlp(txt_modulated2) - encoder_hidden_states = encoder_hidden_states + txt_gate2 * txt_mlp_output - - # Clip to prevent overflow for fp16 - if encoder_hidden_states.dtype == torch.float16: - encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) - if hidden_states.dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - return encoder_hidden_states, hidden_states - - -class QwenImageTransformer2DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin, AttentionMixin -): - """ - The Transformer model introduced in Qwen. - - Args: - patch_size (`int`, defaults to `2`): - Patch size to turn the input data into small patches. - in_channels (`int`, defaults to `64`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `None`): - The number of channels in the output. If not specified, it defaults to `in_channels`. - num_layers (`int`, defaults to `60`): - The number of layers of dual stream DiT blocks to use. - attention_head_dim (`int`, defaults to `128`): - The number of dimensions to use for each attention head. - num_attention_heads (`int`, defaults to `24`): - The number of attention heads to use. - joint_attention_dim (`int`, defaults to `3584`): - The number of dimensions to use for the joint attention (embedding/channel dimension of - `encoder_hidden_states`). - guidance_embeds (`bool`, defaults to `False`): - Whether to use guidance embeddings for guidance-distilled variant of the model. - axes_dims_rope (`tuple[int]`, defaults to `(16, 56, 56)`): - The dimensions to use for the rotary positional embeddings. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["QwenImageTransformerBlock"] - _skip_layerwise_casting_patterns = ["pos_embed", "norm"] - _repeated_blocks = ["QwenImageTransformerBlock"] - # Make CP plan compatible with https://github.com/huggingface/diffusers/pull/12702 - _cp_plan = { - "transformer_blocks.0": { - "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - "encoder_hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - }, - "transformer_blocks.*": { - "modulate_index": ContextParallelInput(split_dim=1, expected_dims=2, split_output=False), - "encoder_hidden_states_mask": ContextParallelInput(split_dim=1, expected_dims=2, split_output=False), - }, - "pos_embed": { - 0: ContextParallelInput(split_dim=0, expected_dims=2, split_output=True), - 1: ContextParallelInput(split_dim=0, expected_dims=2, split_output=True), - }, - "proj_out": ContextParallelOutput(gather_dim=1, expected_dims=3), - } - - @register_to_config - def __init__( - self, - patch_size: int = 2, - in_channels: int = 64, - out_channels: int | None = 16, - num_layers: int = 60, - attention_head_dim: int = 128, - num_attention_heads: int = 24, - joint_attention_dim: int = 3584, - guidance_embeds: bool = False, # TODO: this should probably be removed - axes_dims_rope: tuple[int, int, int] = (16, 56, 56), - zero_cond_t: bool = False, - use_additional_t_cond: bool = False, - use_layer3d_rope: bool = False, - ): - super().__init__() - self.out_channels = out_channels or in_channels - self.inner_dim = num_attention_heads * attention_head_dim - - if not use_layer3d_rope: - self.pos_embed = QwenEmbedRope(theta=10000, axes_dim=list(axes_dims_rope), scale_rope=True) - else: - self.pos_embed = QwenEmbedLayer3DRope(theta=10000, axes_dim=list(axes_dims_rope), scale_rope=True) - - self.time_text_embed = QwenTimestepProjEmbeddings( - embedding_dim=self.inner_dim, use_additional_t_cond=use_additional_t_cond - ) - - self.txt_norm = RMSNorm(joint_attention_dim, eps=1e-6) - - self.img_in = nn.Linear(in_channels, self.inner_dim) - self.txt_in = nn.Linear(joint_attention_dim, self.inner_dim) - - self.transformer_blocks = nn.ModuleList( - [ - QwenImageTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - zero_cond_t=zero_cond_t, - ) - for _ in range(num_layers) - ] - ) - - self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) - - self.gradient_checkpointing = False - self.zero_cond_t = zero_cond_t - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - encoder_hidden_states_mask: torch.Tensor = None, - timestep: torch.LongTensor = None, - img_shapes: list[tuple[int, int, int]] | None = None, - guidance: torch.Tensor = None, # TODO: this should probably be removed - attention_kwargs: dict[str, Any] | None = None, - controlnet_block_samples=None, - additional_t_cond=None, - return_dict: bool = True, - ) -> torch.Tensor | Transformer2DModelOutput: - """ - The [`QwenTransformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, image_sequence_length, in_channels)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, text_sequence_length, joint_attention_dim)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_hidden_states_mask (`torch.Tensor` of shape `(batch_size, text_sequence_length)`, *optional*): - Mask for the encoder hidden states. Expected to have 1.0 for valid tokens and 0.0 for padding tokens. - Used in the attention processor to prevent attending to padding tokens. The mask can have any pattern - (not just contiguous valid tokens followed by padding) since it's applied element-wise in attention. - timestep ( `torch.LongTensor`): - Used to indicate denoising step. - img_shapes (`list[tuple[int, int, int]]`, *optional*): - Image shapes for RoPE computation. - guidance (`torch.Tensor`, *optional*): - Guidance tensor for conditional generation. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - controlnet_block_samples (*optional*): - ControlNet block samples to add to the transformer blocks. - additional_t_cond (`torch.Tensor`, *optional*): - Additional timestep conditioning added to the timestep embedding. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - hidden_states = self.img_in(hidden_states) - - timestep = timestep.to(hidden_states.dtype) - - if self.zero_cond_t: - timestep = torch.cat([timestep, timestep * 0], dim=0) - modulate_index = torch.tensor( - [[0] * prod(sample[0]) + [1] * sum([prod(s) for s in sample[1:]]) for sample in img_shapes], - device=timestep.device, - dtype=torch.int, - ) - else: - modulate_index = None - - encoder_hidden_states = self.txt_norm(encoder_hidden_states) - encoder_hidden_states = self.txt_in(encoder_hidden_states) - - # Use the encoder_hidden_states sequence length for RoPE computation and normalize mask - text_seq_len, _, encoder_hidden_states_mask = compute_text_seq_len_from_mask( - encoder_hidden_states, encoder_hidden_states_mask - ) - - if guidance is not None: - guidance = guidance.to(hidden_states.dtype) * 1000 - - temb = ( - self.time_text_embed(timestep, hidden_states, additional_t_cond) - if guidance is None - else self.time_text_embed(timestep, guidance, hidden_states, additional_t_cond) - ) - - image_rotary_emb = self.pos_embed(img_shapes, max_txt_seq_len=text_seq_len, device=hidden_states.device) - - for index_block, block in enumerate(self.transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - encoder_hidden_states_mask, - temb, - image_rotary_emb, - attention_kwargs, - modulate_index, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - encoder_hidden_states_mask=encoder_hidden_states_mask, - temb=temb, - image_rotary_emb=image_rotary_emb, - joint_attention_kwargs=attention_kwargs, - modulate_index=modulate_index, - ) - - # controlnet residual - if controlnet_block_samples is not None: - interval_control = len(self.transformer_blocks) / len(controlnet_block_samples) - interval_control = int(np.ceil(interval_control)) - hidden_states = hidden_states + controlnet_block_samples[index_block // interval_control] - - if self.zero_cond_t: - temb = temb.chunk(2, dim=0)[0] - # Use only the image part (hidden_states) from the dual-stream blocks - hidden_states = self.norm_out(hidden_states, temb) - output = self.proj_out(hidden_states) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_sana_video.py b/diffusers/models/transformers/transformer_sana_video.py deleted file mode 100644 index db1f08a73a81f892356356f4d85080485ecc3f60..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_sana_video.py +++ /dev/null @@ -1,717 +0,0 @@ -# Copyright 2025 The HuggingFace Team and SANA-Video Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import Any - -import torch -import torch.nn.functional as F -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ..attention import AttentionMixin -from ..attention_dispatch import dispatch_attention_fn -from ..attention_processor import Attention -from ..embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps, get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormSingle, RMSNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class GLUMBTempConv(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - expand_ratio: float = 4, - norm_type: str | None = None, - residual_connection: bool = True, - ) -> None: - super().__init__() - - hidden_channels = int(expand_ratio * in_channels) - self.norm_type = norm_type - self.residual_connection = residual_connection - - self.nonlinearity = nn.SiLU() - self.conv_inverted = nn.Conv2d(in_channels, hidden_channels * 2, 1, 1, 0) - self.conv_depth = nn.Conv2d(hidden_channels * 2, hidden_channels * 2, 3, 1, 1, groups=hidden_channels * 2) - self.conv_point = nn.Conv2d(hidden_channels, out_channels, 1, 1, 0, bias=False) - - self.norm = None - if norm_type == "rms_norm": - self.norm = RMSNorm(out_channels, eps=1e-5, elementwise_affine=True, bias=True) - - self.conv_temp = nn.Conv2d( - out_channels, out_channels, kernel_size=(3, 1), stride=1, padding=(1, 0), bias=False - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.residual_connection: - residual = hidden_states - batch_size, num_frames, height, width, num_channels = hidden_states.shape - hidden_states = hidden_states.view(batch_size * num_frames, height, width, num_channels).permute(0, 3, 1, 2) - - hidden_states = self.conv_inverted(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - - hidden_states = self.conv_depth(hidden_states) - hidden_states, gate = torch.chunk(hidden_states, 2, dim=1) - hidden_states = hidden_states * self.nonlinearity(gate) - - hidden_states = self.conv_point(hidden_states) - - # Temporal aggregation - hidden_states_temporal = hidden_states.view(batch_size, num_frames, num_channels, height * width).permute( - 0, 2, 1, 3 - ) - hidden_states = hidden_states_temporal + self.conv_temp(hidden_states_temporal) - hidden_states = hidden_states.permute(0, 2, 3, 1).view(batch_size, num_frames, height, width, num_channels) - - if self.norm_type == "rms_norm": - # move channel to the last dimension so we apply RMSnorm across channel dimension - hidden_states = self.norm(hidden_states.movedim(1, -1)).movedim(-1, 1) - - if self.residual_connection: - hidden_states = hidden_states + residual - - return hidden_states - - -class SanaLinearAttnProcessor3_0: - r""" - Processor for implementing scaled dot-product linear attention. - """ - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - original_dtype = hidden_states.dtype - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - # B,N,H,C - - query = F.relu(query) - key = F.relu(key) - - if rotary_emb is not None: - - def apply_rotary_emb( - hidden_states: torch.Tensor, - freqs_cos: torch.Tensor, - freqs_sin: torch.Tensor, - ): - x1, x2 = hidden_states.unflatten(-1, (-1, 2)).unbind(-1) - cos = freqs_cos[..., 0::2] - sin = freqs_sin[..., 1::2] - out = torch.empty_like(hidden_states) - out[..., 0::2] = x1 * cos - x2 * sin - out[..., 1::2] = x1 * sin + x2 * cos - return out.type_as(hidden_states) - - query_rotate = apply_rotary_emb(query, *rotary_emb) - key_rotate = apply_rotary_emb(key, *rotary_emb) - - # B,H,C,N - query = query.permute(0, 2, 3, 1) - key = key.permute(0, 2, 3, 1) - query_rotate = query_rotate.permute(0, 2, 3, 1) - key_rotate = key_rotate.permute(0, 2, 3, 1) - value = value.permute(0, 2, 3, 1) - - query_rotate, key_rotate, value = query_rotate.float(), key_rotate.float(), value.float() - - z = 1 / (key.sum(dim=-1, keepdim=True).transpose(-2, -1) @ query + 1e-15) - - scores = torch.matmul(value, key_rotate.transpose(-1, -2)) - hidden_states = torch.matmul(scores, query_rotate) - - hidden_states = hidden_states * z - # B,H,C,N - hidden_states = hidden_states.flatten(1, 2).transpose(1, 2) - hidden_states = hidden_states.to(original_dtype) - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - - return hidden_states - - -class WanRotaryPosEmbed(nn.Module): - def __init__( - self, - attention_head_dim: int, - patch_size: tuple[int, int, int], - max_seq_len: int, - theta: float = 10000.0, - ): - super().__init__() - - self.attention_head_dim = attention_head_dim - self.patch_size = patch_size - self.max_seq_len = max_seq_len - - h_dim = w_dim = 2 * (attention_head_dim // 6) - t_dim = attention_head_dim - h_dim - w_dim - - self.t_dim = t_dim - self.h_dim = h_dim - self.w_dim = w_dim - - freqs_dtype = torch.float32 if torch.backends.mps.is_available() else torch.float64 - - freqs_cos = [] - freqs_sin = [] - - for dim in [t_dim, h_dim, w_dim]: - freq_cos, freq_sin = get_1d_rotary_pos_embed( - dim, - max_seq_len, - theta, - use_real=True, - repeat_interleave_real=True, - freqs_dtype=freqs_dtype, - ) - freqs_cos.append(freq_cos) - freqs_sin.append(freq_sin) - - self.register_buffer("freqs_cos", torch.cat(freqs_cos, dim=1), persistent=False) - self.register_buffer("freqs_sin", torch.cat(freqs_sin, dim=1), persistent=False) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.patch_size - ppf, pph, ppw = num_frames // p_t, height // p_h, width // p_w - - split_sizes = [self.t_dim, self.h_dim, self.w_dim] - - freqs_cos = self.freqs_cos.split(split_sizes, dim=1) - freqs_sin = self.freqs_sin.split(split_sizes, dim=1) - - freqs_cos_f = freqs_cos[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_cos_h = freqs_cos[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_cos_w = freqs_cos[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - - freqs_sin_f = freqs_sin[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_sin_h = freqs_sin[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_sin_w = freqs_sin[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - - freqs_cos = torch.cat([freqs_cos_f, freqs_cos_h, freqs_cos_w], dim=-1).reshape(1, ppf * pph * ppw, 1, -1) - freqs_sin = torch.cat([freqs_sin_f, freqs_sin_h, freqs_sin_w], dim=-1).reshape(1, ppf * pph * ppw, 1, -1) - - return freqs_cos, freqs_sin - - -class SanaModulatedNorm(nn.Module): - def __init__(self, dim: int, elementwise_affine: bool = False, eps: float = 1e-6): - super().__init__() - self.norm = nn.LayerNorm(dim, elementwise_affine=elementwise_affine, eps=eps) - - def forward( - self, hidden_states: torch.Tensor, temb: torch.Tensor, scale_shift_table: torch.Tensor - ) -> torch.Tensor: - hidden_states = self.norm(hidden_states) - shift, scale = (scale_shift_table[None, None] + temb[:, :, None].to(scale_shift_table.device)).unbind(dim=2) - hidden_states = hidden_states * (1 + scale) + shift - return hidden_states - - -class SanaCombinedTimestepGuidanceEmbeddings(nn.Module): - def __init__(self, embedding_dim): - super().__init__() - self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - self.guidance_condition_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) - self.guidance_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - - self.silu = nn.SiLU() - self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=True) - - def forward(self, timestep: torch.Tensor, guidance: torch.Tensor = None, hidden_dtype: torch.dtype = None): - timesteps_proj = self.time_proj(timestep) - timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) - - guidance_proj = self.guidance_condition_proj(guidance) - guidance_emb = self.guidance_embedder(guidance_proj.to(dtype=hidden_dtype)) - conditioning = timesteps_emb + guidance_emb - - return self.linear(self.silu(conditioning)), conditioning - - -class SanaAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("SanaAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim) - key = key.view(batch_size, -1, attn.heads, head_dim) - value = value.view(batch_size, -1, attn.heads, head_dim) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class SanaVideoTransformerBlock(nn.Module): - r""" - Transformer block introduced in [Sana-Video](https://huggingface.co/papers/2509.24695). - """ - - def __init__( - self, - dim: int = 2240, - num_attention_heads: int = 20, - attention_head_dim: int = 112, - dropout: float = 0.0, - num_cross_attention_heads: int | None = 20, - cross_attention_head_dim: int | None = 112, - cross_attention_dim: int | None = 2240, - attention_bias: bool = True, - norm_elementwise_affine: bool = False, - norm_eps: float = 1e-6, - attention_out_bias: bool = True, - mlp_ratio: float = 3.0, - qk_norm: str | None = "rms_norm_across_heads", - rope_max_seq_len: int = 1024, - ) -> None: - super().__init__() - - # 1. Self Attention - self.norm1 = nn.LayerNorm(dim, elementwise_affine=False, eps=norm_eps) - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - kv_heads=num_attention_heads if qk_norm is not None else None, - qk_norm=qk_norm, - dropout=dropout, - bias=attention_bias, - cross_attention_dim=None, - processor=SanaLinearAttnProcessor3_0(), - ) - - # 2. Cross Attention - if cross_attention_dim is not None: - self.norm2 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps) - self.attn2 = Attention( - query_dim=dim, - qk_norm=qk_norm, - kv_heads=num_cross_attention_heads if qk_norm is not None else None, - cross_attention_dim=cross_attention_dim, - heads=num_cross_attention_heads, - dim_head=cross_attention_head_dim, - dropout=dropout, - bias=True, - out_bias=attention_out_bias, - processor=SanaAttnProcessor2_0(), - ) - - # 3. Feed-forward - self.ff = GLUMBTempConv(dim, dim, mlp_ratio, norm_type=None, residual_connection=False) - - self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - timestep: torch.LongTensor | None = None, - frames: int = None, - height: int = None, - width: int = None, - rotary_emb: torch.Tensor | None = None, - ) -> torch.Tensor: - batch_size = hidden_states.shape[0] - - # 1. Modulation - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( - self.scale_shift_table[None, None] + timestep.reshape(batch_size, timestep.shape[1], 6, -1) - ).unbind(dim=2) - - # 2. Self Attention - norm_hidden_states = self.norm1(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_msa) + shift_msa - norm_hidden_states = norm_hidden_states.to(hidden_states.dtype) - - attn_output = self.attn1(norm_hidden_states, rotary_emb=rotary_emb) - hidden_states = hidden_states + gate_msa * attn_output - - # 3. Cross Attention - if self.attn2 is not None: - attn_output = self.attn2( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=encoder_attention_mask, - ) - hidden_states = attn_output + hidden_states - - # 4. Feed-forward - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp - - norm_hidden_states = norm_hidden_states.unflatten(1, (frames, height, width)) - ff_output = self.ff(norm_hidden_states) - ff_output = ff_output.flatten(1, 3) - hidden_states = hidden_states + gate_mlp * ff_output - - return hidden_states - - -class SanaVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, AttentionMixin): - r""" - A 3D Transformer model introduced in [Sana-Video](https://huggingface.co/papers/2509.24695) family of models. - - Args: - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, *optional*, defaults to `16`): - The number of channels in the output. - num_attention_heads (`int`, defaults to `20`): - The number of heads to use for multi-head attention. - attention_head_dim (`int`, defaults to `112`): - The number of channels in each head. - num_layers (`int`, defaults to `20`): - The number of layers of Transformer blocks to use. - num_cross_attention_heads (`int`, *optional*, defaults to `20`): - The number of heads to use for cross-attention. - cross_attention_head_dim (`int`, *optional*, defaults to `112`): - The number of channels in each head for cross-attention. - cross_attention_dim (`int`, *optional*, defaults to `2240`): - The number of channels in the cross-attention output. - caption_channels (`int`, defaults to `2304`): - The number of channels in the caption embeddings. - mlp_ratio (`float`, defaults to `2.5`): - The expansion ratio to use in the GLUMBConv layer. - dropout (`float`, defaults to `0.0`): - The dropout probability. - attention_bias (`bool`, defaults to `False`): - Whether to use bias in the attention layer. - sample_size (`int`, defaults to `32`): - The base size of the input latent. - patch_size (`int`, defaults to `1`): - The size of the patches to use in the patch embedding layer. - norm_elementwise_affine (`bool`, defaults to `False`): - Whether to use elementwise affinity in the normalization layer. - norm_eps (`float`, defaults to `1e-6`): - The epsilon value for the normalization layer. - qk_norm (`str`, *optional*, defaults to `None`): - The normalization to use for the query and key. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["SanaVideoTransformerBlock", "SanaModulatedNorm"] - _skip_layerwise_casting_patterns = ["patch_embedding", "norm"] - - @register_to_config - def __init__( - self, - in_channels: int = 16, - out_channels: int | None = 16, - num_attention_heads: int = 20, - attention_head_dim: int = 112, - num_layers: int = 20, - num_cross_attention_heads: int | None = 20, - cross_attention_head_dim: int | None = 112, - cross_attention_dim: int | None = 2240, - caption_channels: int = 2304, - mlp_ratio: float = 2.5, - dropout: float = 0.0, - attention_bias: bool = False, - sample_size: int = 30, - patch_size: tuple[int, int, int] = (1, 2, 2), - norm_elementwise_affine: bool = False, - norm_eps: float = 1e-6, - interpolation_scale: int | None = None, - guidance_embeds: bool = False, - guidance_embeds_scale: float = 0.1, - qk_norm: str | None = "rms_norm_across_heads", - rope_max_seq_len: int = 1024, - ) -> None: - super().__init__() - - out_channels = out_channels or in_channels - inner_dim = num_attention_heads * attention_head_dim - - # 1. Patch & position embedding - self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len) - self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - - # 2. Additional condition embeddings - if guidance_embeds: - self.time_embed = SanaCombinedTimestepGuidanceEmbeddings(inner_dim) - else: - self.time_embed = AdaLayerNormSingle(inner_dim) - - self.caption_projection = PixArtAlphaTextProjection(in_features=caption_channels, hidden_size=inner_dim) - self.caption_norm = RMSNorm(inner_dim, eps=1e-5, elementwise_affine=True) - - # 3. Transformer blocks - self.transformer_blocks = nn.ModuleList( - [ - SanaVideoTransformerBlock( - inner_dim, - num_attention_heads, - attention_head_dim, - dropout=dropout, - num_cross_attention_heads=num_cross_attention_heads, - cross_attention_head_dim=cross_attention_head_dim, - cross_attention_dim=cross_attention_dim, - attention_bias=attention_bias, - norm_elementwise_affine=norm_elementwise_affine, - norm_eps=norm_eps, - mlp_ratio=mlp_ratio, - qk_norm=qk_norm, - ) - for _ in range(num_layers) - ] - ) - - # 4. Output blocks - self.scale_shift_table = nn.Parameter(torch.randn(2, inner_dim) / inner_dim**0.5) - self.norm_out = SanaModulatedNorm(inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(inner_dim, math.prod(patch_size) * out_channels) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - timestep: torch.Tensor, - guidance: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - attention_kwargs: dict[str, Any] | None = None, - controlnet_block_samples: tuple[torch.Tensor] | None = None, - return_dict: bool = True, - ) -> tuple[torch.Tensor, ...] | Transformer2DModelOutput: - """ - The [`SanaVideoTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, in_channels, num_frames, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - guidance (`torch.Tensor`, *optional*): - Guidance scale embedding. - encoder_attention_mask (`torch.Tensor`, *optional*): - Cross-attention mask applied to `encoder_hidden_states`. - attention_mask (`torch.Tensor`, *optional*): - Self-attention mask applied to `hidden_states`. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - controlnet_block_samples (`tuple` of `torch.Tensor`, *optional*): - A list of tensors that if specified are added to the residuals of transformer blocks. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - # ensure attention_mask is a bias, and give it a singleton query_tokens dimension. - # we may have done this conversion already, e.g. if we came here via UNet2DConditionModel#forward. - # we can tell by counting dims; if ndim == 2: it's a mask rather than a bias. - # expects mask of shape: - # [batch, key_tokens] - # adds singleton query_tokens dimension: - # [batch, 1, key_tokens] - # this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes: - # [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn) - # [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn) - if attention_mask is not None and attention_mask.ndim == 2: - # assume that mask is expressed as: - # (1 = keep, 0 = discard) - # convert mask into a bias that can be added to attention scores: - # (keep = +0, discard = -10000.0) - attention_mask = (1 - attention_mask.to(hidden_states.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # convert encoder_attention_mask to a bias the same way we do for attention_mask - if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2: - encoder_attention_mask = (1 - encoder_attention_mask.to(hidden_states.dtype)) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - # 1. Input - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.config.patch_size - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p_h - post_patch_width = width // p_w - - rotary_emb = self.rope(hidden_states) - - hidden_states = self.patch_embedding(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) - - if guidance is not None: - timestep, embedded_timestep = self.time_embed( - timestep.flatten(), guidance=guidance, hidden_dtype=hidden_states.dtype - ) - else: - timestep, embedded_timestep = self.time_embed( - timestep.flatten(), batch_size=batch_size, hidden_dtype=hidden_states.dtype - ) - - timestep = timestep.view(batch_size, -1, timestep.size(-1)) - embedded_timestep = embedded_timestep.view(batch_size, -1, embedded_timestep.size(-1)) - - encoder_hidden_states = self.caption_projection(encoder_hidden_states) - encoder_hidden_states = encoder_hidden_states.view(batch_size, -1, hidden_states.shape[-1]) - - encoder_hidden_states = self.caption_norm(encoder_hidden_states) - - # 2. Transformer blocks - if torch.is_grad_enabled() and self.gradient_checkpointing: - for index_block, block in enumerate(self.transformer_blocks): - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - attention_mask, - encoder_hidden_states, - encoder_attention_mask, - timestep, - post_patch_num_frames, - post_patch_height, - post_patch_width, - rotary_emb, - ) - if controlnet_block_samples is not None and 0 < index_block <= len(controlnet_block_samples): - hidden_states = hidden_states + controlnet_block_samples[index_block - 1] - - else: - for index_block, block in enumerate(self.transformer_blocks): - hidden_states = block( - hidden_states, - attention_mask, - encoder_hidden_states, - encoder_attention_mask, - timestep, - post_patch_num_frames, - post_patch_height, - post_patch_width, - rotary_emb, - ) - if controlnet_block_samples is not None and 0 < index_block <= len(controlnet_block_samples): - hidden_states = hidden_states + controlnet_block_samples[index_block - 1] - - # 3. Normalization - hidden_states = self.norm_out(hidden_states, embedded_timestep, self.scale_shift_table) - - hidden_states = self.proj_out(hidden_states) - - # 5. Unpatchify - hidden_states = hidden_states.reshape( - batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1 - ) - hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) - output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_sd3.py b/diffusers/models/transformers/transformer_sd3.py deleted file mode 100644 index 9a56ca4e226de34208eaac171d80b7b83d2eefef..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_sd3.py +++ /dev/null @@ -1,347 +0,0 @@ -# Copyright 2025 Stability AI, The HuggingFace Team and The InstantX Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin, SD3Transformer2DLoadersMixin -from ...utils import apply_lora_scale, logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import AttentionMixin, FeedForward, JointTransformerBlock -from ..attention_processor import ( - Attention, - FusedJointAttnProcessor2_0, - JointAttnProcessor2_0, -) -from ..embeddings import CombinedTimestepTextProjEmbeddings, PatchEmbed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import AdaLayerNormContinuous, AdaLayerNormZero - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@maybe_allow_in_graph -class SD3SingleTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - ): - super().__init__() - - self.norm1 = AdaLayerNormZero(dim) - self.attn = Attention( - query_dim=dim, - dim_head=attention_head_dim, - heads=num_attention_heads, - out_dim=dim, - bias=True, - processor=JointAttnProcessor2_0(), - eps=1e-6, - ) - - self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) - self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor): - # 1. Attention - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) - attn_output = self.attn(hidden_states=norm_hidden_states, encoder_hidden_states=None) - attn_output = gate_msa.unsqueeze(1) * attn_output - hidden_states = hidden_states + attn_output - - # 2. Feed Forward - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp.unsqueeze(1)) + shift_mlp.unsqueeze(1) - ff_output = self.ff(norm_hidden_states) - ff_output = gate_mlp.unsqueeze(1) * ff_output - hidden_states = hidden_states + ff_output - - return hidden_states - - -class SD3Transformer2DModel( - ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, SD3Transformer2DLoadersMixin -): - """ - The Transformer model introduced in [Stable Diffusion 3](https://huggingface.co/papers/2403.03206). - - Parameters: - sample_size (`int`, defaults to `128`): - The width/height of the latents. This is fixed during training since it is used to learn a number of - position embeddings. - patch_size (`int`, defaults to `2`): - Patch size to turn the input data into small patches. - in_channels (`int`, defaults to `16`): - The number of latent channels in the input. - num_layers (`int`, defaults to `18`): - The number of layers of transformer blocks to use. - attention_head_dim (`int`, defaults to `64`): - The number of channels in each head. - num_attention_heads (`int`, defaults to `18`): - The number of heads to use for multi-head attention. - joint_attention_dim (`int`, defaults to `4096`): - The embedding dimension to use for joint text-image attention. - caption_projection_dim (`int`, defaults to `1152`): - The embedding dimension of caption embeddings. - pooled_projection_dim (`int`, defaults to `2048`): - The embedding dimension of pooled text projections. - out_channels (`int`, defaults to `16`): - The number of latent channels in the output. - pos_embed_max_size (`int`, defaults to `96`): - The maximum latent height/width of positional embeddings. - dual_attention_layers (`tuple[int, ...]`, defaults to `()`): - The number of dual-stream transformer blocks to use. - qk_norm (`str`, *optional*, defaults to `None`): - The normalization to use for query and key in the attention layer. If `None`, no normalization is used. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["JointTransformerBlock"] - _skip_layerwise_casting_patterns = ["pos_embed", "norm"] - - @register_to_config - def __init__( - self, - sample_size: int = 128, - patch_size: int = 2, - in_channels: int = 16, - num_layers: int = 18, - attention_head_dim: int = 64, - num_attention_heads: int = 18, - joint_attention_dim: int = 4096, - caption_projection_dim: int = 1152, - pooled_projection_dim: int = 2048, - out_channels: int = 16, - pos_embed_max_size: int = 96, - dual_attention_layers: tuple[ - int, ... - ] = (), # () for sd3.0; (0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12) for sd3.5 - qk_norm: str | None = None, - ): - super().__init__() - self.out_channels = out_channels if out_channels is not None else in_channels - self.inner_dim = num_attention_heads * attention_head_dim - - self.pos_embed = PatchEmbed( - height=sample_size, - width=sample_size, - patch_size=patch_size, - in_channels=in_channels, - embed_dim=self.inner_dim, - pos_embed_max_size=pos_embed_max_size, # hard-code for now. - ) - self.time_text_embed = CombinedTimestepTextProjEmbeddings( - embedding_dim=self.inner_dim, pooled_projection_dim=pooled_projection_dim - ) - self.context_embedder = nn.Linear(joint_attention_dim, caption_projection_dim) - - self.transformer_blocks = nn.ModuleList( - [ - JointTransformerBlock( - dim=self.inner_dim, - num_attention_heads=num_attention_heads, - attention_head_dim=attention_head_dim, - context_pre_only=i == num_layers - 1, - qk_norm=qk_norm, - use_dual_attention=True if i in dual_attention_layers else False, - ) - for i in range(num_layers) - ] - ) - - self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) - self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) - - self.gradient_checkpointing = False - - # Copied from diffusers.models.unets.unet_3d_condition.UNet3DConditionModel.enable_forward_chunking - def enable_forward_chunking(self, chunk_size: int | None = None, dim: int = 0) -> None: - """ - Sets the attention processor to use [feed forward - chunking](https://huggingface.co/blog/reformer#2-chunked-feed-forward-layers). - - Parameters: - chunk_size (`int`, *optional*): - The chunk size of the feed-forward layers. If not specified, will run feed-forward layer individually - over each tensor of dim=`dim`. - dim (`int`, *optional*, defaults to `0`): - The dimension over which the feed-forward computation should be chunked. Choose between dim=0 (batch) - or dim=1 (sequence length). - """ - if dim not in [0, 1]: - raise ValueError(f"Make sure to set `dim` to either 0 or 1, not {dim}") - - # By default chunk size is 1 - chunk_size = chunk_size or 1 - - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, chunk_size, dim) - - # Copied from diffusers.models.unets.unet_3d_condition.UNet3DConditionModel.disable_forward_chunking - def disable_forward_chunking(self): - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, None, 0) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections with FusedAttnProcessor2_0->FusedJointAttnProcessor2_0 - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedJointAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) - - @apply_lora_scale("joint_attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor = None, - pooled_projections: torch.Tensor = None, - timestep: torch.LongTensor = None, - block_controlnet_hidden_states: list = None, - joint_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - skip_layers: list[int] | None = None, - ) -> torch.Tensor | Transformer2DModelOutput: - """ - The [`SD3Transformer2DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch size, channel, height, width)`): - Input `hidden_states`. - encoder_hidden_states (`torch.Tensor` of shape `(batch size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - pooled_projections (`torch.Tensor` of shape `(batch_size, projection_dim)`): - Embeddings projected from the embeddings of input conditions. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - block_controlnet_hidden_states (`list` of `torch.Tensor`): - A list of tensors that if specified are added to the residuals of transformer blocks. - joint_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - skip_layers (`list` of `int`, *optional*): - A list of layer indices to skip during the forward pass. - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - - height, width = hidden_states.shape[-2:] - - hidden_states = self.pos_embed(hidden_states) # takes care of adding positional embeddings too. - # pos_embed output is non-contiguous due to flatten+transpose in PatchEmbed (BCHW -> BNC). - hidden_states = hidden_states.contiguous() - temb = self.time_text_embed(timestep, pooled_projections) - encoder_hidden_states = self.context_embedder(encoder_hidden_states) - - if joint_attention_kwargs is not None and "ip_adapter_image_embeds" in joint_attention_kwargs: - ip_adapter_image_embeds = joint_attention_kwargs.pop("ip_adapter_image_embeds") - ip_hidden_states, ip_temb = self.image_proj(ip_adapter_image_embeds, timestep) - - joint_attention_kwargs.update(ip_hidden_states=ip_hidden_states, temb=ip_temb) - - for index_block, block in enumerate(self.transformer_blocks): - # Skip specified layers - is_skip = True if skip_layers is not None and index_block in skip_layers else False - - if torch.is_grad_enabled() and self.gradient_checkpointing and not is_skip: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - temb, - joint_attention_kwargs, - ) - elif not is_skip: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - joint_attention_kwargs=joint_attention_kwargs, - ) - - # controlnet residual - if block_controlnet_hidden_states is not None and block.context_pre_only is False: - interval_control = len(self.transformer_blocks) / len(block_controlnet_hidden_states) - hidden_states = hidden_states + block_controlnet_hidden_states[int(index_block / interval_control)] - - hidden_states = self.norm_out(hidden_states, temb) - hidden_states = self.proj_out(hidden_states) - - # unpatchify - patch_size = self.config.patch_size - height = height // patch_size - width = width // patch_size - - hidden_states = hidden_states.reshape( - shape=(hidden_states.shape[0], height, width, patch_size, patch_size, self.out_channels) - ) - hidden_states = torch.einsum("nhwpqc->nchpwq", hidden_states) - output = hidden_states.reshape( - shape=(hidden_states.shape[0], self.out_channels, height * patch_size, width * patch_size) - ) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_skyreels_v2.py b/diffusers/models/transformers/transformer_skyreels_v2.py deleted file mode 100644 index 81caf6cb71417d6807e499b91709a2ac08fc5b4a..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_skyreels_v2.py +++ /dev/null @@ -1,794 +0,0 @@ -# Copyright 2025 The SkyReels Team, The Wan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, deprecate, logging -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import ( - PixArtAlphaTextProjection, - TimestepEmbedding, - get_1d_rotary_pos_embed, - get_1d_sincos_pos_embed_from_grid, -) -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin, get_parameter_dtype -from ..normalization import FP32LayerNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _get_qkv_projections( - attn: "SkyReelsV2Attention", hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor -): - # encoder_hidden_states is only passed for cross-attention - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - if attn.fused_projections: - if attn.cross_attention_dim_head is None: - # In self-attention layers, we can fuse the entire QKV projection into a single linear - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - else: - # In cross-attention layers, we can only fuse the KV projections into a single linear - query = attn.to_q(hidden_states) - key, value = attn.to_kv(encoder_hidden_states).chunk(2, dim=-1) - else: - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - return query, key, value - - -def _get_added_kv_projections(attn: "SkyReelsV2Attention", encoder_hidden_states_img: torch.Tensor): - if attn.fused_projections: - key_img, value_img = attn.to_added_kv(encoder_hidden_states_img).chunk(2, dim=-1) - else: - key_img = attn.add_k_proj(encoder_hidden_states_img) - value_img = attn.add_v_proj(encoder_hidden_states_img) - return key_img, value_img - - -class SkyReelsV2AttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "SkyReelsV2AttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0." - ) - - def __call__( - self, - attn: "SkyReelsV2Attention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - encoder_hidden_states_img = None - if attn.add_k_proj is not None: - # 512 is the context length of the text encoder, hardcoded for now - image_context_length = encoder_hidden_states.shape[1] - 512 - encoder_hidden_states_img = encoder_hidden_states[:, :image_context_length] - encoder_hidden_states = encoder_hidden_states[:, image_context_length:] - - query, key, value = _get_qkv_projections(attn, hidden_states, encoder_hidden_states) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - if rotary_emb is not None: - - def apply_rotary_emb( - hidden_states: torch.Tensor, - freqs_cos: torch.Tensor, - freqs_sin: torch.Tensor, - ): - x1, x2 = hidden_states.unflatten(-1, (-1, 2)).unbind(-1) - cos = freqs_cos[..., 0::2] - sin = freqs_sin[..., 1::2] - out = torch.empty_like(hidden_states) - out[..., 0::2] = x1 * cos - x2 * sin - out[..., 1::2] = x1 * sin + x2 * cos - return out.type_as(hidden_states) - - query = apply_rotary_emb(query, *rotary_emb) - key = apply_rotary_emb(key, *rotary_emb) - - # I2V task - hidden_states_img = None - if encoder_hidden_states_img is not None: - key_img, value_img = _get_added_kv_projections(attn, encoder_hidden_states_img) - key_img = attn.norm_added_k(key_img) - - key_img = key_img.unflatten(2, (attn.heads, -1)) - value_img = value_img.unflatten(2, (attn.heads, -1)) - - hidden_states_img = dispatch_attention_fn( - query, - key_img, - value_img, - attn_mask=None, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - hidden_states_img = hidden_states_img.flatten(2, 3) - hidden_states_img = hidden_states_img.type_as(query) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - - if hidden_states_img is not None: - hidden_states = hidden_states + hidden_states_img - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class SkyReelsV2AttnProcessor2_0: - def __new__(cls, *args, **kwargs): - deprecation_message = ( - "The SkyReelsV2AttnProcessor2_0 class is deprecated and will be removed in a future version. " - "Please use SkyReelsV2AttnProcessor instead. " - ) - deprecate("SkyReelsV2AttnProcessor2_0", "1.0.0", deprecation_message, standard_warn=False) - return SkyReelsV2AttnProcessor(*args, **kwargs) - - -class SkyReelsV2Attention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = SkyReelsV2AttnProcessor - _available_processors = [SkyReelsV2AttnProcessor] - - def __init__( - self, - dim: int, - heads: int = 8, - dim_head: int = 64, - eps: float = 1e-5, - dropout: float = 0.0, - added_kv_proj_dim: int | None = None, - cross_attention_dim_head: int | None = None, - processor=None, - is_cross_attention=None, - ): - super().__init__() - - self.inner_dim = dim_head * heads - self.heads = heads - self.added_kv_proj_dim = added_kv_proj_dim - self.cross_attention_dim_head = cross_attention_dim_head - self.kv_inner_dim = self.inner_dim if cross_attention_dim_head is None else cross_attention_dim_head * heads - - self.to_q = torch.nn.Linear(dim, self.inner_dim, bias=True) - self.to_k = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) - self.to_v = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) - self.to_out = torch.nn.ModuleList( - [ - torch.nn.Linear(self.inner_dim, dim, bias=True), - torch.nn.Dropout(dropout), - ] - ) - self.norm_q = torch.nn.RMSNorm(dim_head * heads, eps=eps, elementwise_affine=True) - self.norm_k = torch.nn.RMSNorm(dim_head * heads, eps=eps, elementwise_affine=True) - - self.add_k_proj = self.add_v_proj = None - if added_kv_proj_dim is not None: - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=True) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=True) - self.norm_added_k = torch.nn.RMSNorm(dim_head * heads, eps=eps) - - self.is_cross_attention = cross_attention_dim_head is not None - - self.set_processor(processor) - - def fuse_projections(self): - if getattr(self, "fused_projections", False): - return - - if self.cross_attention_dim_head is None: - concatenated_weights = torch.cat([self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data]) - concatenated_bias = torch.cat([self.to_q.bias.data, self.to_k.bias.data, self.to_v.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_qkv = nn.Linear(in_features, out_features, bias=True) - self.to_qkv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - else: - concatenated_weights = torch.cat([self.to_k.weight.data, self.to_v.weight.data]) - concatenated_bias = torch.cat([self.to_k.bias.data, self.to_v.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_kv = nn.Linear(in_features, out_features, bias=True) - self.to_kv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - - if self.added_kv_proj_dim is not None: - concatenated_weights = torch.cat([self.add_k_proj.weight.data, self.add_v_proj.weight.data]) - concatenated_bias = torch.cat([self.add_k_proj.bias.data, self.add_v_proj.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_added_kv = nn.Linear(in_features, out_features, bias=True) - self.to_added_kv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - - self.fused_projections = True - - @torch.no_grad() - def unfuse_projections(self): - if not getattr(self, "fused_projections", False): - return - - if hasattr(self, "to_qkv"): - delattr(self, "to_qkv") - if hasattr(self, "to_kv"): - delattr(self, "to_kv") - if hasattr(self, "to_added_kv"): - delattr(self, "to_added_kv") - - self.fused_projections = False - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - **kwargs, - ) -> torch.Tensor: - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, rotary_emb, **kwargs) - - -class SkyReelsV2ImageEmbedding(torch.nn.Module): - def __init__(self, in_features: int, out_features: int, pos_embed_seq_len=None): - super().__init__() - - self.norm1 = FP32LayerNorm(in_features) - self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu") - self.norm2 = FP32LayerNorm(out_features) - if pos_embed_seq_len is not None: - self.pos_embed = nn.Parameter(torch.zeros(1, pos_embed_seq_len, in_features)) - else: - self.pos_embed = None - - def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor: - if self.pos_embed is not None: - batch_size, seq_len, embed_dim = encoder_hidden_states_image.shape - encoder_hidden_states_image = encoder_hidden_states_image.view(-1, 2 * seq_len, embed_dim) - encoder_hidden_states_image = encoder_hidden_states_image + self.pos_embed - - hidden_states = self.norm1(encoder_hidden_states_image) - hidden_states = self.ff(hidden_states) - hidden_states = self.norm2(hidden_states) - return hidden_states - - -class SkyReelsV2Timesteps(nn.Module): - def __init__(self, num_channels: int, flip_sin_to_cos: bool, output_type: str = "pt"): - super().__init__() - self.num_channels = num_channels - self.output_type = output_type - self.flip_sin_to_cos = flip_sin_to_cos - - def forward(self, timesteps: torch.Tensor) -> torch.Tensor: - original_shape = timesteps.shape - t_emb = get_1d_sincos_pos_embed_from_grid( - self.num_channels, - timesteps, - output_type=self.output_type, - flip_sin_to_cos=self.flip_sin_to_cos, - ) - # Reshape back to maintain batch structure - if len(original_shape) > 1: - t_emb = t_emb.reshape(*original_shape, self.num_channels) - return t_emb - - -class SkyReelsV2TimeTextImageEmbedding(nn.Module): - def __init__( - self, - dim: int, - time_freq_dim: int, - time_proj_dim: int, - text_embed_dim: int, - image_embed_dim: int | None = None, - pos_embed_seq_len: int | None = None, - ): - super().__init__() - - self.timesteps_proj = SkyReelsV2Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True) - self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim) - self.act_fn = nn.SiLU() - self.time_proj = nn.Linear(dim, time_proj_dim) - self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh") - - self.image_embedder = None - if image_embed_dim is not None: - self.image_embedder = SkyReelsV2ImageEmbedding(image_embed_dim, dim, pos_embed_seq_len=pos_embed_seq_len) - - def forward( - self, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: torch.Tensor | None = None, - ): - timestep = self.timesteps_proj(timestep) - - time_embedder_dtype = get_parameter_dtype(self.time_embedder) - if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8: - timestep = timestep.to(time_embedder_dtype) - temb = self.time_embedder(timestep).type_as(encoder_hidden_states) - timestep_proj = self.time_proj(self.act_fn(temb)) - - encoder_hidden_states = self.text_embedder(encoder_hidden_states) - if encoder_hidden_states_image is not None: - encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image) - - return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image - - -class SkyReelsV2RotaryPosEmbed(nn.Module): - def __init__( - self, - attention_head_dim: int, - patch_size: tuple[int, int, int], - max_seq_len: int, - theta: float = 10000.0, - ): - super().__init__() - - self.attention_head_dim = attention_head_dim - self.patch_size = patch_size - self.max_seq_len = max_seq_len - - h_dim = w_dim = 2 * (attention_head_dim // 6) - t_dim = attention_head_dim - h_dim - w_dim - freqs_dtype = torch.float32 if torch.backends.mps.is_available() else torch.float64 - - self.t_dim = t_dim - self.h_dim = h_dim - self.w_dim = w_dim - - freqs_cos = [] - freqs_sin = [] - - for dim in [t_dim, h_dim, w_dim]: - freq_cos, freq_sin = get_1d_rotary_pos_embed( - dim, - max_seq_len, - theta, - use_real=True, - repeat_interleave_real=True, - freqs_dtype=freqs_dtype, - ) - freqs_cos.append(freq_cos) - freqs_sin.append(freq_sin) - - self.register_buffer("freqs_cos", torch.cat(freqs_cos, dim=1), persistent=False) - self.register_buffer("freqs_sin", torch.cat(freqs_sin, dim=1), persistent=False) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.patch_size - ppf, pph, ppw = num_frames // p_t, height // p_h, width // p_w - - split_sizes = [self.t_dim, self.h_dim, self.w_dim] - - freqs_cos = self.freqs_cos.split(split_sizes, dim=1) - freqs_sin = self.freqs_sin.split(split_sizes, dim=1) - - freqs_cos_f = freqs_cos[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_cos_h = freqs_cos[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_cos_w = freqs_cos[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - - freqs_sin_f = freqs_sin[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_sin_h = freqs_sin[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_sin_w = freqs_sin[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - - freqs_cos = torch.cat([freqs_cos_f, freqs_cos_h, freqs_cos_w], dim=-1).reshape(1, ppf * pph * ppw, 1, -1) - freqs_sin = torch.cat([freqs_sin_f, freqs_sin_h, freqs_sin_w], dim=-1).reshape(1, ppf * pph * ppw, 1, -1) - - return freqs_cos, freqs_sin - - -@maybe_allow_in_graph -class SkyReelsV2TransformerBlock(nn.Module): - def __init__( - self, - dim: int, - ffn_dim: int, - num_heads: int, - qk_norm: str = "rms_norm_across_heads", - cross_attn_norm: bool = False, - eps: float = 1e-6, - added_kv_proj_dim: int | None = None, - ): - super().__init__() - - # 1. Self-attention - self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False) - self.attn1 = SkyReelsV2Attention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - cross_attention_dim_head=None, - processor=SkyReelsV2AttnProcessor(), - ) - - # 2. Cross-attention - self.attn2 = SkyReelsV2Attention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - added_kv_proj_dim=added_kv_proj_dim, - cross_attention_dim_head=dim // num_heads, - processor=SkyReelsV2AttnProcessor(), - ) - self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity() - - # 3. Feed-forward - self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate") - self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False) - - self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - rotary_emb: torch.Tensor, - attention_mask: torch.Tensor, - ) -> torch.Tensor: - if temb.dim() == 3: - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table + temb.float() - ).chunk(6, dim=1) - elif temb.dim() == 4: - # For 4D temb in Diffusion Forcing framework, we assume the shape is (b, 6, f * pp_h * pp_w, inner_dim) - e = (self.scale_shift_table.unsqueeze(2) + temb.float()).chunk(6, dim=1) - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = [ei.squeeze(1) for ei in e] - - # 1. Self-attention - norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states) - attn_output = self.attn1(norm_hidden_states, None, attention_mask, rotary_emb) - hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states) - - # 2. Cross-attention - norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states) - attn_output = self.attn2(norm_hidden_states, encoder_hidden_states, None, None) - hidden_states = hidden_states + attn_output - - # 3. Feed-forward - norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as( - hidden_states - ) - ff_output = self.ffn(norm_hidden_states) - hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states) - - return hidden_states - - -class SkyReelsV2Transformer3DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin, AttentionMixin -): - r""" - A Transformer model for video-like data used in the Wan-based SkyReels-V2 model. - - Args: - patch_size (`tuple[int]`, defaults to `(1, 2, 2)`): - 3D patch dimensions for video embedding (t_patch, h_patch, w_patch). - num_attention_heads (`int`, defaults to `16`): - Fixed length for text embeddings. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - text_dim (`int`, defaults to `4096`): - Input dimension for text embeddings. - freq_dim (`int`, defaults to `256`): - Dimension for sinusoidal time embeddings. - ffn_dim (`int`, defaults to `8192`): - Intermediate dimension in feed-forward network. - num_layers (`int`, defaults to `32`): - The number of layers of transformer blocks to use. - window_size (`tuple[int]`, defaults to `(-1, -1)`): - Window size for local attention (-1 indicates global attention). - cross_attn_norm (`bool`, defaults to `True`): - Enable cross-attention normalization. - qk_norm (`str`, *optional*, defaults to `"rms_norm_across_heads"`): - Enable query/key normalization. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - inject_sample_info (`bool`, defaults to `False`): - Whether to inject sample information into the model. - image_dim (`int`, *optional*): - The dimension of the image embeddings. - added_kv_proj_dim (`int`, *optional*): - The dimension of the added key/value projection. - rope_max_seq_len (`int`, defaults to `1024`): - The maximum sequence length for the rotary embeddings. - pos_embed_seq_len (`int`, *optional*): - The sequence length for the positional embeddings. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"] - _no_split_modules = ["SkyReelsV2TransformerBlock"] - _keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"] - _keys_to_ignore_on_load_unexpected = ["norm_added_q"] - _repeated_blocks = ["SkyReelsV2TransformerBlock"] - - @register_to_config - def __init__( - self, - patch_size: tuple[int] = (1, 2, 2), - num_attention_heads: int = 16, - attention_head_dim: int = 128, - in_channels: int = 16, - out_channels: int = 16, - text_dim: int = 4096, - freq_dim: int = 256, - ffn_dim: int = 8192, - num_layers: int = 32, - cross_attn_norm: bool = True, - qk_norm: str | None = "rms_norm_across_heads", - eps: float = 1e-6, - image_dim: int | None = None, - added_kv_proj_dim: int | None = None, - rope_max_seq_len: int = 1024, - pos_embed_seq_len: int | None = None, - inject_sample_info: bool = False, - num_frame_per_block: int = 1, - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - # 1. Patch & position embedding - self.rope = SkyReelsV2RotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len) - self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - - # 2. Condition embeddings - # image_embedding_dim=1280 for I2V model - self.condition_embedder = SkyReelsV2TimeTextImageEmbedding( - dim=inner_dim, - time_freq_dim=freq_dim, - time_proj_dim=inner_dim * 6, - text_embed_dim=text_dim, - image_embed_dim=image_dim, - pos_embed_seq_len=pos_embed_seq_len, - ) - - # 3. Transformer blocks - self.blocks = nn.ModuleList( - [ - SkyReelsV2TransformerBlock( - inner_dim, ffn_dim, num_attention_heads, qk_norm, cross_attn_norm, eps, added_kv_proj_dim - ) - for _ in range(num_layers) - ] - ) - - # 4. Output norm & projection - self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size)) - self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5) - - if inject_sample_info: - self.fps_embedding = nn.Embedding(2, inner_dim) - self.fps_projection = FeedForward(inner_dim, inner_dim * 6, mult=1, activation_fn="linear-silu") - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: torch.Tensor | None = None, - enable_diffusion_forcing: bool = False, - fps: torch.Tensor | None = None, - return_dict: bool = True, - attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor | dict[str, torch.Tensor]: - """ - The [`SkyReelsV2Transformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_hidden_states_image (`torch.Tensor`, *optional*): - Conditional image embeddings for image-conditioned generation. - enable_diffusion_forcing (`bool`, *optional*, defaults to `False`): - Whether to enable diffusion forcing (per-block causal masking). - fps (`torch.Tensor`, *optional*): - FPS conditioning embedding. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.config.patch_size - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p_h - post_patch_width = width // p_w - - rotary_emb = self.rope(hidden_states) - - hidden_states = self.patch_embedding(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) - - causal_mask = None - if self.config.num_frame_per_block > 1: - block_num = post_patch_num_frames // self.config.num_frame_per_block - range_tensor = torch.arange(block_num, device=hidden_states.device).repeat_interleave( - self.config.num_frame_per_block - ) - causal_mask = range_tensor.unsqueeze(0) <= range_tensor.unsqueeze(1) # f, f - causal_mask = causal_mask.view(post_patch_num_frames, 1, 1, post_patch_num_frames, 1, 1) - causal_mask = causal_mask.repeat( - 1, post_patch_height, post_patch_width, 1, post_patch_height, post_patch_width - ) - causal_mask = causal_mask.reshape( - post_patch_num_frames * post_patch_height * post_patch_width, - post_patch_num_frames * post_patch_height * post_patch_width, - ) - causal_mask = causal_mask.unsqueeze(0).unsqueeze(0) - - temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder( - timestep, encoder_hidden_states, encoder_hidden_states_image - ) - - timestep_proj = timestep_proj.unflatten(-1, (6, -1)) - - if encoder_hidden_states_image is not None: - encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) - - if self.config.inject_sample_info: - fps = torch.tensor(fps, dtype=torch.long, device=hidden_states.device) - - fps_emb = self.fps_embedding(fps) - if enable_diffusion_forcing: - timestep_proj = timestep_proj + self.fps_projection(fps_emb).unflatten(1, (6, -1)).repeat( - timestep.shape[1], 1, 1 - ) - else: - timestep_proj = timestep_proj + self.fps_projection(fps_emb).unflatten(1, (6, -1)) - - if enable_diffusion_forcing: - b, f = timestep.shape - temb = temb.view(b, f, 1, 1, -1) - timestep_proj = timestep_proj.view(b, f, 1, 1, 6, -1) # (b, f, 1, 1, 6, inner_dim) - temb = temb.repeat(1, 1, post_patch_height, post_patch_width, 1).flatten(1, 3) - timestep_proj = timestep_proj.repeat(1, 1, post_patch_height, post_patch_width, 1, 1).flatten( - 1, 3 - ) # (b, f, pp_h, pp_w, 6, inner_dim) -> (b, f * pp_h * pp_w, 6, inner_dim) - timestep_proj = timestep_proj.transpose(1, 2).contiguous() # (b, 6, f * pp_h * pp_w, inner_dim) - - # 4. Transformer blocks - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.blocks: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - timestep_proj, - rotary_emb, - causal_mask, - ) - else: - for block in self.blocks: - hidden_states = block( - hidden_states, - encoder_hidden_states, - timestep_proj, - rotary_emb, - causal_mask, - ) - - if temb.dim() == 2: - # If temb is 2D, we assume it has time 1-D time embedding values for each batch. - # For models: - # - Skywork/SkyReels-V2-T2V-14B-540P-Diffusers - # - Skywork/SkyReels-V2-T2V-14B-720P-Diffusers - # - Skywork/SkyReels-V2-I2V-1.3B-540P-Diffusers - # - Skywork/SkyReels-V2-I2V-14B-540P-Diffusers - # - Skywork/SkyReels-V2-I2V-14B-720P-Diffusers - shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1) - elif temb.dim() == 3: - # If temb is 3D, we assume it has 2-D time embedding values for each batch. - # Each time embedding tensor includes values for each latent frame; thus Diffusion Forcing. - # For models: - # - Skywork/SkyReels-V2-DF-1.3B-540P-Diffusers - # - Skywork/SkyReels-V2-DF-14B-540P-Diffusers - # - Skywork/SkyReels-V2-DF-14B-720P-Diffusers - shift, scale = (self.scale_shift_table.unsqueeze(2) + temb.unsqueeze(1)).chunk(2, dim=1) - shift, scale = shift.squeeze(1), scale.squeeze(1) - - # Move the shift and scale tensors to the same device as hidden_states. - # When using multi-GPU inference via accelerate these will be on the - # first device rather than the last device, which hidden_states ends up - # on. - shift = shift.to(hidden_states.device) - scale = scale.to(hidden_states.device) - - hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states) - - hidden_states = self.proj_out(hidden_states) - - hidden_states = hidden_states.reshape( - batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1 - ) - hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) - output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) - - def _set_ar_attention(self, causal_block_size: int): - self.register_to_config(num_frame_per_block=causal_block_size) diff --git a/diffusers/models/transformers/transformer_temporal.py b/diffusers/models/transformers/transformer_temporal.py deleted file mode 100644 index 1cc42aa98ce4fa0383ac71251c0ae7b8c9414029..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_temporal.py +++ /dev/null @@ -1,375 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from dataclasses import dataclass -from typing import Any - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput -from ..attention import BasicTransformerBlock, TemporalBasicTransformerBlock -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from ..resnet import AlphaBlender - - -@dataclass -class TransformerTemporalModelOutput(BaseOutput): - """ - The output of [`TransformerTemporalModel`]. - - Args: - sample (`torch.Tensor` of shape `(batch_size x num_frames, num_channels, height, width)`): - The hidden states output conditioned on `encoder_hidden_states` input. - """ - - sample: torch.Tensor - - -class TransformerTemporalModel(ModelMixin, ConfigMixin): - """ - A Transformer model for video-like data. - - Parameters: - num_attention_heads (`int`, *optional*, defaults to 16): The number of heads to use for multi-head attention. - attention_head_dim (`int`, *optional*, defaults to 88): The number of channels in each head. - in_channels (`int`, *optional*): - The number of channels in the input and output (specify if the input is **continuous**). - num_layers (`int`, *optional*, defaults to 1): The number of layers of Transformer blocks to use. - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - cross_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use. - attention_bias (`bool`, *optional*): - Configure if the `TransformerBlock` attention should contain a bias parameter. - sample_size (`int`, *optional*): The width of the latent images (specify if the input is **discrete**). - This is fixed during training since it is used to learn a number of position embeddings. - activation_fn (`str`, *optional*, defaults to `"geglu"`): - Activation function to use in feed-forward. See `diffusers.models.activations.get_activation` for supported - activation functions. - norm_elementwise_affine (`bool`, *optional*): - Configure if the `TransformerBlock` should use learnable elementwise affine parameters for normalization. - double_self_attention (`bool`, *optional*): - Configure if each `TransformerBlock` should contain two self-attention layers. - positional_embeddings: (`str`, *optional*): - The type of positional embeddings to apply to the sequence input before passing use. - num_positional_embeddings: (`int`, *optional*): - The maximum length of the sequence over which to apply positional embeddings. - """ - - _skip_layerwise_casting_patterns = ["norm"] - - @register_to_config - def __init__( - self, - num_attention_heads: int = 16, - attention_head_dim: int = 88, - in_channels: int | None = None, - out_channels: int | None = None, - num_layers: int = 1, - dropout: float = 0.0, - norm_num_groups: int = 32, - cross_attention_dim: int | None = None, - attention_bias: bool = False, - sample_size: int | None = None, - activation_fn: str = "geglu", - norm_elementwise_affine: bool = True, - double_self_attention: bool = True, - positional_embeddings: str | None = None, - num_positional_embeddings: int | None = None, - ): - super().__init__() - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - inner_dim = num_attention_heads * attention_head_dim - - self.in_channels = in_channels - - self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True) - self.proj_in = nn.Linear(in_channels, inner_dim) - - # 3. Define transformers blocks - self.transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlock( - inner_dim, - num_attention_heads, - attention_head_dim, - dropout=dropout, - cross_attention_dim=cross_attention_dim, - activation_fn=activation_fn, - attention_bias=attention_bias, - double_self_attention=double_self_attention, - norm_elementwise_affine=norm_elementwise_affine, - positional_embeddings=positional_embeddings, - num_positional_embeddings=num_positional_embeddings, - ) - for d in range(num_layers) - ] - ) - - self.proj_out = nn.Linear(inner_dim, in_channels) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.LongTensor | None = None, - timestep: torch.LongTensor | None = None, - class_labels: torch.LongTensor = None, - num_frames: int = 1, - cross_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> TransformerTemporalModelOutput: - """ - The [`TransformerTemporal`] forward method. - - Args: - hidden_states (`torch.LongTensor` of shape `(batch size, num latent pixels)` if discrete, `torch.Tensor` of shape `(batch size, channel, height, width)` if continuous): - Input hidden_states. - encoder_hidden_states ( `torch.LongTensor` of shape `(batch size, encoder_hidden_states dim)`, *optional*): - Conditional embeddings for cross attention layer. If not given, cross-attention defaults to - self-attention. - timestep ( `torch.LongTensor`, *optional*): - Used to indicate denoising step. Optional timestep to be applied as an embedding in `AdaLayerNorm`. - class_labels ( `torch.LongTensor` of shape `(batch size, num classes)`, *optional*): - Used to indicate class labels conditioning. Optional class labels to be applied as an embedding in - `AdaLayerZeroNorm`. - num_frames (`int`, *optional*, defaults to 1): - The number of frames to be processed per batch. This is used to reshape the hidden states. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`] - instead of a plain tuple. - - Returns: - [`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`] or `tuple`: - If `return_dict` is True, an - [`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - # 1. Input - batch_frames, channel, height, width = hidden_states.shape - batch_size = batch_frames // num_frames - - residual = hidden_states - - hidden_states = hidden_states[None, :].reshape(batch_size, num_frames, channel, height, width) - hidden_states = hidden_states.permute(0, 2, 1, 3, 4) - - hidden_states = self.norm(hidden_states) - hidden_states = hidden_states.permute(0, 3, 4, 2, 1).reshape(batch_size * height * width, num_frames, channel) - - hidden_states = self.proj_in(hidden_states) - - # 2. Blocks - for block in self.transformer_blocks: - hidden_states = block( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - timestep=timestep, - cross_attention_kwargs=cross_attention_kwargs, - class_labels=class_labels, - ) - - # 3. Output - hidden_states = self.proj_out(hidden_states) - hidden_states = ( - hidden_states[None, None, :] - .reshape(batch_size, height, width, num_frames, channel) - .permute(0, 3, 4, 1, 2) - .contiguous() - ) - hidden_states = hidden_states.reshape(batch_frames, channel, height, width) - - output = hidden_states + residual - - if not return_dict: - return (output,) - - return TransformerTemporalModelOutput(sample=output) - - -class TransformerSpatioTemporalModel(nn.Module): - """ - A Transformer model for video-like data. - - Parameters: - num_attention_heads (`int`, *optional*, defaults to 16): The number of heads to use for multi-head attention. - attention_head_dim (`int`, *optional*, defaults to 88): The number of channels in each head. - in_channels (`int`, *optional*): - The number of channels in the input and output (specify if the input is **continuous**). - out_channels (`int`, *optional*): - The number of channels in the output (specify if the input is **continuous**). - num_layers (`int`, *optional*, defaults to 1): The number of layers of Transformer blocks to use. - cross_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use. - """ - - def __init__( - self, - num_attention_heads: int = 16, - attention_head_dim: int = 88, - in_channels: int = 320, - out_channels: int | None = None, - num_layers: int = 1, - cross_attention_dim: int | None = None, - ): - super().__init__() - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - - inner_dim = num_attention_heads * attention_head_dim - self.inner_dim = inner_dim - - # 2. Define input layers - self.in_channels = in_channels - self.norm = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6) - self.proj_in = nn.Linear(in_channels, inner_dim) - - # 3. Define transformers blocks - self.transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlock( - inner_dim, - num_attention_heads, - attention_head_dim, - cross_attention_dim=cross_attention_dim, - ) - for d in range(num_layers) - ] - ) - - time_mix_inner_dim = inner_dim - self.temporal_transformer_blocks = nn.ModuleList( - [ - TemporalBasicTransformerBlock( - inner_dim, - time_mix_inner_dim, - num_attention_heads, - attention_head_dim, - cross_attention_dim=cross_attention_dim, - ) - for _ in range(num_layers) - ] - ) - - time_embed_dim = in_channels * 4 - self.time_pos_embed = TimestepEmbedding(in_channels, time_embed_dim, out_dim=in_channels) - self.time_proj = Timesteps(in_channels, True, 0) - self.time_mixer = AlphaBlender(alpha=0.5, merge_strategy="learned_with_images") - - # 4. Define output layers - self.out_channels = in_channels if out_channels is None else out_channels - # TODO: should use out_channels for continuous projections - self.proj_out = nn.Linear(inner_dim, in_channels) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - image_only_indicator: torch.Tensor | None = None, - return_dict: bool = True, - ): - """ - Args: - hidden_states (`torch.Tensor` of shape `(batch size, channel, height, width)`): - Input hidden_states. - num_frames (`int`): - The number of frames to be processed per batch. This is used to reshape the hidden states. - encoder_hidden_states ( `torch.LongTensor` of shape `(batch size, encoder_hidden_states dim)`, *optional*): - Conditional embeddings for cross attention layer. If not given, cross-attention defaults to - self-attention. - image_only_indicator (`torch.LongTensor` of shape `(batch size, num_frames)`, *optional*): - A tensor indicating whether the input contains only images. 1 indicates that the input contains only - images, 0 indicates that the input contains video frames. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`] - instead of a plain tuple. - - Returns: - [`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`] or `tuple`: - If `return_dict` is True, an - [`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - # 1. Input - batch_frames, _, height, width = hidden_states.shape - num_frames = image_only_indicator.shape[-1] - batch_size = batch_frames // num_frames - - time_context = encoder_hidden_states - time_context_first_timestep = time_context[None, :].reshape( - batch_size, num_frames, -1, time_context.shape[-1] - )[:, 0] - time_context = time_context_first_timestep[:, None].broadcast_to( - batch_size, height * width, time_context.shape[-2], time_context.shape[-1] - ) - time_context = time_context.reshape(batch_size * height * width, -1, time_context.shape[-1]) - - residual = hidden_states - - hidden_states = self.norm(hidden_states) - inner_dim = hidden_states.shape[1] - hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch_frames, height * width, inner_dim) - hidden_states = self.proj_in(hidden_states) - - num_frames_emb = torch.arange(num_frames, device=hidden_states.device) - num_frames_emb = num_frames_emb.repeat(batch_size, 1) - num_frames_emb = num_frames_emb.reshape(-1) - t_emb = self.time_proj(num_frames_emb) - - # `Timesteps` does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=hidden_states.dtype) - - emb = self.time_pos_embed(t_emb) - emb = emb[:, None, :] - - # 2. Blocks - for block, temporal_block in zip(self.transformer_blocks, self.temporal_transformer_blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, hidden_states, None, encoder_hidden_states, None - ) - else: - hidden_states = block(hidden_states, encoder_hidden_states=encoder_hidden_states) - - hidden_states_mix = hidden_states - hidden_states_mix = hidden_states_mix + emb - - hidden_states_mix = temporal_block( - hidden_states_mix, - num_frames=num_frames, - encoder_hidden_states=time_context, - ) - hidden_states = self.time_mixer( - x_spatial=hidden_states, - x_temporal=hidden_states_mix, - image_only_indicator=image_only_indicator, - ) - - # 3. Output - hidden_states = self.proj_out(hidden_states) - hidden_states = hidden_states.reshape(batch_frames, height, width, inner_dim).permute(0, 3, 1, 2).contiguous() - - output = hidden_states + residual - - if not return_dict: - return (output,) - - return TransformerTemporalModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_wan.py b/diffusers/models/transformers/transformer_wan.py deleted file mode 100644 index cf1b4ecc5d78073709039c7b040ee36e5c0c689c..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_wan.py +++ /dev/null @@ -1,735 +0,0 @@ -# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, deprecate, logging -from ...utils.torch_utils import maybe_allow_in_graph -from .._modeling_parallel import ContextParallelInput, ContextParallelOutput -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps, get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import FP32LayerNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _get_qkv_projections(attn: "WanAttention", hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor): - # encoder_hidden_states is only passed for cross-attention - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - if attn.fused_projections: - if not attn.is_cross_attention: - # In self-attention layers, we can fuse the entire QKV projection into a single linear - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - else: - # In cross-attention layers, we can only fuse the KV projections into a single linear - query = attn.to_q(hidden_states) - key, value = attn.to_kv(encoder_hidden_states).chunk(2, dim=-1) - else: - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - return query, key, value - - -def _get_added_kv_projections(attn: "WanAttention", encoder_hidden_states_img: torch.Tensor): - if attn.fused_projections: - key_img, value_img = attn.to_added_kv(encoder_hidden_states_img).chunk(2, dim=-1) - else: - key_img = attn.add_k_proj(encoder_hidden_states_img) - value_img = attn.add_v_proj(encoder_hidden_states_img) - return key_img, value_img - - -class WanAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "WanAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to version 2.0 or higher." - ) - - def __call__( - self, - attn: "WanAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - encoder_hidden_states_img = None - if attn.add_k_proj is not None: - # 512 is the context length of the text encoder, hardcoded for now - image_context_length = encoder_hidden_states.shape[1] - 512 - encoder_hidden_states_img = encoder_hidden_states[:, :image_context_length] - encoder_hidden_states = encoder_hidden_states[:, image_context_length:] - - query, key, value = _get_qkv_projections(attn, hidden_states, encoder_hidden_states) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - if rotary_emb is not None: - - def apply_rotary_emb( - hidden_states: torch.Tensor, - freqs_cos: torch.Tensor, - freqs_sin: torch.Tensor, - ): - x1, x2 = hidden_states.unflatten(-1, (-1, 2)).unbind(-1) - cos = freqs_cos[..., 0::2] - sin = freqs_sin[..., 1::2] - out = torch.empty_like(hidden_states) - out[..., 0::2] = x1 * cos - x2 * sin - out[..., 1::2] = x1 * sin + x2 * cos - return out.type_as(hidden_states) - - query = apply_rotary_emb(query, *rotary_emb) - key = apply_rotary_emb(key, *rotary_emb) - - # I2V task - hidden_states_img = None - if encoder_hidden_states_img is not None: - key_img, value_img = _get_added_kv_projections(attn, encoder_hidden_states_img) - key_img = attn.norm_added_k(key_img) - - key_img = key_img.unflatten(2, (attn.heads, -1)) - value_img = value_img.unflatten(2, (attn.heads, -1)) - - hidden_states_img = dispatch_attention_fn( - query, - key_img, - value_img, - attn_mask=None, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - # Reference: https://github.com/huggingface/diffusers/pull/12909 - parallel_config=None, - ) - hidden_states_img = hidden_states_img.flatten(2, 3) - hidden_states_img = hidden_states_img.type_as(query) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - # Reference: https://github.com/huggingface/diffusers/pull/12909 - parallel_config=(self._parallel_config if encoder_hidden_states is None else None), - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - - if hidden_states_img is not None: - hidden_states = hidden_states + hidden_states_img - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -class WanAttnProcessor2_0: - def __new__(cls, *args, **kwargs): - deprecation_message = ( - "The WanAttnProcessor2_0 class is deprecated and will be removed in a future version. " - "Please use WanAttnProcessor instead. " - ) - deprecate("WanAttnProcessor2_0", "1.0.0", deprecation_message, standard_warn=False) - return WanAttnProcessor(*args, **kwargs) - - -class WanAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = WanAttnProcessor - _available_processors = [WanAttnProcessor] - - def __init__( - self, - dim: int, - heads: int = 8, - dim_head: int = 64, - eps: float = 1e-5, - dropout: float = 0.0, - added_kv_proj_dim: int | None = None, - cross_attention_dim_head: int | None = None, - processor=None, - is_cross_attention=None, - ): - super().__init__() - - self.inner_dim = dim_head * heads - self.heads = heads - self.added_kv_proj_dim = added_kv_proj_dim - self.cross_attention_dim_head = cross_attention_dim_head - self.kv_inner_dim = self.inner_dim if cross_attention_dim_head is None else cross_attention_dim_head * heads - - self.to_q = torch.nn.Linear(dim, self.inner_dim, bias=True) - self.to_k = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) - self.to_v = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) - self.to_out = torch.nn.ModuleList( - [ - torch.nn.Linear(self.inner_dim, dim, bias=True), - torch.nn.Dropout(dropout), - ] - ) - self.norm_q = torch.nn.RMSNorm(dim_head * heads, eps=eps, elementwise_affine=True) - self.norm_k = torch.nn.RMSNorm(dim_head * heads, eps=eps, elementwise_affine=True) - - self.add_k_proj = self.add_v_proj = None - if added_kv_proj_dim is not None: - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=True) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=True) - self.norm_added_k = torch.nn.RMSNorm(dim_head * heads, eps=eps) - - if is_cross_attention is not None: - self.is_cross_attention = is_cross_attention - else: - self.is_cross_attention = cross_attention_dim_head is not None - - self.set_processor(processor) - - def fuse_projections(self): - if getattr(self, "fused_projections", False): - return - - if not self.is_cross_attention: - concatenated_weights = torch.cat([self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data]) - concatenated_bias = torch.cat([self.to_q.bias.data, self.to_k.bias.data, self.to_v.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_qkv = nn.Linear(in_features, out_features, bias=True) - self.to_qkv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - else: - concatenated_weights = torch.cat([self.to_k.weight.data, self.to_v.weight.data]) - concatenated_bias = torch.cat([self.to_k.bias.data, self.to_v.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_kv = nn.Linear(in_features, out_features, bias=True) - self.to_kv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - - if self.added_kv_proj_dim is not None: - concatenated_weights = torch.cat([self.add_k_proj.weight.data, self.add_v_proj.weight.data]) - concatenated_bias = torch.cat([self.add_k_proj.bias.data, self.add_v_proj.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_added_kv = nn.Linear(in_features, out_features, bias=True) - self.to_added_kv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - - self.fused_projections = True - - @torch.no_grad() - def unfuse_projections(self): - if not getattr(self, "fused_projections", False): - return - - if hasattr(self, "to_qkv"): - delattr(self, "to_qkv") - if hasattr(self, "to_kv"): - delattr(self, "to_kv") - if hasattr(self, "to_added_kv"): - delattr(self, "to_added_kv") - - self.fused_projections = False - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - **kwargs, - ) -> torch.Tensor: - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, rotary_emb, **kwargs) - - -class WanImageEmbedding(torch.nn.Module): - def __init__(self, in_features: int, out_features: int, pos_embed_seq_len=None): - super().__init__() - - self.norm1 = FP32LayerNorm(in_features) - self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu") - self.norm2 = FP32LayerNorm(out_features) - if pos_embed_seq_len is not None: - self.pos_embed = nn.Parameter(torch.zeros(1, pos_embed_seq_len, in_features)) - else: - self.pos_embed = None - - def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor: - if self.pos_embed is not None: - batch_size, seq_len, embed_dim = encoder_hidden_states_image.shape - encoder_hidden_states_image = encoder_hidden_states_image.view(-1, 2 * seq_len, embed_dim) - encoder_hidden_states_image = encoder_hidden_states_image + self.pos_embed - - hidden_states = self.norm1(encoder_hidden_states_image) - hidden_states = self.ff(hidden_states) - hidden_states = self.norm2(hidden_states) - return hidden_states - - -class WanTimeTextImageEmbedding(nn.Module): - def __init__( - self, - dim: int, - time_freq_dim: int, - time_proj_dim: int, - text_embed_dim: int, - image_embed_dim: int | None = None, - pos_embed_seq_len: int | None = None, - ): - super().__init__() - - self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim) - self.act_fn = nn.SiLU() - self.time_proj = nn.Linear(dim, time_proj_dim) - self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh") - - self.image_embedder = None - if image_embed_dim is not None: - self.image_embedder = WanImageEmbedding(image_embed_dim, dim, pos_embed_seq_len=pos_embed_seq_len) - - def forward( - self, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: torch.Tensor | None = None, - timestep_seq_len: int | None = None, - ): - timestep = self.timesteps_proj(timestep) - if timestep_seq_len is not None: - timestep = timestep.unflatten(0, (-1, timestep_seq_len)) - - time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype - if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8: - timestep = timestep.to(time_embedder_dtype) - temb = self.time_embedder(timestep).type_as(encoder_hidden_states) - timestep_proj = self.time_proj(self.act_fn(temb)) - - encoder_hidden_states = self.text_embedder(encoder_hidden_states) - if encoder_hidden_states_image is not None: - encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image) - - return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image - - -class WanRotaryPosEmbed(nn.Module): - def __init__( - self, - attention_head_dim: int, - patch_size: tuple[int, int, int], - max_seq_len: int, - theta: float = 10000.0, - ): - super().__init__() - - self.attention_head_dim = attention_head_dim - self.patch_size = patch_size - self.max_seq_len = max_seq_len - - h_dim = w_dim = 2 * (attention_head_dim // 6) - t_dim = attention_head_dim - h_dim - w_dim - - self.t_dim = t_dim - self.h_dim = h_dim - self.w_dim = w_dim - - freqs_dtype = torch.float32 if torch.backends.mps.is_available() else torch.float64 - - freqs_cos = [] - freqs_sin = [] - - for dim in [t_dim, h_dim, w_dim]: - freq_cos, freq_sin = get_1d_rotary_pos_embed( - dim, - max_seq_len, - theta, - use_real=True, - repeat_interleave_real=True, - freqs_dtype=freqs_dtype, - ) - freqs_cos.append(freq_cos) - freqs_sin.append(freq_sin) - - self.register_buffer("freqs_cos", torch.cat(freqs_cos, dim=1), persistent=False) - self.register_buffer("freqs_sin", torch.cat(freqs_sin, dim=1), persistent=False) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.patch_size - ppf, pph, ppw = num_frames // p_t, height // p_h, width // p_w - - split_sizes = [self.t_dim, self.h_dim, self.w_dim] - - freqs_cos = self.freqs_cos.split(split_sizes, dim=1) - freqs_sin = self.freqs_sin.split(split_sizes, dim=1) - - freqs_cos_f = freqs_cos[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_cos_h = freqs_cos[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_cos_w = freqs_cos[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - - freqs_sin_f = freqs_sin[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_sin_h = freqs_sin[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_sin_w = freqs_sin[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - - freqs_cos = torch.cat([freqs_cos_f, freqs_cos_h, freqs_cos_w], dim=-1).reshape(1, ppf * pph * ppw, 1, -1) - freqs_sin = torch.cat([freqs_sin_f, freqs_sin_h, freqs_sin_w], dim=-1).reshape(1, ppf * pph * ppw, 1, -1) - - return freqs_cos, freqs_sin - - -@maybe_allow_in_graph -class WanTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - ffn_dim: int, - num_heads: int, - qk_norm: str = "rms_norm_across_heads", - cross_attn_norm: bool = False, - eps: float = 1e-6, - added_kv_proj_dim: int | None = None, - ): - super().__init__() - - # 1. Self-attention - self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False) - self.attn1 = WanAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - cross_attention_dim_head=None, - processor=WanAttnProcessor(), - ) - - # 2. Cross-attention - self.attn2 = WanAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - added_kv_proj_dim=added_kv_proj_dim, - cross_attention_dim_head=dim // num_heads, - processor=WanAttnProcessor(), - ) - self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity() - - # 3. Feed-forward - self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate") - self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False) - - self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - rotary_emb: torch.Tensor, - ) -> torch.Tensor: - if temb.ndim == 4: - # temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v) - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table.unsqueeze(0) + temb.float() - ).chunk(6, dim=2) - # batch_size, seq_len, 1, inner_dim - shift_msa = shift_msa.squeeze(2) - scale_msa = scale_msa.squeeze(2) - gate_msa = gate_msa.squeeze(2) - c_shift_msa = c_shift_msa.squeeze(2) - c_scale_msa = c_scale_msa.squeeze(2) - c_gate_msa = c_gate_msa.squeeze(2) - else: - # temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B) - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table + temb.float() - ).chunk(6, dim=1) - - # 1. Self-attention - norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states) - attn_output = self.attn1(norm_hidden_states, None, None, rotary_emb) - hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states) - - # 2. Cross-attention - norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states) - attn_output = self.attn2(norm_hidden_states, encoder_hidden_states, None, None) - hidden_states = hidden_states + attn_output - - # 3. Feed-forward - norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as( - hidden_states - ) - ff_output = self.ffn(norm_hidden_states) - hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states) - - return hidden_states - - -class WanTransformer3DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin, AttentionMixin -): - r""" - A Transformer model for video-like data used in the Wan model. - - Args: - patch_size (`tuple[int]`, defaults to `(1, 2, 2)`): - 3D patch dimensions for video embedding (t_patch, h_patch, w_patch). - num_attention_heads (`int`, defaults to `40`): - Fixed length for text embeddings. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - text_dim (`int`, defaults to `512`): - Input dimension for text embeddings. - freq_dim (`int`, defaults to `256`): - Dimension for sinusoidal time embeddings. - ffn_dim (`int`, defaults to `13824`): - Intermediate dimension in feed-forward network. - num_layers (`int`, defaults to `40`): - The number of layers of transformer blocks to use. - window_size (`tuple[int]`, defaults to `(-1, -1)`): - Window size for local attention (-1 indicates global attention). - cross_attn_norm (`bool`, defaults to `True`): - Enable cross-attention normalization. - qk_norm (`bool`, defaults to `True`): - Enable query/key normalization. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - add_img_emb (`bool`, defaults to `False`): - Whether to use img_emb. - added_kv_proj_dim (`int`, *optional*, defaults to `None`): - The number of channels to use for the added key and value projections. If `None`, no projection is used. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"] - _no_split_modules = ["WanTransformerBlock"] - _keep_in_fp32_modules = ["rope", "time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"] - _keys_to_ignore_on_load_unexpected = ["norm_added_q"] - _repeated_blocks = ["WanTransformerBlock"] - _cp_plan = { - "rope": { - 0: ContextParallelInput(split_dim=1, expected_dims=4, split_output=True), - 1: ContextParallelInput(split_dim=1, expected_dims=4, split_output=True), - }, - "blocks.0": { - "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), - }, - # Reference: https://github.com/huggingface/diffusers/pull/12909 - # We need to disable the splitting of encoder_hidden_states because the image_encoder - # (Wan 2.1 I2V) consistently generates 257 tokens for image_embed. This causes the shape - # of encoder_hidden_states—whose token count is always 769 (512 + 257) after concatenation - # —to be indivisible by the number of devices in the CP. - "proj_out": ContextParallelOutput(gather_dim=1, expected_dims=3), - "": { - "timestep": ContextParallelInput(split_dim=1, expected_dims=2, split_output=False), - }, - } - - @register_to_config - def __init__( - self, - patch_size: tuple[int, ...] = (1, 2, 2), - num_attention_heads: int = 40, - attention_head_dim: int = 128, - in_channels: int = 16, - out_channels: int = 16, - text_dim: int = 4096, - freq_dim: int = 256, - ffn_dim: int = 13824, - num_layers: int = 40, - cross_attn_norm: bool = True, - qk_norm: str | None = "rms_norm_across_heads", - eps: float = 1e-6, - image_dim: int | None = None, - added_kv_proj_dim: int | None = None, - rope_max_seq_len: int = 1024, - pos_embed_seq_len: int | None = None, - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - # 1. Patch & position embedding - self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len) - self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - - # 2. Condition embeddings - # image_embedding_dim=1280 for I2V model - self.condition_embedder = WanTimeTextImageEmbedding( - dim=inner_dim, - time_freq_dim=freq_dim, - time_proj_dim=inner_dim * 6, - text_embed_dim=text_dim, - image_embed_dim=image_dim, - pos_embed_seq_len=pos_embed_seq_len, - ) - - # 3. Transformer blocks - self.blocks = nn.ModuleList( - [ - WanTransformerBlock( - inner_dim, ffn_dim, num_attention_heads, qk_norm, cross_attn_norm, eps, added_kv_proj_dim - ) - for _ in range(num_layers) - ] - ) - - # 4. Output norm & projection - self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size)) - self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: torch.Tensor | None = None, - return_dict: bool = True, - attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor | dict[str, torch.Tensor]: - """ - The [`WanTransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_hidden_states_image (`torch.Tensor`, *optional*): - Conditional image embeddings for image-conditioned generation. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.config.patch_size - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p_h - post_patch_width = width // p_w - - rotary_emb = self.rope(hidden_states) - - hidden_states = self.patch_embedding(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) - - # flatten+transpose produces a non-contiguous tensor; make it contiguous before the block loop. - hidden_states = hidden_states.contiguous() - - # timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v) - if timestep.ndim == 2: - ts_seq_len = timestep.shape[1] - timestep = timestep.flatten() # batch_size * seq_len - else: - ts_seq_len = None - - temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder( - timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len - ) - if ts_seq_len is not None: - # batch_size, seq_len, 6, inner_dim - timestep_proj = timestep_proj.unflatten(2, (6, -1)) - else: - # batch_size, 6, inner_dim - timestep_proj = timestep_proj.unflatten(1, (6, -1)) - - if encoder_hidden_states_image is not None: - encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) - - # 4. Transformer blocks - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.blocks: - hidden_states = self._gradient_checkpointing_func( - block, hidden_states, encoder_hidden_states, timestep_proj, rotary_emb - ) - else: - for block in self.blocks: - hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb) - - # 5. Output norm, projection & unpatchify - if temb.ndim == 3: - # batch_size, seq_len, inner_dim (wan 2.2 ti2v) - shift, scale = (self.scale_shift_table.unsqueeze(0).to(temb.device) + temb.unsqueeze(2)).chunk(2, dim=2) - shift = shift.squeeze(2) - scale = scale.squeeze(2) - else: - # batch_size, inner_dim - shift, scale = (self.scale_shift_table.to(temb.device) + temb.unsqueeze(1)).chunk(2, dim=1) - - # Move the shift and scale tensors to the same device as hidden_states. - # When using multi-GPU inference via accelerate these will be on the - # first device rather than the last device, which hidden_states ends up - # on. - shift = shift.to(hidden_states.device) - scale = scale.to(hidden_states.device) - - hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states) - hidden_states = self.proj_out(hidden_states) - - hidden_states = hidden_states.reshape( - batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1 - ) - hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) - output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_wan_animate.py b/diffusers/models/transformers/transformer_wan_animate.py deleted file mode 100644 index 084c3a2aed7dc7198115bccf6196d9619aa91748..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_wan_animate.py +++ /dev/null @@ -1,1306 +0,0 @@ -# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin -from ..embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps, get_1d_rotary_pos_embed -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import FP32LayerNorm - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -WAN_ANIMATE_MOTION_ENCODER_CHANNEL_SIZES = { - "4": 512, - "8": 512, - "16": 512, - "32": 512, - "64": 256, - "128": 128, - "256": 64, - "512": 32, - "1024": 16, -} - - -# Copied from diffusers.models.transformers.transformer_wan._get_qkv_projections -def _get_qkv_projections(attn: "WanAttention", hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor): - # encoder_hidden_states is only passed for cross-attention - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - - if attn.fused_projections: - if not attn.is_cross_attention: - # In self-attention layers, we can fuse the entire QKV projection into a single linear - query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) - else: - # In cross-attention layers, we can only fuse the KV projections into a single linear - query = attn.to_q(hidden_states) - key, value = attn.to_kv(encoder_hidden_states).chunk(2, dim=-1) - else: - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - return query, key, value - - -# Copied from diffusers.models.transformers.transformer_wan._get_added_kv_projections -def _get_added_kv_projections(attn: "WanAttention", encoder_hidden_states_img: torch.Tensor): - if attn.fused_projections: - key_img, value_img = attn.to_added_kv(encoder_hidden_states_img).chunk(2, dim=-1) - else: - key_img = attn.add_k_proj(encoder_hidden_states_img) - value_img = attn.add_v_proj(encoder_hidden_states_img) - return key_img, value_img - - -class FusedLeakyReLU(nn.Module): - """ - Fused LeakyRelu with scale factor and channel-wise bias. - """ - - def __init__(self, negative_slope: float = 0.2, scale: float = 2**0.5, bias_channels: int | None = None): - super().__init__() - self.negative_slope = negative_slope - self.scale = scale - self.channels = bias_channels - - if self.channels is not None: - self.bias = nn.Parameter( - torch.zeros( - self.channels, - ) - ) - else: - self.bias = None - - def forward(self, x: torch.Tensor, channel_dim: int = 1) -> torch.Tensor: - if self.bias is not None: - # Expand self.bias to have all singleton dims except at self.channel_dim - expanded_shape = [1] * x.ndim - expanded_shape[channel_dim] = self.bias.shape[0] - bias = self.bias.reshape(*expanded_shape) - x = x + bias - return F.leaky_relu(x, self.negative_slope) * self.scale - - -class MotionConv2d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int, - stride: int = 1, - padding: int = 0, - bias: bool = True, - blur_kernel: tuple[int, ...] | None = None, - blur_upsample_factor: int = 1, - use_activation: bool = True, - ): - super().__init__() - self.use_activation = use_activation - self.in_channels = in_channels - - # Handle blurring (applying a FIR filter with the given kernel) if available - self.blur = False - if blur_kernel is not None: - p = (len(blur_kernel) - stride) + (kernel_size - 1) - self.blur_padding = ((p + 1) // 2, p // 2) - - kernel = torch.tensor(blur_kernel) - # Convert kernel to 2D if necessary - if kernel.ndim == 1: - kernel = kernel[None, :] * kernel[:, None] - # Normalize kernel - kernel = kernel / kernel.sum() - if blur_upsample_factor > 1: - kernel = kernel * (blur_upsample_factor**2) - self.register_buffer("blur_kernel", kernel, persistent=False) - self.blur = True - - # Main Conv2d parameters (with scale factor) - self.weight = nn.Parameter(torch.randn(out_channels, in_channels, kernel_size, kernel_size)) - self.scale = 1 / math.sqrt(in_channels * kernel_size**2) - - self.stride = stride - self.padding = padding - - # If using an activation function, the bias will be fused into the activation - if bias and not self.use_activation: - self.bias = nn.Parameter(torch.zeros(out_channels)) - else: - self.bias = None - - if self.use_activation: - self.act_fn = FusedLeakyReLU(bias_channels=out_channels) - else: - self.act_fn = None - - def forward(self, x: torch.Tensor, channel_dim: int = 1) -> torch.Tensor: - # Apply blur if using - if self.blur: - # NOTE: the original implementation uses a 2D upfirdn operation with the upsampling and downsampling rates - # set to 1, which should be equivalent to a 2D convolution - expanded_kernel = self.blur_kernel[None, None, :, :].expand(self.in_channels, 1, -1, -1) - x = F.conv2d(x, expanded_kernel.to(x.dtype), padding=self.blur_padding, groups=self.in_channels) - - # Main Conv2D with scaling - x = x.to(self.weight.dtype) - x = F.conv2d(x, self.weight * self.scale, bias=self.bias, stride=self.stride, padding=self.padding) - - # Activation with fused bias, if using - if self.use_activation: - x = self.act_fn(x, channel_dim=channel_dim) - return x - - def __repr__(self): - return ( - f"{self.__class__.__name__}({self.weight.shape[1]}, {self.weight.shape[0]}," - f" kernel_size={self.weight.shape[2]}, stride={self.stride}, padding={self.padding})" - ) - - -class MotionLinear(nn.Module): - def __init__( - self, - in_dim: int, - out_dim: int, - bias: bool = True, - use_activation: bool = False, - ): - super().__init__() - self.use_activation = use_activation - - # Linear weight with scale factor - self.weight = nn.Parameter(torch.randn(out_dim, in_dim)) - self.scale = 1 / math.sqrt(in_dim) - - # If an activation is present, the bias will be fused to it - if bias and not self.use_activation: - self.bias = nn.Parameter(torch.zeros(out_dim)) - else: - self.bias = None - - if self.use_activation: - self.act_fn = FusedLeakyReLU(bias_channels=out_dim) - else: - self.act_fn = None - - def forward(self, input: torch.Tensor, channel_dim: int = 1) -> torch.Tensor: - out = F.linear(input, self.weight * self.scale, bias=self.bias) - if self.use_activation: - out = self.act_fn(out, channel_dim=channel_dim) - return out - - def __repr__(self): - return ( - f"{self.__class__.__name__}(in_features={self.weight.shape[1]}, out_features={self.weight.shape[0]}," - f" bias={self.bias is not None})" - ) - - -class MotionEncoderResBlock(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int = 3, - kernel_size_skip: int = 1, - blur_kernel: tuple[int, ...] = (1, 3, 3, 1), - downsample_factor: int = 2, - ): - super().__init__() - self.downsample_factor = downsample_factor - - # 3 x 3 Conv + fused leaky ReLU - self.conv1 = MotionConv2d( - in_channels, - in_channels, - kernel_size, - stride=1, - padding=kernel_size // 2, - use_activation=True, - ) - - # 3 x 3 Conv that downsamples 2x + fused leaky ReLU - self.conv2 = MotionConv2d( - in_channels, - out_channels, - kernel_size=kernel_size, - stride=self.downsample_factor, - padding=0, - blur_kernel=blur_kernel, - use_activation=True, - ) - - # 1 x 1 Conv that downsamples 2x in skip connection - self.conv_skip = MotionConv2d( - in_channels, - out_channels, - kernel_size=kernel_size_skip, - stride=self.downsample_factor, - padding=0, - bias=False, - blur_kernel=blur_kernel, - use_activation=False, - ) - - def forward(self, x: torch.Tensor, channel_dim: int = 1) -> torch.Tensor: - x_out = self.conv1(x, channel_dim) - x_out = self.conv2(x_out, channel_dim) - - x_skip = self.conv_skip(x, channel_dim) - - x_out = (x_out + x_skip) / math.sqrt(2) - return x_out - - -class WanAnimateMotionEncoder(nn.Module): - def __init__( - self, - size: int = 512, - style_dim: int = 512, - motion_dim: int = 20, - out_dim: int = 512, - motion_blocks: int = 5, - channels: dict[str, int] | None = None, - ): - super().__init__() - self.size = size - - # Appearance encoder: conv layers - if channels is None: - channels = WAN_ANIMATE_MOTION_ENCODER_CHANNEL_SIZES - - self.conv_in = MotionConv2d(3, channels[str(size)], 1, use_activation=True) - - self.res_blocks = nn.ModuleList() - in_channels = channels[str(size)] - log_size = int(math.log(size, 2)) - for i in range(log_size, 2, -1): - out_channels = channels[str(2 ** (i - 1))] - self.res_blocks.append(MotionEncoderResBlock(in_channels, out_channels)) - in_channels = out_channels - - self.conv_out = MotionConv2d(in_channels, style_dim, 4, padding=0, bias=False, use_activation=False) - - # Motion encoder: linear layers - # NOTE: there are no activations in between the linear layers here, which is weird but I believe matches the - # original code. - linears = [MotionLinear(style_dim, style_dim) for _ in range(motion_blocks - 1)] - linears.append(MotionLinear(style_dim, motion_dim)) - self.motion_network = nn.ModuleList(linears) - - self.motion_synthesis_weight = nn.Parameter(torch.randn(out_dim, motion_dim)) - - def forward(self, face_image: torch.Tensor, channel_dim: int = 1) -> torch.Tensor: - if (face_image.shape[-2] != self.size) or (face_image.shape[-1] != self.size): - raise ValueError( - f"Face pixel values has resolution ({face_image.shape[-1]}, {face_image.shape[-2]}) but is expected" - f" to have resolution ({self.size}, {self.size})" - ) - - # Appearance encoding through convs - face_image = self.conv_in(face_image, channel_dim) - for block in self.res_blocks: - face_image = block(face_image, channel_dim) - face_image = self.conv_out(face_image, channel_dim) - motion_feat = face_image.squeeze(-1).squeeze(-1) - - # Motion feature extraction - for linear_layer in self.motion_network: - motion_feat = linear_layer(motion_feat, channel_dim=channel_dim) - - # Motion synthesis via Linear Motion Decomposition - weight = self.motion_synthesis_weight + 1e-8 - # Upcast the QR orthogonalization operation to FP32 - original_motion_dtype = motion_feat.dtype - motion_feat = motion_feat.to(torch.float32) - weight = weight.to(torch.float32) - - Q = torch.linalg.qr(weight)[0].to(device=motion_feat.device) - - motion_feat_diag = torch.diag_embed(motion_feat) # Alpha, diagonal matrix - motion_decomposition = torch.matmul(motion_feat_diag, Q.T) - motion_vec = torch.sum(motion_decomposition, dim=1) - - motion_vec = motion_vec.to(dtype=original_motion_dtype) - - return motion_vec - - -class WanAnimateFaceEncoder(nn.Module): - def __init__( - self, - in_dim: int, - out_dim: int, - hidden_dim: int = 1024, - num_heads: int = 4, - kernel_size: int = 3, - eps: float = 1e-6, - pad_mode: str = "replicate", - ): - super().__init__() - self.num_heads = num_heads - self.time_causal_padding = (kernel_size - 1, 0) - self.pad_mode = pad_mode - - self.act = nn.SiLU() - - self.conv1_local = nn.Conv1d(in_dim, hidden_dim * num_heads, kernel_size=kernel_size, stride=1) - self.conv2 = nn.Conv1d(hidden_dim, hidden_dim, kernel_size, stride=2) - self.conv3 = nn.Conv1d(hidden_dim, hidden_dim, kernel_size, stride=2) - - self.norm1 = nn.LayerNorm(hidden_dim, eps, elementwise_affine=False) - self.norm2 = nn.LayerNorm(hidden_dim, eps, elementwise_affine=False) - self.norm3 = nn.LayerNorm(hidden_dim, eps, elementwise_affine=False) - - self.out_proj = nn.Linear(hidden_dim, out_dim) - - self.padding_tokens = nn.Parameter(torch.zeros(1, 1, 1, out_dim)) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - batch_size = x.shape[0] - - # Reshape to channels-first to apply causal Conv1d over frame dim - x = x.permute(0, 2, 1) - x = F.pad(x, self.time_causal_padding, mode=self.pad_mode) - x = self.conv1_local(x) # [B, C, T_padded] --> [B, N * C, T] - x = x.unflatten(1, (self.num_heads, -1)).flatten(0, 1) # [B, N * C, T] --> [B * N, C, T] - # Reshape back to channels-last to apply LayerNorm over channel dim - x = x.permute(0, 2, 1) - x = self.norm1(x) - x = self.act(x) - - x = x.permute(0, 2, 1) - x = F.pad(x, self.time_causal_padding, mode=self.pad_mode) - x = self.conv2(x) - x = x.permute(0, 2, 1) - x = self.norm2(x) - x = self.act(x) - - x = x.permute(0, 2, 1) - x = F.pad(x, self.time_causal_padding, mode=self.pad_mode) - x = self.conv3(x) - x = x.permute(0, 2, 1) - x = self.norm3(x) - x = self.act(x) - - x = self.out_proj(x) - x = x.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3) # [B * N, T, C_out] --> [B, T, N, C_out] - - padding = self.padding_tokens.repeat(batch_size, x.shape[1], 1, 1).to(device=x.device) - x = torch.cat([x, padding], dim=-2) # [B, T, N, C_out] --> [B, T, N + 1, C_out] - - return x - - -class WanAnimateFaceBlockAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - f"{self.__class__.__name__} requires PyTorch 2.0. To use it, please upgrade PyTorch to version 2.0 or" - f" higher." - ) - - def __call__( - self, - attn: "WanAnimateFaceBlockCrossAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - # encoder_hidden_states corresponds to the motion vec - # attention_mask corresponds to the motion mask (if any) - hidden_states = attn.pre_norm_q(hidden_states) - encoder_hidden_states = attn.pre_norm_kv(encoder_hidden_states) - - # B --> batch_size, T --> reduced inference segment len, N --> face_encoder_num_heads + 1, C --> attn.dim - B, T, N, C = encoder_hidden_states.shape - - # Flatten T and N so the K/V projections see a 3D tensor; BnB int8 matmul only - # accepts 2D/3D inputs and would otherwise fail on this 4D activation. - encoder_hidden_states = encoder_hidden_states.flatten(1, 2) # [B, T, N, C] --> [B, T * N, C] - - query, key, value = _get_qkv_projections(attn, hidden_states, encoder_hidden_states) - - query = query.unflatten(2, (attn.heads, -1)) # [B, S, H * D] --> [B, S, H, D] - key = key.view(B, T, N, attn.heads, -1) # [B, T * N, H * D_kv] --> [B, T, N, H, D_kv] - value = value.view(B, T, N, attn.heads, -1) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - # NOTE: the below line (which follows the official code) means that in practice, the number of frames T in - # encoder_hidden_states (the motion vector after applying the face encoder) must evenly divide the - # post-patchify sequence length S of the transformer hidden_states. Is it possible to remove this dependency? - query = query.unflatten(1, (T, -1)).flatten(0, 1) # [B, S, H, D] --> [B * T, S / T, H, D] - key = key.flatten(0, 1) # [B, T, N, H, D_kv] --> [B * T, N, H, D_kv] - value = value.flatten(0, 1) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=None, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - hidden_states = hidden_states.unflatten(0, (B, T)).flatten(1, 2) - - hidden_states = attn.to_out(hidden_states) - - if attention_mask is not None: - # NOTE: attention_mask is assumed to be a multiplicative mask - attention_mask = attention_mask.flatten(start_dim=1) - hidden_states = hidden_states * attention_mask - - return hidden_states - - -class WanAnimateFaceBlockCrossAttention(nn.Module, AttentionModuleMixin): - """ - Temporally-aligned cross attention with the face motion signal in the Wan Animate Face Blocks. - """ - - _default_processor_cls = WanAnimateFaceBlockAttnProcessor - _available_processors = [WanAnimateFaceBlockAttnProcessor] - - def __init__( - self, - dim: int, - heads: int = 8, - dim_head: int = 64, - eps: float = 1e-6, - cross_attention_dim_head: int | None = None, - bias: bool = True, - processor=None, - ): - super().__init__() - self.inner_dim = dim_head * heads - self.heads = heads - self.cross_attention_dim_head = cross_attention_dim_head - self.kv_inner_dim = self.inner_dim if cross_attention_dim_head is None else cross_attention_dim_head * heads - self.use_bias = bias - self.is_cross_attention = cross_attention_dim_head is not None - - # 1. Pre-Attention Norms for the hidden_states (video latents) and encoder_hidden_states (motion vector). - # NOTE: this is not used in "vanilla" WanAttention - self.pre_norm_q = nn.LayerNorm(dim, eps, elementwise_affine=False) - self.pre_norm_kv = nn.LayerNorm(dim, eps, elementwise_affine=False) - - # 2. QKV and Output Projections - self.to_q = torch.nn.Linear(dim, self.inner_dim, bias=bias) - self.to_k = torch.nn.Linear(dim, self.kv_inner_dim, bias=bias) - self.to_v = torch.nn.Linear(dim, self.kv_inner_dim, bias=bias) - self.to_out = torch.nn.Linear(self.inner_dim, dim, bias=bias) - - # 3. QK Norm - # NOTE: this is applied after the reshape, so only over dim_head rather than dim_head * heads - self.norm_q = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=True) - self.norm_k = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=True) - - # 4. Set attention processor - if processor is None: - processor = self._default_processor_cls() - self.set_processor(processor) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - **kwargs, - ) -> torch.Tensor: - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask) - - -# Copied from diffusers.models.transformers.transformer_wan.WanAttnProcessor -class WanAttnProcessor: - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "WanAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to version 2.0 or higher." - ) - - def __call__( - self, - attn: "WanAttention", - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - ) -> torch.Tensor: - encoder_hidden_states_img = None - if attn.add_k_proj is not None: - # 512 is the context length of the text encoder, hardcoded for now - image_context_length = encoder_hidden_states.shape[1] - 512 - encoder_hidden_states_img = encoder_hidden_states[:, :image_context_length] - encoder_hidden_states = encoder_hidden_states[:, image_context_length:] - - query, key, value = _get_qkv_projections(attn, hidden_states, encoder_hidden_states) - - query = attn.norm_q(query) - key = attn.norm_k(key) - - query = query.unflatten(2, (attn.heads, -1)) - key = key.unflatten(2, (attn.heads, -1)) - value = value.unflatten(2, (attn.heads, -1)) - - if rotary_emb is not None: - - def apply_rotary_emb( - hidden_states: torch.Tensor, - freqs_cos: torch.Tensor, - freqs_sin: torch.Tensor, - ): - x1, x2 = hidden_states.unflatten(-1, (-1, 2)).unbind(-1) - cos = freqs_cos[..., 0::2] - sin = freqs_sin[..., 1::2] - out = torch.empty_like(hidden_states) - out[..., 0::2] = x1 * cos - x2 * sin - out[..., 1::2] = x1 * sin + x2 * cos - return out.type_as(hidden_states) - - query = apply_rotary_emb(query, *rotary_emb) - key = apply_rotary_emb(key, *rotary_emb) - - # I2V task - hidden_states_img = None - if encoder_hidden_states_img is not None: - key_img, value_img = _get_added_kv_projections(attn, encoder_hidden_states_img) - key_img = attn.norm_added_k(key_img) - - key_img = key_img.unflatten(2, (attn.heads, -1)) - value_img = value_img.unflatten(2, (attn.heads, -1)) - - hidden_states_img = dispatch_attention_fn( - query, - key_img, - value_img, - attn_mask=None, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - # Reference: https://github.com/huggingface/diffusers/pull/12909 - parallel_config=None, - ) - hidden_states_img = hidden_states_img.flatten(2, 3) - hidden_states_img = hidden_states_img.type_as(query) - - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - # Reference: https://github.com/huggingface/diffusers/pull/12909 - parallel_config=(self._parallel_config if encoder_hidden_states is None else None), - ) - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.type_as(query) - - if hidden_states_img is not None: - hidden_states = hidden_states + hidden_states_img - - hidden_states = attn.to_out[0](hidden_states) - hidden_states = attn.to_out[1](hidden_states) - return hidden_states - - -# Copied from diffusers.models.transformers.transformer_wan.WanAttention -class WanAttention(torch.nn.Module, AttentionModuleMixin): - _default_processor_cls = WanAttnProcessor - _available_processors = [WanAttnProcessor] - - def __init__( - self, - dim: int, - heads: int = 8, - dim_head: int = 64, - eps: float = 1e-5, - dropout: float = 0.0, - added_kv_proj_dim: int | None = None, - cross_attention_dim_head: int | None = None, - processor=None, - is_cross_attention=None, - ): - super().__init__() - - self.inner_dim = dim_head * heads - self.heads = heads - self.added_kv_proj_dim = added_kv_proj_dim - self.cross_attention_dim_head = cross_attention_dim_head - self.kv_inner_dim = self.inner_dim if cross_attention_dim_head is None else cross_attention_dim_head * heads - - self.to_q = torch.nn.Linear(dim, self.inner_dim, bias=True) - self.to_k = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) - self.to_v = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) - self.to_out = torch.nn.ModuleList( - [ - torch.nn.Linear(self.inner_dim, dim, bias=True), - torch.nn.Dropout(dropout), - ] - ) - self.norm_q = torch.nn.RMSNorm(dim_head * heads, eps=eps, elementwise_affine=True) - self.norm_k = torch.nn.RMSNorm(dim_head * heads, eps=eps, elementwise_affine=True) - - self.add_k_proj = self.add_v_proj = None - if added_kv_proj_dim is not None: - self.add_k_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=True) - self.add_v_proj = torch.nn.Linear(added_kv_proj_dim, self.inner_dim, bias=True) - self.norm_added_k = torch.nn.RMSNorm(dim_head * heads, eps=eps) - - if is_cross_attention is not None: - self.is_cross_attention = is_cross_attention - else: - self.is_cross_attention = cross_attention_dim_head is not None - - self.set_processor(processor) - - def fuse_projections(self): - if getattr(self, "fused_projections", False): - return - - if not self.is_cross_attention: - concatenated_weights = torch.cat([self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data]) - concatenated_bias = torch.cat([self.to_q.bias.data, self.to_k.bias.data, self.to_v.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_qkv = nn.Linear(in_features, out_features, bias=True) - self.to_qkv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - else: - concatenated_weights = torch.cat([self.to_k.weight.data, self.to_v.weight.data]) - concatenated_bias = torch.cat([self.to_k.bias.data, self.to_v.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_kv = nn.Linear(in_features, out_features, bias=True) - self.to_kv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - - if self.added_kv_proj_dim is not None: - concatenated_weights = torch.cat([self.add_k_proj.weight.data, self.add_v_proj.weight.data]) - concatenated_bias = torch.cat([self.add_k_proj.bias.data, self.add_v_proj.bias.data]) - out_features, in_features = concatenated_weights.shape - with torch.device("meta"): - self.to_added_kv = nn.Linear(in_features, out_features, bias=True) - self.to_added_kv.load_state_dict( - {"weight": concatenated_weights, "bias": concatenated_bias}, strict=True, assign=True - ) - - self.fused_projections = True - - @torch.no_grad() - def unfuse_projections(self): - if not getattr(self, "fused_projections", False): - return - - if hasattr(self, "to_qkv"): - delattr(self, "to_qkv") - if hasattr(self, "to_kv"): - delattr(self, "to_kv") - if hasattr(self, "to_added_kv"): - delattr(self, "to_added_kv") - - self.fused_projections = False - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, - **kwargs, - ) -> torch.Tensor: - return self.processor(self, hidden_states, encoder_hidden_states, attention_mask, rotary_emb, **kwargs) - - -# Copied from diffusers.models.transformers.transformer_wan.WanImageEmbedding -class WanImageEmbedding(torch.nn.Module): - def __init__(self, in_features: int, out_features: int, pos_embed_seq_len=None): - super().__init__() - - self.norm1 = FP32LayerNorm(in_features) - self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu") - self.norm2 = FP32LayerNorm(out_features) - if pos_embed_seq_len is not None: - self.pos_embed = nn.Parameter(torch.zeros(1, pos_embed_seq_len, in_features)) - else: - self.pos_embed = None - - def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor: - if self.pos_embed is not None: - batch_size, seq_len, embed_dim = encoder_hidden_states_image.shape - encoder_hidden_states_image = encoder_hidden_states_image.view(-1, 2 * seq_len, embed_dim) - encoder_hidden_states_image = encoder_hidden_states_image + self.pos_embed - - hidden_states = self.norm1(encoder_hidden_states_image) - hidden_states = self.ff(hidden_states) - hidden_states = self.norm2(hidden_states) - return hidden_states - - -# Modified from diffusers.models.transformers.transformer_wan.WanTimeTextImageEmbedding -class WanTimeTextImageEmbedding(nn.Module): - def __init__( - self, - dim: int, - time_freq_dim: int, - time_proj_dim: int, - text_embed_dim: int, - image_embed_dim: int | None = None, - pos_embed_seq_len: int | None = None, - ): - super().__init__() - - self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0) - self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim) - self.act_fn = nn.SiLU() - self.time_proj = nn.Linear(dim, time_proj_dim) - self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh") - - self.image_embedder = None - if image_embed_dim is not None: - self.image_embedder = WanImageEmbedding(image_embed_dim, dim, pos_embed_seq_len=pos_embed_seq_len) - - def forward( - self, - timestep: torch.Tensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: torch.Tensor | None = None, - timestep_seq_len: int | None = None, - ): - timestep = self.timesteps_proj(timestep) - if timestep_seq_len is not None: - timestep = timestep.unflatten(0, (-1, timestep_seq_len)) - - if self.time_embedder.linear_1.weight.dtype.is_floating_point: - time_embedder_dtype = self.time_embedder.linear_1.weight.dtype - else: - time_embedder_dtype = encoder_hidden_states.dtype - - temb = self.time_embedder(timestep.to(time_embedder_dtype)).type_as(encoder_hidden_states) - timestep_proj = self.time_proj(self.act_fn(temb)) - - encoder_hidden_states = self.text_embedder(encoder_hidden_states) - if encoder_hidden_states_image is not None: - encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image) - - return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image - - -# Copied from diffusers.models.transformers.transformer_wan.WanRotaryPosEmbed -class WanRotaryPosEmbed(nn.Module): - def __init__( - self, - attention_head_dim: int, - patch_size: tuple[int, int, int], - max_seq_len: int, - theta: float = 10000.0, - ): - super().__init__() - - self.attention_head_dim = attention_head_dim - self.patch_size = patch_size - self.max_seq_len = max_seq_len - - h_dim = w_dim = 2 * (attention_head_dim // 6) - t_dim = attention_head_dim - h_dim - w_dim - - self.t_dim = t_dim - self.h_dim = h_dim - self.w_dim = w_dim - - freqs_dtype = torch.float32 if torch.backends.mps.is_available() else torch.float64 - - freqs_cos = [] - freqs_sin = [] - - for dim in [t_dim, h_dim, w_dim]: - freq_cos, freq_sin = get_1d_rotary_pos_embed( - dim, - max_seq_len, - theta, - use_real=True, - repeat_interleave_real=True, - freqs_dtype=freqs_dtype, - ) - freqs_cos.append(freq_cos) - freqs_sin.append(freq_sin) - - self.register_buffer("freqs_cos", torch.cat(freqs_cos, dim=1), persistent=False) - self.register_buffer("freqs_sin", torch.cat(freqs_sin, dim=1), persistent=False) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.patch_size - ppf, pph, ppw = num_frames // p_t, height // p_h, width // p_w - - split_sizes = [self.t_dim, self.h_dim, self.w_dim] - - freqs_cos = self.freqs_cos.split(split_sizes, dim=1) - freqs_sin = self.freqs_sin.split(split_sizes, dim=1) - - freqs_cos_f = freqs_cos[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_cos_h = freqs_cos[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_cos_w = freqs_cos[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - - freqs_sin_f = freqs_sin[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1) - freqs_sin_h = freqs_sin[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1) - freqs_sin_w = freqs_sin[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) - - freqs_cos = torch.cat([freqs_cos_f, freqs_cos_h, freqs_cos_w], dim=-1).reshape(1, ppf * pph * ppw, 1, -1) - freqs_sin = torch.cat([freqs_sin_f, freqs_sin_h, freqs_sin_w], dim=-1).reshape(1, ppf * pph * ppw, 1, -1) - - return freqs_cos, freqs_sin - - -# Copied from diffusers.models.transformers.transformer_wan.WanTransformerBlock -class WanTransformerBlock(nn.Module): - def __init__( - self, - dim: int, - ffn_dim: int, - num_heads: int, - qk_norm: str = "rms_norm_across_heads", - cross_attn_norm: bool = False, - eps: float = 1e-6, - added_kv_proj_dim: int | None = None, - ): - super().__init__() - - # 1. Self-attention - self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False) - self.attn1 = WanAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - cross_attention_dim_head=None, - processor=WanAttnProcessor(), - ) - - # 2. Cross-attention - self.attn2 = WanAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - added_kv_proj_dim=added_kv_proj_dim, - cross_attention_dim_head=dim // num_heads, - processor=WanAttnProcessor(), - ) - self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity() - - # 3. Feed-forward - self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate") - self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False) - - self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - temb: torch.Tensor, - rotary_emb: torch.Tensor, - ) -> torch.Tensor: - if temb.ndim == 4: - # temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v) - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table.unsqueeze(0) + temb.float() - ).chunk(6, dim=2) - # batch_size, seq_len, 1, inner_dim - shift_msa = shift_msa.squeeze(2) - scale_msa = scale_msa.squeeze(2) - gate_msa = gate_msa.squeeze(2) - c_shift_msa = c_shift_msa.squeeze(2) - c_scale_msa = c_scale_msa.squeeze(2) - c_gate_msa = c_gate_msa.squeeze(2) - else: - # temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B) - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table + temb.float() - ).chunk(6, dim=1) - - # 1. Self-attention - norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states) - attn_output = self.attn1(norm_hidden_states, None, None, rotary_emb) - hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states) - - # 2. Cross-attention - norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states) - attn_output = self.attn2(norm_hidden_states, encoder_hidden_states, None, None) - hidden_states = hidden_states + attn_output - - # 3. Feed-forward - norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as( - hidden_states - ) - ff_output = self.ffn(norm_hidden_states) - hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states) - - return hidden_states - - -class WanAnimateTransformer3DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin, AttentionMixin -): - r""" - A Transformer model for video-like data used in the WanAnimate model. - - Args: - patch_size (`tuple[int]`, defaults to `(1, 2, 2)`): - 3D patch dimensions for video embedding (t_patch, h_patch, w_patch). - num_attention_heads (`int`, defaults to `40`): - Fixed length for text embeddings. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - text_dim (`int`, defaults to `512`): - Input dimension for text embeddings. - freq_dim (`int`, defaults to `256`): - Dimension for sinusoidal time embeddings. - ffn_dim (`int`, defaults to `13824`): - Intermediate dimension in feed-forward network. - num_layers (`int`, defaults to `40`): - The number of layers of transformer blocks to use. - window_size (`tuple[int]`, defaults to `(-1, -1)`): - Window size for local attention (-1 indicates global attention). - cross_attn_norm (`bool`, defaults to `True`): - Enable cross-attention normalization. - qk_norm (`bool`, defaults to `True`): - Enable query/key normalization. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - image_dim (`int`, *optional*, defaults to `1280`): - The number of channels to use for the image embedding. If `None`, no projection is used. - added_kv_proj_dim (`int`, *optional*, defaults to `5120`): - The number of channels to use for the added key and value projections. If `None`, no projection is used. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"] - _no_split_modules = ["WanTransformerBlock", "MotionEncoderResBlock"] - _keep_in_fp32_modules = [ - "time_embedder", - "scale_shift_table", - "norm1", - "norm2", - "norm3", - "motion_synthesis_weight", - "rope", - ] - _keys_to_ignore_on_load_unexpected = ["norm_added_q"] - _repeated_blocks = ["WanTransformerBlock"] - - @register_to_config - def __init__( - self, - patch_size: tuple[int] = (1, 2, 2), - num_attention_heads: int = 40, - attention_head_dim: int = 128, - in_channels: int | None = 36, - latent_channels: int | None = 16, - out_channels: int | None = 16, - text_dim: int = 4096, - freq_dim: int = 256, - ffn_dim: int = 13824, - num_layers: int = 40, - cross_attn_norm: bool = True, - qk_norm: str | None = "rms_norm_across_heads", - eps: float = 1e-6, - image_dim: int | None = 1280, - added_kv_proj_dim: int | None = None, - rope_max_seq_len: int = 1024, - pos_embed_seq_len: int | None = None, - motion_encoder_channel_sizes: dict[str, int] | None = None, # Start of Wan Animate-specific args - motion_encoder_size: int = 512, - motion_style_dim: int = 512, - motion_dim: int = 20, - motion_encoder_dim: int = 512, - face_encoder_hidden_dim: int = 1024, - face_encoder_num_heads: int = 4, - inject_face_latents_blocks: int = 5, - motion_encoder_batch_size: int = 8, - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - # Allow either only in_channels or only latent_channels to be set for convenience - if in_channels is None and latent_channels is not None: - in_channels = 2 * latent_channels + 4 - elif in_channels is not None and latent_channels is None: - latent_channels = (in_channels - 4) // 2 - elif in_channels is not None and latent_channels is not None: - # TODO: should this always be true? - assert in_channels == 2 * latent_channels + 4, "in_channels should be 2 * latent_channels + 4" - else: - raise ValueError("At least one of `in_channels` and `latent_channels` must be supplied.") - out_channels = out_channels or latent_channels - - # 1. Patch & position embedding - self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len) - self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - self.pose_patch_embedding = nn.Conv3d(latent_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - - # 2. Condition embeddings - self.condition_embedder = WanTimeTextImageEmbedding( - dim=inner_dim, - time_freq_dim=freq_dim, - time_proj_dim=inner_dim * 6, - text_embed_dim=text_dim, - image_embed_dim=image_dim, - pos_embed_seq_len=pos_embed_seq_len, - ) - - # Motion encoder - self.motion_encoder = WanAnimateMotionEncoder( - size=motion_encoder_size, - style_dim=motion_style_dim, - motion_dim=motion_dim, - out_dim=motion_encoder_dim, - channels=motion_encoder_channel_sizes, - ) - - # Face encoder - self.face_encoder = WanAnimateFaceEncoder( - in_dim=motion_encoder_dim, - out_dim=inner_dim, - hidden_dim=face_encoder_hidden_dim, - num_heads=face_encoder_num_heads, - ) - - # 3. Transformer blocks - self.blocks = nn.ModuleList( - [ - WanTransformerBlock( - dim=inner_dim, - ffn_dim=ffn_dim, - num_heads=num_attention_heads, - qk_norm=qk_norm, - cross_attn_norm=cross_attn_norm, - eps=eps, - added_kv_proj_dim=added_kv_proj_dim, - ) - for _ in range(num_layers) - ] - ) - - self.face_adapter = nn.ModuleList( - [ - WanAnimateFaceBlockCrossAttention( - dim=inner_dim, - heads=num_attention_heads, - dim_head=inner_dim // num_attention_heads, - eps=eps, - cross_attention_dim_head=inner_dim // num_attention_heads, - processor=WanAnimateFaceBlockAttnProcessor(), - ) - for _ in range(num_layers // inject_face_latents_blocks) - ] - ) - - # 4. Output norm & projection - self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size)) - self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: torch.Tensor | None = None, - pose_hidden_states: torch.Tensor | None = None, - face_pixel_values: torch.Tensor | None = None, - motion_encode_batch_size: int | None = None, - return_dict: bool = True, - attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor | dict[str, torch.Tensor]: - """ - Forward pass of Wan2.2-Animate transformer model. - - Args: - hidden_states (`torch.Tensor` of shape `(B, 2C + 4, T + 1, H, W)`): - Input noisy video latents of shape `(B, 2C + 4, T + 1, H, W)`, where B is the batch size, C is the - number of latent channels (16 for Wan VAE), T is the number of latent frames in an inference segment, H - is the latent height, and W is the latent width. - timestep: (`torch.LongTensor`): - The current timestep in the denoising loop. - encoder_hidden_states (`torch.Tensor`): - Text embeddings from the text encoder (umT5 for Wan Animate). - encoder_hidden_states_image (`torch.Tensor`): - CLIP visual features of the reference (character) image. - pose_hidden_states (`torch.Tensor` of shape `(B, C, T, H, W)`): - Pose video latents. TODO: description - face_pixel_values (`torch.Tensor` of shape `(B, C', S, H', W')`): - Face video in pixel space (not latent space). Typically C' = 3 and H' and W' are the height/width of - the face video in pixels. Here S is the inference segment length, usually set to 77. - motion_encode_batch_size (`int`, *optional*): - The batch size for batched encoding of the face video via the motion encoder. Will default to - `self.config.motion_encoder_batch_size` if not set. - return_dict (`bool`, *optional*, defaults to `True`): - Whether to return the output as a dict or tuple. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - - Returns: - [`~models.transformer_2d.Transformer2DModelOutput`] or `tuple`: - If `return_dict` is True, a [`~models.transformer_2d.Transformer2DModelOutput`] whose `sample` is the - denoised video latent is returned, otherwise a plain `tuple` whose first element is that tensor is - returned. - """ - - # Check that shapes match up - if pose_hidden_states is not None and pose_hidden_states.shape[2] + 1 != hidden_states.shape[2]: - raise ValueError( - f"pose_hidden_states frame dim (dim 2) is {pose_hidden_states.shape[2]} but must be one less than the" - f" hidden_states's corresponding frame dim: {hidden_states.shape[2]}" - ) - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.config.patch_size - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p_h - post_patch_width = width // p_w - - # 1. Rotary position embedding - rotary_emb = self.rope(hidden_states) - - # 2. Patch embedding - hidden_states = self.patch_embedding(hidden_states) - pose_hidden_states = self.pose_patch_embedding(pose_hidden_states) - # Add pose embeddings to hidden states - hidden_states[:, :, 1:] = hidden_states[:, :, 1:] + pose_hidden_states - # Calling contiguous() here is important so that we don't recompile when performing regional compilation - hidden_states = hidden_states.flatten(2).transpose(1, 2).contiguous() - - # 3. Condition embeddings (time, text, image) - # Wan Animate is based on Wan 2.1 and thus uses Wan 2.1's timestep logic - temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder( - timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=None - ) - - # batch_size, 6, inner_dim - timestep_proj = timestep_proj.unflatten(1, (6, -1)) - - if encoder_hidden_states_image is not None: - encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) - - # 4. Get motion features from the face video - # Motion vector computation from face pixel values - batch_size, channels, num_face_frames, height, width = face_pixel_values.shape - # Rearrange from (B, C, T, H, W) to (B*T, C, H, W) - face_pixel_values = face_pixel_values.permute(0, 2, 1, 3, 4).reshape(-1, channels, height, width) - - # Extract motion features using motion encoder - # Perform batched motion encoder inference to allow trading off inference speed for memory usage - motion_encode_batch_size = motion_encode_batch_size or self.config.motion_encoder_batch_size - face_batches = torch.split(face_pixel_values, motion_encode_batch_size) - motion_vec_batches = [] - for face_batch in face_batches: - motion_vec_batch = self.motion_encoder(face_batch) - motion_vec_batches.append(motion_vec_batch) - motion_vec = torch.cat(motion_vec_batches) - motion_vec = motion_vec.view(batch_size, num_face_frames, -1) - - # Now get face features from the motion vector - motion_vec = self.face_encoder(motion_vec) - - # Add padding at the beginning (prepend zeros) - pad_face = torch.zeros_like(motion_vec[:, :1]) - motion_vec = torch.cat([pad_face, motion_vec], dim=1) - - # 5. Transformer blocks with face adapter integration - for block_idx, block in enumerate(self.blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - block, hidden_states, encoder_hidden_states, timestep_proj, rotary_emb - ) - else: - hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb) - - # Face adapter integration: apply after every 5th block (0, 5, 10, 15, ...) - if block_idx % self.config.inject_face_latents_blocks == 0: - face_adapter_block_idx = block_idx // self.config.inject_face_latents_blocks - face_adapter_output = self.face_adapter[face_adapter_block_idx](hidden_states, motion_vec) - # In case the face adapter and main transformer blocks are on different devices, which can happen when - # using model parallelism - face_adapter_output = face_adapter_output.to(device=hidden_states.device) - hidden_states = face_adapter_output + hidden_states - - # 6. Output norm, projection & unpatchify - # batch_size, inner_dim - shift, scale = (self.scale_shift_table.to(temb.device) + temb.unsqueeze(1)).chunk(2, dim=1) - - hidden_states_original_dtype = hidden_states.dtype - hidden_states = self.norm_out(hidden_states.float()) - # Move the shift and scale tensors to the same device as hidden_states. - # When using multi-GPU inference via accelerate these will be on the - # first device rather than the last device, which hidden_states ends up - # on. - shift = shift.to(hidden_states.device) - scale = scale.to(hidden_states.device) - hidden_states = (hidden_states * (1 + scale) + shift).to(dtype=hidden_states_original_dtype) - - hidden_states = self.proj_out(hidden_states) - - hidden_states = hidden_states.reshape( - batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1 - ) - hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) - output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_wan_vace.py b/diffusers/models/transformers/transformer_wan_vace.py deleted file mode 100644 index af40c7545d20e037f67237e7b44aff8eefc58792..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_wan_vace.py +++ /dev/null @@ -1,401 +0,0 @@ -# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...utils import apply_lora_scale, logging -from ..attention import AttentionMixin, FeedForward -from ..cache_utils import CacheMixin -from ..modeling_outputs import Transformer2DModelOutput -from ..modeling_utils import ModelMixin -from ..normalization import FP32LayerNorm -from .transformer_wan import ( - WanAttention, - WanAttnProcessor, - WanRotaryPosEmbed, - WanTimeTextImageEmbedding, - WanTransformerBlock, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class WanVACETransformerBlock(nn.Module): - def __init__( - self, - dim: int, - ffn_dim: int, - num_heads: int, - qk_norm: str = "rms_norm_across_heads", - cross_attn_norm: bool = False, - eps: float = 1e-6, - added_kv_proj_dim: int | None = None, - apply_input_projection: bool = False, - apply_output_projection: bool = False, - ): - super().__init__() - - # 1. Input projection - self.proj_in = None - if apply_input_projection: - self.proj_in = nn.Linear(dim, dim) - - # 2. Self-attention - self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False) - self.attn1 = WanAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - processor=WanAttnProcessor(), - ) - - # 3. Cross-attention - self.attn2 = WanAttention( - dim=dim, - heads=num_heads, - dim_head=dim // num_heads, - eps=eps, - added_kv_proj_dim=added_kv_proj_dim, - processor=WanAttnProcessor(), - is_cross_attention=True, - ) - self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity() - - # 4. Feed-forward - self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate") - self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False) - - # 5. Output projection - self.proj_out = None - if apply_output_projection: - self.proj_out = nn.Linear(dim, dim) - - self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor, - control_hidden_states: torch.Tensor, - temb: torch.Tensor, - rotary_emb: torch.Tensor, - ) -> torch.Tensor: - if self.proj_in is not None: - control_hidden_states = self.proj_in(control_hidden_states) - control_hidden_states = control_hidden_states + hidden_states - - shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( - self.scale_shift_table.to(temb.device) + temb.float() - ).chunk(6, dim=1) - - # 1. Self-attention - norm_hidden_states = (self.norm1(control_hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as( - control_hidden_states - ) - attn_output = self.attn1(norm_hidden_states, None, None, rotary_emb) - control_hidden_states = (control_hidden_states.float() + attn_output * gate_msa).type_as(control_hidden_states) - - # 2. Cross-attention - norm_hidden_states = self.norm2(control_hidden_states.float()).type_as(control_hidden_states) - attn_output = self.attn2(norm_hidden_states, encoder_hidden_states, None, None) - control_hidden_states = control_hidden_states + attn_output - - # 3. Feed-forward - norm_hidden_states = (self.norm3(control_hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as( - control_hidden_states - ) - ff_output = self.ffn(norm_hidden_states) - control_hidden_states = (control_hidden_states.float() + ff_output.float() * c_gate_msa).type_as( - control_hidden_states - ) - - conditioning_states = None - if self.proj_out is not None: - conditioning_states = self.proj_out(control_hidden_states) - - return conditioning_states, control_hidden_states - - -class WanVACETransformer3DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin, AttentionMixin -): - r""" - A Transformer model for video-like data used in the Wan model. - - Args: - patch_size (`tuple[int]`, defaults to `(1, 2, 2)`): - 3D patch dimensions for video embedding (t_patch, h_patch, w_patch). - num_attention_heads (`int`, defaults to `40`): - Fixed length for text embeddings. - attention_head_dim (`int`, defaults to `128`): - The number of channels in each head. - in_channels (`int`, defaults to `16`): - The number of channels in the input. - out_channels (`int`, defaults to `16`): - The number of channels in the output. - text_dim (`int`, defaults to `512`): - Input dimension for text embeddings. - freq_dim (`int`, defaults to `256`): - Dimension for sinusoidal time embeddings. - ffn_dim (`int`, defaults to `13824`): - Intermediate dimension in feed-forward network. - num_layers (`int`, defaults to `40`): - The number of layers of transformer blocks to use. - window_size (`tuple[int]`, defaults to `(-1, -1)`): - Window size for local attention (-1 indicates global attention). - cross_attn_norm (`bool`, defaults to `True`): - Enable cross-attention normalization. - qk_norm (`bool`, defaults to `True`): - Enable query/key normalization. - eps (`float`, defaults to `1e-6`): - Epsilon value for normalization layers. - add_img_emb (`bool`, defaults to `False`): - Whether to use img_emb. - added_kv_proj_dim (`int`, *optional*, defaults to `None`): - The number of channels to use for the added key and value projections. If `None`, no projection is used. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["patch_embedding", "vace_patch_embedding", "condition_embedder", "norm"] - _no_split_modules = ["WanTransformerBlock", "WanVACETransformerBlock"] - _keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"] - _keys_to_ignore_on_load_unexpected = ["norm_added_q"] - _repeated_blocks = ["WanTransformerBlock", "WanVACETransformerBlock"] - - @register_to_config - def __init__( - self, - patch_size: tuple[int, ...] = (1, 2, 2), - num_attention_heads: int = 40, - attention_head_dim: int = 128, - in_channels: int = 16, - out_channels: int = 16, - text_dim: int = 4096, - freq_dim: int = 256, - ffn_dim: int = 13824, - num_layers: int = 40, - cross_attn_norm: bool = True, - qk_norm: str | None = "rms_norm_across_heads", - eps: float = 1e-6, - image_dim: int | None = None, - added_kv_proj_dim: int | None = None, - rope_max_seq_len: int = 1024, - pos_embed_seq_len: int | None = None, - vace_layers: list[int] = [0, 5, 10, 15, 20, 25, 30, 35], - vace_in_channels: int = 96, - ) -> None: - super().__init__() - - inner_dim = num_attention_heads * attention_head_dim - out_channels = out_channels or in_channels - - if max(vace_layers) >= num_layers: - raise ValueError(f"VACE layers {vace_layers} exceed the number of transformer layers {num_layers}.") - if 0 not in vace_layers: - raise ValueError("VACE layers must include layer 0.") - - # 1. Patch & position embedding - self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len) - self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - self.vace_patch_embedding = nn.Conv3d(vace_in_channels, inner_dim, kernel_size=patch_size, stride=patch_size) - - # 2. Condition embeddings - # image_embedding_dim=1280 for I2V model - self.condition_embedder = WanTimeTextImageEmbedding( - dim=inner_dim, - time_freq_dim=freq_dim, - time_proj_dim=inner_dim * 6, - text_embed_dim=text_dim, - image_embed_dim=image_dim, - pos_embed_seq_len=pos_embed_seq_len, - ) - - # 3. Transformer blocks - self.blocks = nn.ModuleList( - [ - WanTransformerBlock( - inner_dim, ffn_dim, num_attention_heads, qk_norm, cross_attn_norm, eps, added_kv_proj_dim - ) - for _ in range(num_layers) - ] - ) - - self.vace_blocks = nn.ModuleList( - [ - WanVACETransformerBlock( - inner_dim, - ffn_dim, - num_attention_heads, - qk_norm, - cross_attn_norm, - eps, - added_kv_proj_dim, - apply_input_projection=i == 0, # Layer 0 always has input projection and is in vace_layers - apply_output_projection=True, - ) - for i in range(len(vace_layers)) - ] - ) - - # 4. Output norm & projection - self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False) - self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size)) - self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5) - - self.gradient_checkpointing = False - - @apply_lora_scale("attention_kwargs") - def forward( - self, - hidden_states: torch.Tensor, - timestep: torch.LongTensor, - encoder_hidden_states: torch.Tensor, - encoder_hidden_states_image: torch.Tensor | None = None, - control_hidden_states: torch.Tensor = None, - control_hidden_states_scale: torch.Tensor = None, - return_dict: bool = True, - attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor | dict[str, torch.Tensor]: - """ - The [`WanVACETransformer3DModel`] forward method. - - Args: - hidden_states (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - Input `hidden_states`. - timestep (`torch.LongTensor`): - Used to indicate denoising step. - encoder_hidden_states (`torch.Tensor` of shape `(batch_size, sequence_len, embed_dims)`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_hidden_states_image (`torch.Tensor`, *optional*): - Conditional image embeddings for image-conditioned generation. - control_hidden_states (`torch.Tensor`, *optional*): - Control latents used by the VACE control branch. - control_hidden_states_scale (`torch.Tensor`, *optional*): - Per-VACE-layer scale applied to the control hidden states. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - - Returns: - If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a - `tuple` where the first element is the sample tensor. - """ - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.config.patch_size - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p_h - post_patch_width = width // p_w - - if control_hidden_states_scale is None: - control_hidden_states_scale = control_hidden_states.new_ones(len(self.config.vace_layers)) - control_hidden_states_scale = torch.unbind(control_hidden_states_scale) - if len(control_hidden_states_scale) != len(self.config.vace_layers): - raise ValueError( - f"Length of `control_hidden_states_scale` {len(control_hidden_states_scale)} should be " - f"equal to {len(self.config.vace_layers)}." - ) - - # 1. Rotary position embedding - rotary_emb = self.rope(hidden_states) - - # 2. Patch embedding - hidden_states = self.patch_embedding(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) - - control_hidden_states = self.vace_patch_embedding(control_hidden_states) - control_hidden_states = control_hidden_states.flatten(2).transpose(1, 2) - control_hidden_states_padding = control_hidden_states.new_zeros( - batch_size, hidden_states.size(1) - control_hidden_states.size(1), control_hidden_states.size(2) - ) - control_hidden_states = torch.cat([control_hidden_states, control_hidden_states_padding], dim=1) - - # 3. Time embedding - temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder( - timestep, encoder_hidden_states, encoder_hidden_states_image - ) - timestep_proj = timestep_proj.unflatten(1, (6, -1)) - - # 4. Image embedding - if encoder_hidden_states_image is not None: - encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) - - # 5. Transformer blocks - if torch.is_grad_enabled() and self.gradient_checkpointing: - # Prepare VACE hints - control_hidden_states_list = [] - for i, block in enumerate(self.vace_blocks): - conditioning_states, control_hidden_states = self._gradient_checkpointing_func( - block, hidden_states, encoder_hidden_states, control_hidden_states, timestep_proj, rotary_emb - ) - control_hidden_states_list.append((conditioning_states, control_hidden_states_scale[i])) - control_hidden_states_list = control_hidden_states_list[::-1] - - for i, block in enumerate(self.blocks): - hidden_states = self._gradient_checkpointing_func( - block, hidden_states, encoder_hidden_states, timestep_proj, rotary_emb - ) - if i in self.config.vace_layers: - control_hint, scale = control_hidden_states_list.pop() - hidden_states = hidden_states + control_hint.to(hidden_states.device) * scale - else: - # Prepare VACE hints - control_hidden_states_list = [] - for i, block in enumerate(self.vace_blocks): - conditioning_states, control_hidden_states = block( - hidden_states, encoder_hidden_states, control_hidden_states, timestep_proj, rotary_emb - ) - control_hidden_states_list.append((conditioning_states, control_hidden_states_scale[i])) - control_hidden_states_list = control_hidden_states_list[::-1] - - for i, block in enumerate(self.blocks): - hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb) - if i in self.config.vace_layers: - control_hint, scale = control_hidden_states_list.pop() - hidden_states = hidden_states + control_hint.to(hidden_states.device) * scale - - # 6. Output norm, projection & unpatchify - shift, scale = (self.scale_shift_table.to(temb.device) + temb.unsqueeze(1)).chunk(2, dim=1) - - # Move the shift and scale tensors to the same device as hidden_states. - # When using multi-GPU inference via accelerate these will be on the - # first device rather than the last device, which hidden_states ends up - # on. - shift = shift.to(hidden_states.device) - scale = scale.to(hidden_states.device) - - hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states) - hidden_states = self.proj_out(hidden_states) - - hidden_states = hidden_states.reshape( - batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1 - ) - hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) - output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) diff --git a/diffusers/models/transformers/transformer_z_image.py b/diffusers/models/transformers/transformer_z_image.py deleted file mode 100644 index 4cea745e5ed5f36c8231109248752809c0f690f6..0000000000000000000000000000000000000000 --- a/diffusers/models/transformers/transformer_z_image.py +++ /dev/null @@ -1,1070 +0,0 @@ -# Copyright 2025 Alibaba Z-Image Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math - -import torch -import torch.nn as nn -import torch.nn.functional as F -from torch.nn.utils.rnn import pad_sequence - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin -from ...models.attention_processor import Attention -from ...models.modeling_utils import ModelMixin -from ...models.normalization import RMSNorm -from ...utils.torch_utils import maybe_allow_in_graph -from ..attention_dispatch import dispatch_attention_fn -from ..modeling_outputs import Transformer2DModelOutput - - -ADALN_EMBED_DIM = 256 -SEQ_MULTI_OF = 32 -X_PAD_DIM = 64 - - -class TimestepEmbedder(nn.Module): - def __init__(self, out_size, mid_size=None, frequency_embedding_size=256): - super().__init__() - if mid_size is None: - mid_size = out_size - self.mlp = nn.Sequential( - nn.Linear(frequency_embedding_size, mid_size, bias=True), - nn.SiLU(), - nn.Linear(mid_size, out_size, bias=True), - ) - - self.frequency_embedding_size = frequency_embedding_size - - @staticmethod - def timestep_embedding(t, dim, max_period=10000): - with torch.amp.autocast("cuda", enabled=False): - half = dim // 2 - freqs = torch.exp( - -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32, device=t.device) / half - ) - args = t[:, None].float() * freqs[None] - embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) - if dim % 2: - embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) - return embedding - - def forward(self, t): - t_freq = self.timestep_embedding(t, self.frequency_embedding_size) - weight_dtype = self.mlp[0].weight.dtype - compute_dtype = getattr(self.mlp[0], "compute_dtype", None) - if weight_dtype.is_floating_point: - t_freq = t_freq.to(weight_dtype) - elif compute_dtype is not None: - t_freq = t_freq.to(compute_dtype) - t_emb = self.mlp(t_freq) - return t_emb - - -class ZSingleStreamAttnProcessor: - """ - Processor for Z-Image single stream attention that adapts the existing Attention class to match the behavior of the - original Z-ImageAttention module. - """ - - _attention_backend = None - _parallel_config = None - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError( - "ZSingleStreamAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to version 2.0 or higher." - ) - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - freqs_cis: torch.Tensor | None = None, - ) -> torch.Tensor: - query = attn.to_q(hidden_states) - key = attn.to_k(hidden_states) - value = attn.to_v(hidden_states) - - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) - - # Apply Norms - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - # Apply RoPE - def apply_rotary_emb(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: - with torch.amp.autocast("cuda", enabled=False): - x = torch.view_as_complex(x_in.float().reshape(*x_in.shape[:-1], -1, 2)) - freqs_cis = freqs_cis.unsqueeze(2) - x_out = torch.view_as_real(x * freqs_cis).flatten(3) - return x_out.type_as(x_in) # todo - - if freqs_cis is not None: - query = apply_rotary_emb(query, freqs_cis) - key = apply_rotary_emb(key, freqs_cis) - - # Cast to correct dtype - dtype = query.dtype - query, key = query.to(dtype), key.to(dtype) - - # From [batch, seq_len] to [batch, 1, 1, seq_len] -> broadcast to [batch, heads, seq_len, seq_len] - if attention_mask is not None and attention_mask.ndim == 2: - attention_mask = attention_mask[:, None, None, :] - - # Compute joint attention - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - dropout_p=0.0, - is_causal=False, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - # Reshape back - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(dtype) - - output = attn.to_out[0](hidden_states) - if len(attn.to_out) > 1: # dropout - output = attn.to_out[1](output) - - return output - - -def select_per_token( - value_noisy: torch.Tensor, - value_clean: torch.Tensor, - noise_mask: torch.Tensor, - seq_len: int, -) -> torch.Tensor: - noise_mask_expanded = noise_mask.unsqueeze(-1) # (batch, seq_len, 1) - return torch.where( - noise_mask_expanded == 1, - value_noisy.unsqueeze(1).expand(-1, seq_len, -1), - value_clean.unsqueeze(1).expand(-1, seq_len, -1), - ) - - -class FeedForward(nn.Module): - def __init__(self, dim: int, hidden_dim: int): - super().__init__() - self.w1 = nn.Linear(dim, hidden_dim, bias=False) - self.w2 = nn.Linear(hidden_dim, dim, bias=False) - self.w3 = nn.Linear(dim, hidden_dim, bias=False) - - def _forward_silu_gating(self, x1, x3): - return F.silu(x1) * x3 - - def forward(self, x): - return self.w2(self._forward_silu_gating(self.w1(x), self.w3(x))) - - -@maybe_allow_in_graph -class ZImageTransformerBlock(nn.Module): - def __init__( - self, - layer_id: int, - dim: int, - n_heads: int, - n_kv_heads: int, - norm_eps: float, - qk_norm: bool, - modulation=True, - ): - super().__init__() - self.dim = dim - self.head_dim = dim // n_heads - - # Refactored to use diffusers Attention with custom processor - # Original Z-Image params: dim, n_heads, n_kv_heads, qk_norm - self.attention = Attention( - query_dim=dim, - cross_attention_dim=None, - dim_head=dim // n_heads, - heads=n_heads, - qk_norm="rms_norm" if qk_norm else None, - eps=1e-5, - bias=False, - out_bias=False, - processor=ZSingleStreamAttnProcessor(), - ) - - self.feed_forward = FeedForward(dim=dim, hidden_dim=int(dim / 3 * 8)) - self.layer_id = layer_id - - self.attention_norm1 = RMSNorm(dim, eps=norm_eps) - self.ffn_norm1 = RMSNorm(dim, eps=norm_eps) - - self.attention_norm2 = RMSNorm(dim, eps=norm_eps) - self.ffn_norm2 = RMSNorm(dim, eps=norm_eps) - - self.modulation = modulation - if modulation: - self.adaLN_modulation = nn.Sequential(nn.Linear(min(dim, ADALN_EMBED_DIM), 4 * dim, bias=True)) - - def forward( - self, - x: torch.Tensor, - attn_mask: torch.Tensor, - freqs_cis: torch.Tensor, - adaln_input: torch.Tensor | None = None, - noise_mask: torch.Tensor | None = None, - adaln_noisy: torch.Tensor | None = None, - adaln_clean: torch.Tensor | None = None, - ): - if self.modulation: - seq_len = x.shape[1] - - if noise_mask is not None: - # Per-token modulation: different modulation for noisy/clean tokens - mod_noisy = self.adaLN_modulation(adaln_noisy) - mod_clean = self.adaLN_modulation(adaln_clean) - - scale_msa_noisy, gate_msa_noisy, scale_mlp_noisy, gate_mlp_noisy = mod_noisy.chunk(4, dim=1) - scale_msa_clean, gate_msa_clean, scale_mlp_clean, gate_mlp_clean = mod_clean.chunk(4, dim=1) - - gate_msa_noisy, gate_mlp_noisy = gate_msa_noisy.tanh(), gate_mlp_noisy.tanh() - gate_msa_clean, gate_mlp_clean = gate_msa_clean.tanh(), gate_mlp_clean.tanh() - - scale_msa_noisy, scale_mlp_noisy = 1.0 + scale_msa_noisy, 1.0 + scale_mlp_noisy - scale_msa_clean, scale_mlp_clean = 1.0 + scale_msa_clean, 1.0 + scale_mlp_clean - - scale_msa = select_per_token(scale_msa_noisy, scale_msa_clean, noise_mask, seq_len) - scale_mlp = select_per_token(scale_mlp_noisy, scale_mlp_clean, noise_mask, seq_len) - gate_msa = select_per_token(gate_msa_noisy, gate_msa_clean, noise_mask, seq_len) - gate_mlp = select_per_token(gate_mlp_noisy, gate_mlp_clean, noise_mask, seq_len) - else: - # Global modulation: same modulation for all tokens (avoid double select) - mod = self.adaLN_modulation(adaln_input) - scale_msa, gate_msa, scale_mlp, gate_mlp = mod.unsqueeze(1).chunk(4, dim=2) - gate_msa, gate_mlp = gate_msa.tanh(), gate_mlp.tanh() - scale_msa, scale_mlp = 1.0 + scale_msa, 1.0 + scale_mlp - - # Attention block - attn_out = self.attention( - self.attention_norm1(x) * scale_msa, attention_mask=attn_mask, freqs_cis=freqs_cis - ) - x = x + gate_msa * self.attention_norm2(attn_out) - - # FFN block - x = x + gate_mlp * self.ffn_norm2(self.feed_forward(self.ffn_norm1(x) * scale_mlp)) - else: - # Attention block - attn_out = self.attention(self.attention_norm1(x), attention_mask=attn_mask, freqs_cis=freqs_cis) - x = x + self.attention_norm2(attn_out) - - # FFN block - x = x + self.ffn_norm2(self.feed_forward(self.ffn_norm1(x))) - - return x - - -class FinalLayer(nn.Module): - def __init__(self, hidden_size, out_channels): - super().__init__() - self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.linear = nn.Linear(hidden_size, out_channels, bias=True) - - self.adaLN_modulation = nn.Sequential( - nn.SiLU(), - nn.Linear(min(hidden_size, ADALN_EMBED_DIM), hidden_size, bias=True), - ) - - def forward(self, x, c=None, noise_mask=None, c_noisy=None, c_clean=None): - seq_len = x.shape[1] - - if noise_mask is not None: - # Per-token modulation - scale_noisy = 1.0 + self.adaLN_modulation(c_noisy) - scale_clean = 1.0 + self.adaLN_modulation(c_clean) - scale = select_per_token(scale_noisy, scale_clean, noise_mask, seq_len) - else: - # Original global modulation - assert c is not None, "Either c or (c_noisy, c_clean) must be provided" - scale = 1.0 + self.adaLN_modulation(c) - scale = scale.unsqueeze(1) - - x = self.norm_final(x) * scale - x = self.linear(x) - return x - - -class RopeEmbedder: - def __init__( - self, - theta: float = 256.0, - axes_dims: list[int] = (16, 56, 56), - axes_lens: list[int] = (64, 128, 128), - ): - self.theta = theta - self.axes_dims = axes_dims - self.axes_lens = axes_lens - assert len(axes_dims) == len(axes_lens), "axes_dims and axes_lens must have the same length" - self.freqs_cis = None - - @staticmethod - def precompute_freqs_cis(dim: list[int], end: list[int], theta: float = 256.0): - with torch.device("cpu"): - freqs_cis = [] - for i, (d, e) in enumerate(zip(dim, end)): - freqs = 1.0 / (theta ** (torch.arange(0, d, 2, dtype=torch.float64, device="cpu") / d)) - timestep = torch.arange(e, device=freqs.device, dtype=torch.float64) - freqs = torch.outer(timestep, freqs).float() - freqs_cis_i = torch.polar(torch.ones_like(freqs), freqs).to(torch.complex64) # complex64 - freqs_cis.append(freqs_cis_i) - - return freqs_cis - - def __call__(self, ids: torch.Tensor): - assert ids.ndim == 2 - assert ids.shape[-1] == len(self.axes_dims) - device = ids.device - - if self.freqs_cis is None: - self.freqs_cis = self.precompute_freqs_cis(self.axes_dims, self.axes_lens, theta=self.theta) - self.freqs_cis = [freqs_cis.to(device) for freqs_cis in self.freqs_cis] - else: - # Ensure freqs_cis are on the same device as ids - if self.freqs_cis[0].device != device: - self.freqs_cis = [freqs_cis.to(device) for freqs_cis in self.freqs_cis] - - result = [] - for i in range(len(self.axes_dims)): - index = ids[:, i] - result.append(self.freqs_cis[i][index]) - return torch.cat(result, dim=-1) - - -class ZImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): - _supports_gradient_checkpointing = True - _no_split_modules = ["ZImageTransformerBlock"] - _repeated_blocks = ["ZImageTransformerBlock"] - _skip_layerwise_casting_patterns = ["t_embedder", "cap_embedder"] # precision sensitive layers - - @register_to_config - def __init__( - self, - all_patch_size=(2,), - all_f_patch_size=(1,), - in_channels=16, - dim=3840, - n_layers=30, - n_refiner_layers=2, - n_heads=30, - n_kv_heads=30, - norm_eps=1e-5, - qk_norm=True, - cap_feat_dim=2560, - siglip_feat_dim=None, # Optional: set to enable SigLIP support for Omni - rope_theta=256.0, - t_scale=1000.0, - axes_dims=[32, 48, 48], - axes_lens=[1024, 512, 512], - ) -> None: - super().__init__() - self.in_channels = in_channels - self.out_channels = in_channels - self.all_patch_size = all_patch_size - self.all_f_patch_size = all_f_patch_size - self.dim = dim - self.n_heads = n_heads - - self.rope_theta = rope_theta - self.t_scale = t_scale - self.gradient_checkpointing = False - - assert len(all_patch_size) == len(all_f_patch_size) - - all_x_embedder = {} - all_final_layer = {} - for patch_idx, (patch_size, f_patch_size) in enumerate(zip(all_patch_size, all_f_patch_size)): - x_embedder = nn.Linear(f_patch_size * patch_size * patch_size * in_channels, dim, bias=True) - all_x_embedder[f"{patch_size}-{f_patch_size}"] = x_embedder - - final_layer = FinalLayer(dim, patch_size * patch_size * f_patch_size * self.out_channels) - all_final_layer[f"{patch_size}-{f_patch_size}"] = final_layer - - self.all_x_embedder = nn.ModuleDict(all_x_embedder) - self.all_final_layer = nn.ModuleDict(all_final_layer) - self.noise_refiner = nn.ModuleList( - [ - ZImageTransformerBlock( - 1000 + layer_id, - dim, - n_heads, - n_kv_heads, - norm_eps, - qk_norm, - modulation=True, - ) - for layer_id in range(n_refiner_layers) - ] - ) - self.context_refiner = nn.ModuleList( - [ - ZImageTransformerBlock( - layer_id, - dim, - n_heads, - n_kv_heads, - norm_eps, - qk_norm, - modulation=False, - ) - for layer_id in range(n_refiner_layers) - ] - ) - self.t_embedder = TimestepEmbedder(min(dim, ADALN_EMBED_DIM), mid_size=1024) - self.cap_embedder = nn.Sequential(RMSNorm(cap_feat_dim, eps=norm_eps), nn.Linear(cap_feat_dim, dim, bias=True)) - - # Optional SigLIP components (for Omni variant) - if siglip_feat_dim is not None: - self.siglip_embedder = nn.Sequential( - RMSNorm(siglip_feat_dim, eps=norm_eps), nn.Linear(siglip_feat_dim, dim, bias=True) - ) - self.siglip_refiner = nn.ModuleList( - [ - ZImageTransformerBlock( - 2000 + layer_id, - dim, - n_heads, - n_kv_heads, - norm_eps, - qk_norm, - modulation=False, - ) - for layer_id in range(n_refiner_layers) - ] - ) - self.siglip_pad_token = nn.Parameter(torch.zeros((1, dim))) - else: - self.siglip_embedder = None - self.siglip_refiner = None - self.siglip_pad_token = None - - self.x_pad_token = nn.Parameter(torch.zeros((1, dim))) - self.cap_pad_token = nn.Parameter(torch.zeros((1, dim))) - - self.layers = nn.ModuleList( - [ - ZImageTransformerBlock(layer_id, dim, n_heads, n_kv_heads, norm_eps, qk_norm) - for layer_id in range(n_layers) - ] - ) - head_dim = dim // n_heads - assert head_dim == sum(axes_dims) - self.axes_dims = axes_dims - self.axes_lens = axes_lens - - self.rope_embedder = RopeEmbedder(theta=rope_theta, axes_dims=axes_dims, axes_lens=axes_lens) - - def unpatchify( - self, - x: list[torch.Tensor], - size: list[tuple], - patch_size, - f_patch_size, - x_pos_offsets: list[tuple[int, int]] | None = None, - ) -> list[torch.Tensor]: - pH = pW = patch_size - pF = f_patch_size - bsz = len(x) - assert len(size) == bsz - - if x_pos_offsets is not None: - # Omni: extract target image from unified sequence (cond_images + target) - result = [] - for i in range(bsz): - unified_x = x[i][x_pos_offsets[i][0] : x_pos_offsets[i][1]] - cu_len = 0 - x_item = None - for j in range(len(size[i])): - if size[i][j] is None: - ori_len = 0 - pad_len = SEQ_MULTI_OF - cu_len += pad_len + ori_len - else: - F, H, W = size[i][j] - ori_len = (F // pF) * (H // pH) * (W // pW) - pad_len = (-ori_len) % SEQ_MULTI_OF - x_item = ( - unified_x[cu_len : cu_len + ori_len] - .view(F // pF, H // pH, W // pW, pF, pH, pW, self.out_channels) - .permute(6, 0, 3, 1, 4, 2, 5) - .reshape(self.out_channels, F, H, W) - ) - cu_len += ori_len + pad_len - result.append(x_item) # Return only the last (target) image - return result - else: - # Original mode: simple unpatchify - for i in range(bsz): - F, H, W = size[i] - ori_len = (F // pF) * (H // pH) * (W // pW) - # "f h w pf ph pw c -> c (f pf) (h ph) (w pw)" - x[i] = ( - x[i][:ori_len] - .view(F // pF, H // pH, W // pW, pF, pH, pW, self.out_channels) - .permute(6, 0, 3, 1, 4, 2, 5) - .reshape(self.out_channels, F, H, W) - ) - return x - - @staticmethod - def create_coordinate_grid(size, start=None, device=None): - if start is None: - start = (0 for _ in size) - axes = [torch.arange(x0, x0 + span, dtype=torch.int32, device=device) for x0, span in zip(start, size)] - grids = torch.meshgrid(axes, indexing="ij") - return torch.stack(grids, dim=-1) - - def _patchify_image(self, image: torch.Tensor, patch_size: int, f_patch_size: int): - """Patchify a single image tensor: (C, F, H, W) -> (num_patches, patch_dim).""" - pH, pW, pF = patch_size, patch_size, f_patch_size - C, F, H, W = image.size() - F_tokens, H_tokens, W_tokens = F // pF, H // pH, W // pW - image = image.view(C, F_tokens, pF, H_tokens, pH, W_tokens, pW) - image = image.permute(1, 3, 5, 2, 4, 6, 0).reshape(F_tokens * H_tokens * W_tokens, pF * pH * pW * C) - return image, (F, H, W), (F_tokens, H_tokens, W_tokens) - - def _pad_with_ids( - self, - feat: torch.Tensor, - pos_grid_size: tuple, - pos_start: tuple, - device: torch.device, - noise_mask_val: int | None = None, - ): - """Pad feature to SEQ_MULTI_OF, create position IDs and pad mask.""" - ori_len = len(feat) - pad_len = (-ori_len) % SEQ_MULTI_OF - total_len = ori_len + pad_len - - # Pos IDs - ori_pos_ids = self.create_coordinate_grid(size=pos_grid_size, start=pos_start, device=device).flatten(0, 2) - if pad_len > 0: - pad_pos_ids = ( - self.create_coordinate_grid(size=(1, 1, 1), start=(0, 0, 0), device=device) - .flatten(0, 2) - .repeat(pad_len, 1) - ) - pos_ids = torch.cat([ori_pos_ids, pad_pos_ids], dim=0) - padded_feat = torch.cat([feat, feat[-1:].repeat(pad_len, 1)], dim=0) - pad_mask = torch.cat( - [ - torch.zeros(ori_len, dtype=torch.bool, device=device), - torch.ones(pad_len, dtype=torch.bool, device=device), - ] - ) - else: - pos_ids = ori_pos_ids - padded_feat = feat - pad_mask = torch.zeros(ori_len, dtype=torch.bool, device=device) - - noise_mask = [noise_mask_val] * total_len if noise_mask_val is not None else None # token level - return padded_feat, pos_ids, pad_mask, total_len, noise_mask - - def patchify_and_embed( - self, all_image: list[torch.Tensor], all_cap_feats: list[torch.Tensor], patch_size: int, f_patch_size: int - ): - """Patchify for basic mode: single image per batch item.""" - device = all_image[0].device - all_img_out, all_img_size, all_img_pos_ids, all_img_pad_mask = [], [], [], [] - all_cap_out, all_cap_pos_ids, all_cap_pad_mask = [], [], [] - - for image, cap_feat in zip(all_image, all_cap_feats): - # Caption - cap_out, cap_pos_ids, cap_pad_mask, cap_len, _ = self._pad_with_ids( - cap_feat, (len(cap_feat) + (-len(cap_feat)) % SEQ_MULTI_OF, 1, 1), (1, 0, 0), device - ) - all_cap_out.append(cap_out) - all_cap_pos_ids.append(cap_pos_ids) - all_cap_pad_mask.append(cap_pad_mask) - - # Image - img_patches, size, (F_t, H_t, W_t) = self._patchify_image(image, patch_size, f_patch_size) - img_out, img_pos_ids, img_pad_mask, _, _ = self._pad_with_ids( - img_patches, (F_t, H_t, W_t), (cap_len + 1, 0, 0), device - ) - all_img_out.append(img_out) - all_img_size.append(size) - all_img_pos_ids.append(img_pos_ids) - all_img_pad_mask.append(img_pad_mask) - - return ( - all_img_out, - all_cap_out, - all_img_size, - all_img_pos_ids, - all_cap_pos_ids, - all_img_pad_mask, - all_cap_pad_mask, - ) - - def patchify_and_embed_omni( - self, - all_x: list[list[torch.Tensor]], - all_cap_feats: list[list[torch.Tensor]], - all_siglip_feats: list[list[torch.Tensor]], - patch_size: int, - f_patch_size: int, - images_noise_mask: list[list[int]], - ): - """Patchify for omni mode: multiple images per batch item with noise masks.""" - bsz = len(all_x) - device = all_x[0][-1].device - dtype = all_x[0][-1].dtype - - all_x_out, all_x_size, all_x_pos_ids, all_x_pad_mask, all_x_len, all_x_noise_mask = [], [], [], [], [], [] - all_cap_out, all_cap_pos_ids, all_cap_pad_mask, all_cap_len, all_cap_noise_mask = [], [], [], [], [] - all_sig_out, all_sig_pos_ids, all_sig_pad_mask, all_sig_len, all_sig_noise_mask = [], [], [], [], [] - - for i in range(bsz): - num_images = len(all_x[i]) - cap_feats_list, cap_pos_list, cap_mask_list, cap_lens, cap_noise = [], [], [], [], [] - cap_end_pos = [] - cap_cu_len = 1 - - # Process captions - for j, cap_item in enumerate(all_cap_feats[i]): - noise_val = images_noise_mask[i][j] if j < len(images_noise_mask[i]) else 1 - cap_out, cap_pos, cap_mask, cap_len, cap_nm = self._pad_with_ids( - cap_item, - (len(cap_item) + (-len(cap_item)) % SEQ_MULTI_OF, 1, 1), - (cap_cu_len, 0, 0), - device, - noise_val, - ) - cap_feats_list.append(cap_out) - cap_pos_list.append(cap_pos) - cap_mask_list.append(cap_mask) - cap_lens.append(cap_len) - cap_noise.extend(cap_nm) - cap_cu_len += len(cap_item) - cap_end_pos.append(cap_cu_len) - cap_cu_len += 2 # for image vae and siglip tokens - - all_cap_out.append(torch.cat(cap_feats_list, dim=0)) - all_cap_pos_ids.append(torch.cat(cap_pos_list, dim=0)) - all_cap_pad_mask.append(torch.cat(cap_mask_list, dim=0)) - all_cap_len.append(cap_lens) - all_cap_noise_mask.append(cap_noise) - - # Process images - x_feats_list, x_pos_list, x_mask_list, x_lens, x_size, x_noise = [], [], [], [], [], [] - for j, x_item in enumerate(all_x[i]): - noise_val = images_noise_mask[i][j] - if x_item is not None: - x_patches, size, (F_t, H_t, W_t) = self._patchify_image(x_item, patch_size, f_patch_size) - x_out, x_pos, x_mask, x_len, x_nm = self._pad_with_ids( - x_patches, (F_t, H_t, W_t), (cap_end_pos[j], 0, 0), device, noise_val - ) - x_size.append(size) - else: - x_len = SEQ_MULTI_OF - x_out = torch.zeros((x_len, X_PAD_DIM), dtype=dtype, device=device) - x_pos = self.create_coordinate_grid((1, 1, 1), (0, 0, 0), device).flatten(0, 2).repeat(x_len, 1) - x_mask = torch.ones(x_len, dtype=torch.bool, device=device) - x_nm = [noise_val] * x_len - x_size.append(None) - x_feats_list.append(x_out) - x_pos_list.append(x_pos) - x_mask_list.append(x_mask) - x_lens.append(x_len) - x_noise.extend(x_nm) - - all_x_out.append(torch.cat(x_feats_list, dim=0)) - all_x_pos_ids.append(torch.cat(x_pos_list, dim=0)) - all_x_pad_mask.append(torch.cat(x_mask_list, dim=0)) - all_x_size.append(x_size) - all_x_len.append(x_lens) - all_x_noise_mask.append(x_noise) - - # Process siglip - if all_siglip_feats[i] is None: - all_sig_len.append([0] * num_images) - all_sig_out.append(None) - else: - sig_feats_list, sig_pos_list, sig_mask_list, sig_lens, sig_noise = [], [], [], [], [] - for j, sig_item in enumerate(all_siglip_feats[i]): - noise_val = images_noise_mask[i][j] - if sig_item is not None: - sig_H, sig_W, sig_C = sig_item.size() - sig_flat = sig_item.permute(2, 0, 1).reshape(sig_H * sig_W, sig_C) - sig_out, sig_pos, sig_mask, sig_len, sig_nm = self._pad_with_ids( - sig_flat, (1, sig_H, sig_W), (cap_end_pos[j] + 1, 0, 0), device, noise_val - ) - # Scale position IDs to match x resolution - if x_size[j] is not None: - sig_pos = sig_pos.float() - sig_pos[..., 1] = sig_pos[..., 1] / max(sig_H - 1, 1) * (x_size[j][1] - 1) - sig_pos[..., 2] = sig_pos[..., 2] / max(sig_W - 1, 1) * (x_size[j][2] - 1) - sig_pos = sig_pos.to(torch.int32) - else: - sig_len = SEQ_MULTI_OF - sig_out = torch.zeros((sig_len, self.config.siglip_feat_dim), dtype=dtype, device=device) - sig_pos = ( - self.create_coordinate_grid((1, 1, 1), (0, 0, 0), device).flatten(0, 2).repeat(sig_len, 1) - ) - sig_mask = torch.ones(sig_len, dtype=torch.bool, device=device) - sig_nm = [noise_val] * sig_len - sig_feats_list.append(sig_out) - sig_pos_list.append(sig_pos) - sig_mask_list.append(sig_mask) - sig_lens.append(sig_len) - sig_noise.extend(sig_nm) - - all_sig_out.append(torch.cat(sig_feats_list, dim=0)) - all_sig_pos_ids.append(torch.cat(sig_pos_list, dim=0)) - all_sig_pad_mask.append(torch.cat(sig_mask_list, dim=0)) - all_sig_len.append(sig_lens) - all_sig_noise_mask.append(sig_noise) - - # Compute x position offsets - all_x_pos_offsets = [(sum(all_cap_len[i]), sum(all_cap_len[i]) + sum(all_x_len[i])) for i in range(bsz)] - - return ( - all_x_out, - all_cap_out, - all_sig_out, - all_x_size, - all_x_pos_ids, - all_cap_pos_ids, - all_sig_pos_ids, - all_x_pad_mask, - all_cap_pad_mask, - all_sig_pad_mask, - all_x_pos_offsets, - all_x_noise_mask, - all_cap_noise_mask, - all_sig_noise_mask, - ) - - def _prepare_sequence( - self, - feats: list[torch.Tensor], - pos_ids: list[torch.Tensor], - inner_pad_mask: list[torch.Tensor], - pad_token: torch.nn.Parameter, - noise_mask: list[list[int]] | None = None, - device: torch.device = None, - ): - """Prepare sequence: apply pad token, RoPE embed, pad to batch, create attention mask.""" - item_seqlens = [len(f) for f in feats] - max_seqlen = max(item_seqlens) - bsz = len(feats) - - # Pad token - feats_cat = torch.cat(feats, dim=0) - mask = torch.cat(inner_pad_mask).unsqueeze(-1) - feats_cat = torch.where(mask, pad_token, feats_cat) - feats = list(feats_cat.split(item_seqlens, dim=0)) - - # RoPE - freqs_cis = list(self.rope_embedder(torch.cat(pos_ids, dim=0)).split([len(p) for p in pos_ids], dim=0)) - - # Pad to batch - feats = pad_sequence(feats, batch_first=True, padding_value=0.0) - freqs_cis = pad_sequence(freqs_cis, batch_first=True, padding_value=0.0)[:, : feats.shape[1]] - - # Attention mask - if all(seq == max_seqlen for seq in item_seqlens): - attn_mask = None - else: - attn_mask = torch.zeros((bsz, max_seqlen), dtype=torch.bool, device=device) - for i, seq_len in enumerate(item_seqlens): - attn_mask[i, :seq_len] = 1 - - # Noise mask - noise_mask_tensor = None - if noise_mask is not None: - noise_mask_tensor = pad_sequence( - [torch.tensor(m, dtype=torch.long, device=device) for m in noise_mask], - batch_first=True, - padding_value=0, - )[:, : feats.shape[1]] - - return feats, freqs_cis, attn_mask, item_seqlens, noise_mask_tensor - - def _build_unified_sequence( - self, - x: torch.Tensor, - x_freqs: torch.Tensor, - x_seqlens: list[int], - x_noise_mask: list[list[int]] | None, - cap: torch.Tensor, - cap_freqs: torch.Tensor, - cap_seqlens: list[int], - cap_noise_mask: list[list[int]] | None, - siglip: torch.Tensor | None, - siglip_freqs: torch.Tensor | None, - siglip_seqlens: list[int] | None, - siglip_noise_mask: list[list[int]] | None, - omni_mode: bool, - device: torch.device, - ): - """Build unified sequence: x, cap, and optionally siglip. - Basic mode order: [x, cap]; Omni mode order: [cap, x, siglip] - """ - bsz = len(x_seqlens) - unified = [] - unified_freqs = [] - unified_noise_mask = [] - - for i in range(bsz): - x_len, cap_len = x_seqlens[i], cap_seqlens[i] - - if omni_mode: - # Omni: [cap, x, siglip] - if siglip is not None and siglip_seqlens is not None: - sig_len = siglip_seqlens[i] - unified.append(torch.cat([cap[i][:cap_len], x[i][:x_len], siglip[i][:sig_len]])) - unified_freqs.append( - torch.cat([cap_freqs[i][:cap_len], x_freqs[i][:x_len], siglip_freqs[i][:sig_len]]) - ) - unified_noise_mask.append( - torch.tensor( - cap_noise_mask[i] + x_noise_mask[i] + siglip_noise_mask[i], dtype=torch.long, device=device - ) - ) - else: - unified.append(torch.cat([cap[i][:cap_len], x[i][:x_len]])) - unified_freqs.append(torch.cat([cap_freqs[i][:cap_len], x_freqs[i][:x_len]])) - unified_noise_mask.append( - torch.tensor(cap_noise_mask[i] + x_noise_mask[i], dtype=torch.long, device=device) - ) - else: - # Basic: [x, cap] - unified.append(torch.cat([x[i][:x_len], cap[i][:cap_len]])) - unified_freqs.append(torch.cat([x_freqs[i][:x_len], cap_freqs[i][:cap_len]])) - - # Compute unified seqlens - if omni_mode: - if siglip is not None and siglip_seqlens is not None: - unified_seqlens = [a + b + c for a, b, c in zip(cap_seqlens, x_seqlens, siglip_seqlens)] - else: - unified_seqlens = [a + b for a, b in zip(cap_seqlens, x_seqlens)] - else: - unified_seqlens = [a + b for a, b in zip(x_seqlens, cap_seqlens)] - - max_seqlen = max(unified_seqlens) - - # Pad to batch - unified = pad_sequence(unified, batch_first=True, padding_value=0.0) - unified_freqs = pad_sequence(unified_freqs, batch_first=True, padding_value=0.0) - - # Attention mask - if all(seq == max_seqlen for seq in unified_seqlens): - attn_mask = None - else: - attn_mask = torch.zeros((bsz, max_seqlen), dtype=torch.bool, device=device) - for i, seq_len in enumerate(unified_seqlens): - attn_mask[i, :seq_len] = 1 - - # Noise mask - noise_mask_tensor = None - if omni_mode: - noise_mask_tensor = pad_sequence(unified_noise_mask, batch_first=True, padding_value=0)[ - :, : unified.shape[1] - ] - - return unified, unified_freqs, attn_mask, noise_mask_tensor - - def forward( - self, - x: list[torch.Tensor, list[list[torch.Tensor]]], - t, - cap_feats: list[torch.Tensor, list[list[torch.Tensor]]], - return_dict: bool = True, - controlnet_block_samples: dict[int, torch.Tensor] | None = None, - siglip_feats: list[list[torch.Tensor]] | None = None, - image_noise_mask: list[list[int]] | None = None, - patch_size: int = 2, - f_patch_size: int = 1, - ): - """ - The [`ZImageTransformer2DModel`] forward method. - - Flow: patchify -> t_embed -> x_embed -> x_refine -> cap_embed -> cap_refine - -> [siglip_embed -> siglip_refine] -> build_unified -> main_layers -> final_layer -> unpatchify - - Args: - x (`list` of `torch.Tensor` or nested `list` of `torch.Tensor`): - Input latents. A flat list when running in standard mode, or a nested list when running in omni mode. - t (`torch.Tensor`): - Used to indicate denoising step. - cap_feats (`list` of `torch.Tensor` or nested `list` of `torch.Tensor`): - Conditional caption embeddings (embeddings computed from the input conditions such as prompts) to use. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain - tuple. - controlnet_block_samples (`dict` of `int` to `torch.Tensor`, *optional*): - A mapping from block index to tensor that if specified are added to the residuals of transformer - blocks. - siglip_feats (`list` of `list` of `torch.Tensor`, *optional*): - Optional SigLIP image features used as additional conditioning. - image_noise_mask (`list` of `list` of `int`, *optional*): - Per-image noise masks indicating noisy vs. clean tokens in omni mode. - patch_size (`int`, *optional*, defaults to 2): - Spatial patch size used to patchify the input latents. - f_patch_size (`int`, *optional*, defaults to 1): - Temporal patch size used to patchify the input latents. - """ - assert patch_size in self.all_patch_size and f_patch_size in self.all_f_patch_size - omni_mode = isinstance(x[0], list) - device = x[0][-1].device if omni_mode else x[0].device - - if omni_mode: - # Dual embeddings: noisy (t) and clean (t=1) - t_noisy = self.t_embedder(t * self.t_scale).type_as(x[0][-1]) - t_clean = self.t_embedder(torch.ones_like(t) * self.t_scale).type_as(x[0][-1]) - adaln_input = None - else: - # Single embedding for all tokens - adaln_input = self.t_embedder(t * self.t_scale).type_as(x[0]) - t_noisy = t_clean = None - - # Patchify - if omni_mode: - ( - x, - cap_feats, - siglip_feats, - x_size, - x_pos_ids, - cap_pos_ids, - siglip_pos_ids, - x_pad_mask, - cap_pad_mask, - siglip_pad_mask, - x_pos_offsets, - x_noise_mask, - cap_noise_mask, - siglip_noise_mask, - ) = self.patchify_and_embed_omni(x, cap_feats, siglip_feats, patch_size, f_patch_size, image_noise_mask) - else: - ( - x, - cap_feats, - x_size, - x_pos_ids, - cap_pos_ids, - x_pad_mask, - cap_pad_mask, - ) = self.patchify_and_embed(x, cap_feats, patch_size, f_patch_size) - x_pos_offsets = x_noise_mask = cap_noise_mask = siglip_noise_mask = None - - # X embed & refine - x_seqlens = [len(xi) for xi in x] - x = self.all_x_embedder[f"{patch_size}-{f_patch_size}"](torch.cat(x, dim=0)) # embed - x, x_freqs, x_mask, _, x_noise_tensor = self._prepare_sequence( - list(x.split(x_seqlens, dim=0)), x_pos_ids, x_pad_mask, self.x_pad_token, x_noise_mask, device - ) - - for layer in self.noise_refiner: - x = ( - self._gradient_checkpointing_func( - layer, x, x_mask, x_freqs, adaln_input, x_noise_tensor, t_noisy, t_clean - ) - if torch.is_grad_enabled() and self.gradient_checkpointing - else layer(x, x_mask, x_freqs, adaln_input, x_noise_tensor, t_noisy, t_clean) - ) - - # Cap embed & refine - cap_seqlens = [len(ci) for ci in cap_feats] - cap_feats = self.cap_embedder(torch.cat(cap_feats, dim=0)) # embed - cap_feats, cap_freqs, cap_mask, _, _ = self._prepare_sequence( - list(cap_feats.split(cap_seqlens, dim=0)), cap_pos_ids, cap_pad_mask, self.cap_pad_token, None, device - ) - - for layer in self.context_refiner: - cap_feats = ( - self._gradient_checkpointing_func(layer, cap_feats, cap_mask, cap_freqs) - if torch.is_grad_enabled() and self.gradient_checkpointing - else layer(cap_feats, cap_mask, cap_freqs) - ) - - # Siglip embed & refine - siglip_seqlens = siglip_freqs = None - if omni_mode and siglip_feats[0] is not None and self.siglip_embedder is not None: - siglip_seqlens = [len(si) for si in siglip_feats] - siglip_feats = self.siglip_embedder(torch.cat(siglip_feats, dim=0)) # embed - siglip_feats, siglip_freqs, siglip_mask, _, _ = self._prepare_sequence( - list(siglip_feats.split(siglip_seqlens, dim=0)), - siglip_pos_ids, - siglip_pad_mask, - self.siglip_pad_token, - None, - device, - ) - - for layer in self.siglip_refiner: - siglip_feats = ( - self._gradient_checkpointing_func(layer, siglip_feats, siglip_mask, siglip_freqs) - if torch.is_grad_enabled() and self.gradient_checkpointing - else layer(siglip_feats, siglip_mask, siglip_freqs) - ) - - # Unified sequence - unified, unified_freqs, unified_mask, unified_noise_tensor = self._build_unified_sequence( - x, - x_freqs, - x_seqlens, - x_noise_mask, - cap_feats, - cap_freqs, - cap_seqlens, - cap_noise_mask, - siglip_feats, - siglip_freqs, - siglip_seqlens, - siglip_noise_mask, - omni_mode, - device, - ) - - # Main transformer layers - for layer_idx, layer in enumerate(self.layers): - unified = ( - self._gradient_checkpointing_func( - layer, unified, unified_mask, unified_freqs, adaln_input, unified_noise_tensor, t_noisy, t_clean - ) - if torch.is_grad_enabled() and self.gradient_checkpointing - else layer(unified, unified_mask, unified_freqs, adaln_input, unified_noise_tensor, t_noisy, t_clean) - ) - if controlnet_block_samples is not None and layer_idx in controlnet_block_samples: - unified = unified + controlnet_block_samples[layer_idx] - - unified = ( - self.all_final_layer[f"{patch_size}-{f_patch_size}"]( - unified, noise_mask=unified_noise_tensor, c_noisy=t_noisy, c_clean=t_clean - ) - if omni_mode - else self.all_final_layer[f"{patch_size}-{f_patch_size}"](unified, c=adaln_input) - ) - - # Unpatchify - x = self.unpatchify(list(unified.unbind(dim=0)), x_size, patch_size, f_patch_size, x_pos_offsets) - - return (x,) if not return_dict else Transformer2DModelOutput(sample=x) diff --git a/diffusers/models/unets/__init__.py b/diffusers/models/unets/__init__.py deleted file mode 100644 index d3b69d6d5e8c6cdc7f5f45da486090ec841a72ce..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/__init__.py +++ /dev/null @@ -1,15 +0,0 @@ -from ...utils import is_torch_available - - -if is_torch_available(): - from .unet_1d import UNet1DModel - from .unet_2d import UNet2DModel - from .unet_2d_condition import UNet2DConditionModel - from .unet_3d_condition import UNet3DConditionModel - from .unet_dreamlite import DreamLiteUNetModel - from .unet_i2vgen_xl import I2VGenXLUNet - from .unet_kandinsky3 import Kandinsky3UNet - from .unet_motion_model import MotionAdapter, UNetMotionModel - from .unet_spatio_temporal_condition import UNetSpatioTemporalConditionModel - from .unet_stable_cascade import StableCascadeUNet - from .uvit_2d import UVit2DModel diff --git a/diffusers/models/unets/unet_1d.py b/diffusers/models/unets/unet_1d.py deleted file mode 100644 index 959e82e9d7cd42badc5150bdc080618859fb877e..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_1d.py +++ /dev/null @@ -1,265 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput -from ..embeddings import GaussianFourierProjection, TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from .unet_1d_blocks import get_down_block, get_mid_block, get_out_block, get_up_block - - -@dataclass -class UNet1DOutput(BaseOutput): - """ - The output of [`UNet1DModel`]. - - Args: - sample (`torch.Tensor` of shape `(batch_size, num_channels, sample_size)`): - The hidden states output from the last layer of the model. - """ - - sample: torch.Tensor - - -class UNet1DModel(ModelMixin, ConfigMixin): - r""" - A 1D UNet model that takes a noisy sample and a timestep and returns a sample shaped output. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - sample_size (`int`, *optional*): Default length of sample. Should be adaptable at runtime. - in_channels (`int`, *optional*, defaults to 2): Number of channels in the input sample. - out_channels (`int`, *optional*, defaults to 2): Number of channels in the output. - extra_in_channels (`int`, *optional*, defaults to 0): - Number of additional channels to be added to the input of the first down block. Useful for cases where the - input data has more channels than what the model was initially designed for. - time_embedding_type (`str`, *optional*, defaults to `"fourier"`): Type of time embedding to use. - freq_shift (`float`, *optional*, defaults to 0.0): Frequency shift for Fourier time embedding. - flip_sin_to_cos (`bool`, *optional*, defaults to `False`): - Whether to flip sin to cos for Fourier time embedding. - down_block_types (`tuple[str]`, *optional*, defaults to `("DownBlock1DNoSkip", "DownBlock1D", "AttnDownBlock1D")`): - tuple of downsample block types. - up_block_types (`tuple[str]`, *optional*, defaults to `("AttnUpBlock1D", "UpBlock1D", "UpBlock1DNoSkip")`): - tuple of upsample block types. - block_out_channels (`tuple[int]`, *optional*, defaults to `(32, 32, 64)`): - tuple of block output channels. - mid_block_type (`str`, *optional*, defaults to `"UNetMidBlock1D"`): Block type for middle of UNet. - out_block_type (`str`, *optional*, defaults to `None`): Optional output processing block of UNet. - act_fn (`str`, *optional*, defaults to `None`): Optional activation function in UNet blocks. - norm_num_groups (`int`, *optional*, defaults to 8): The number of groups for normalization. - layers_per_block (`int`, *optional*, defaults to 1): The number of layers per block. - downsample_each_block (`int`, *optional*, defaults to `False`): - Experimental feature for using a UNet without upsampling. - """ - - _skip_layerwise_casting_patterns = ["norm"] - - @register_to_config - def __init__( - self, - sample_size: int = 65536, - sample_rate: int | None = None, - in_channels: int = 2, - out_channels: int = 2, - extra_in_channels: int = 0, - time_embedding_type: str = "fourier", - time_embedding_dim: int | None = None, - flip_sin_to_cos: bool = True, - use_timestep_embedding: bool = False, - freq_shift: float = 0.0, - down_block_types: tuple[str, ...] = ("DownBlock1DNoSkip", "DownBlock1D", "AttnDownBlock1D"), - up_block_types: tuple[str, ...] = ("AttnUpBlock1D", "UpBlock1D", "UpBlock1DNoSkip"), - mid_block_type: str = "UNetMidBlock1D", - out_block_type: str = None, - block_out_channels: tuple[int, ...] = (32, 32, 64), - act_fn: str = None, - norm_num_groups: int = 8, - layers_per_block: int = 1, - downsample_each_block: bool = False, - ): - super().__init__() - self.sample_size = sample_size - - # time - if time_embedding_type == "fourier": - time_embed_dim = time_embedding_dim or block_out_channels[0] * 2 - if time_embed_dim % 2 != 0: - raise ValueError(f"`time_embed_dim` should be divisible by 2, but is {time_embed_dim}.") - self.time_proj = GaussianFourierProjection( - embedding_size=time_embed_dim // 2, set_W_to_weight=False, log=False, flip_sin_to_cos=flip_sin_to_cos - ) - timestep_input_dim = time_embed_dim - elif time_embedding_type == "positional": - time_embed_dim = time_embedding_dim or block_out_channels[0] * 4 - self.time_proj = Timesteps( - block_out_channels[0], flip_sin_to_cos=flip_sin_to_cos, downscale_freq_shift=freq_shift - ) - timestep_input_dim = block_out_channels[0] - else: - raise ValueError( - f"{time_embedding_type} does not exist. Please make sure to use one of `fourier` or `positional`." - ) - - if use_timestep_embedding: - time_embed_dim = block_out_channels[0] * 4 - self.time_mlp = TimestepEmbedding( - in_channels=timestep_input_dim, - time_embed_dim=time_embed_dim, - act_fn=act_fn, - out_dim=block_out_channels[0], - ) - - self.down_blocks = nn.ModuleList([]) - self.mid_block = None - self.up_blocks = nn.ModuleList([]) - self.out_block = None - - # down - output_channel = in_channels - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - - if i == 0: - input_channel += extra_in_channels - - is_final_block = i == len(block_out_channels) - 1 - - down_block = get_down_block( - down_block_type, - num_layers=layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - temb_channels=block_out_channels[0], - add_downsample=not is_final_block or downsample_each_block, - ) - self.down_blocks.append(down_block) - - # mid - self.mid_block = get_mid_block( - mid_block_type, - in_channels=block_out_channels[-1], - mid_channels=block_out_channels[-1], - out_channels=block_out_channels[-1], - embed_dim=block_out_channels[0], - num_layers=layers_per_block, - add_downsample=downsample_each_block, - ) - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - output_channel = reversed_block_out_channels[0] - if out_block_type is None: - final_upsample_channels = out_channels - else: - final_upsample_channels = block_out_channels[0] - - for i, up_block_type in enumerate(up_block_types): - prev_output_channel = output_channel - output_channel = ( - reversed_block_out_channels[i + 1] if i < len(up_block_types) - 1 else final_upsample_channels - ) - - is_final_block = i == len(block_out_channels) - 1 - - up_block = get_up_block( - up_block_type, - num_layers=layers_per_block, - in_channels=prev_output_channel, - out_channels=output_channel, - temb_channels=block_out_channels[0], - add_upsample=not is_final_block, - ) - self.up_blocks.append(up_block) - prev_output_channel = output_channel - - # out - num_groups_out = norm_num_groups if norm_num_groups is not None else min(block_out_channels[0] // 4, 32) - self.out_block = get_out_block( - out_block_type=out_block_type, - num_groups_out=num_groups_out, - embed_dim=block_out_channels[0], - out_channels=out_channels, - act_fn=act_fn, - fc_dim=block_out_channels[-1] // 4, - ) - - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - return_dict: bool = True, - ) -> UNet1DOutput | tuple: - r""" - The [`UNet1DModel`] forward method. - - Args: - sample (`torch.Tensor`): - The noisy input tensor with the following shape `(batch_size, num_channels, sample_size)`. - timestep (`torch.Tensor` or `float` or `int`): The number of timesteps to denoise an input. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unets.unet_1d.UNet1DOutput`] instead of a plain tuple. - - Returns: - [`~models.unets.unet_1d.UNet1DOutput`] or `tuple`: - If `return_dict` is True, an [`~models.unets.unet_1d.UNet1DOutput`] is returned, otherwise a `tuple` is - returned where the first element is the sample tensor. - """ - - # 1. time - timesteps = timestep - if not torch.is_tensor(timesteps): - timesteps = torch.tensor([timesteps], dtype=torch.long, device=sample.device) - elif torch.is_tensor(timesteps) and len(timesteps.shape) == 0: - timesteps = timesteps[None].to(sample.device) - - timestep_embed = self.time_proj(timesteps) - if self.config.use_timestep_embedding: - timestep_embed = self.time_mlp(timestep_embed.to(sample.dtype)) - else: - timestep_embed = timestep_embed[..., None] - timestep_embed = timestep_embed.repeat([1, 1, sample.shape[2]]).to(sample.dtype) - timestep_embed = timestep_embed.broadcast_to((sample.shape[:1] + timestep_embed.shape[1:])) - - # 2. down - down_block_res_samples = () - for downsample_block in self.down_blocks: - sample, res_samples = downsample_block(hidden_states=sample, temb=timestep_embed) - down_block_res_samples += res_samples - - # 3. mid - if self.mid_block: - sample = self.mid_block(sample, timestep_embed) - - # 4. up - for i, upsample_block in enumerate(self.up_blocks): - res_samples = down_block_res_samples[-1:] - down_block_res_samples = down_block_res_samples[:-1] - sample = upsample_block(sample, res_hidden_states_tuple=res_samples, temb=timestep_embed) - - # 5. post-process - if self.out_block: - sample = self.out_block(sample, timestep_embed) - - if not return_dict: - return (sample,) - - return UNet1DOutput(sample=sample) diff --git a/diffusers/models/unets/unet_1d_blocks.py b/diffusers/models/unets/unet_1d_blocks.py deleted file mode 100644 index f4d5c7d93a1278bf0409e3c98869eb60fabd3c58..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_1d_blocks.py +++ /dev/null @@ -1,701 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import math - -import torch -import torch.nn.functional as F -from torch import nn - -from ..activations import get_activation -from ..resnet import Downsample1D, ResidualTemporalBlock1D, Upsample1D, rearrange_dims - - -class DownResnetBlock1D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - num_layers: int = 1, - conv_shortcut: bool = False, - temb_channels: int = 32, - groups: int = 32, - groups_out: int | None = None, - non_linearity: str | None = None, - time_embedding_norm: str = "default", - output_scale_factor: float = 1.0, - add_downsample: bool = True, - ): - super().__init__() - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - self.use_conv_shortcut = conv_shortcut - self.time_embedding_norm = time_embedding_norm - self.add_downsample = add_downsample - self.output_scale_factor = output_scale_factor - - if groups_out is None: - groups_out = groups - - # there will always be at least one resnet - resnets = [ResidualTemporalBlock1D(in_channels, out_channels, embed_dim=temb_channels)] - - for _ in range(num_layers): - resnets.append(ResidualTemporalBlock1D(out_channels, out_channels, embed_dim=temb_channels)) - - self.resnets = nn.ModuleList(resnets) - - if non_linearity is None: - self.nonlinearity = None - else: - self.nonlinearity = get_activation(non_linearity) - - self.downsample = None - if add_downsample: - self.downsample = Downsample1D(out_channels, use_conv=True, padding=1) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - output_states = () - - hidden_states = self.resnets[0](hidden_states, temb) - for resnet in self.resnets[1:]: - hidden_states = resnet(hidden_states, temb) - - output_states += (hidden_states,) - - if self.nonlinearity is not None: - hidden_states = self.nonlinearity(hidden_states) - - if self.downsample is not None: - hidden_states = self.downsample(hidden_states) - - return hidden_states, output_states - - -class UpResnetBlock1D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int | None = None, - num_layers: int = 1, - temb_channels: int = 32, - groups: int = 32, - groups_out: int | None = None, - non_linearity: str | None = None, - time_embedding_norm: str = "default", - output_scale_factor: float = 1.0, - add_upsample: bool = True, - ): - super().__init__() - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - self.time_embedding_norm = time_embedding_norm - self.add_upsample = add_upsample - self.output_scale_factor = output_scale_factor - - if groups_out is None: - groups_out = groups - - # there will always be at least one resnet - resnets = [ResidualTemporalBlock1D(2 * in_channels, out_channels, embed_dim=temb_channels)] - - for _ in range(num_layers): - resnets.append(ResidualTemporalBlock1D(out_channels, out_channels, embed_dim=temb_channels)) - - self.resnets = nn.ModuleList(resnets) - - if non_linearity is None: - self.nonlinearity = None - else: - self.nonlinearity = get_activation(non_linearity) - - self.upsample = None - if add_upsample: - self.upsample = Upsample1D(out_channels, use_conv_transpose=True) - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...] | None = None, - temb: torch.Tensor | None = None, - ) -> torch.Tensor: - if res_hidden_states_tuple is not None: - res_hidden_states = res_hidden_states_tuple[-1] - hidden_states = torch.cat((hidden_states, res_hidden_states), dim=1) - - hidden_states = self.resnets[0](hidden_states, temb) - for resnet in self.resnets[1:]: - hidden_states = resnet(hidden_states, temb) - - if self.nonlinearity is not None: - hidden_states = self.nonlinearity(hidden_states) - - if self.upsample is not None: - hidden_states = self.upsample(hidden_states) - - return hidden_states - - -class ValueFunctionMidBlock1D(nn.Module): - def __init__(self, in_channels: int, out_channels: int, embed_dim: int): - super().__init__() - self.in_channels = in_channels - self.out_channels = out_channels - self.embed_dim = embed_dim - - self.res1 = ResidualTemporalBlock1D(in_channels, in_channels // 2, embed_dim=embed_dim) - self.down1 = Downsample1D(out_channels // 2, use_conv=True) - self.res2 = ResidualTemporalBlock1D(in_channels // 2, in_channels // 4, embed_dim=embed_dim) - self.down2 = Downsample1D(out_channels // 4, use_conv=True) - - def forward(self, x: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - x = self.res1(x, temb) - x = self.down1(x) - x = self.res2(x, temb) - x = self.down2(x) - return x - - -class MidResTemporalBlock1D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - embed_dim: int, - num_layers: int = 1, - add_downsample: bool = False, - add_upsample: bool = False, - non_linearity: str | None = None, - ): - super().__init__() - self.in_channels = in_channels - self.out_channels = out_channels - self.add_downsample = add_downsample - - # there will always be at least one resnet - resnets = [ResidualTemporalBlock1D(in_channels, out_channels, embed_dim=embed_dim)] - - for _ in range(num_layers): - resnets.append(ResidualTemporalBlock1D(out_channels, out_channels, embed_dim=embed_dim)) - - self.resnets = nn.ModuleList(resnets) - - if non_linearity is None: - self.nonlinearity = None - else: - self.nonlinearity = get_activation(non_linearity) - - self.upsample = None - if add_upsample: - self.upsample = Upsample1D(out_channels, use_conv=True) - - self.downsample = None - if add_downsample: - self.downsample = Downsample1D(out_channels, use_conv=True) - - if self.upsample and self.downsample: - raise ValueError("Block cannot downsample and upsample") - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor) -> torch.Tensor: - hidden_states = self.resnets[0](hidden_states, temb) - for resnet in self.resnets[1:]: - hidden_states = resnet(hidden_states, temb) - - if self.upsample: - hidden_states = self.upsample(hidden_states) - if self.downsample: - hidden_states = self.downsample(hidden_states) - - return hidden_states - - -class OutConv1DBlock(nn.Module): - def __init__(self, num_groups_out: int, out_channels: int, embed_dim: int, act_fn: str): - super().__init__() - self.final_conv1d_1 = nn.Conv1d(embed_dim, embed_dim, 5, padding=2) - self.final_conv1d_gn = nn.GroupNorm(num_groups_out, embed_dim) - self.final_conv1d_act = get_activation(act_fn) - self.final_conv1d_2 = nn.Conv1d(embed_dim, out_channels, 1) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - hidden_states = self.final_conv1d_1(hidden_states) - hidden_states = rearrange_dims(hidden_states) - hidden_states = self.final_conv1d_gn(hidden_states) - hidden_states = rearrange_dims(hidden_states) - hidden_states = self.final_conv1d_act(hidden_states) - hidden_states = self.final_conv1d_2(hidden_states) - return hidden_states - - -class OutValueFunctionBlock(nn.Module): - def __init__(self, fc_dim: int, embed_dim: int, act_fn: str = "mish"): - super().__init__() - self.final_block = nn.ModuleList( - [ - nn.Linear(fc_dim + embed_dim, fc_dim // 2), - get_activation(act_fn), - nn.Linear(fc_dim // 2, 1), - ] - ) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor) -> torch.Tensor: - hidden_states = hidden_states.view(hidden_states.shape[0], -1) - hidden_states = torch.cat((hidden_states, temb), dim=-1) - for layer in self.final_block: - hidden_states = layer(hidden_states) - - return hidden_states - - -_kernels = { - "linear": [1 / 8, 3 / 8, 3 / 8, 1 / 8], - "cubic": [-0.01171875, -0.03515625, 0.11328125, 0.43359375, 0.43359375, 0.11328125, -0.03515625, -0.01171875], - "lanczos3": [ - 0.003689131001010537, - 0.015056144446134567, - -0.03399861603975296, - -0.066637322306633, - 0.13550527393817902, - 0.44638532400131226, - 0.44638532400131226, - 0.13550527393817902, - -0.066637322306633, - -0.03399861603975296, - 0.015056144446134567, - 0.003689131001010537, - ], -} - - -class Downsample1d(nn.Module): - def __init__(self, kernel: str = "linear", pad_mode: str = "reflect"): - super().__init__() - self.pad_mode = pad_mode - kernel_1d = torch.tensor(_kernels[kernel]) - self.pad = kernel_1d.shape[0] // 2 - 1 - self.register_buffer("kernel", kernel_1d) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = F.pad(hidden_states, (self.pad,) * 2, self.pad_mode) - weight = hidden_states.new_zeros([hidden_states.shape[1], hidden_states.shape[1], self.kernel.shape[0]]) - indices = torch.arange(hidden_states.shape[1], device=hidden_states.device) - kernel = self.kernel.to(weight)[None, :].expand(hidden_states.shape[1], -1) - weight[indices, indices] = kernel - return F.conv1d(hidden_states, weight, stride=2) - - -class Upsample1d(nn.Module): - def __init__(self, kernel: str = "linear", pad_mode: str = "reflect"): - super().__init__() - self.pad_mode = pad_mode - kernel_1d = torch.tensor(_kernels[kernel]) * 2 - self.pad = kernel_1d.shape[0] // 2 - 1 - self.register_buffer("kernel", kernel_1d) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - hidden_states = F.pad(hidden_states, ((self.pad + 1) // 2,) * 2, self.pad_mode) - weight = hidden_states.new_zeros([hidden_states.shape[1], hidden_states.shape[1], self.kernel.shape[0]]) - indices = torch.arange(hidden_states.shape[1], device=hidden_states.device) - kernel = self.kernel.to(weight)[None, :].expand(hidden_states.shape[1], -1) - weight[indices, indices] = kernel - return F.conv_transpose1d(hidden_states, weight, stride=2, padding=self.pad * 2 + 1) - - -class SelfAttention1d(nn.Module): - def __init__(self, in_channels: int, n_head: int = 1, dropout_rate: float = 0.0): - super().__init__() - self.channels = in_channels - self.group_norm = nn.GroupNorm(1, num_channels=in_channels) - self.num_heads = n_head - - self.query = nn.Linear(self.channels, self.channels) - self.key = nn.Linear(self.channels, self.channels) - self.value = nn.Linear(self.channels, self.channels) - - self.proj_attn = nn.Linear(self.channels, self.channels, bias=True) - - self.dropout = nn.Dropout(dropout_rate, inplace=True) - - def transpose_for_scores(self, projection: torch.Tensor) -> torch.Tensor: - new_projection_shape = projection.size()[:-1] + (self.num_heads, -1) - # move heads to 2nd position (B, T, H * D) -> (B, T, H, D) -> (B, H, T, D) - new_projection = projection.view(new_projection_shape).permute(0, 2, 1, 3) - return new_projection - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - residual = hidden_states - batch, channel_dim, seq = hidden_states.shape - - hidden_states = self.group_norm(hidden_states) - hidden_states = hidden_states.transpose(1, 2) - - query_proj = self.query(hidden_states) - key_proj = self.key(hidden_states) - value_proj = self.value(hidden_states) - - query_states = self.transpose_for_scores(query_proj) - key_states = self.transpose_for_scores(key_proj) - value_states = self.transpose_for_scores(value_proj) - - scale = 1 / math.sqrt(math.sqrt(key_states.shape[-1])) - - attention_scores = torch.matmul(query_states * scale, key_states.transpose(-1, -2) * scale) - attention_probs = torch.softmax(attention_scores, dim=-1) - - # compute attention output - hidden_states = torch.matmul(attention_probs, value_states) - - hidden_states = hidden_states.permute(0, 2, 1, 3).contiguous() - new_hidden_states_shape = hidden_states.size()[:-2] + (self.channels,) - hidden_states = hidden_states.view(new_hidden_states_shape) - - # compute next hidden_states - hidden_states = self.proj_attn(hidden_states) - hidden_states = hidden_states.transpose(1, 2) - hidden_states = self.dropout(hidden_states) - - output = hidden_states + residual - - return output - - -class ResConvBlock(nn.Module): - def __init__(self, in_channels: int, mid_channels: int, out_channels: int, is_last: bool = False): - super().__init__() - self.is_last = is_last - self.has_conv_skip = in_channels != out_channels - - if self.has_conv_skip: - self.conv_skip = nn.Conv1d(in_channels, out_channels, 1, bias=False) - - self.conv_1 = nn.Conv1d(in_channels, mid_channels, 5, padding=2) - self.group_norm_1 = nn.GroupNorm(1, mid_channels) - self.gelu_1 = nn.GELU() - self.conv_2 = nn.Conv1d(mid_channels, out_channels, 5, padding=2) - - if not self.is_last: - self.group_norm_2 = nn.GroupNorm(1, out_channels) - self.gelu_2 = nn.GELU() - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - residual = self.conv_skip(hidden_states) if self.has_conv_skip else hidden_states - - hidden_states = self.conv_1(hidden_states) - hidden_states = self.group_norm_1(hidden_states) - hidden_states = self.gelu_1(hidden_states) - hidden_states = self.conv_2(hidden_states) - - if not self.is_last: - hidden_states = self.group_norm_2(hidden_states) - hidden_states = self.gelu_2(hidden_states) - - output = hidden_states + residual - return output - - -class UNetMidBlock1D(nn.Module): - def __init__(self, mid_channels: int, in_channels: int, out_channels: int | None = None): - super().__init__() - - out_channels = in_channels if out_channels is None else out_channels - - # there is always at least one resnet - self.down = Downsample1d("cubic") - resnets = [ - ResConvBlock(in_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, out_channels), - ] - attentions = [ - SelfAttention1d(mid_channels, mid_channels // 32), - SelfAttention1d(mid_channels, mid_channels // 32), - SelfAttention1d(mid_channels, mid_channels // 32), - SelfAttention1d(mid_channels, mid_channels // 32), - SelfAttention1d(mid_channels, mid_channels // 32), - SelfAttention1d(out_channels, out_channels // 32), - ] - self.up = Upsample1d(kernel="cubic") - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - hidden_states = self.down(hidden_states) - for attn, resnet in zip(self.attentions, self.resnets): - hidden_states = resnet(hidden_states) - hidden_states = attn(hidden_states) - - hidden_states = self.up(hidden_states) - - return hidden_states - - -class AttnDownBlock1D(nn.Module): - def __init__(self, out_channels: int, in_channels: int, mid_channels: int | None = None): - super().__init__() - mid_channels = out_channels if mid_channels is None else mid_channels - - self.down = Downsample1d("cubic") - resnets = [ - ResConvBlock(in_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, out_channels), - ] - attentions = [ - SelfAttention1d(mid_channels, mid_channels // 32), - SelfAttention1d(mid_channels, mid_channels // 32), - SelfAttention1d(out_channels, out_channels // 32), - ] - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - hidden_states = self.down(hidden_states) - - for resnet, attn in zip(self.resnets, self.attentions): - hidden_states = resnet(hidden_states) - hidden_states = attn(hidden_states) - - return hidden_states, (hidden_states,) - - -class DownBlock1D(nn.Module): - def __init__(self, out_channels: int, in_channels: int, mid_channels: int | None = None): - super().__init__() - mid_channels = out_channels if mid_channels is None else mid_channels - - self.down = Downsample1d("cubic") - resnets = [ - ResConvBlock(in_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, out_channels), - ] - - self.resnets = nn.ModuleList(resnets) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - hidden_states = self.down(hidden_states) - - for resnet in self.resnets: - hidden_states = resnet(hidden_states) - - return hidden_states, (hidden_states,) - - -class DownBlock1DNoSkip(nn.Module): - def __init__(self, out_channels: int, in_channels: int, mid_channels: int | None = None): - super().__init__() - mid_channels = out_channels if mid_channels is None else mid_channels - - resnets = [ - ResConvBlock(in_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, out_channels), - ] - - self.resnets = nn.ModuleList(resnets) - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - hidden_states = torch.cat([hidden_states, temb], dim=1) - for resnet in self.resnets: - hidden_states = resnet(hidden_states) - - return hidden_states, (hidden_states,) - - -class AttnUpBlock1D(nn.Module): - def __init__(self, in_channels: int, out_channels: int, mid_channels: int | None = None): - super().__init__() - mid_channels = out_channels if mid_channels is None else mid_channels - - resnets = [ - ResConvBlock(2 * in_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, out_channels), - ] - attentions = [ - SelfAttention1d(mid_channels, mid_channels // 32), - SelfAttention1d(mid_channels, mid_channels // 32), - SelfAttention1d(out_channels, out_channels // 32), - ] - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - self.up = Upsample1d(kernel="cubic") - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - ) -> torch.Tensor: - res_hidden_states = res_hidden_states_tuple[-1] - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - for resnet, attn in zip(self.resnets, self.attentions): - hidden_states = resnet(hidden_states) - hidden_states = attn(hidden_states) - - hidden_states = self.up(hidden_states) - - return hidden_states - - -class UpBlock1D(nn.Module): - def __init__(self, in_channels: int, out_channels: int, mid_channels: int | None = None): - super().__init__() - mid_channels = in_channels if mid_channels is None else mid_channels - - resnets = [ - ResConvBlock(2 * in_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, out_channels), - ] - - self.resnets = nn.ModuleList(resnets) - self.up = Upsample1d(kernel="cubic") - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - ) -> torch.Tensor: - res_hidden_states = res_hidden_states_tuple[-1] - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - for resnet in self.resnets: - hidden_states = resnet(hidden_states) - - hidden_states = self.up(hidden_states) - - return hidden_states - - -class UpBlock1DNoSkip(nn.Module): - def __init__(self, in_channels: int, out_channels: int, mid_channels: int | None = None): - super().__init__() - mid_channels = in_channels if mid_channels is None else mid_channels - - resnets = [ - ResConvBlock(2 * in_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, mid_channels), - ResConvBlock(mid_channels, mid_channels, out_channels, is_last=True), - ] - - self.resnets = nn.ModuleList(resnets) - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - ) -> torch.Tensor: - res_hidden_states = res_hidden_states_tuple[-1] - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - for resnet in self.resnets: - hidden_states = resnet(hidden_states) - - return hidden_states - - -DownBlockType = DownResnetBlock1D | DownBlock1D | AttnDownBlock1D | DownBlock1DNoSkip -MidBlockType = MidResTemporalBlock1D | ValueFunctionMidBlock1D | UNetMidBlock1D -OutBlockType = OutConv1DBlock | OutValueFunctionBlock -UpBlockType = UpResnetBlock1D | UpBlock1D | AttnUpBlock1D | UpBlock1DNoSkip - - -def get_down_block( - down_block_type: str, - num_layers: int, - in_channels: int, - out_channels: int, - temb_channels: int, - add_downsample: bool, -) -> DownBlockType: - if down_block_type == "DownResnetBlock1D": - return DownResnetBlock1D( - in_channels=in_channels, - num_layers=num_layers, - out_channels=out_channels, - temb_channels=temb_channels, - add_downsample=add_downsample, - ) - elif down_block_type == "DownBlock1D": - return DownBlock1D(out_channels=out_channels, in_channels=in_channels) - elif down_block_type == "AttnDownBlock1D": - return AttnDownBlock1D(out_channels=out_channels, in_channels=in_channels) - elif down_block_type == "DownBlock1DNoSkip": - return DownBlock1DNoSkip(out_channels=out_channels, in_channels=in_channels) - raise ValueError(f"{down_block_type} does not exist.") - - -def get_up_block( - up_block_type: str, num_layers: int, in_channels: int, out_channels: int, temb_channels: int, add_upsample: bool -) -> UpBlockType: - if up_block_type == "UpResnetBlock1D": - return UpResnetBlock1D( - in_channels=in_channels, - num_layers=num_layers, - out_channels=out_channels, - temb_channels=temb_channels, - add_upsample=add_upsample, - ) - elif up_block_type == "UpBlock1D": - return UpBlock1D(in_channels=in_channels, out_channels=out_channels) - elif up_block_type == "AttnUpBlock1D": - return AttnUpBlock1D(in_channels=in_channels, out_channels=out_channels) - elif up_block_type == "UpBlock1DNoSkip": - return UpBlock1DNoSkip(in_channels=in_channels, out_channels=out_channels) - raise ValueError(f"{up_block_type} does not exist.") - - -def get_mid_block( - mid_block_type: str, - num_layers: int, - in_channels: int, - mid_channels: int, - out_channels: int, - embed_dim: int, - add_downsample: bool, -) -> MidBlockType: - if mid_block_type == "MidResTemporalBlock1D": - return MidResTemporalBlock1D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - embed_dim=embed_dim, - add_downsample=add_downsample, - ) - elif mid_block_type == "ValueFunctionMidBlock1D": - return ValueFunctionMidBlock1D(in_channels=in_channels, out_channels=out_channels, embed_dim=embed_dim) - elif mid_block_type == "UNetMidBlock1D": - return UNetMidBlock1D(in_channels=in_channels, mid_channels=mid_channels, out_channels=out_channels) - raise ValueError(f"{mid_block_type} does not exist.") - - -def get_out_block( - *, out_block_type: str, num_groups_out: int, embed_dim: int, out_channels: int, act_fn: str, fc_dim: int -) -> OutBlockType | None: - if out_block_type == "OutConv1DBlock": - return OutConv1DBlock(num_groups_out, out_channels, embed_dim, act_fn) - elif out_block_type == "ValueFunction": - return OutValueFunctionBlock(fc_dim, embed_dim, act_fn) - return None diff --git a/diffusers/models/unets/unet_2d.py b/diffusers/models/unets/unet_2d.py deleted file mode 100644 index 4bbe0535e94aea3bcc4bbc704dcebad241e369a3..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_2d.py +++ /dev/null @@ -1,353 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from dataclasses import dataclass - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput -from ..embeddings import GaussianFourierProjection, TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from .unet_2d_blocks import UNetMidBlock2D, get_down_block, get_up_block - - -@dataclass -class UNet2DOutput(BaseOutput): - """ - The output of [`UNet2DModel`]. - - Args: - sample (`torch.Tensor` of shape `(batch_size, num_channels, height, width)`): - The hidden states output from the last layer of the model. - """ - - sample: torch.Tensor - - -class UNet2DModel(ModelMixin, ConfigMixin): - r""" - A 2D UNet model that takes a noisy sample and a timestep and returns a sample shaped output. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - sample_size (`int` or `tuple[int, int]`, *optional*, defaults to `None`): - Height and width of input/output sample. Dimensions must be a multiple of `2 ** (len(block_out_channels) - - 1)`. - in_channels (`int`, *optional*, defaults to 3): Number of channels in the input sample. - out_channels (`int`, *optional*, defaults to 3): Number of channels in the output. - center_input_sample (`bool`, *optional*, defaults to `False`): Whether to center the input sample. - time_embedding_type (`str`, *optional*, defaults to `"positional"`): Type of time embedding to use. - freq_shift (`int`, *optional*, defaults to 0): Frequency shift for Fourier time embedding. - flip_sin_to_cos (`bool`, *optional*, defaults to `True`): - Whether to flip sin to cos for Fourier time embedding. - down_block_types (`tuple[str]`, *optional*, defaults to `("DownBlock2D", "AttnDownBlock2D", "AttnDownBlock2D", "AttnDownBlock2D")`): - tuple of downsample block types. - mid_block_type (`str`, *optional*, defaults to `"UNetMidBlock2D"`): - Block type for middle of UNet, it can be either `UNetMidBlock2D` or `None`. - up_block_types (`tuple[str]`, *optional*, defaults to `("AttnUpBlock2D", "AttnUpBlock2D", "AttnUpBlock2D", "UpBlock2D")`): - tuple of upsample block types. - block_out_channels (`tuple[int]`, *optional*, defaults to `(224, 448, 672, 896)`): - tuple of block output channels. - layers_per_block (`int`, *optional*, defaults to `2`): The number of layers per block. - mid_block_scale_factor (`float`, *optional*, defaults to `1`): The scale factor for the mid block. - downsample_padding (`int`, *optional*, defaults to `1`): The padding for the downsample convolution. - downsample_type (`str`, *optional*, defaults to `conv`): - The downsample type for downsampling layers. Choose between "conv" and "resnet" - upsample_type (`str`, *optional*, defaults to `conv`): - The upsample type for upsampling layers. Choose between "conv" and "resnet" - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use. - attention_head_dim (`int`, *optional*, defaults to `8`): The attention head dimension. - norm_num_groups (`int`, *optional*, defaults to `32`): The number of groups for normalization. - attn_norm_num_groups (`int`, *optional*, defaults to `None`): - If set to an integer, a group norm layer will be created in the mid block's [`Attention`] layer with the - given number of groups. If left as `None`, the group norm layer will only be created if - `resnet_time_scale_shift` is set to `default`, and if created will have `norm_num_groups` groups. - norm_eps (`float`, *optional*, defaults to `1e-5`): The epsilon for normalization. - resnet_time_scale_shift (`str`, *optional*, defaults to `"default"`): Time scale shift config - for ResNet blocks (see [`~models.resnet.ResnetBlock2D`]). Choose from `default` or `scale_shift`. - class_embed_type (`str`, *optional*, defaults to `None`): - The type of class embedding to use which is ultimately summed with the time embeddings. Choose from `None`, - `"timestep"`, or `"identity"`. - num_class_embeds (`int`, *optional*, defaults to `None`): - Input dimension of the learnable embedding matrix to be projected to `time_embed_dim` when performing class - conditioning with `class_embed_type` equal to `None`. - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["norm"] - - @register_to_config - def __init__( - self, - sample_size: int | tuple[int, int] | None = None, - in_channels: int = 3, - out_channels: int = 3, - center_input_sample: bool = False, - time_embedding_type: str = "positional", - time_embedding_dim: int | None = None, - freq_shift: int = 0, - flip_sin_to_cos: bool = True, - down_block_types: tuple[str, ...] = ("DownBlock2D", "AttnDownBlock2D", "AttnDownBlock2D", "AttnDownBlock2D"), - mid_block_type: str | None = "UNetMidBlock2D", - up_block_types: tuple[str, ...] = ("AttnUpBlock2D", "AttnUpBlock2D", "AttnUpBlock2D", "UpBlock2D"), - block_out_channels: tuple[int, ...] = (224, 448, 672, 896), - layers_per_block: int = 2, - mid_block_scale_factor: float = 1, - downsample_padding: int = 1, - downsample_type: str = "conv", - upsample_type: str = "conv", - dropout: float = 0.0, - act_fn: str = "silu", - attention_head_dim: int | None = 8, - norm_num_groups: int = 32, - attn_norm_num_groups: int | None = None, - norm_eps: float = 1e-5, - resnet_time_scale_shift: str = "default", - add_attention: bool = True, - class_embed_type: str | None = None, - num_class_embeds: int | None = None, - num_train_timesteps: int | None = None, - ): - super().__init__() - - self.sample_size = sample_size - time_embed_dim = time_embedding_dim or block_out_channels[0] * 4 - - # Check inputs - if len(down_block_types) != len(up_block_types): - raise ValueError( - f"Must provide the same number of `down_block_types` as `up_block_types`. `down_block_types`: {down_block_types}. `up_block_types`: {up_block_types}." - ) - - if len(block_out_channels) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." - ) - - # input - self.conv_in = nn.Conv2d(in_channels, block_out_channels[0], kernel_size=3, padding=(1, 1)) - - # time - if time_embedding_type == "fourier": - self.time_proj = GaussianFourierProjection(embedding_size=block_out_channels[0], scale=16) - timestep_input_dim = 2 * block_out_channels[0] - elif time_embedding_type == "positional": - self.time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos, freq_shift) - timestep_input_dim = block_out_channels[0] - elif time_embedding_type == "learned": - self.time_proj = nn.Embedding(num_train_timesteps, block_out_channels[0]) - timestep_input_dim = block_out_channels[0] - - self.time_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim) - - # class embedding - if class_embed_type is None and num_class_embeds is not None: - self.class_embedding = nn.Embedding(num_class_embeds, time_embed_dim) - elif class_embed_type == "timestep": - self.class_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim) - elif class_embed_type == "identity": - self.class_embedding = nn.Identity(time_embed_dim, time_embed_dim) - else: - self.class_embedding = None - - self.down_blocks = nn.ModuleList([]) - self.mid_block = None - self.up_blocks = nn.ModuleList([]) - - # down - output_channel = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - down_block = get_down_block( - down_block_type, - num_layers=layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - temb_channels=time_embed_dim, - add_downsample=not is_final_block, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - attention_head_dim=attention_head_dim if attention_head_dim is not None else output_channel, - downsample_padding=downsample_padding, - resnet_time_scale_shift=resnet_time_scale_shift, - downsample_type=downsample_type, - dropout=dropout, - ) - self.down_blocks.append(down_block) - - # mid - if mid_block_type is None: - self.mid_block = None - else: - self.mid_block = UNetMidBlock2D( - in_channels=block_out_channels[-1], - temb_channels=time_embed_dim, - dropout=dropout, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - output_scale_factor=mid_block_scale_factor, - resnet_time_scale_shift=resnet_time_scale_shift, - attention_head_dim=attention_head_dim if attention_head_dim is not None else block_out_channels[-1], - resnet_groups=norm_num_groups, - attn_groups=attn_norm_num_groups, - add_attention=add_attention, - ) - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - output_channel = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)] - - is_final_block = i == len(block_out_channels) - 1 - - up_block = get_up_block( - up_block_type, - num_layers=layers_per_block + 1, - in_channels=input_channel, - out_channels=output_channel, - prev_output_channel=prev_output_channel, - temb_channels=time_embed_dim, - add_upsample=not is_final_block, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - attention_head_dim=attention_head_dim if attention_head_dim is not None else output_channel, - resnet_time_scale_shift=resnet_time_scale_shift, - upsample_type=upsample_type, - dropout=dropout, - ) - self.up_blocks.append(up_block) - - # out - num_groups_out = norm_num_groups if norm_num_groups is not None else min(block_out_channels[0] // 4, 32) - self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=num_groups_out, eps=norm_eps) - self.conv_act = nn.SiLU() - self.conv_out = nn.Conv2d(block_out_channels[0], out_channels, kernel_size=3, padding=1) - - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - class_labels: torch.Tensor | None = None, - return_dict: bool = True, - ) -> UNet2DOutput | tuple: - r""" - The [`UNet2DModel`] forward method. - - Args: - sample (`torch.Tensor`): - The noisy input tensor with the following shape `(batch, channel, height, width)`. - timestep (`torch.Tensor` or `float` or `int`): The number of timesteps to denoise an input. - class_labels (`torch.Tensor`, *optional*, defaults to `None`): - Optional class labels for conditioning. Their embeddings will be summed with the timestep embeddings. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unets.unet_2d.UNet2DOutput`] instead of a plain tuple. - - Returns: - [`~models.unets.unet_2d.UNet2DOutput`] or `tuple`: - If `return_dict` is True, an [`~models.unets.unet_2d.UNet2DOutput`] is returned, otherwise a `tuple` is - returned where the first element is the sample tensor. - """ - # 0. center input if necessary - if self.config.center_input_sample: - sample = 2 * sample - 1.0 - - # 1. time - timesteps = timestep - if not torch.is_tensor(timesteps): - timesteps = torch.tensor([timesteps], dtype=torch.long, device=sample.device) - elif torch.is_tensor(timesteps) and len(timesteps.shape) == 0: - timesteps = timesteps[None].to(sample.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - timesteps = timesteps * torch.ones(sample.shape[0], dtype=timesteps.dtype, device=timesteps.device) - - t_emb = self.time_proj(timesteps) - - # timesteps does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=self.dtype) - emb = self.time_embedding(t_emb) - - if self.class_embedding is not None: - if class_labels is None: - raise ValueError("class_labels should be provided when doing class conditioning") - - if self.config.class_embed_type == "timestep": - class_labels = self.time_proj(class_labels) - - class_emb = self.class_embedding(class_labels).to(dtype=self.dtype) - emb = emb + class_emb - elif self.class_embedding is None and class_labels is not None: - raise ValueError("class_embedding needs to be initialized in order to use class conditioning") - - # 2. pre-process - skip_sample = sample - sample = self.conv_in(sample) - - # 3. down - down_block_res_samples = (sample,) - for downsample_block in self.down_blocks: - if hasattr(downsample_block, "skip_conv"): - sample, res_samples, skip_sample = downsample_block( - hidden_states=sample, temb=emb, skip_sample=skip_sample - ) - else: - sample, res_samples = downsample_block(hidden_states=sample, temb=emb) - - down_block_res_samples += res_samples - - # 4. mid - if self.mid_block is not None: - sample = self.mid_block(sample, emb) - - # 5. up - skip_sample = None - for upsample_block in self.up_blocks: - res_samples = down_block_res_samples[-len(upsample_block.resnets) :] - down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)] - - if hasattr(upsample_block, "skip_conv"): - sample, skip_sample = upsample_block(sample, res_samples, emb, skip_sample) - else: - sample = upsample_block(sample, res_samples, emb) - - # 6. post-process - sample = self.conv_norm_out(sample) - sample = self.conv_act(sample) - sample = self.conv_out(sample) - - if skip_sample is not None: - sample += skip_sample - - if self.config.time_embedding_type == "fourier": - timesteps = timesteps.reshape((sample.shape[0], *([1] * len(sample.shape[1:])))) - sample = sample / timesteps - - if not return_dict: - return (sample,) - - return UNet2DOutput(sample=sample) diff --git a/diffusers/models/unets/unet_2d_blocks.py b/diffusers/models/unets/unet_2d_blocks.py deleted file mode 100644 index 611d0113e174af97dd9db61827b91155dc21cffd..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_2d_blocks.py +++ /dev/null @@ -1,3583 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from typing import Any - -import numpy as np -import torch -import torch.nn.functional as F -from torch import nn - -from ...utils import deprecate, logging -from ...utils.torch_utils import apply_freeu -from ..activations import get_activation -from ..attention_processor import Attention, AttnAddedKVProcessor, AttnAddedKVProcessor2_0 -from ..normalization import AdaGroupNorm -from ..resnet import ( - Downsample2D, - FirDownsample2D, - FirUpsample2D, - KDownsample2D, - KUpsample2D, - ResnetBlock2D, - ResnetBlockCondNorm2D, - Upsample2D, -) -from ..transformers.dual_transformer_2d import DualTransformer2DModel -from ..transformers.transformer_2d import Transformer2DModel - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def get_down_block( - down_block_type: str, - num_layers: int, - in_channels: int, - out_channels: int, - temb_channels: int, - add_downsample: bool, - resnet_eps: float, - resnet_act_fn: str, - transformer_layers_per_block: int = 1, - num_attention_heads: int | None = None, - resnet_groups: int | None = None, - cross_attention_dim: int | None = None, - downsample_padding: int | None = None, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - resnet_time_scale_shift: str = "default", - attention_type: str = "default", - resnet_skip_time_act: bool = False, - resnet_out_scale_factor: float = 1.0, - cross_attention_norm: str | None = None, - attention_head_dim: int | None = None, - downsample_type: str | None = None, - dropout: float = 0.0, -): - # If attn head dim is not defined, we default it to the number of heads - if attention_head_dim is None: - logger.warning( - f"It is recommended to provide `attention_head_dim` when calling `get_down_block`. Defaulting `attention_head_dim` to {num_attention_heads}." - ) - attention_head_dim = num_attention_heads - - down_block_type = down_block_type[7:] if down_block_type.startswith("UNetRes") else down_block_type - if down_block_type == "DownBlock2D": - return DownBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - downsample_padding=downsample_padding, - resnet_time_scale_shift=resnet_time_scale_shift, - ) - elif down_block_type == "ResnetDownsampleBlock2D": - return ResnetDownsampleBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - resnet_time_scale_shift=resnet_time_scale_shift, - skip_time_act=resnet_skip_time_act, - output_scale_factor=resnet_out_scale_factor, - ) - elif down_block_type == "AttnDownBlock2D": - if add_downsample is False: - downsample_type = None - else: - downsample_type = downsample_type or "conv" # default to 'conv' - return AttnDownBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - dropout=dropout, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - downsample_padding=downsample_padding, - attention_head_dim=attention_head_dim, - resnet_time_scale_shift=resnet_time_scale_shift, - downsample_type=downsample_type, - ) - elif down_block_type == "CrossAttnDownBlock2D": - if cross_attention_dim is None: - raise ValueError("cross_attention_dim must be specified for CrossAttnDownBlock2D") - return CrossAttnDownBlock2D( - num_layers=num_layers, - transformer_layers_per_block=transformer_layers_per_block, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - downsample_padding=downsample_padding, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads, - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - attention_type=attention_type, - ) - elif down_block_type == "SimpleCrossAttnDownBlock2D": - if cross_attention_dim is None: - raise ValueError("cross_attention_dim must be specified for SimpleCrossAttnDownBlock2D") - return SimpleCrossAttnDownBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - cross_attention_dim=cross_attention_dim, - attention_head_dim=attention_head_dim, - resnet_time_scale_shift=resnet_time_scale_shift, - skip_time_act=resnet_skip_time_act, - output_scale_factor=resnet_out_scale_factor, - only_cross_attention=only_cross_attention, - cross_attention_norm=cross_attention_norm, - ) - elif down_block_type == "SkipDownBlock2D": - return SkipDownBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - downsample_padding=downsample_padding, - resnet_time_scale_shift=resnet_time_scale_shift, - ) - elif down_block_type == "AttnSkipDownBlock2D": - return AttnSkipDownBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - attention_head_dim=attention_head_dim, - resnet_time_scale_shift=resnet_time_scale_shift, - ) - elif down_block_type == "DownEncoderBlock2D": - return DownEncoderBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - downsample_padding=downsample_padding, - resnet_time_scale_shift=resnet_time_scale_shift, - ) - elif down_block_type == "AttnDownEncoderBlock2D": - return AttnDownEncoderBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - downsample_padding=downsample_padding, - attention_head_dim=attention_head_dim, - resnet_time_scale_shift=resnet_time_scale_shift, - ) - elif down_block_type == "KDownBlock2D": - return KDownBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - ) - elif down_block_type == "KCrossAttnDownBlock2D": - return KCrossAttnDownBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - cross_attention_dim=cross_attention_dim, - attention_head_dim=attention_head_dim, - add_self_attention=True if not add_downsample else False, - ) - raise ValueError(f"{down_block_type} does not exist.") - - -def get_mid_block( - mid_block_type: str, - temb_channels: int, - in_channels: int, - resnet_eps: float, - resnet_act_fn: str, - resnet_groups: int, - output_scale_factor: float = 1.0, - transformer_layers_per_block: int = 1, - num_attention_heads: int | None = None, - cross_attention_dim: int | None = None, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - mid_block_only_cross_attention: bool = False, - upcast_attention: bool = False, - resnet_time_scale_shift: str = "default", - attention_type: str = "default", - resnet_skip_time_act: bool = False, - cross_attention_norm: str | None = None, - attention_head_dim: int | None = 1, - dropout: float = 0.0, -): - if mid_block_type == "UNetMidBlock2DCrossAttn": - return UNetMidBlock2DCrossAttn( - transformer_layers_per_block=transformer_layers_per_block, - in_channels=in_channels, - temb_channels=temb_channels, - dropout=dropout, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - output_scale_factor=output_scale_factor, - resnet_time_scale_shift=resnet_time_scale_shift, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads, - resnet_groups=resnet_groups, - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - attention_type=attention_type, - ) - elif mid_block_type == "UNetMidBlock2DSimpleCrossAttn": - return UNetMidBlock2DSimpleCrossAttn( - in_channels=in_channels, - temb_channels=temb_channels, - dropout=dropout, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - output_scale_factor=output_scale_factor, - cross_attention_dim=cross_attention_dim, - attention_head_dim=attention_head_dim, - resnet_groups=resnet_groups, - resnet_time_scale_shift=resnet_time_scale_shift, - skip_time_act=resnet_skip_time_act, - only_cross_attention=mid_block_only_cross_attention, - cross_attention_norm=cross_attention_norm, - ) - elif mid_block_type == "UNetMidBlock2D": - return UNetMidBlock2D( - in_channels=in_channels, - temb_channels=temb_channels, - dropout=dropout, - num_layers=0, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - output_scale_factor=output_scale_factor, - resnet_groups=resnet_groups, - resnet_time_scale_shift=resnet_time_scale_shift, - add_attention=False, - ) - elif mid_block_type is None: - return None - else: - raise ValueError(f"unknown mid_block_type : {mid_block_type}") - - -def get_up_block( - up_block_type: str, - num_layers: int, - in_channels: int, - out_channels: int, - prev_output_channel: int, - temb_channels: int, - add_upsample: bool, - resnet_eps: float, - resnet_act_fn: str, - resolution_idx: int | None = None, - transformer_layers_per_block: int = 1, - num_attention_heads: int | None = None, - resnet_groups: int | None = None, - cross_attention_dim: int | None = None, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - resnet_time_scale_shift: str = "default", - attention_type: str = "default", - resnet_skip_time_act: bool = False, - resnet_out_scale_factor: float = 1.0, - cross_attention_norm: str | None = None, - attention_head_dim: int | None = None, - upsample_type: str | None = None, - dropout: float = 0.0, -) -> nn.Module: - # If attn head dim is not defined, we default it to the number of heads - if attention_head_dim is None: - logger.warning( - f"It is recommended to provide `attention_head_dim` when calling `get_up_block`. Defaulting `attention_head_dim` to {num_attention_heads}." - ) - attention_head_dim = num_attention_heads - - up_block_type = up_block_type[7:] if up_block_type.startswith("UNetRes") else up_block_type - if up_block_type == "UpBlock2D": - return UpBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - resnet_time_scale_shift=resnet_time_scale_shift, - ) - elif up_block_type == "ResnetUpsampleBlock2D": - return ResnetUpsampleBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - resnet_time_scale_shift=resnet_time_scale_shift, - skip_time_act=resnet_skip_time_act, - output_scale_factor=resnet_out_scale_factor, - ) - elif up_block_type == "CrossAttnUpBlock2D": - if cross_attention_dim is None: - raise ValueError("cross_attention_dim must be specified for CrossAttnUpBlock2D") - return CrossAttnUpBlock2D( - num_layers=num_layers, - transformer_layers_per_block=transformer_layers_per_block, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads, - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - attention_type=attention_type, - ) - elif up_block_type == "SimpleCrossAttnUpBlock2D": - if cross_attention_dim is None: - raise ValueError("cross_attention_dim must be specified for SimpleCrossAttnUpBlock2D") - return SimpleCrossAttnUpBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - cross_attention_dim=cross_attention_dim, - attention_head_dim=attention_head_dim, - resnet_time_scale_shift=resnet_time_scale_shift, - skip_time_act=resnet_skip_time_act, - output_scale_factor=resnet_out_scale_factor, - only_cross_attention=only_cross_attention, - cross_attention_norm=cross_attention_norm, - ) - elif up_block_type == "AttnUpBlock2D": - if add_upsample is False: - upsample_type = None - else: - upsample_type = upsample_type or "conv" # default to 'conv' - - return AttnUpBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - dropout=dropout, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - attention_head_dim=attention_head_dim, - resnet_time_scale_shift=resnet_time_scale_shift, - upsample_type=upsample_type, - ) - elif up_block_type == "SkipUpBlock2D": - return SkipUpBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_time_scale_shift=resnet_time_scale_shift, - ) - elif up_block_type == "AttnSkipUpBlock2D": - return AttnSkipUpBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - attention_head_dim=attention_head_dim, - resnet_time_scale_shift=resnet_time_scale_shift, - ) - elif up_block_type == "UpDecoderBlock2D": - return UpDecoderBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - resnet_time_scale_shift=resnet_time_scale_shift, - temb_channels=temb_channels, - ) - elif up_block_type == "AttnUpDecoderBlock2D": - return AttnUpDecoderBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - attention_head_dim=attention_head_dim, - resnet_time_scale_shift=resnet_time_scale_shift, - temb_channels=temb_channels, - ) - elif up_block_type == "KUpBlock2D": - return KUpBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - ) - elif up_block_type == "KCrossAttnUpBlock2D": - return KCrossAttnUpBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - cross_attention_dim=cross_attention_dim, - attention_head_dim=attention_head_dim, - ) - - raise ValueError(f"{up_block_type} does not exist.") - - -class AutoencoderTinyBlock(nn.Module): - """ - Tiny Autoencoder block used in [`AutoencoderTiny`]. It is a mini residual module consisting of plain conv + ReLU - blocks. - - Args: - in_channels (`int`): The number of input channels. - out_channels (`int`): The number of output channels. - act_fn (`str`): - ` The activation function to use. Supported values are `"swish"`, `"mish"`, `"gelu"`, and `"relu"`. - - Returns: - `torch.Tensor`: A tensor with the same shape as the input tensor, but with the number of channels equal to - `out_channels`. - """ - - def __init__(self, in_channels: int, out_channels: int, act_fn: str): - super().__init__() - act_fn = get_activation(act_fn) - self.conv = nn.Sequential( - nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1), - act_fn, - nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1), - act_fn, - nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1), - ) - self.skip = ( - nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False) - if in_channels != out_channels - else nn.Identity() - ) - self.fuse = nn.ReLU() - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.fuse(self.conv(x) + self.skip(x)) - - -class UNetMidBlock2D(nn.Module): - """ - A 2D UNet mid-block [`UNetMidBlock2D`] with multiple residual blocks and optional attention blocks. - - Args: - in_channels (`int`): The number of input channels. - temb_channels (`int`): The number of temporal embedding channels. - dropout (`float`, *optional*, defaults to 0.0): The dropout rate. - num_layers (`int`, *optional*, defaults to 1): The number of residual blocks. - resnet_eps (`float`, *optional*, 1e-6 ): The epsilon value for the resnet blocks. - resnet_time_scale_shift (`str`, *optional*, defaults to `default`): - The type of normalization to apply to the time embeddings. This can help to improve the performance of the - model on tasks with long-range temporal dependencies. - resnet_act_fn (`str`, *optional*, defaults to `swish`): The activation function for the resnet blocks. - resnet_groups (`int`, *optional*, defaults to 32): - The number of groups to use in the group normalization layers of the resnet blocks. - attn_groups (`int | None`, *optional*, defaults to None): The number of groups for the attention blocks. - resnet_pre_norm (`bool`, *optional*, defaults to `True`): - Whether to use pre-normalization for the resnet blocks. - add_attention (`bool`, *optional*, defaults to `True`): Whether to add attention blocks. - attention_head_dim (`int`, *optional*, defaults to 1): - Dimension of a single attention head. The number of attention heads is determined based on this value and - the number of input channels. - output_scale_factor (`float`, *optional*, defaults to 1.0): The output scale factor. - - Returns: - `torch.Tensor`: The output of the last residual block, which is a tensor of shape `(batch_size, in_channels, - height, width)`. - - """ - - def __init__( - self, - in_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", # default, spatial - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - attn_groups: int | None = None, - resnet_pre_norm: bool = True, - add_attention: bool = True, - attention_head_dim: int = 1, - output_scale_factor: float = 1.0, - ): - super().__init__() - resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32) - self.add_attention = add_attention - - if attn_groups is None: - attn_groups = resnet_groups if resnet_time_scale_shift == "default" else None - - # there is always at least one resnet - if resnet_time_scale_shift == "spatial": - resnets = [ - ResnetBlockCondNorm2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm="spatial", - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - ) - ] - else: - resnets = [ - ResnetBlock2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ] - attentions = [] - - if attention_head_dim is None: - logger.warning( - f"It is not recommend to pass `attention_head_dim=None`. Defaulting `attention_head_dim` to `in_channels`: {in_channels}." - ) - attention_head_dim = in_channels - - for _ in range(num_layers): - if self.add_attention: - attentions.append( - Attention( - in_channels, - heads=in_channels // attention_head_dim, - dim_head=attention_head_dim, - rescale_output_factor=output_scale_factor, - eps=resnet_eps, - norm_num_groups=attn_groups, - spatial_norm_dim=temb_channels if resnet_time_scale_shift == "spatial" else None, - residual_connection=True, - bias=True, - upcast_softmax=True, - _from_deprecated_attn_block=True, - ) - ) - else: - attentions.append(None) - - if resnet_time_scale_shift == "spatial": - resnets.append( - ResnetBlockCondNorm2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm="spatial", - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - ) - ) - else: - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - hidden_states = self.resnets[0](hidden_states, temb) - for attn, resnet in zip(self.attentions, self.resnets[1:]): - if torch.is_grad_enabled() and self.gradient_checkpointing: - if attn is not None: - hidden_states = attn(hidden_states, temb=temb) - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - if attn is not None: - hidden_states = attn(hidden_states, temb=temb) - hidden_states = resnet(hidden_states, temb) - - return hidden_states - - -class UNetMidBlock2DCrossAttn(nn.Module): - def __init__( - self, - in_channels: int, - temb_channels: int, - out_channels: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_groups_out: int | None = None, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - output_scale_factor: float = 1.0, - cross_attention_dim: int = 1280, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - upcast_attention: bool = False, - attention_type: str = "default", - ): - super().__init__() - - out_channels = out_channels or in_channels - self.in_channels = in_channels - self.out_channels = out_channels - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32) - - # support for variable transformer layers per block - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - resnet_groups_out = resnet_groups_out or resnet_groups - - # there is always at least one resnet - resnets = [ - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - groups_out=resnet_groups_out, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ] - attentions = [] - - for i in range(num_layers): - if not dual_cross_attention: - attentions.append( - Transformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups_out, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - attention_type=attention_type, - ) - ) - else: - attentions.append( - DualTransformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - resnets.append( - ResnetBlock2D( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups_out, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - hidden_states = self.resnets[0](hidden_states, temb) - for attn, resnet in zip(self.attentions, self.resnets[1:]): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - hidden_states = resnet(hidden_states, temb) - - return hidden_states - - -class UNetMidBlock2DSimpleCrossAttn(nn.Module): - def __init__( - self, - in_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - attention_head_dim: int = 1, - output_scale_factor: float = 1.0, - cross_attention_dim: int = 1280, - skip_time_act: bool = False, - only_cross_attention: bool = False, - cross_attention_norm: str | None = None, - ): - super().__init__() - - self.has_cross_attention = True - - self.attention_head_dim = attention_head_dim - resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32) - - self.num_heads = in_channels // self.attention_head_dim - - # there is always at least one resnet - resnets = [ - ResnetBlock2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - skip_time_act=skip_time_act, - ) - ] - attentions = [] - - for _ in range(num_layers): - processor = ( - AttnAddedKVProcessor2_0() if hasattr(F, "scaled_dot_product_attention") else AttnAddedKVProcessor() - ) - - attentions.append( - Attention( - query_dim=in_channels, - cross_attention_dim=in_channels, - heads=self.num_heads, - dim_head=self.attention_head_dim, - added_kv_proj_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - bias=True, - upcast_softmax=True, - only_cross_attention=only_cross_attention, - cross_attention_norm=cross_attention_norm, - processor=processor, - ) - ) - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - skip_time_act=skip_time_act, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - if attention_mask is None: - # if encoder_hidden_states is defined: we are doing cross-attn, so we should use cross-attn mask. - mask = None if encoder_hidden_states is None else encoder_attention_mask - else: - # when attention_mask is defined: we don't even check for encoder_attention_mask. - # this is to maintain compatibility with UnCLIP, which uses 'attention_mask' param for cross-attn masks. - # TODO: UnCLIP should express cross-attn mask via encoder_attention_mask param instead of via attention_mask. - # then we can simplify this whole if/else block to: - # mask = attention_mask if encoder_hidden_states is None else encoder_attention_mask - mask = attention_mask - - hidden_states = self.resnets[0](hidden_states, temb) - for attn, resnet in zip(self.attentions, self.resnets[1:]): - # attn - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=mask, - **cross_attention_kwargs, - ) - - # resnet - hidden_states = resnet(hidden_states, temb) - - return hidden_states - - -class AttnDownBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - attention_head_dim: int = 1, - output_scale_factor: float = 1.0, - downsample_padding: int = 1, - downsample_type: str = "conv", - ): - super().__init__() - resnets = [] - attentions = [] - self.downsample_type = downsample_type - - if attention_head_dim is None: - logger.warning( - f"It is not recommend to pass `attention_head_dim=None`. Defaulting `attention_head_dim` to `in_channels`: {out_channels}." - ) - attention_head_dim = out_channels - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - attentions.append( - Attention( - out_channels, - heads=out_channels // attention_head_dim, - dim_head=attention_head_dim, - rescale_output_factor=output_scale_factor, - eps=resnet_eps, - norm_num_groups=resnet_groups, - residual_connection=True, - bias=True, - upcast_softmax=True, - _from_deprecated_attn_block=True, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if downsample_type == "conv": - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op" - ) - ] - ) - elif downsample_type == "resnet": - self.downsamplers = nn.ModuleList( - [ - ResnetBlock2D( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - down=True, - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - upsample_size: int | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - output_states = () - - for resnet, attn in zip(self.resnets, self.attentions): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - hidden_states = attn(hidden_states, **cross_attention_kwargs) - output_states = output_states + (hidden_states,) - else: - hidden_states = resnet(hidden_states, temb) - hidden_states = attn(hidden_states, **cross_attention_kwargs) - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - if self.downsample_type == "resnet": - hidden_states = downsampler(hidden_states, temb=temb) - else: - hidden_states = downsampler(hidden_states) - - output_states += (hidden_states,) - - return hidden_states, output_states - - -class CrossAttnDownBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - downsample_padding: int = 1, - add_downsample: bool = True, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - attention_type: str = "default", - ): - super().__init__() - resnets = [] - attentions = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - if not dual_cross_attention: - attentions.append( - Transformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - attention_type=attention_type, - ) - ) - else: - attentions.append( - DualTransformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op" - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - additional_residuals: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - output_states = () - - blocks = list(zip(self.resnets, self.attentions)) - - for i, (resnet, attn) in enumerate(blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - else: - hidden_states = resnet(hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - # apply additional residuals to the output of the last pair of resnet and attention blocks - if i == len(blocks) - 1 and additional_residuals is not None: - hidden_states = hidden_states + additional_residuals - - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - output_states = output_states + (hidden_states,) - - return hidden_states, output_states - - -class DownBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_downsample: bool = True, - downsample_padding: int = 1, - ): - super().__init__() - resnets = [] - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op" - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None, *args, **kwargs - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - output_states = () - - for resnet in self.resnets: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - output_states = output_states + (hidden_states,) - - return hidden_states, output_states - - -class DownEncoderBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_downsample: bool = True, - downsample_padding: int = 1, - ): - super().__init__() - resnets = [] - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - if resnet_time_scale_shift == "spatial": - resnets.append( - ResnetBlockCondNorm2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=None, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm="spatial", - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - ) - ) - else: - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=None, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op" - ) - ] - ) - else: - self.downsamplers = None - - def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - for resnet in self.resnets: - hidden_states = resnet(hidden_states, temb=None) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - return hidden_states - - -class AttnDownEncoderBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - attention_head_dim: int = 1, - output_scale_factor: float = 1.0, - add_downsample: bool = True, - downsample_padding: int = 1, - ): - super().__init__() - resnets = [] - attentions = [] - - if attention_head_dim is None: - logger.warning( - f"It is not recommend to pass `attention_head_dim=None`. Defaulting `attention_head_dim` to `in_channels`: {out_channels}." - ) - attention_head_dim = out_channels - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - if resnet_time_scale_shift == "spatial": - resnets.append( - ResnetBlockCondNorm2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=None, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm="spatial", - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - ) - ) - else: - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=None, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - attentions.append( - Attention( - out_channels, - heads=out_channels // attention_head_dim, - dim_head=attention_head_dim, - rescale_output_factor=output_scale_factor, - eps=resnet_eps, - norm_num_groups=resnet_groups, - residual_connection=True, - bias=True, - upcast_softmax=True, - _from_deprecated_attn_block=True, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op" - ) - ] - ) - else: - self.downsamplers = None - - def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - for resnet, attn in zip(self.resnets, self.attentions): - hidden_states = resnet(hidden_states, temb=None) - hidden_states = attn(hidden_states) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - return hidden_states - - -class AttnSkipDownBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_pre_norm: bool = True, - attention_head_dim: int = 1, - output_scale_factor: float = np.sqrt(2.0), - add_downsample: bool = True, - ): - super().__init__() - self.attentions = nn.ModuleList([]) - self.resnets = nn.ModuleList([]) - - if attention_head_dim is None: - logger.warning( - f"It is not recommend to pass `attention_head_dim=None`. Defaulting `attention_head_dim` to `in_channels`: {out_channels}." - ) - attention_head_dim = out_channels - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - self.resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=min(in_channels // 4, 32), - groups_out=min(out_channels // 4, 32), - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - self.attentions.append( - Attention( - out_channels, - heads=out_channels // attention_head_dim, - dim_head=attention_head_dim, - rescale_output_factor=output_scale_factor, - eps=resnet_eps, - norm_num_groups=32, - residual_connection=True, - bias=True, - upcast_softmax=True, - _from_deprecated_attn_block=True, - ) - ) - - if add_downsample: - self.resnet_down = ResnetBlock2D( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=min(out_channels // 4, 32), - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_in_shortcut=True, - down=True, - kernel="fir", - ) - self.downsamplers = nn.ModuleList([FirDownsample2D(out_channels, out_channels=out_channels)]) - self.skip_conv = nn.Conv2d(3, out_channels, kernel_size=(1, 1), stride=(1, 1)) - else: - self.resnet_down = None - self.downsamplers = None - self.skip_conv = None - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - skip_sample: torch.Tensor | None = None, - *args, - **kwargs, - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...], torch.Tensor]: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - output_states = () - - for resnet, attn in zip(self.resnets, self.attentions): - hidden_states = resnet(hidden_states, temb) - hidden_states = attn(hidden_states) - output_states += (hidden_states,) - - if self.downsamplers is not None: - hidden_states = self.resnet_down(hidden_states, temb) - for downsampler in self.downsamplers: - skip_sample = downsampler(skip_sample) - - hidden_states = self.skip_conv(skip_sample) + hidden_states - - output_states += (hidden_states,) - - return hidden_states, output_states, skip_sample - - -class SkipDownBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_pre_norm: bool = True, - output_scale_factor: float = np.sqrt(2.0), - add_downsample: bool = True, - downsample_padding: int = 1, - ): - super().__init__() - self.resnets = nn.ModuleList([]) - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - self.resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=min(in_channels // 4, 32), - groups_out=min(out_channels // 4, 32), - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - if add_downsample: - self.resnet_down = ResnetBlock2D( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=min(out_channels // 4, 32), - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_in_shortcut=True, - down=True, - kernel="fir", - ) - self.downsamplers = nn.ModuleList([FirDownsample2D(out_channels, out_channels=out_channels)]) - self.skip_conv = nn.Conv2d(3, out_channels, kernel_size=(1, 1), stride=(1, 1)) - else: - self.resnet_down = None - self.downsamplers = None - self.skip_conv = None - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - skip_sample: torch.Tensor | None = None, - *args, - **kwargs, - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...], torch.Tensor]: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - output_states = () - - for resnet in self.resnets: - hidden_states = resnet(hidden_states, temb) - output_states += (hidden_states,) - - if self.downsamplers is not None: - hidden_states = self.resnet_down(hidden_states, temb) - for downsampler in self.downsamplers: - skip_sample = downsampler(skip_sample) - - hidden_states = self.skip_conv(skip_sample) + hidden_states - - output_states += (hidden_states,) - - return hidden_states, output_states, skip_sample - - -class ResnetDownsampleBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_downsample: bool = True, - skip_time_act: bool = False, - ): - super().__init__() - resnets = [] - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - skip_time_act=skip_time_act, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - ResnetBlock2D( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - skip_time_act=skip_time_act, - down=True, - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None, *args, **kwargs - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - output_states = () - - for resnet in self.resnets: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states, temb) - - output_states = output_states + (hidden_states,) - - return hidden_states, output_states - - -class SimpleCrossAttnDownBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - attention_head_dim: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - add_downsample: bool = True, - skip_time_act: bool = False, - only_cross_attention: bool = False, - cross_attention_norm: str | None = None, - ): - super().__init__() - - self.has_cross_attention = True - - resnets = [] - attentions = [] - - self.attention_head_dim = attention_head_dim - self.num_heads = out_channels // self.attention_head_dim - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - skip_time_act=skip_time_act, - ) - ) - - processor = ( - AttnAddedKVProcessor2_0() if hasattr(F, "scaled_dot_product_attention") else AttnAddedKVProcessor() - ) - - attentions.append( - Attention( - query_dim=out_channels, - cross_attention_dim=out_channels, - heads=self.num_heads, - dim_head=attention_head_dim, - added_kv_proj_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - bias=True, - upcast_softmax=True, - only_cross_attention=only_cross_attention, - cross_attention_norm=cross_attention_norm, - processor=processor, - ) - ) - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - ResnetBlock2D( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - skip_time_act=skip_time_act, - down=True, - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - output_states = () - - if attention_mask is None: - # if encoder_hidden_states is defined: we are doing cross-attn, so we should use cross-attn mask. - mask = None if encoder_hidden_states is None else encoder_attention_mask - else: - # when attention_mask is defined: we don't even check for encoder_attention_mask. - # this is to maintain compatibility with UnCLIP, which uses 'attention_mask' param for cross-attn masks. - # TODO: UnCLIP should express cross-attn mask via encoder_attention_mask param instead of via attention_mask. - # then we can simplify this whole if/else block to: - # mask = attention_mask if encoder_hidden_states is None else encoder_attention_mask - mask = attention_mask - - for resnet, attn in zip(self.resnets, self.attentions): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=mask, - **cross_attention_kwargs, - ) - else: - hidden_states = resnet(hidden_states, temb) - - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=mask, - **cross_attention_kwargs, - ) - - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states, temb) - - output_states = output_states + (hidden_states,) - - return hidden_states, output_states - - -class KDownBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 4, - resnet_eps: float = 1e-5, - resnet_act_fn: str = "gelu", - resnet_group_size: int = 32, - add_downsample: bool = False, - ): - super().__init__() - resnets = [] - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - groups = in_channels // resnet_group_size - groups_out = out_channels // resnet_group_size - - resnets.append( - ResnetBlockCondNorm2D( - in_channels=in_channels, - out_channels=out_channels, - dropout=dropout, - temb_channels=temb_channels, - groups=groups, - groups_out=groups_out, - eps=resnet_eps, - non_linearity=resnet_act_fn, - time_embedding_norm="ada_group", - conv_shortcut_bias=False, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - # YiYi's comments- might be able to use FirDownsample2D, look into details later - self.downsamplers = nn.ModuleList([KDownsample2D()]) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None, *args, **kwargs - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - output_states = () - - for resnet in self.resnets: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - - output_states += (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - return hidden_states, output_states - - -class KCrossAttnDownBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - cross_attention_dim: int, - dropout: float = 0.0, - num_layers: int = 4, - resnet_group_size: int = 32, - add_downsample: bool = True, - attention_head_dim: int = 64, - add_self_attention: bool = False, - resnet_eps: float = 1e-5, - resnet_act_fn: str = "gelu", - ): - super().__init__() - resnets = [] - attentions = [] - - self.has_cross_attention = True - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - groups = in_channels // resnet_group_size - groups_out = out_channels // resnet_group_size - - resnets.append( - ResnetBlockCondNorm2D( - in_channels=in_channels, - out_channels=out_channels, - dropout=dropout, - temb_channels=temb_channels, - groups=groups, - groups_out=groups_out, - eps=resnet_eps, - non_linearity=resnet_act_fn, - time_embedding_norm="ada_group", - conv_shortcut_bias=False, - ) - ) - attentions.append( - KAttentionBlock( - out_channels, - out_channels // attention_head_dim, - attention_head_dim, - cross_attention_dim=cross_attention_dim, - temb_channels=temb_channels, - attention_bias=True, - add_self_attention=add_self_attention, - cross_attention_norm="layer_norm", - group_size=resnet_group_size, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.attentions = nn.ModuleList(attentions) - - if add_downsample: - self.downsamplers = nn.ModuleList([KDownsample2D()]) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - output_states = () - - for resnet, attn in zip(self.resnets, self.attentions): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - resnet, - hidden_states, - temb, - ) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - emb=temb, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - encoder_attention_mask=encoder_attention_mask, - ) - else: - hidden_states = resnet(hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - emb=temb, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - encoder_attention_mask=encoder_attention_mask, - ) - - if self.downsamplers is None: - output_states += (None,) - else: - output_states += (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - return hidden_states, output_states - - -class AttnUpBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - prev_output_channel: int, - out_channels: int, - temb_channels: int, - resolution_idx: int = None, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - attention_head_dim: int = 1, - output_scale_factor: float = 1.0, - upsample_type: str = "conv", - ): - super().__init__() - resnets = [] - attentions = [] - - self.upsample_type = upsample_type - - if attention_head_dim is None: - logger.warning( - f"It is not recommend to pass `attention_head_dim=None`. Defaulting `attention_head_dim` to `in_channels`: {out_channels}." - ) - attention_head_dim = out_channels - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - attentions.append( - Attention( - out_channels, - heads=out_channels // attention_head_dim, - dim_head=attention_head_dim, - rescale_output_factor=output_scale_factor, - eps=resnet_eps, - norm_num_groups=resnet_groups, - residual_connection=True, - bias=True, - upcast_softmax=True, - _from_deprecated_attn_block=True, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if upsample_type == "conv": - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - elif upsample_type == "resnet": - self.upsamplers = nn.ModuleList( - [ - ResnetBlock2D( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - up=True, - ) - ] - ) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - upsample_size: int | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - for resnet, attn in zip(self.resnets, self.attentions): - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - hidden_states = attn(hidden_states) - else: - hidden_states = resnet(hidden_states, temb) - hidden_states = attn(hidden_states) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - if self.upsample_type == "resnet": - hidden_states = upsampler(hidden_states, temb=temb) - else: - hidden_states = upsampler(hidden_states) - - return hidden_states - - -class CrossAttnUpBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - prev_output_channel: int, - temb_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - attention_type: str = "default", - ): - super().__init__() - resnets = [] - attentions = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - if not dual_cross_attention: - attentions.append( - Transformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - attention_type=attention_type, - ) - ) - else: - attentions.append( - DualTransformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - upsample_size: int | None = None, - attention_mask: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - is_freeu_enabled = ( - getattr(self, "s1", None) - and getattr(self, "s2", None) - and getattr(self, "b1", None) - and getattr(self, "b2", None) - ) - - for resnet, attn in zip(self.resnets, self.attentions): - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - - # FreeU: Only operate on the first two stages - if is_freeu_enabled: - hidden_states, res_hidden_states = apply_freeu( - self.resolution_idx, - hidden_states, - res_hidden_states, - s1=self.s1, - s2=self.s2, - b1=self.b1, - b2=self.b2, - ) - - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - else: - hidden_states = resnet(hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size) - - return hidden_states - - -class UpBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - prev_output_channel: int, - out_channels: int, - temb_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - ): - super().__init__() - resnets = [] - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - upsample_size: int | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - is_freeu_enabled = ( - getattr(self, "s1", None) - and getattr(self, "s2", None) - and getattr(self, "b1", None) - and getattr(self, "b2", None) - ) - - for resnet in self.resnets: - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - - # FreeU: Only operate on the first two stages - if is_freeu_enabled: - hidden_states, res_hidden_states = apply_freeu( - self.resolution_idx, - hidden_states, - res_hidden_states, - s1=self.s1, - s2=self.s2, - b1=self.b1, - b2=self.b2, - ) - - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size) - - return hidden_states - - -class UpDecoderBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", # default, spatial - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - temb_channels: int | None = None, - ): - super().__init__() - resnets = [] - - for i in range(num_layers): - input_channels = in_channels if i == 0 else out_channels - - if resnet_time_scale_shift == "spatial": - resnets.append( - ResnetBlockCondNorm2D( - in_channels=input_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm="spatial", - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - ) - ) - else: - resnets.append( - ResnetBlock2D( - in_channels=input_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.resolution_idx = resolution_idx - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - for resnet in self.resnets: - hidden_states = resnet(hidden_states, temb=temb) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - return hidden_states - - -class AttnUpDecoderBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - attention_head_dim: int = 1, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - temb_channels: int | None = None, - ): - super().__init__() - resnets = [] - attentions = [] - - if attention_head_dim is None: - logger.warning( - f"It is not recommend to pass `attention_head_dim=None`. Defaulting `attention_head_dim` to `out_channels`: {out_channels}." - ) - attention_head_dim = out_channels - - for i in range(num_layers): - input_channels = in_channels if i == 0 else out_channels - - if resnet_time_scale_shift == "spatial": - resnets.append( - ResnetBlockCondNorm2D( - in_channels=input_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm="spatial", - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - ) - ) - else: - resnets.append( - ResnetBlock2D( - in_channels=input_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - attentions.append( - Attention( - out_channels, - heads=out_channels // attention_head_dim, - dim_head=attention_head_dim, - rescale_output_factor=output_scale_factor, - eps=resnet_eps, - norm_num_groups=resnet_groups if resnet_time_scale_shift != "spatial" else None, - spatial_norm_dim=temb_channels if resnet_time_scale_shift == "spatial" else None, - residual_connection=True, - bias=True, - upcast_softmax=True, - _from_deprecated_attn_block=True, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.resolution_idx = resolution_idx - - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor: - for resnet, attn in zip(self.resnets, self.attentions): - hidden_states = resnet(hidden_states, temb=temb) - hidden_states = attn(hidden_states, temb=temb) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - return hidden_states - - -class AttnSkipUpBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - prev_output_channel: int, - out_channels: int, - temb_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_pre_norm: bool = True, - attention_head_dim: int = 1, - output_scale_factor: float = np.sqrt(2.0), - add_upsample: bool = True, - ): - super().__init__() - self.attentions = nn.ModuleList([]) - self.resnets = nn.ModuleList([]) - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - self.resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=min(resnet_in_channels + res_skip_channels // 4, 32), - groups_out=min(out_channels // 4, 32), - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - if attention_head_dim is None: - logger.warning( - f"It is not recommend to pass `attention_head_dim=None`. Defaulting `attention_head_dim` to `out_channels`: {out_channels}." - ) - attention_head_dim = out_channels - - self.attentions.append( - Attention( - out_channels, - heads=out_channels // attention_head_dim, - dim_head=attention_head_dim, - rescale_output_factor=output_scale_factor, - eps=resnet_eps, - norm_num_groups=32, - residual_connection=True, - bias=True, - upcast_softmax=True, - _from_deprecated_attn_block=True, - ) - ) - - self.upsampler = FirUpsample2D(in_channels, out_channels=out_channels) - if add_upsample: - self.resnet_up = ResnetBlock2D( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=min(out_channels // 4, 32), - groups_out=min(out_channels // 4, 32), - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_in_shortcut=True, - up=True, - kernel="fir", - ) - self.skip_conv = nn.Conv2d(out_channels, 3, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) - self.skip_norm = torch.nn.GroupNorm( - num_groups=min(out_channels // 4, 32), num_channels=out_channels, eps=resnet_eps, affine=True - ) - self.act = nn.SiLU() - else: - self.resnet_up = None - self.skip_conv = None - self.skip_norm = None - self.act = None - - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - skip_sample=None, - *args, - **kwargs, - ) -> tuple[torch.Tensor, torch.Tensor]: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - for resnet in self.resnets: - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - hidden_states = resnet(hidden_states, temb) - - hidden_states = self.attentions[0](hidden_states) - - if skip_sample is not None: - skip_sample = self.upsampler(skip_sample) - else: - skip_sample = 0 - - if self.resnet_up is not None: - skip_sample_states = self.skip_norm(hidden_states) - skip_sample_states = self.act(skip_sample_states) - skip_sample_states = self.skip_conv(skip_sample_states) - - skip_sample = skip_sample + skip_sample_states - - hidden_states = self.resnet_up(hidden_states, temb) - - return hidden_states, skip_sample - - -class SkipUpBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - prev_output_channel: int, - out_channels: int, - temb_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_pre_norm: bool = True, - output_scale_factor: float = np.sqrt(2.0), - add_upsample: bool = True, - upsample_padding: int = 1, - ): - super().__init__() - self.resnets = nn.ModuleList([]) - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - self.resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=min((resnet_in_channels + res_skip_channels) // 4, 32), - groups_out=min(out_channels // 4, 32), - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - self.upsampler = FirUpsample2D(in_channels, out_channels=out_channels) - if add_upsample: - self.resnet_up = ResnetBlock2D( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=min(out_channels // 4, 32), - groups_out=min(out_channels // 4, 32), - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_in_shortcut=True, - up=True, - kernel="fir", - ) - self.skip_conv = nn.Conv2d(out_channels, 3, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) - self.skip_norm = torch.nn.GroupNorm( - num_groups=min(out_channels // 4, 32), num_channels=out_channels, eps=resnet_eps, affine=True - ) - self.act = nn.SiLU() - else: - self.resnet_up = None - self.skip_conv = None - self.skip_norm = None - self.act = None - - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - skip_sample=None, - *args, - **kwargs, - ) -> tuple[torch.Tensor, torch.Tensor]: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - for resnet in self.resnets: - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - hidden_states = resnet(hidden_states, temb) - - if skip_sample is not None: - skip_sample = self.upsampler(skip_sample) - else: - skip_sample = 0 - - if self.resnet_up is not None: - skip_sample_states = self.skip_norm(hidden_states) - skip_sample_states = self.act(skip_sample_states) - skip_sample_states = self.skip_conv(skip_sample_states) - - skip_sample = skip_sample + skip_sample_states - - hidden_states = self.resnet_up(hidden_states, temb) - - return hidden_states, skip_sample - - -class ResnetUpsampleBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - prev_output_channel: int, - out_channels: int, - temb_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - skip_time_act: bool = False, - ): - super().__init__() - resnets = [] - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - skip_time_act=skip_time_act, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList( - [ - ResnetBlock2D( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - skip_time_act=skip_time_act, - up=True, - ) - ] - ) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - upsample_size: int | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - for resnet in self.resnets: - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, temb) - - return hidden_states - - -class SimpleCrossAttnUpBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - prev_output_channel: int, - temb_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - attention_head_dim: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - skip_time_act: bool = False, - only_cross_attention: bool = False, - cross_attention_norm: str | None = None, - ): - super().__init__() - resnets = [] - attentions = [] - - self.has_cross_attention = True - self.attention_head_dim = attention_head_dim - - self.num_heads = out_channels // self.attention_head_dim - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - skip_time_act=skip_time_act, - ) - ) - - processor = ( - AttnAddedKVProcessor2_0() if hasattr(F, "scaled_dot_product_attention") else AttnAddedKVProcessor() - ) - - attentions.append( - Attention( - query_dim=out_channels, - cross_attention_dim=out_channels, - heads=self.num_heads, - dim_head=self.attention_head_dim, - added_kv_proj_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - bias=True, - upcast_softmax=True, - only_cross_attention=only_cross_attention, - cross_attention_norm=cross_attention_norm, - processor=processor, - ) - ) - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList( - [ - ResnetBlock2D( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - skip_time_act=skip_time_act, - up=True, - ) - ] - ) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - upsample_size: int | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - if attention_mask is None: - # if encoder_hidden_states is defined: we are doing cross-attn, so we should use cross-attn mask. - mask = None if encoder_hidden_states is None else encoder_attention_mask - else: - # when attention_mask is defined: we don't even check for encoder_attention_mask. - # this is to maintain compatibility with UnCLIP, which uses 'attention_mask' param for cross-attn masks. - # TODO: UnCLIP should express cross-attn mask via encoder_attention_mask param instead of via attention_mask. - # then we can simplify this whole if/else block to: - # mask = attention_mask if encoder_hidden_states is None else encoder_attention_mask - mask = attention_mask - - for resnet, attn in zip(self.resnets, self.attentions): - # resnet - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=mask, - **cross_attention_kwargs, - ) - else: - hidden_states = resnet(hidden_states, temb) - - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=mask, - **cross_attention_kwargs, - ) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, temb) - - return hidden_states - - -class KUpBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - resolution_idx: int, - dropout: float = 0.0, - num_layers: int = 5, - resnet_eps: float = 1e-5, - resnet_act_fn: str = "gelu", - resnet_group_size: int | None = 32, - add_upsample: bool = True, - ): - super().__init__() - resnets = [] - k_in_channels = 2 * out_channels - k_out_channels = in_channels - num_layers = num_layers - 1 - - for i in range(num_layers): - in_channels = k_in_channels if i == 0 else out_channels - groups = in_channels // resnet_group_size - groups_out = out_channels // resnet_group_size - - resnets.append( - ResnetBlockCondNorm2D( - in_channels=in_channels, - out_channels=k_out_channels if (i == num_layers - 1) else out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=groups, - groups_out=groups_out, - dropout=dropout, - non_linearity=resnet_act_fn, - time_embedding_norm="ada_group", - conv_shortcut_bias=False, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList([KUpsample2D()]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - upsample_size: int | None = None, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - res_hidden_states_tuple = res_hidden_states_tuple[-1] - if res_hidden_states_tuple is not None: - hidden_states = torch.cat([hidden_states, res_hidden_states_tuple], dim=1) - - for resnet in self.resnets: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - return hidden_states - - -class KCrossAttnUpBlock2D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - resolution_idx: int, - dropout: float = 0.0, - num_layers: int = 4, - resnet_eps: float = 1e-5, - resnet_act_fn: str = "gelu", - resnet_group_size: int = 32, - attention_head_dim: int = 1, # attention dim_head - cross_attention_dim: int = 768, - add_upsample: bool = True, - upcast_attention: bool = False, - ): - super().__init__() - resnets = [] - attentions = [] - - is_first_block = in_channels == out_channels == temb_channels - is_middle_block = in_channels != out_channels - add_self_attention = True if is_first_block else False - - self.has_cross_attention = True - self.attention_head_dim = attention_head_dim - - # in_channels, and out_channels for the block (k-unet) - k_in_channels = out_channels if is_first_block else 2 * out_channels - k_out_channels = in_channels - - num_layers = num_layers - 1 - - for i in range(num_layers): - in_channels = k_in_channels if i == 0 else out_channels - groups = in_channels // resnet_group_size - groups_out = out_channels // resnet_group_size - - if is_middle_block and (i == num_layers - 1): - conv_2d_out_channels = k_out_channels - else: - conv_2d_out_channels = None - - resnets.append( - ResnetBlockCondNorm2D( - in_channels=in_channels, - out_channels=out_channels, - conv_2d_out_channels=conv_2d_out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=groups, - groups_out=groups_out, - dropout=dropout, - non_linearity=resnet_act_fn, - time_embedding_norm="ada_group", - conv_shortcut_bias=False, - ) - ) - attentions.append( - KAttentionBlock( - k_out_channels if (i == num_layers - 1) else out_channels, - k_out_channels // attention_head_dim - if (i == num_layers - 1) - else out_channels // attention_head_dim, - attention_head_dim, - cross_attention_dim=cross_attention_dim, - temb_channels=temb_channels, - attention_bias=True, - add_self_attention=add_self_attention, - cross_attention_norm="layer_norm", - upcast_attention=upcast_attention, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.attentions = nn.ModuleList(attentions) - - if add_upsample: - self.upsamplers = nn.ModuleList([KUpsample2D()]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - upsample_size: int | None = None, - attention_mask: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - res_hidden_states_tuple = res_hidden_states_tuple[-1] - if res_hidden_states_tuple is not None: - hidden_states = torch.cat([hidden_states, res_hidden_states_tuple], dim=1) - - for resnet, attn in zip(self.resnets, self.attentions): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - resnet, - hidden_states, - temb, - ) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - emb=temb, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - encoder_attention_mask=encoder_attention_mask, - ) - else: - hidden_states = resnet(hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - emb=temb, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - encoder_attention_mask=encoder_attention_mask, - ) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - return hidden_states - - -# can potentially later be renamed to `No-feed-forward` attention -class KAttentionBlock(nn.Module): - r""" - A basic Transformer block. - - Parameters: - dim (`int`): The number of channels in the input and output. - num_attention_heads (`int`): The number of heads to use for multi-head attention. - attention_head_dim (`int`): The number of channels in each head. - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention. - attention_bias (`bool`, *optional*, defaults to `False`): - Configure if the attention layers should contain a bias parameter. - upcast_attention (`bool`, *optional*, defaults to `False`): - Set to `True` to upcast the attention computation to `float32`. - temb_channels (`int`, *optional*, defaults to 768): - The number of channels in the token embedding. - add_self_attention (`bool`, *optional*, defaults to `False`): - Set to `True` to add self-attention to the block. - cross_attention_norm (`str`, *optional*, defaults to `None`): - The type of normalization to use for the cross attention. Can be `None`, `layer_norm`, or `group_norm`. - group_size (`int`, *optional*, defaults to 32): - The number of groups to separate the channels into for group normalization. - """ - - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - dropout: float = 0.0, - cross_attention_dim: int | None = None, - attention_bias: bool = False, - upcast_attention: bool = False, - temb_channels: int = 768, # for ada_group_norm - add_self_attention: bool = False, - cross_attention_norm: str | None = None, - group_size: int = 32, - ): - super().__init__() - self.add_self_attention = add_self_attention - - # 1. Self-Attn - if add_self_attention: - self.norm1 = AdaGroupNorm(temb_channels, dim, max(1, dim // group_size)) - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - cross_attention_dim=None, - cross_attention_norm=None, - ) - - # 2. Cross-Attn - self.norm2 = AdaGroupNorm(temb_channels, dim, max(1, dim // group_size)) - self.attn2 = Attention( - query_dim=dim, - cross_attention_dim=cross_attention_dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=attention_bias, - upcast_attention=upcast_attention, - cross_attention_norm=cross_attention_norm, - ) - - def _to_3d(self, hidden_states: torch.Tensor, height: int, weight: int) -> torch.Tensor: - return hidden_states.permute(0, 2, 3, 1).reshape(hidden_states.shape[0], height * weight, -1) - - def _to_4d(self, hidden_states: torch.Tensor, height: int, weight: int) -> torch.Tensor: - return hidden_states.permute(0, 2, 1).reshape(hidden_states.shape[0], -1, height, weight) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - # TODO: mark emb as non-optional (self.norm2 requires it). - # requires assessing impact of change to positional param interface. - emb: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - # 1. Self-Attention - if self.add_self_attention: - norm_hidden_states = self.norm1(hidden_states, emb) - - height, weight = norm_hidden_states.shape[2:] - norm_hidden_states = self._to_3d(norm_hidden_states, height, weight) - - attn_output = self.attn1( - norm_hidden_states, - encoder_hidden_states=None, - attention_mask=attention_mask, - **cross_attention_kwargs, - ) - attn_output = self._to_4d(attn_output, height, weight) - - hidden_states = attn_output + hidden_states - - # 2. Cross-Attention/None - norm_hidden_states = self.norm2(hidden_states, emb) - - height, weight = norm_hidden_states.shape[2:] - norm_hidden_states = self._to_3d(norm_hidden_states, height, weight) - attn_output = self.attn2( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask if encoder_hidden_states is None else encoder_attention_mask, - **cross_attention_kwargs, - ) - attn_output = self._to_4d(attn_output, height, weight) - - hidden_states = attn_output + hidden_states - - return hidden_states diff --git a/diffusers/models/unets/unet_2d_condition.py b/diffusers/models/unets/unet_2d_condition.py deleted file mode 100644 index af44f0e9d2cb003ba01bbe8f11a7988c30573359..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_2d_condition.py +++ /dev/null @@ -1,1235 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from dataclasses import dataclass -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin, UNet2DConditionLoadersMixin -from ...loaders.single_file_model import FromOriginalModelMixin -from ...utils import ( - BaseOutput, - apply_lora_scale, - deprecate, - logging, -) -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..activations import get_activation -from ..attention import AttentionMixin -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - Attention, - AttnAddedKVProcessor, - AttnProcessor, - FusedAttnProcessor2_0, -) -from ..embeddings import ( - GaussianFourierProjection, - GLIGENTextBoundingboxProjection, - ImageHintTimeEmbedding, - ImageProjection, - ImageTimeEmbedding, - TextImageProjection, - TextImageTimeEmbedding, - TextTimeEmbedding, - TimestepEmbedding, - Timesteps, -) -from ..modeling_utils import ModelMixin -from .unet_2d_blocks import ( - get_down_block, - get_mid_block, - get_up_block, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class UNet2DConditionOutput(BaseOutput): - """ - The output of [`UNet2DConditionModel`]. - - Args: - sample (`torch.Tensor` of shape `(batch_size, num_channels, height, width)`): - The hidden states output conditioned on `encoder_hidden_states` input. Output of last layer of model. - """ - - sample: torch.Tensor = None - - -class UNet2DConditionModel( - ModelMixin, AttentionMixin, ConfigMixin, FromOriginalModelMixin, UNet2DConditionLoadersMixin, PeftAdapterMixin -): - r""" - A conditional 2D UNet model that takes a noisy sample, conditional state, and a timestep and returns a sample - shaped output. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - sample_size (`int` or `tuple[int, int]`, *optional*, defaults to `None`): - Height and width of input/output sample. - in_channels (`int`, *optional*, defaults to 4): Number of channels in the input sample. - out_channels (`int`, *optional*, defaults to 4): Number of channels in the output. - center_input_sample (`bool`, *optional*, defaults to `False`): Whether to center the input sample. - flip_sin_to_cos (`bool`, *optional*, defaults to `True`): - Whether to flip the sin to cos in the time embedding. - freq_shift (`int`, *optional*, defaults to 0): The frequency shift to apply to the time embedding. - down_block_types (`tuple[str]`, *optional*, defaults to `("CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "DownBlock2D")`): - The tuple of downsample blocks to use. - mid_block_type (`str`, *optional*, defaults to `"UNetMidBlock2DCrossAttn"`): - Block type for middle of UNet, it can be one of `UNetMidBlock2DCrossAttn`, `UNetMidBlock2D`, or - `UNetMidBlock2DSimpleCrossAttn`. If `None`, the mid block layer is skipped. - up_block_types (`tuple[str]`, *optional*, defaults to `("UpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D")`): - The tuple of upsample blocks to use. - only_cross_attention(`bool` or `tuple[bool]`, *optional*, default to `False`): - Whether to include self-attention in the basic transformer blocks, see - [`~models.attention.BasicTransformerBlock`]. - block_out_channels (`tuple[int]`, *optional*, defaults to `(320, 640, 1280, 1280)`): - The tuple of output channels for each block. - layers_per_block (`int`, *optional*, defaults to 2): The number of layers per block. - downsample_padding (`int`, *optional*, defaults to 1): The padding to use for the downsampling convolution. - mid_block_scale_factor (`float`, *optional*, defaults to 1.0): The scale factor to use for the mid block. - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use. - norm_num_groups (`int`, *optional*, defaults to 32): The number of groups to use for the normalization. - If `None`, normalization and activation layers is skipped in post-processing. - norm_eps (`float`, *optional*, defaults to 1e-5): The epsilon to use for the normalization. - cross_attention_dim (`int` or `tuple[int]`, *optional*, defaults to 1280): - The dimension of the cross attention features. - transformer_layers_per_block (`int`, `tuple[int]`, or `tuple[tuple]` , *optional*, defaults to 1): - The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`]. Only relevant for - [`~models.unets.unet_2d_blocks.CrossAttnDownBlock2D`], [`~models.unets.unet_2d_blocks.CrossAttnUpBlock2D`], - [`~models.unets.unet_2d_blocks.UNetMidBlock2DCrossAttn`]. - reverse_transformer_layers_per_block : (`tuple[tuple]`, *optional*, defaults to None): - The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`], in the upsampling - blocks of the U-Net. Only relevant if `transformer_layers_per_block` is of type `tuple[tuple]` and for - [`~models.unets.unet_2d_blocks.CrossAttnDownBlock2D`], [`~models.unets.unet_2d_blocks.CrossAttnUpBlock2D`], - [`~models.unets.unet_2d_blocks.UNetMidBlock2DCrossAttn`]. - encoder_hid_dim (`int`, *optional*, defaults to None): - If `encoder_hid_dim_type` is defined, `encoder_hidden_states` will be projected from `encoder_hid_dim` - dimension to `cross_attention_dim`. - encoder_hid_dim_type (`str`, *optional*, defaults to `None`): - If given, the `encoder_hidden_states` and potentially other embeddings are down-projected to text - embeddings of dimension `cross_attention` according to `encoder_hid_dim_type`. - attention_head_dim (`int`, *optional*, defaults to 8): The dimension of the attention heads. - num_attention_heads (`int`, *optional*): - The number of attention heads. If not defined, defaults to `attention_head_dim` - resnet_time_scale_shift (`str`, *optional*, defaults to `"default"`): Time scale shift config - for ResNet blocks (see [`~models.resnet.ResnetBlock2D`]). Choose from `default` or `scale_shift`. - class_embed_type (`str`, *optional*, defaults to `None`): - The type of class embedding to use which is ultimately summed with the time embeddings. Choose from `None`, - `"timestep"`, `"identity"`, `"projection"`, or `"simple_projection"`. - addition_embed_type (`str`, *optional*, defaults to `None`): - Configures an optional embedding which will be summed with the time embeddings. Choose from `None` or - "text". "text" will use the `TextTimeEmbedding` layer. - addition_time_embed_dim: (`int`, *optional*, defaults to `None`): - Dimension for the timestep embeddings. - num_class_embeds (`int`, *optional*, defaults to `None`): - Input dimension of the learnable embedding matrix to be projected to `time_embed_dim`, when performing - class conditioning with `class_embed_type` equal to `None`. - time_embedding_type (`str`, *optional*, defaults to `positional`): - The type of position embedding to use for timesteps. Choose from `positional` or `fourier`. - time_embedding_dim (`int`, *optional*, defaults to `None`): - An optional override for the dimension of the projected time embedding. - time_embedding_act_fn (`str`, *optional*, defaults to `None`): - Optional activation function to use only once on the time embeddings before they are passed to the rest of - the UNet. Choose from `silu`, `mish`, `gelu`, and `swish`. - timestep_post_act (`str`, *optional*, defaults to `None`): - The second activation function to use in timestep embedding. Choose from `silu`, `mish` and `gelu`. - time_cond_proj_dim (`int`, *optional*, defaults to `None`): - The dimension of `cond_proj` layer in the timestep embedding. - conv_in_kernel (`int`, *optional*, default to `3`): The kernel size of `conv_in` layer. - conv_out_kernel (`int`, *optional*, default to `3`): The kernel size of `conv_out` layer. - projection_class_embeddings_input_dim (`int`, *optional*): The dimension of the `class_labels` input when - `class_embed_type="projection"`. Required when `class_embed_type="projection"`. - class_embeddings_concat (`bool`, *optional*, defaults to `False`): Whether to concatenate the time - embeddings with the class embeddings. - mid_block_only_cross_attention (`bool`, *optional*, defaults to `None`): - Whether to use cross attention with the mid block when using the `UNetMidBlock2DSimpleCrossAttn`. If - `only_cross_attention` is given as a single boolean and `mid_block_only_cross_attention` is `None`, the - `only_cross_attention` value is used as the value for `mid_block_only_cross_attention`. Default to `False` - otherwise. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = ["BasicTransformerBlock", "ResnetBlock2D", "CrossAttnUpBlock2D", "UpBlock2D"] - _skip_layerwise_casting_patterns = ["norm"] - _repeated_blocks = ["BasicTransformerBlock"] - - @register_to_config - def __init__( - self, - sample_size: int | tuple[int, int] | None = None, - in_channels: int = 4, - out_channels: int = 4, - center_input_sample: bool = False, - flip_sin_to_cos: bool = True, - freq_shift: int = 0, - down_block_types: tuple[str, ...] = ( - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "DownBlock2D", - ), - mid_block_type: str | None = "UNetMidBlock2DCrossAttn", - up_block_types: tuple[str, ...] = ( - "UpBlock2D", - "CrossAttnUpBlock2D", - "CrossAttnUpBlock2D", - "CrossAttnUpBlock2D", - ), - only_cross_attention: bool | tuple[bool] = False, - block_out_channels: tuple[int, ...] = (320, 640, 1280, 1280), - layers_per_block: int | tuple[int] = 2, - downsample_padding: int = 1, - mid_block_scale_factor: float = 1, - dropout: float = 0.0, - act_fn: str = "silu", - norm_num_groups: int | None = 32, - norm_eps: float = 1e-5, - cross_attention_dim: int | tuple[int] = 1280, - transformer_layers_per_block: int | tuple[int] | tuple[tuple] = 1, - reverse_transformer_layers_per_block: tuple[tuple[int]] | None = None, - encoder_hid_dim: int | None = None, - encoder_hid_dim_type: str | None = None, - attention_head_dim: int | tuple[int] = 8, - num_attention_heads: int | tuple[int] | None = None, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - class_embed_type: str | None = None, - addition_embed_type: str | None = None, - addition_time_embed_dim: int | None = None, - num_class_embeds: int | None = None, - upcast_attention: bool = False, - resnet_time_scale_shift: str = "default", - resnet_skip_time_act: bool = False, - resnet_out_scale_factor: float = 1.0, - time_embedding_type: str = "positional", - time_embedding_dim: int | None = None, - time_embedding_act_fn: str | None = None, - timestep_post_act: str | None = None, - time_cond_proj_dim: int | None = None, - conv_in_kernel: int = 3, - conv_out_kernel: int = 3, - projection_class_embeddings_input_dim: int | None = None, - attention_type: str = "default", - class_embeddings_concat: bool = False, - mid_block_only_cross_attention: bool | None = None, - cross_attention_norm: str | None = None, - addition_embed_type_num_heads: int = 64, - ): - super().__init__() - - self.sample_size = sample_size - - if num_attention_heads is not None: - raise ValueError( - "At the moment it is not possible to define the number of attention heads via `num_attention_heads` because of a naming issue as described in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131. Passing `num_attention_heads` will only be supported in diffusers v0.19." - ) - - # If `num_attention_heads` is not defined (which is the case for most models) - # it will default to `attention_head_dim`. This looks weird upon first reading it and it is. - # The reason for this behavior is to correct for incorrectly named variables that were introduced - # when this library was created. The incorrect naming was only discovered much later in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131 - # Changing `attention_head_dim` to `num_attention_heads` for 40,000+ configurations is too backwards breaking - # which is why we correct for the naming here. - num_attention_heads = num_attention_heads or attention_head_dim - - # Check inputs - self._check_config( - down_block_types=down_block_types, - up_block_types=up_block_types, - only_cross_attention=only_cross_attention, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - cross_attention_dim=cross_attention_dim, - transformer_layers_per_block=transformer_layers_per_block, - reverse_transformer_layers_per_block=reverse_transformer_layers_per_block, - attention_head_dim=attention_head_dim, - num_attention_heads=num_attention_heads, - ) - - # input - conv_in_padding = (conv_in_kernel - 1) // 2 - self.conv_in = nn.Conv2d( - in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding - ) - - # time - time_embed_dim, timestep_input_dim = self._set_time_proj( - time_embedding_type, - block_out_channels=block_out_channels, - flip_sin_to_cos=flip_sin_to_cos, - freq_shift=freq_shift, - time_embedding_dim=time_embedding_dim, - ) - - self.time_embedding = TimestepEmbedding( - timestep_input_dim, - time_embed_dim, - act_fn=act_fn, - post_act_fn=timestep_post_act, - cond_proj_dim=time_cond_proj_dim, - ) - - self._set_encoder_hid_proj( - encoder_hid_dim_type, - cross_attention_dim=cross_attention_dim, - encoder_hid_dim=encoder_hid_dim, - ) - - # class embedding - self._set_class_embedding( - class_embed_type, - act_fn=act_fn, - num_class_embeds=num_class_embeds, - projection_class_embeddings_input_dim=projection_class_embeddings_input_dim, - time_embed_dim=time_embed_dim, - timestep_input_dim=timestep_input_dim, - ) - - self._set_add_embedding( - addition_embed_type, - addition_embed_type_num_heads=addition_embed_type_num_heads, - addition_time_embed_dim=addition_time_embed_dim, - cross_attention_dim=cross_attention_dim, - encoder_hid_dim=encoder_hid_dim, - flip_sin_to_cos=flip_sin_to_cos, - freq_shift=freq_shift, - projection_class_embeddings_input_dim=projection_class_embeddings_input_dim, - time_embed_dim=time_embed_dim, - ) - - if time_embedding_act_fn is None: - self.time_embed_act = None - else: - self.time_embed_act = get_activation(time_embedding_act_fn) - - self.down_blocks = nn.ModuleList([]) - self.up_blocks = nn.ModuleList([]) - - if isinstance(only_cross_attention, bool): - if mid_block_only_cross_attention is None: - mid_block_only_cross_attention = only_cross_attention - - only_cross_attention = [only_cross_attention] * len(down_block_types) - - if mid_block_only_cross_attention is None: - mid_block_only_cross_attention = False - - if isinstance(num_attention_heads, int): - num_attention_heads = (num_attention_heads,) * len(down_block_types) - - if isinstance(attention_head_dim, int): - attention_head_dim = (attention_head_dim,) * len(down_block_types) - - if isinstance(cross_attention_dim, int): - cross_attention_dim = (cross_attention_dim,) * len(down_block_types) - - if isinstance(layers_per_block, int): - layers_per_block = [layers_per_block] * len(down_block_types) - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types) - - if class_embeddings_concat: - # The time embeddings are concatenated with the class embeddings. The dimension of the - # time embeddings passed to the down, middle, and up blocks is twice the dimension of the - # regular time embeddings - blocks_time_embed_dim = time_embed_dim * 2 - else: - blocks_time_embed_dim = time_embed_dim - - # down - output_channel = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - down_block = get_down_block( - down_block_type, - num_layers=layers_per_block[i], - transformer_layers_per_block=transformer_layers_per_block[i], - in_channels=input_channel, - out_channels=output_channel, - temb_channels=blocks_time_embed_dim, - add_downsample=not is_final_block, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - cross_attention_dim=cross_attention_dim[i], - num_attention_heads=num_attention_heads[i], - downsample_padding=downsample_padding, - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention[i], - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - attention_type=attention_type, - resnet_skip_time_act=resnet_skip_time_act, - resnet_out_scale_factor=resnet_out_scale_factor, - cross_attention_norm=cross_attention_norm, - attention_head_dim=attention_head_dim[i] if attention_head_dim[i] is not None else output_channel, - dropout=dropout, - ) - self.down_blocks.append(down_block) - - # mid - self.mid_block = get_mid_block( - mid_block_type, - temb_channels=blocks_time_embed_dim, - in_channels=block_out_channels[-1], - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - output_scale_factor=mid_block_scale_factor, - transformer_layers_per_block=transformer_layers_per_block[-1], - num_attention_heads=num_attention_heads[-1], - cross_attention_dim=cross_attention_dim[-1], - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - mid_block_only_cross_attention=mid_block_only_cross_attention, - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - attention_type=attention_type, - resnet_skip_time_act=resnet_skip_time_act, - cross_attention_norm=cross_attention_norm, - attention_head_dim=attention_head_dim[-1], - dropout=dropout, - ) - - # count how many layers upsample the images - self.num_upsamplers = 0 - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - reversed_num_attention_heads = list(reversed(num_attention_heads)) - reversed_layers_per_block = list(reversed(layers_per_block)) - reversed_cross_attention_dim = list(reversed(cross_attention_dim)) - reversed_transformer_layers_per_block = ( - list(reversed(transformer_layers_per_block)) - if reverse_transformer_layers_per_block is None - else reverse_transformer_layers_per_block - ) - only_cross_attention = list(reversed(only_cross_attention)) - - output_channel = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - is_final_block = i == len(block_out_channels) - 1 - - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)] - - # add upsample block for all BUT final layer - if not is_final_block: - add_upsample = True - self.num_upsamplers += 1 - else: - add_upsample = False - - up_block = get_up_block( - up_block_type, - num_layers=reversed_layers_per_block[i] + 1, - transformer_layers_per_block=reversed_transformer_layers_per_block[i], - in_channels=input_channel, - out_channels=output_channel, - prev_output_channel=prev_output_channel, - temb_channels=blocks_time_embed_dim, - add_upsample=add_upsample, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resolution_idx=i, - resnet_groups=norm_num_groups, - cross_attention_dim=reversed_cross_attention_dim[i], - num_attention_heads=reversed_num_attention_heads[i], - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention[i], - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - attention_type=attention_type, - resnet_skip_time_act=resnet_skip_time_act, - resnet_out_scale_factor=resnet_out_scale_factor, - cross_attention_norm=cross_attention_norm, - attention_head_dim=attention_head_dim[i] if attention_head_dim[i] is not None else output_channel, - dropout=dropout, - ) - self.up_blocks.append(up_block) - - # out - if norm_num_groups is not None: - self.conv_norm_out = nn.GroupNorm( - num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=norm_eps - ) - - self.conv_act = get_activation(act_fn) - - else: - self.conv_norm_out = None - self.conv_act = None - - conv_out_padding = (conv_out_kernel - 1) // 2 - self.conv_out = nn.Conv2d( - block_out_channels[0], out_channels, kernel_size=conv_out_kernel, padding=conv_out_padding - ) - - self._set_pos_net_if_use_gligen(attention_type=attention_type, cross_attention_dim=cross_attention_dim) - - def _check_config( - self, - down_block_types: tuple[str, ...], - up_block_types: tuple[str, ...], - only_cross_attention: bool | tuple[bool], - block_out_channels: tuple[int, ...], - layers_per_block: int | tuple[int], - cross_attention_dim: int | tuple[int], - transformer_layers_per_block: int | tuple[int, tuple[tuple[int]]], - reverse_transformer_layers_per_block: bool, - attention_head_dim: int, - num_attention_heads: int | tuple[int] | None, - ): - if len(down_block_types) != len(up_block_types): - raise ValueError( - f"Must provide the same number of `down_block_types` as `up_block_types`. `down_block_types`: {down_block_types}. `up_block_types`: {up_block_types}." - ) - - if len(block_out_channels) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(only_cross_attention, bool) and len(only_cross_attention) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `only_cross_attention` as `down_block_types`. `only_cross_attention`: {only_cross_attention}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(attention_head_dim, int) and len(attention_head_dim) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `attention_head_dim` as `down_block_types`. `attention_head_dim`: {attention_head_dim}. `down_block_types`: {down_block_types}." - ) - - if isinstance(cross_attention_dim, list) and len(cross_attention_dim) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `cross_attention_dim` as `down_block_types`. `cross_attention_dim`: {cross_attention_dim}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(layers_per_block, int) and len(layers_per_block) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `layers_per_block` as `down_block_types`. `layers_per_block`: {layers_per_block}. `down_block_types`: {down_block_types}." - ) - if isinstance(transformer_layers_per_block, list) and reverse_transformer_layers_per_block is None: - for layer_number_per_block in transformer_layers_per_block: - if isinstance(layer_number_per_block, list): - raise ValueError("Must provide 'reverse_transformer_layers_per_block` if using asymmetrical UNet.") - - def _set_time_proj( - self, - time_embedding_type: str, - block_out_channels: int, - flip_sin_to_cos: bool, - freq_shift: float, - time_embedding_dim: int, - ) -> tuple[int, int]: - if time_embedding_type == "fourier": - time_embed_dim = time_embedding_dim or block_out_channels[0] * 2 - if time_embed_dim % 2 != 0: - raise ValueError(f"`time_embed_dim` should be divisible by 2, but is {time_embed_dim}.") - self.time_proj = GaussianFourierProjection( - time_embed_dim // 2, set_W_to_weight=False, log=False, flip_sin_to_cos=flip_sin_to_cos - ) - timestep_input_dim = time_embed_dim - elif time_embedding_type == "positional": - time_embed_dim = time_embedding_dim or block_out_channels[0] * 4 - - self.time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos, freq_shift) - timestep_input_dim = block_out_channels[0] - else: - raise ValueError( - f"{time_embedding_type} does not exist. Please make sure to use one of `fourier` or `positional`." - ) - - return time_embed_dim, timestep_input_dim - - def _set_encoder_hid_proj( - self, - encoder_hid_dim_type: str | None, - cross_attention_dim: int | tuple[int], - encoder_hid_dim: int | None, - ): - if encoder_hid_dim_type is None and encoder_hid_dim is not None: - encoder_hid_dim_type = "text_proj" - self.register_to_config(encoder_hid_dim_type=encoder_hid_dim_type) - logger.info("encoder_hid_dim_type defaults to 'text_proj' as `encoder_hid_dim` is defined.") - - if encoder_hid_dim is None and encoder_hid_dim_type is not None: - raise ValueError( - f"`encoder_hid_dim` has to be defined when `encoder_hid_dim_type` is set to {encoder_hid_dim_type}." - ) - - if encoder_hid_dim_type == "text_proj": - self.encoder_hid_proj = nn.Linear(encoder_hid_dim, cross_attention_dim) - elif encoder_hid_dim_type == "text_image_proj": - # image_embed_dim DOESN'T have to be `cross_attention_dim`. To not clutter the __init__ too much - # they are set to `cross_attention_dim` here as this is exactly the required dimension for the currently only use - # case when `addition_embed_type == "text_image_proj"` (Kandinsky 2.1)` - self.encoder_hid_proj = TextImageProjection( - text_embed_dim=encoder_hid_dim, - image_embed_dim=cross_attention_dim, - cross_attention_dim=cross_attention_dim, - ) - elif encoder_hid_dim_type == "image_proj": - # Kandinsky 2.2 - self.encoder_hid_proj = ImageProjection( - image_embed_dim=encoder_hid_dim, - cross_attention_dim=cross_attention_dim, - ) - elif encoder_hid_dim_type is not None: - raise ValueError( - f"`encoder_hid_dim_type`: {encoder_hid_dim_type} must be None, 'text_proj', 'text_image_proj', or 'image_proj'." - ) - else: - self.encoder_hid_proj = None - - def _set_class_embedding( - self, - class_embed_type: str | None, - act_fn: str, - num_class_embeds: int | None, - projection_class_embeddings_input_dim: int | None, - time_embed_dim: int, - timestep_input_dim: int, - ): - if class_embed_type is None and num_class_embeds is not None: - self.class_embedding = nn.Embedding(num_class_embeds, time_embed_dim) - elif class_embed_type == "timestep": - self.class_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim, act_fn=act_fn) - elif class_embed_type == "identity": - self.class_embedding = nn.Identity(time_embed_dim, time_embed_dim) - elif class_embed_type == "projection": - if projection_class_embeddings_input_dim is None: - raise ValueError( - "`class_embed_type`: 'projection' requires `projection_class_embeddings_input_dim` be set" - ) - # The projection `class_embed_type` is the same as the timestep `class_embed_type` except - # 1. the `class_labels` inputs are not first converted to sinusoidal embeddings - # 2. it projects from an arbitrary input dimension. - # - # Note that `TimestepEmbedding` is quite general, being mainly linear layers and activations. - # When used for embedding actual timesteps, the timesteps are first converted to sinusoidal embeddings. - # As a result, `TimestepEmbedding` can be passed arbitrary vectors. - self.class_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) - elif class_embed_type == "simple_projection": - if projection_class_embeddings_input_dim is None: - raise ValueError( - "`class_embed_type`: 'simple_projection' requires `projection_class_embeddings_input_dim` be set" - ) - self.class_embedding = nn.Linear(projection_class_embeddings_input_dim, time_embed_dim) - else: - self.class_embedding = None - - def _set_add_embedding( - self, - addition_embed_type: str, - addition_embed_type_num_heads: int, - addition_time_embed_dim: int | None, - flip_sin_to_cos: bool, - freq_shift: float, - cross_attention_dim: int | None, - encoder_hid_dim: int | None, - projection_class_embeddings_input_dim: int | None, - time_embed_dim: int, - ): - if addition_embed_type == "text": - if encoder_hid_dim is not None: - text_time_embedding_from_dim = encoder_hid_dim - else: - text_time_embedding_from_dim = cross_attention_dim - - self.add_embedding = TextTimeEmbedding( - text_time_embedding_from_dim, time_embed_dim, num_heads=addition_embed_type_num_heads - ) - elif addition_embed_type == "text_image": - # text_embed_dim and image_embed_dim DON'T have to be `cross_attention_dim`. To not clutter the __init__ too much - # they are set to `cross_attention_dim` here as this is exactly the required dimension for the currently only use - # case when `addition_embed_type == "text_image"` (Kandinsky 2.1)` - self.add_embedding = TextImageTimeEmbedding( - text_embed_dim=cross_attention_dim, image_embed_dim=cross_attention_dim, time_embed_dim=time_embed_dim - ) - elif addition_embed_type == "text_time": - self.add_time_proj = Timesteps(addition_time_embed_dim, flip_sin_to_cos, freq_shift) - self.add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) - elif addition_embed_type == "image": - # Kandinsky 2.2 - self.add_embedding = ImageTimeEmbedding(image_embed_dim=encoder_hid_dim, time_embed_dim=time_embed_dim) - elif addition_embed_type == "image_hint": - # Kandinsky 2.2 ControlNet - self.add_embedding = ImageHintTimeEmbedding(image_embed_dim=encoder_hid_dim, time_embed_dim=time_embed_dim) - elif addition_embed_type is not None: - raise ValueError( - f"`addition_embed_type`: {addition_embed_type} must be None, 'text', 'text_image', 'text_time', 'image', or 'image_hint'." - ) - - def _set_pos_net_if_use_gligen(self, attention_type: str, cross_attention_dim: int): - if attention_type in ["gated", "gated-text-image"]: - positive_len = 768 - if isinstance(cross_attention_dim, int): - positive_len = cross_attention_dim - elif isinstance(cross_attention_dim, (list, tuple)): - positive_len = cross_attention_dim[0] - - feature_type = "text-only" if attention_type == "gated" else "text-image" - self.position_net = GLIGENTextBoundingboxProjection( - positive_len=positive_len, out_dim=cross_attention_dim, feature_type=feature_type - ) - - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - def set_attention_slice(self, slice_size: str | int | list[int] = "auto"): - r""" - Enable sliced attention computation. - - When this option is enabled, the attention module splits the input tensor in slices to compute attention in - several steps. This is useful for saving some memory in exchange for a small decrease in speed. - - Args: - slice_size (`str` or `int` or `list(int)`, *optional*, defaults to `"auto"`): - When `"auto"`, input to the attention heads is halved, so attention is computed in two steps. If - `"max"`, maximum amount of memory is saved by running only one slice at a time. If a number is - provided, uses as many slices as `attention_head_dim // slice_size`. In this case, `attention_head_dim` - must be a multiple of `slice_size`. - """ - sliceable_head_dims = [] - - def fn_recursive_retrieve_sliceable_dims(module: torch.nn.Module): - if hasattr(module, "set_attention_slice"): - sliceable_head_dims.append(module.sliceable_head_dim) - - for child in module.children(): - fn_recursive_retrieve_sliceable_dims(child) - - # retrieve number of attention layers - for module in self.children(): - fn_recursive_retrieve_sliceable_dims(module) - - num_sliceable_layers = len(sliceable_head_dims) - - if slice_size == "auto": - # half the attention head size is usually a good trade-off between - # speed and memory - slice_size = [dim // 2 for dim in sliceable_head_dims] - elif slice_size == "max": - # make smallest slice possible - slice_size = num_sliceable_layers * [1] - - slice_size = num_sliceable_layers * [slice_size] if not isinstance(slice_size, list) else slice_size - - if len(slice_size) != len(sliceable_head_dims): - raise ValueError( - f"You have provided {len(slice_size)}, but {self.config} has {len(sliceable_head_dims)} different" - f" attention layers. Make sure to match `len(slice_size)` to be {len(sliceable_head_dims)}." - ) - - for i in range(len(slice_size)): - size = slice_size[i] - dim = sliceable_head_dims[i] - if size is not None and size > dim: - raise ValueError(f"size {size} has to be smaller or equal to {dim}.") - - # Recursively walk through all the children. - # Any children which exposes the set_attention_slice method - # gets the message - def fn_recursive_set_attention_slice(module: torch.nn.Module, slice_size: list[int]): - if hasattr(module, "set_attention_slice"): - module.set_attention_slice(slice_size.pop()) - - for child in module.children(): - fn_recursive_set_attention_slice(child, slice_size) - - reversed_slice_size = list(reversed(slice_size)) - for module in self.children(): - fn_recursive_set_attention_slice(module, reversed_slice_size) - - def enable_freeu(self, s1: float, s2: float, b1: float, b2: float): - r"""Enables the FreeU mechanism from https://huggingface.co/papers/2309.11497. - - The suffixes after the scaling factors represent the stage blocks where they are being applied. - - Please refer to the [official repository](https://github.com/ChenyangSi/FreeU) for combinations of values that - are known to work well for different pipelines such as Stable Diffusion v1, v2, and Stable Diffusion XL. - - Args: - s1 (`float`): - Scaling factor for stage 1 to attenuate the contributions of the skip features. This is done to - mitigate the "oversmoothing effect" in the enhanced denoising process. - s2 (`float`): - Scaling factor for stage 2 to attenuate the contributions of the skip features. This is done to - mitigate the "oversmoothing effect" in the enhanced denoising process. - b1 (`float`): Scaling factor for stage 1 to amplify the contributions of backbone features. - b2 (`float`): Scaling factor for stage 2 to amplify the contributions of backbone features. - """ - for i, upsample_block in enumerate(self.up_blocks): - setattr(upsample_block, "s1", s1) - setattr(upsample_block, "s2", s2) - setattr(upsample_block, "b1", b1) - setattr(upsample_block, "b2", b2) - - def disable_freeu(self): - """Disables the FreeU mechanism.""" - freeu_keys = {"s1", "s2", "b1", "b2"} - for i, upsample_block in enumerate(self.up_blocks): - for k in freeu_keys: - if hasattr(upsample_block, k) or getattr(upsample_block, k, None) is not None: - setattr(upsample_block, k, None) - - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedAttnProcessor2_0()) - - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) - - def get_time_embed(self, sample: torch.Tensor, timestep: torch.Tensor | float | int) -> torch.Tensor | None: - timesteps = timestep - if not torch.is_tensor(timesteps): - # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can - # This would be a good case for the `match` statement (Python 3.10+) - dtype = maybe_adjust_dtype_for_device( - torch.float64 if isinstance(timestep, float) else torch.int64, sample.device - ) - timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) - elif len(timesteps.shape) == 0: - timesteps = timesteps[None].to(sample.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - timesteps = timesteps.expand(sample.shape[0]) - - t_emb = self.time_proj(timesteps) - # `Timesteps` does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=sample.dtype) - return t_emb - - def get_class_embed(self, sample: torch.Tensor, class_labels: torch.Tensor | None) -> torch.Tensor | None: - class_emb = None - if self.class_embedding is not None: - if class_labels is None: - raise ValueError("class_labels should be provided when num_class_embeds > 0") - - if self.config.class_embed_type == "timestep": - class_labels = self.time_proj(class_labels) - - # `Timesteps` does not contain any weights and will always return f32 tensors - # there might be better ways to encapsulate this. - class_labels = class_labels.to(dtype=sample.dtype) - - class_emb = self.class_embedding(class_labels).to(dtype=sample.dtype) - return class_emb - - def get_aug_embed( - self, emb: torch.Tensor, encoder_hidden_states: torch.Tensor, added_cond_kwargs: dict[str, Any] - ) -> torch.Tensor | None: - aug_emb = None - if self.config.addition_embed_type == "text": - aug_emb = self.add_embedding(encoder_hidden_states) - elif self.config.addition_embed_type == "text_image": - # Kandinsky 2.1 - style - if "image_embeds" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'text_image' which requires the keyword argument `image_embeds` to be passed in `added_cond_kwargs`" - ) - - image_embs = added_cond_kwargs.get("image_embeds") - text_embs = added_cond_kwargs.get("text_embeds", encoder_hidden_states) - aug_emb = self.add_embedding(text_embs, image_embs) - elif self.config.addition_embed_type == "text_time": - # SDXL - style - if "text_embeds" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `text_embeds` to be passed in `added_cond_kwargs`" - ) - text_embeds = added_cond_kwargs.get("text_embeds") - if "time_ids" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `time_ids` to be passed in `added_cond_kwargs`" - ) - time_ids = added_cond_kwargs.get("time_ids") - time_embeds = self.add_time_proj(time_ids.flatten()) - time_embeds = time_embeds.reshape((text_embeds.shape[0], -1)) - add_embeds = torch.concat([text_embeds, time_embeds], dim=-1) - add_embeds = add_embeds.to(emb.dtype) - aug_emb = self.add_embedding(add_embeds) - elif self.config.addition_embed_type == "image": - # Kandinsky 2.2 - style - if "image_embeds" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'image' which requires the keyword argument `image_embeds` to be passed in `added_cond_kwargs`" - ) - image_embs = added_cond_kwargs.get("image_embeds") - aug_emb = self.add_embedding(image_embs) - elif self.config.addition_embed_type == "image_hint": - # Kandinsky 2.2 ControlNet - style - if "image_embeds" not in added_cond_kwargs or "hint" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'image_hint' which requires the keyword arguments `image_embeds` and `hint` to be passed in `added_cond_kwargs`" - ) - image_embs = added_cond_kwargs.get("image_embeds") - hint = added_cond_kwargs.get("hint") - aug_emb = self.add_embedding(image_embs, hint) - return aug_emb - - def process_encoder_hidden_states( - self, encoder_hidden_states: torch.Tensor, added_cond_kwargs: dict[str, Any] - ) -> torch.Tensor: - if self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "text_proj": - encoder_hidden_states = self.encoder_hid_proj(encoder_hidden_states) - elif self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "text_image_proj": - # Kandinsky 2.1 - style - if "image_embeds" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'text_image_proj' which requires the keyword argument `image_embeds` to be passed in `added_cond_kwargs`" - ) - - image_embeds = added_cond_kwargs.get("image_embeds") - encoder_hidden_states = self.encoder_hid_proj(encoder_hidden_states, image_embeds) - elif self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "image_proj": - # Kandinsky 2.2 - style - if "image_embeds" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'image_proj' which requires the keyword argument `image_embeds` to be passed in `added_cond_kwargs`" - ) - image_embeds = added_cond_kwargs.get("image_embeds") - encoder_hidden_states = self.encoder_hid_proj(image_embeds) - elif self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "ip_image_proj": - if "image_embeds" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'ip_image_proj' which requires the keyword argument `image_embeds` to be passed in `added_cond_kwargs`" - ) - - if hasattr(self, "text_encoder_hid_proj") and self.text_encoder_hid_proj is not None: - encoder_hidden_states = self.text_encoder_hid_proj(encoder_hidden_states) - - image_embeds = added_cond_kwargs.get("image_embeds") - image_embeds = self.encoder_hid_proj(image_embeds) - encoder_hidden_states = (encoder_hidden_states, image_embeds) - return encoder_hidden_states - - @apply_lora_scale("cross_attention_kwargs") - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - encoder_hidden_states: torch.Tensor, - class_labels: torch.Tensor | None = None, - timestep_cond: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - added_cond_kwargs: dict[str, torch.Tensor] | None = None, - down_block_additional_residuals: tuple[torch.Tensor] | None = None, - mid_block_additional_residual: torch.Tensor | None = None, - down_intrablock_additional_residuals: tuple[torch.Tensor] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - return_dict: bool = True, - ) -> UNet2DConditionOutput | tuple: - r""" - The [`UNet2DConditionModel`] forward method. - - Args: - sample (`torch.Tensor`): - The noisy input tensor with the following shape `(batch, channel, height, width)`. - timestep (`torch.Tensor` or `float` or `int`): The number of timesteps to denoise an input. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states with shape `(batch, sequence_length, feature_dim)`. - class_labels (`torch.Tensor`, *optional*, defaults to `None`): - Optional class labels for conditioning. Their embeddings will be summed with the timestep embeddings. - timestep_cond: (`torch.Tensor`, *optional*, defaults to `None`): - Conditional embeddings for timestep. If provided, the embeddings will be summed with the samples passed - through the `self.time_embedding` layer to obtain the timestep embeddings. - attention_mask (`torch.Tensor`, *optional*, defaults to `None`): - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask - is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large - negative values to the attention scores corresponding to "discard" tokens. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - added_cond_kwargs: (`dict`, *optional*): - A kwargs dictionary containing additional embeddings that if specified are added to the embeddings that - are passed along to the UNet blocks. - down_block_additional_residuals: (`tuple` of `torch.Tensor`, *optional*): - A tuple of tensors that if specified are added to the residuals of down unet blocks. - mid_block_additional_residual: (`torch.Tensor`, *optional*): - A tensor that if specified is added to the residual of the middle unet block. - down_intrablock_additional_residuals (`tuple` of `torch.Tensor`, *optional*): - additional residuals to be added within UNet down blocks, for example from T2I-Adapter side model(s) - encoder_attention_mask (`torch.Tensor`): - A cross-attention mask of shape `(batch, sequence_length)` is applied to `encoder_hidden_states`. If - `True` the mask is kept, otherwise if `False` it is discarded. Mask will be converted into a bias, - which adds large negative values to the attention scores corresponding to "discard" tokens. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unets.unet_2d_condition.UNet2DConditionOutput`] instead of a plain - tuple. - - Returns: - [`~models.unets.unet_2d_condition.UNet2DConditionOutput`] or `tuple`: - If `return_dict` is True, an [`~models.unets.unet_2d_condition.UNet2DConditionOutput`] is returned, - otherwise a `tuple` is returned where the first element is the sample tensor. - """ - # By default samples have to be AT least a multiple of the overall upsampling factor. - # The overall upsampling factor is equal to 2 ** (# num of upsampling layers). - # However, the upsampling interpolation output size can be forced to fit any upsampling size - # on the fly if necessary. - default_overall_up_factor = 2**self.num_upsamplers - - # upsample size should be forwarded when sample is not a multiple of `default_overall_up_factor` - forward_upsample_size = False - upsample_size = None - - for dim in sample.shape[-2:]: - if dim % default_overall_up_factor != 0: - # Forward upsample size to force interpolation output size. - forward_upsample_size = True - break - - # ensure attention_mask is a bias, and give it a singleton query_tokens dimension - # expects mask of shape: - # [batch, key_tokens] - # adds singleton query_tokens dimension: - # [batch, 1, key_tokens] - # this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes: - # [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn) - # [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn) - if attention_mask is not None: - # assume that mask is expressed as: - # (1 = keep, 0 = discard) - # convert mask into a bias that can be added to attention scores: - # (keep = +0, discard = -10000.0) - attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # convert encoder_attention_mask to a bias the same way we do for attention_mask - if encoder_attention_mask is not None: - encoder_attention_mask = (1 - encoder_attention_mask.to(sample.dtype)) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - # 0. center input if necessary - if self.config.center_input_sample: - sample = 2 * sample - 1.0 - - # 1. time - t_emb = self.get_time_embed(sample=sample, timestep=timestep) - emb = self.time_embedding(t_emb, timestep_cond) - - class_emb = self.get_class_embed(sample=sample, class_labels=class_labels) - if class_emb is not None: - if self.config.class_embeddings_concat: - emb = torch.cat([emb, class_emb], dim=-1) - else: - emb = emb + class_emb - - aug_emb = self.get_aug_embed( - emb=emb, encoder_hidden_states=encoder_hidden_states, added_cond_kwargs=added_cond_kwargs - ) - if self.config.addition_embed_type == "image_hint": - aug_emb, hint = aug_emb - sample = torch.cat([sample, hint], dim=1) - - emb = emb + aug_emb if aug_emb is not None else emb - - if self.time_embed_act is not None: - emb = self.time_embed_act(emb) - - encoder_hidden_states = self.process_encoder_hidden_states( - encoder_hidden_states=encoder_hidden_states, added_cond_kwargs=added_cond_kwargs - ) - - # 2. pre-process - sample = self.conv_in(sample) - - # 2.5 GLIGEN position net - if cross_attention_kwargs is not None and cross_attention_kwargs.get("gligen", None) is not None: - cross_attention_kwargs = cross_attention_kwargs.copy() - gligen_args = cross_attention_kwargs.pop("gligen") - cross_attention_kwargs["gligen"] = {"objs": self.position_net(**gligen_args)} - - # 3. down - is_controlnet = mid_block_additional_residual is not None and down_block_additional_residuals is not None - # using new arg down_intrablock_additional_residuals for T2I-Adapters, to distinguish from controlnets - is_adapter = down_intrablock_additional_residuals is not None - # maintain backward compatibility for legacy usage, where - # T2I-Adapter and ControlNet both use down_block_additional_residuals arg - # but can only use one or the other - if not is_adapter and mid_block_additional_residual is None and down_block_additional_residuals is not None: - deprecate( - "T2I should not use down_block_additional_residuals", - "1.3.0", - "Passing intrablock residual connections with `down_block_additional_residuals` is deprecated \ - and will be removed in diffusers 1.3.0. `down_block_additional_residuals` should only be used \ - for ControlNet. Please make sure use `down_intrablock_additional_residuals` instead. ", - standard_warn=False, - ) - down_intrablock_additional_residuals = down_block_additional_residuals - is_adapter = True - - down_block_res_samples = (sample,) - for downsample_block in self.down_blocks: - if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention: - # For t2i-adapter CrossAttnDownBlock2D - additional_residuals = {} - if is_adapter and len(down_intrablock_additional_residuals) > 0: - additional_residuals["additional_residuals"] = down_intrablock_additional_residuals.pop(0) - - sample, res_samples = downsample_block( - hidden_states=sample, - temb=emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - encoder_attention_mask=encoder_attention_mask, - **additional_residuals, - ) - else: - sample, res_samples = downsample_block(hidden_states=sample, temb=emb) - if is_adapter and len(down_intrablock_additional_residuals) > 0: - sample += down_intrablock_additional_residuals.pop(0) - - down_block_res_samples += res_samples - - if is_controlnet: - new_down_block_res_samples = () - - for down_block_res_sample, down_block_additional_residual in zip( - down_block_res_samples, down_block_additional_residuals - ): - down_block_res_sample = down_block_res_sample + down_block_additional_residual - new_down_block_res_samples = new_down_block_res_samples + (down_block_res_sample,) - - down_block_res_samples = new_down_block_res_samples - - # 4. mid - if self.mid_block is not None: - if hasattr(self.mid_block, "has_cross_attention") and self.mid_block.has_cross_attention: - sample = self.mid_block( - sample, - emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - encoder_attention_mask=encoder_attention_mask, - ) - else: - sample = self.mid_block(sample, emb) - - # To support T2I-Adapter-XL - if ( - is_adapter - and len(down_intrablock_additional_residuals) > 0 - and sample.shape == down_intrablock_additional_residuals[0].shape - ): - sample += down_intrablock_additional_residuals.pop(0) - - if is_controlnet: - sample = sample + mid_block_additional_residual - - # 5. up - for i, upsample_block in enumerate(self.up_blocks): - is_final_block = i == len(self.up_blocks) - 1 - - res_samples = down_block_res_samples[-len(upsample_block.resnets) :] - down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)] - - # if we have not reached the final block and need to forward the - # upsample size, we do it here - if not is_final_block and forward_upsample_size: - upsample_size = down_block_res_samples[-1].shape[2:] - - if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention: - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - upsample_size=upsample_size, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - ) - else: - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - upsample_size=upsample_size, - ) - - # 6. post-process - if self.conv_norm_out: - sample = self.conv_norm_out(sample) - sample = self.conv_act(sample) - sample = self.conv_out(sample) - - if not return_dict: - return (sample,) - - return UNet2DConditionOutput(sample=sample) diff --git a/diffusers/models/unets/unet_3d_blocks.py b/diffusers/models/unets/unet_3d_blocks.py deleted file mode 100644 index e0d7f03bea3a5153dd6703ac5b69a766a35c3995..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_3d_blocks.py +++ /dev/null @@ -1,1419 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -from typing import Any - -import torch -from torch import nn - -from ...utils import deprecate, logging -from ...utils.torch_utils import apply_freeu -from ..attention import Attention -from ..resnet import ( - Downsample2D, - ResnetBlock2D, - SpatioTemporalResBlock, - TemporalConvLayer, - Upsample2D, -) -from ..transformers.transformer_2d import Transformer2DModel -from ..transformers.transformer_temporal import ( - TransformerSpatioTemporalModel, - TransformerTemporalModel, -) -from .unet_motion_model import ( - CrossAttnDownBlockMotion, - CrossAttnUpBlockMotion, - DownBlockMotion, - UNetMidBlockCrossAttnMotion, - UpBlockMotion, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class DownBlockMotion(DownBlockMotion): - def __init__(self, *args, **kwargs): - deprecation_message = "Importing `DownBlockMotion` from `diffusers.models.unets.unet_3d_blocks` is deprecated and this will be removed in a future version. Please use `from diffusers.models.unets.unet_motion_model import DownBlockMotion` instead." - deprecate("DownBlockMotion", "1.0.0", deprecation_message) - super().__init__(*args, **kwargs) - - -class CrossAttnDownBlockMotion(CrossAttnDownBlockMotion): - def __init__(self, *args, **kwargs): - deprecation_message = "Importing `CrossAttnDownBlockMotion` from `diffusers.models.unets.unet_3d_blocks` is deprecated and this will be removed in a future version. Please use `from diffusers.models.unets.unet_motion_model import CrossAttnDownBlockMotion` instead." - deprecate("CrossAttnDownBlockMotion", "1.0.0", deprecation_message) - super().__init__(*args, **kwargs) - - -class UpBlockMotion(UpBlockMotion): - def __init__(self, *args, **kwargs): - deprecation_message = "Importing `UpBlockMotion` from `diffusers.models.unets.unet_3d_blocks` is deprecated and this will be removed in a future version. Please use `from diffusers.models.unets.unet_motion_model import UpBlockMotion` instead." - deprecate("UpBlockMotion", "1.0.0", deprecation_message) - super().__init__(*args, **kwargs) - - -class CrossAttnUpBlockMotion(CrossAttnUpBlockMotion): - def __init__(self, *args, **kwargs): - deprecation_message = "Importing `CrossAttnUpBlockMotion` from `diffusers.models.unets.unet_3d_blocks` is deprecated and this will be removed in a future version. Please use `from diffusers.models.unets.unet_motion_model import CrossAttnUpBlockMotion` instead." - deprecate("CrossAttnUpBlockMotion", "1.0.0", deprecation_message) - super().__init__(*args, **kwargs) - - -class UNetMidBlockCrossAttnMotion(UNetMidBlockCrossAttnMotion): - def __init__(self, *args, **kwargs): - deprecation_message = "Importing `UNetMidBlockCrossAttnMotion` from `diffusers.models.unets.unet_3d_blocks` is deprecated and this will be removed in a future version. Please use `from diffusers.models.unets.unet_motion_model import UNetMidBlockCrossAttnMotion` instead." - deprecate("UNetMidBlockCrossAttnMotion", "1.0.0", deprecation_message) - super().__init__(*args, **kwargs) - - -def get_down_block( - down_block_type: str, - num_layers: int, - in_channels: int, - out_channels: int, - temb_channels: int, - add_downsample: bool, - resnet_eps: float, - resnet_act_fn: str, - num_attention_heads: int, - resnet_groups: int | None = None, - cross_attention_dim: int | None = None, - downsample_padding: int | None = None, - dual_cross_attention: bool = False, - use_linear_projection: bool = True, - only_cross_attention: bool = False, - upcast_attention: bool = False, - resnet_time_scale_shift: str = "default", - temporal_num_attention_heads: int = 8, - temporal_max_seq_length: int = 32, - transformer_layers_per_block: int | tuple[int] = 1, - temporal_transformer_layers_per_block: int | tuple[int] = 1, - dropout: float = 0.0, -) -> "DownBlock3D" | "CrossAttnDownBlock3D" | "DownBlockSpatioTemporal" | "CrossAttnDownBlockSpatioTemporal": - if down_block_type == "DownBlock3D": - return DownBlock3D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - downsample_padding=downsample_padding, - resnet_time_scale_shift=resnet_time_scale_shift, - dropout=dropout, - ) - elif down_block_type == "CrossAttnDownBlock3D": - if cross_attention_dim is None: - raise ValueError("cross_attention_dim must be specified for CrossAttnDownBlock3D") - return CrossAttnDownBlock3D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - downsample_padding=downsample_padding, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads, - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - dropout=dropout, - ) - elif down_block_type == "DownBlockSpatioTemporal": - # added for SDV - return DownBlockSpatioTemporal( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - add_downsample=add_downsample, - ) - elif down_block_type == "CrossAttnDownBlockSpatioTemporal": - # added for SDV - if cross_attention_dim is None: - raise ValueError("cross_attention_dim must be specified for CrossAttnDownBlockSpatioTemporal") - return CrossAttnDownBlockSpatioTemporal( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - num_layers=num_layers, - transformer_layers_per_block=transformer_layers_per_block, - add_downsample=add_downsample, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads, - ) - - raise ValueError(f"{down_block_type} does not exist.") - - -def get_up_block( - up_block_type: str, - num_layers: int, - in_channels: int, - out_channels: int, - prev_output_channel: int, - temb_channels: int, - add_upsample: bool, - resnet_eps: float, - resnet_act_fn: str, - num_attention_heads: int, - resolution_idx: int | None = None, - resnet_groups: int | None = None, - cross_attention_dim: int | None = None, - dual_cross_attention: bool = False, - use_linear_projection: bool = True, - only_cross_attention: bool = False, - upcast_attention: bool = False, - resnet_time_scale_shift: str = "default", - temporal_num_attention_heads: int = 8, - temporal_cross_attention_dim: int | None = None, - temporal_max_seq_length: int = 32, - transformer_layers_per_block: int | tuple[int] = 1, - temporal_transformer_layers_per_block: int | tuple[int] = 1, - dropout: float = 0.0, -) -> "UpBlock3D" | "CrossAttnUpBlock3D" | "UpBlockSpatioTemporal" | "CrossAttnUpBlockSpatioTemporal": - if up_block_type == "UpBlock3D": - return UpBlock3D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - resnet_time_scale_shift=resnet_time_scale_shift, - resolution_idx=resolution_idx, - dropout=dropout, - ) - elif up_block_type == "CrossAttnUpBlock3D": - if cross_attention_dim is None: - raise ValueError("cross_attention_dim must be specified for CrossAttnUpBlock3D") - return CrossAttnUpBlock3D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads, - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - resolution_idx=resolution_idx, - dropout=dropout, - ) - elif up_block_type == "UpBlockSpatioTemporal": - # added for SDV - return UpBlockSpatioTemporal( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - add_upsample=add_upsample, - ) - elif up_block_type == "CrossAttnUpBlockSpatioTemporal": - # added for SDV - if cross_attention_dim is None: - raise ValueError("cross_attention_dim must be specified for CrossAttnUpBlockSpatioTemporal") - return CrossAttnUpBlockSpatioTemporal( - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - num_layers=num_layers, - transformer_layers_per_block=transformer_layers_per_block, - add_upsample=add_upsample, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads, - resolution_idx=resolution_idx, - ) - - raise ValueError(f"{up_block_type} does not exist.") - - -class UNetMidBlock3DCrossAttn(nn.Module): - def __init__( - self, - in_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - output_scale_factor: float = 1.0, - cross_attention_dim: int = 1280, - dual_cross_attention: bool = False, - use_linear_projection: bool = True, - upcast_attention: bool = False, - ): - super().__init__() - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32) - - # there is always at least one resnet - resnets = [ - ResnetBlock2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ] - temp_convs = [ - TemporalConvLayer( - in_channels, - in_channels, - dropout=0.1, - norm_num_groups=resnet_groups, - ) - ] - attentions = [] - temp_attentions = [] - - for _ in range(num_layers): - attentions.append( - Transformer2DModel( - in_channels // num_attention_heads, - num_attention_heads, - in_channels=in_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - ) - ) - temp_attentions.append( - TransformerTemporalModel( - in_channels // num_attention_heads, - num_attention_heads, - in_channels=in_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - temp_convs.append( - TemporalConvLayer( - in_channels, - in_channels, - dropout=0.1, - norm_num_groups=resnet_groups, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.temp_convs = nn.ModuleList(temp_convs) - self.attentions = nn.ModuleList(attentions) - self.temp_attentions = nn.ModuleList(temp_attentions) - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - num_frames: int = 1, - cross_attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor: - hidden_states = self.resnets[0](hidden_states, temb) - hidden_states = self.temp_convs[0](hidden_states, num_frames=num_frames) - for attn, temp_attn, resnet, temp_conv in zip( - self.attentions, self.temp_attentions, self.resnets[1:], self.temp_convs[1:] - ): - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - return_dict=False, - )[0] - hidden_states = temp_attn( - hidden_states, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - return_dict=False, - )[0] - hidden_states = resnet(hidden_states, temb) - hidden_states = temp_conv(hidden_states, num_frames=num_frames) - - return hidden_states - - -class CrossAttnDownBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - downsample_padding: int = 1, - add_downsample: bool = True, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - ): - super().__init__() - resnets = [] - attentions = [] - temp_attentions = [] - temp_convs = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - temp_convs.append( - TemporalConvLayer( - out_channels, - out_channels, - dropout=0.1, - norm_num_groups=resnet_groups, - ) - ) - attentions.append( - Transformer2DModel( - out_channels // num_attention_heads, - num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - ) - ) - temp_attentions.append( - TransformerTemporalModel( - out_channels // num_attention_heads, - num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - self.resnets = nn.ModuleList(resnets) - self.temp_convs = nn.ModuleList(temp_convs) - self.attentions = nn.ModuleList(attentions) - self.temp_attentions = nn.ModuleList(temp_attentions) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, - use_conv=True, - out_channels=out_channels, - padding=downsample_padding, - name="op", - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - num_frames: int = 1, - cross_attention_kwargs: dict[str, Any] = None, - ) -> torch.Tensor | tuple[torch.Tensor, ...]: - # TODO(Patrick, William) - attention mask is not used - output_states = () - - for resnet, temp_conv, attn, temp_attn in zip( - self.resnets, self.temp_convs, self.attentions, self.temp_attentions - ): - hidden_states = resnet(hidden_states, temb) - hidden_states = temp_conv(hidden_states, num_frames=num_frames) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - return_dict=False, - )[0] - hidden_states = temp_attn( - hidden_states, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - return_dict=False, - )[0] - - output_states += (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - output_states += (hidden_states,) - - return hidden_states, output_states - - -class DownBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_downsample: bool = True, - downsample_padding: int = 1, - ): - super().__init__() - resnets = [] - temp_convs = [] - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - temp_convs.append( - TemporalConvLayer( - out_channels, - out_channels, - dropout=0.1, - norm_num_groups=resnet_groups, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.temp_convs = nn.ModuleList(temp_convs) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, - use_conv=True, - out_channels=out_channels, - padding=downsample_padding, - name="op", - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - num_frames: int = 1, - ) -> torch.Tensor | tuple[torch.Tensor, ...]: - output_states = () - - for resnet, temp_conv in zip(self.resnets, self.temp_convs): - hidden_states = resnet(hidden_states, temb) - hidden_states = temp_conv(hidden_states, num_frames=num_frames) - - output_states += (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - output_states += (hidden_states,) - - return hidden_states, output_states - - -class CrossAttnUpBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - prev_output_channel: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - resolution_idx: int | None = None, - ): - super().__init__() - resnets = [] - temp_convs = [] - attentions = [] - temp_attentions = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - temp_convs.append( - TemporalConvLayer( - out_channels, - out_channels, - dropout=0.1, - norm_num_groups=resnet_groups, - ) - ) - attentions.append( - Transformer2DModel( - out_channels // num_attention_heads, - num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - ) - ) - temp_attentions.append( - TransformerTemporalModel( - out_channels // num_attention_heads, - num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - self.resnets = nn.ModuleList(resnets) - self.temp_convs = nn.ModuleList(temp_convs) - self.attentions = nn.ModuleList(attentions) - self.temp_attentions = nn.ModuleList(temp_attentions) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - upsample_size: int | None = None, - attention_mask: torch.Tensor | None = None, - num_frames: int = 1, - cross_attention_kwargs: dict[str, Any] = None, - ) -> torch.Tensor: - is_freeu_enabled = ( - getattr(self, "s1", None) - and getattr(self, "s2", None) - and getattr(self, "b1", None) - and getattr(self, "b2", None) - ) - - # TODO(Patrick, William) - attention mask is not used - for resnet, temp_conv, attn, temp_attn in zip( - self.resnets, self.temp_convs, self.attentions, self.temp_attentions - ): - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - - # FreeU: Only operate on the first two stages - if is_freeu_enabled: - hidden_states, res_hidden_states = apply_freeu( - self.resolution_idx, - hidden_states, - res_hidden_states, - s1=self.s1, - s2=self.s2, - b1=self.b1, - b2=self.b2, - ) - - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - hidden_states = resnet(hidden_states, temb) - hidden_states = temp_conv(hidden_states, num_frames=num_frames) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - return_dict=False, - )[0] - hidden_states = temp_attn( - hidden_states, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - return_dict=False, - )[0] - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size) - - return hidden_states - - -class UpBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - prev_output_channel: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - resolution_idx: int | None = None, - ): - super().__init__() - resnets = [] - temp_convs = [] - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - temp_convs.append( - TemporalConvLayer( - out_channels, - out_channels, - dropout=0.1, - norm_num_groups=resnet_groups, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.temp_convs = nn.ModuleList(temp_convs) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - upsample_size: int | None = None, - num_frames: int = 1, - ) -> torch.Tensor: - is_freeu_enabled = ( - getattr(self, "s1", None) - and getattr(self, "s2", None) - and getattr(self, "b1", None) - and getattr(self, "b2", None) - ) - for resnet, temp_conv in zip(self.resnets, self.temp_convs): - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - - # FreeU: Only operate on the first two stages - if is_freeu_enabled: - hidden_states, res_hidden_states = apply_freeu( - self.resolution_idx, - hidden_states, - res_hidden_states, - s1=self.s1, - s2=self.s2, - b1=self.b1, - b2=self.b2, - ) - - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - hidden_states = resnet(hidden_states, temb) - hidden_states = temp_conv(hidden_states, num_frames=num_frames) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size) - - return hidden_states - - -class MidBlockTemporalDecoder(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - attention_head_dim: int = 512, - num_layers: int = 1, - upcast_attention: bool = False, - ): - super().__init__() - - resnets = [] - attentions = [] - for i in range(num_layers): - input_channels = in_channels if i == 0 else out_channels - resnets.append( - SpatioTemporalResBlock( - in_channels=input_channels, - out_channels=out_channels, - temb_channels=None, - eps=1e-6, - temporal_eps=1e-5, - merge_factor=0.0, - merge_strategy="learned", - switch_spatial_to_temporal_mix=True, - ) - ) - - attentions.append( - Attention( - query_dim=in_channels, - heads=in_channels // attention_head_dim, - dim_head=attention_head_dim, - eps=1e-6, - upcast_attention=upcast_attention, - norm_num_groups=32, - bias=True, - residual_connection=True, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - def forward( - self, - hidden_states: torch.Tensor, - image_only_indicator: torch.Tensor, - ): - hidden_states = self.resnets[0]( - hidden_states, - image_only_indicator=image_only_indicator, - ) - for resnet, attn in zip(self.resnets[1:], self.attentions): - hidden_states = attn(hidden_states) - hidden_states = resnet( - hidden_states, - image_only_indicator=image_only_indicator, - ) - - return hidden_states - - -class UpBlockTemporalDecoder(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - num_layers: int = 1, - add_upsample: bool = True, - ): - super().__init__() - resnets = [] - for i in range(num_layers): - input_channels = in_channels if i == 0 else out_channels - - resnets.append( - SpatioTemporalResBlock( - in_channels=input_channels, - out_channels=out_channels, - temb_channels=None, - eps=1e-6, - temporal_eps=1e-5, - merge_factor=0.0, - merge_strategy="learned", - switch_spatial_to_temporal_mix=True, - ) - ) - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - def forward( - self, - hidden_states: torch.Tensor, - image_only_indicator: torch.Tensor, - ) -> torch.Tensor: - for resnet in self.resnets: - hidden_states = resnet( - hidden_states, - image_only_indicator=image_only_indicator, - ) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states) - - return hidden_states - - -class UNetMidBlockSpatioTemporal(nn.Module): - def __init__( - self, - in_channels: int, - temb_channels: int, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - ): - super().__init__() - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - - # support for variable transformer layers per block - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - # there is always at least one resnet - resnets = [ - SpatioTemporalResBlock( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=1e-5, - ) - ] - attentions = [] - - for i in range(num_layers): - attentions.append( - TransformerSpatioTemporalModel( - num_attention_heads, - in_channels // num_attention_heads, - in_channels=in_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - ) - ) - - resnets.append( - SpatioTemporalResBlock( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=1e-5, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - image_only_indicator: torch.Tensor | None = None, - ) -> torch.Tensor: - hidden_states = self.resnets[0]( - hidden_states, - temb, - image_only_indicator=image_only_indicator, - ) - - for attn, resnet in zip(self.attentions, self.resnets[1:]): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - image_only_indicator=image_only_indicator, - return_dict=False, - )[0] - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, image_only_indicator) - else: - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - image_only_indicator=image_only_indicator, - return_dict=False, - )[0] - hidden_states = resnet(hidden_states, temb, image_only_indicator=image_only_indicator) - - return hidden_states - - -class DownBlockSpatioTemporal(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - num_layers: int = 1, - add_downsample: bool = True, - ): - super().__init__() - resnets = [] - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - SpatioTemporalResBlock( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=1e-5, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, - use_conv=True, - out_channels=out_channels, - name="op", - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - image_only_indicator: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - output_states = () - for resnet in self.resnets: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, image_only_indicator) - else: - hidden_states = resnet(hidden_states, temb, image_only_indicator=image_only_indicator) - - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - output_states = output_states + (hidden_states,) - - return hidden_states, output_states - - -class CrossAttnDownBlockSpatioTemporal(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - add_downsample: bool = True, - ): - super().__init__() - resnets = [] - attentions = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - SpatioTemporalResBlock( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=1e-6, - ) - ) - attentions.append( - TransformerSpatioTemporalModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, - use_conv=True, - out_channels=out_channels, - padding=1, - name="op", - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - image_only_indicator: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - output_states = () - - blocks = list(zip(self.resnets, self.attentions)) - for resnet, attn in blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, image_only_indicator) - - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - image_only_indicator=image_only_indicator, - return_dict=False, - )[0] - else: - hidden_states = resnet(hidden_states, temb, image_only_indicator=image_only_indicator) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - image_only_indicator=image_only_indicator, - return_dict=False, - )[0] - - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - - output_states = output_states + (hidden_states,) - - return hidden_states, output_states - - -class UpBlockSpatioTemporal(nn.Module): - def __init__( - self, - in_channels: int, - prev_output_channel: int, - out_channels: int, - temb_channels: int, - resolution_idx: int | None = None, - num_layers: int = 1, - resnet_eps: float = 1e-6, - add_upsample: bool = True, - ): - super().__init__() - resnets = [] - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - SpatioTemporalResBlock( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - image_only_indicator: torch.Tensor | None = None, - upsample_size: int | None = None, - ) -> torch.Tensor: - for resnet in self.resnets: - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, image_only_indicator) - else: - hidden_states = resnet(hidden_states, temb, image_only_indicator=image_only_indicator) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size) - - return hidden_states - - -class CrossAttnUpBlockSpatioTemporal(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - prev_output_channel: int, - temb_channels: int, - resolution_idx: int | None = None, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - add_upsample: bool = True, - ): - super().__init__() - resnets = [] - attentions = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - SpatioTemporalResBlock( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - ) - ) - attentions.append( - TransformerSpatioTemporalModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - image_only_indicator: torch.Tensor | None = None, - upsample_size: int | None = None, - ) -> torch.Tensor: - for resnet, attn in zip(self.resnets, self.attentions): - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb, image_only_indicator) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - image_only_indicator=image_only_indicator, - return_dict=False, - )[0] - else: - hidden_states = resnet(hidden_states, temb, image_only_indicator=image_only_indicator) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - image_only_indicator=image_only_indicator, - return_dict=False, - )[0] - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size) - - return hidden_states diff --git a/diffusers/models/unets/unet_3d_condition.py b/diffusers/models/unets/unet_3d_condition.py deleted file mode 100644 index 0d15e93da68f89509ad68f9c81b6d9963fb2093a..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_3d_condition.py +++ /dev/null @@ -1,673 +0,0 @@ -# Copyright 2025 Alibaba DAMO-VILAB and The HuggingFace Team. All rights reserved. -# Copyright 2025 The ModelScope Team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import UNet2DConditionLoadersMixin -from ...utils import BaseOutput, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..activations import get_activation -from ..attention import AttentionMixin -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - Attention, - AttnAddedKVProcessor, - AttnProcessor, - FusedAttnProcessor2_0, -) -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from ..transformers.transformer_temporal import TransformerTemporalModel -from .unet_3d_blocks import ( - UNetMidBlock3DCrossAttn, - get_down_block, - get_up_block, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class UNet3DConditionOutput(BaseOutput): - """ - The output of [`UNet3DConditionModel`]. - - Args: - sample (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - The hidden states output conditioned on `encoder_hidden_states` input. Output of last layer of model. - """ - - sample: torch.Tensor - - -class UNet3DConditionModel(ModelMixin, AttentionMixin, ConfigMixin, UNet2DConditionLoadersMixin): - r""" - A conditional 3D UNet model that takes a noisy sample, conditional state, and a timestep and returns a sample - shaped output. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - sample_size (`int` or `tuple[int, int]`, *optional*, defaults to `None`): - Height and width of input/output sample. - in_channels (`int`, *optional*, defaults to 4): The number of channels in the input sample. - out_channels (`int`, *optional*, defaults to 4): The number of channels in the output. - down_block_types (`tuple[str]`, *optional*, defaults to `("CrossAttnDownBlock3D", "CrossAttnDownBlock3D", "CrossAttnDownBlock3D", "DownBlock3D")`): - The tuple of downsample blocks to use. - up_block_types (`tuple[str]`, *optional*, defaults to `("UpBlock3D", "CrossAttnUpBlock3D", "CrossAttnUpBlock3D", "CrossAttnUpBlock3D")`): - The tuple of upsample blocks to use. - block_out_channels (`tuple[int]`, *optional*, defaults to `(320, 640, 1280, 1280)`): - The tuple of output channels for each block. - layers_per_block (`int`, *optional*, defaults to 2): The number of layers per block. - downsample_padding (`int`, *optional*, defaults to 1): The padding to use for the downsampling convolution. - mid_block_scale_factor (`float`, *optional*, defaults to 1.0): The scale factor to use for the mid block. - act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use. - norm_num_groups (`int`, *optional*, defaults to 32): The number of groups to use for the normalization. - If `None`, normalization and activation layers is skipped in post-processing. - norm_eps (`float`, *optional*, defaults to 1e-5): The epsilon to use for the normalization. - cross_attention_dim (`int`, *optional*, defaults to 1024): The dimension of the cross attention features. - attention_head_dim (`int`, *optional*, defaults to 64): The dimension of the attention heads. - num_attention_heads (`int`, *optional*): The number of attention heads. - time_cond_proj_dim (`int`, *optional*, defaults to `None`): - The dimension of `cond_proj` layer in the timestep embedding. - """ - - _supports_gradient_checkpointing = False - _skip_layerwise_casting_patterns = ["norm", "time_embedding"] - - @register_to_config - def __init__( - self, - sample_size: int | None = None, - in_channels: int = 4, - out_channels: int = 4, - down_block_types: tuple[str, ...] = ( - "CrossAttnDownBlock3D", - "CrossAttnDownBlock3D", - "CrossAttnDownBlock3D", - "DownBlock3D", - ), - up_block_types: tuple[str, ...] = ( - "UpBlock3D", - "CrossAttnUpBlock3D", - "CrossAttnUpBlock3D", - "CrossAttnUpBlock3D", - ), - block_out_channels: tuple[int, ...] = (320, 640, 1280, 1280), - layers_per_block: int = 2, - downsample_padding: int = 1, - mid_block_scale_factor: float = 1, - act_fn: str = "silu", - norm_num_groups: int | None = 32, - norm_eps: float = 1e-5, - cross_attention_dim: int = 1024, - attention_head_dim: int | tuple[int] = 64, - num_attention_heads: int | tuple[int] | None = None, - time_cond_proj_dim: int | None = None, - ): - super().__init__() - - self.sample_size = sample_size - - if num_attention_heads is not None: - raise NotImplementedError( - "At the moment it is not possible to define the number of attention heads via `num_attention_heads` because of a naming issue as described in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131. Passing `num_attention_heads` will only be supported in diffusers v0.19." - ) - - # If `num_attention_heads` is not defined (which is the case for most models) - # it will default to `attention_head_dim`. This looks weird upon first reading it and it is. - # The reason for this behavior is to correct for incorrectly named variables that were introduced - # when this library was created. The incorrect naming was only discovered much later in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131 - # Changing `attention_head_dim` to `num_attention_heads` for 40,000+ configurations is too backwards breaking - # which is why we correct for the naming here. - num_attention_heads = num_attention_heads or attention_head_dim - - # Check inputs - if len(down_block_types) != len(up_block_types): - raise ValueError( - f"Must provide the same number of `down_block_types` as `up_block_types`. `down_block_types`: {down_block_types}. `up_block_types`: {up_block_types}." - ) - - if len(block_out_channels) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}." - ) - - # input - conv_in_kernel = 3 - conv_out_kernel = 3 - conv_in_padding = (conv_in_kernel - 1) // 2 - self.conv_in = nn.Conv2d( - in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding - ) - - # time - time_embed_dim = block_out_channels[0] * 4 - self.time_proj = Timesteps(block_out_channels[0], True, 0) - timestep_input_dim = block_out_channels[0] - - self.time_embedding = TimestepEmbedding( - timestep_input_dim, - time_embed_dim, - act_fn=act_fn, - cond_proj_dim=time_cond_proj_dim, - ) - - self.transformer_in = TransformerTemporalModel( - num_attention_heads=8, - attention_head_dim=attention_head_dim, - in_channels=block_out_channels[0], - num_layers=1, - norm_num_groups=norm_num_groups, - ) - - # class embedding - self.down_blocks = nn.ModuleList([]) - self.up_blocks = nn.ModuleList([]) - - if isinstance(num_attention_heads, int): - num_attention_heads = (num_attention_heads,) * len(down_block_types) - - # down - output_channel = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - down_block = get_down_block( - down_block_type, - num_layers=layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - temb_channels=time_embed_dim, - add_downsample=not is_final_block, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads[i], - downsample_padding=downsample_padding, - dual_cross_attention=False, - ) - self.down_blocks.append(down_block) - - # mid - self.mid_block = UNetMidBlock3DCrossAttn( - in_channels=block_out_channels[-1], - temb_channels=time_embed_dim, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - output_scale_factor=mid_block_scale_factor, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads[-1], - resnet_groups=norm_num_groups, - dual_cross_attention=False, - ) - - # count how many layers upsample the images - self.num_upsamplers = 0 - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - reversed_num_attention_heads = list(reversed(num_attention_heads)) - - output_channel = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - is_final_block = i == len(block_out_channels) - 1 - - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)] - - # add upsample block for all BUT final layer - if not is_final_block: - add_upsample = True - self.num_upsamplers += 1 - else: - add_upsample = False - - up_block = get_up_block( - up_block_type, - num_layers=layers_per_block + 1, - in_channels=input_channel, - out_channels=output_channel, - prev_output_channel=prev_output_channel, - temb_channels=time_embed_dim, - add_upsample=add_upsample, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - cross_attention_dim=cross_attention_dim, - num_attention_heads=reversed_num_attention_heads[i], - dual_cross_attention=False, - resolution_idx=i, - ) - self.up_blocks.append(up_block) - prev_output_channel = output_channel - - # out - if norm_num_groups is not None: - self.conv_norm_out = nn.GroupNorm( - num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=norm_eps - ) - self.conv_act = get_activation("silu") - else: - self.conv_norm_out = None - self.conv_act = None - - conv_out_padding = (conv_out_kernel - 1) // 2 - self.conv_out = nn.Conv2d( - block_out_channels[0], out_channels, kernel_size=conv_out_kernel, padding=conv_out_padding - ) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attention_slice - def set_attention_slice(self, slice_size: str | int | list[int]) -> None: - r""" - Enable sliced attention computation. - - When this option is enabled, the attention module splits the input tensor in slices to compute attention in - several steps. This is useful for saving some memory in exchange for a small decrease in speed. - - Args: - slice_size (`str` or `int` or `list(int)`, *optional*, defaults to `"auto"`): - When `"auto"`, input to the attention heads is halved, so attention is computed in two steps. If - `"max"`, maximum amount of memory is saved by running only one slice at a time. If a number is - provided, uses as many slices as `attention_head_dim // slice_size`. In this case, `attention_head_dim` - must be a multiple of `slice_size`. - """ - sliceable_head_dims = [] - - def fn_recursive_retrieve_sliceable_dims(module: torch.nn.Module): - if hasattr(module, "set_attention_slice"): - sliceable_head_dims.append(module.sliceable_head_dim) - - for child in module.children(): - fn_recursive_retrieve_sliceable_dims(child) - - # retrieve number of attention layers - for module in self.children(): - fn_recursive_retrieve_sliceable_dims(module) - - num_sliceable_layers = len(sliceable_head_dims) - - if slice_size == "auto": - # half the attention head size is usually a good trade-off between - # speed and memory - slice_size = [dim // 2 for dim in sliceable_head_dims] - elif slice_size == "max": - # make smallest slice possible - slice_size = num_sliceable_layers * [1] - - slice_size = num_sliceable_layers * [slice_size] if not isinstance(slice_size, list) else slice_size - - if len(slice_size) != len(sliceable_head_dims): - raise ValueError( - f"You have provided {len(slice_size)}, but {self.config} has {len(sliceable_head_dims)} different" - f" attention layers. Make sure to match `len(slice_size)` to be {len(sliceable_head_dims)}." - ) - - for i in range(len(slice_size)): - size = slice_size[i] - dim = sliceable_head_dims[i] - if size is not None and size > dim: - raise ValueError(f"size {size} has to be smaller or equal to {dim}.") - - # Recursively walk through all the children. - # Any children which exposes the set_attention_slice method - # gets the message - def fn_recursive_set_attention_slice(module: torch.nn.Module, slice_size: list[int]): - if hasattr(module, "set_attention_slice"): - module.set_attention_slice(slice_size.pop()) - - for child in module.children(): - fn_recursive_set_attention_slice(child, slice_size) - - reversed_slice_size = list(reversed(slice_size)) - for module in self.children(): - fn_recursive_set_attention_slice(module, reversed_slice_size) - - def enable_forward_chunking(self, chunk_size: int | None = None, dim: int = 0) -> None: - """ - Sets the attention processor to use [feed forward - chunking](https://huggingface.co/blog/reformer#2-chunked-feed-forward-layers). - - Parameters: - chunk_size (`int`, *optional*): - The chunk size of the feed-forward layers. If not specified, will run feed-forward layer individually - over each tensor of dim=`dim`. - dim (`int`, *optional*, defaults to `0`): - The dimension over which the feed-forward computation should be chunked. Choose between dim=0 (batch) - or dim=1 (sequence length). - """ - if dim not in [0, 1]: - raise ValueError(f"Make sure to set `dim` to either 0 or 1, not {dim}") - - # By default chunk size is 1 - chunk_size = chunk_size or 1 - - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, chunk_size, dim) - - def disable_forward_chunking(self): - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, None, 0) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.enable_freeu - def enable_freeu(self, s1, s2, b1, b2): - r"""Enables the FreeU mechanism from https://huggingface.co/papers/2309.11497. - - The suffixes after the scaling factors represent the stage blocks where they are being applied. - - Please refer to the [official repository](https://github.com/ChenyangSi/FreeU) for combinations of values that - are known to work well for different pipelines such as Stable Diffusion v1, v2, and Stable Diffusion XL. - - Args: - s1 (`float`): - Scaling factor for stage 1 to attenuate the contributions of the skip features. This is done to - mitigate the "oversmoothing effect" in the enhanced denoising process. - s2 (`float`): - Scaling factor for stage 2 to attenuate the contributions of the skip features. This is done to - mitigate the "oversmoothing effect" in the enhanced denoising process. - b1 (`float`): Scaling factor for stage 1 to amplify the contributions of backbone features. - b2 (`float`): Scaling factor for stage 2 to amplify the contributions of backbone features. - """ - for i, upsample_block in enumerate(self.up_blocks): - setattr(upsample_block, "s1", s1) - setattr(upsample_block, "s2", s2) - setattr(upsample_block, "b1", b1) - setattr(upsample_block, "b2", b2) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.disable_freeu - def disable_freeu(self): - """Disables the FreeU mechanism.""" - freeu_keys = {"s1", "s2", "b1", "b2"} - for i, upsample_block in enumerate(self.up_blocks): - for k in freeu_keys: - if hasattr(upsample_block, k) or getattr(upsample_block, k, None) is not None: - setattr(upsample_block, k, None) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) - - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - encoder_hidden_states: torch.Tensor, - class_labels: torch.Tensor | None = None, - timestep_cond: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - down_block_additional_residuals: tuple[torch.Tensor] | None = None, - mid_block_additional_residual: torch.Tensor | None = None, - return_dict: bool = True, - ) -> UNet3DConditionOutput | tuple[torch.Tensor]: - r""" - The [`UNet3DConditionModel`] forward method. - - Args: - sample (`torch.Tensor`): - The noisy input tensor with the following shape `(batch, num_channels, num_frames, height, width`. - timestep (`torch.Tensor` or `float` or `int`): The number of timesteps to denoise an input. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states with shape `(batch, sequence_length, feature_dim)`. - class_labels (`torch.Tensor`, *optional*, defaults to `None`): - Optional class labels for conditioning. Their embeddings will be summed with the timestep embeddings. - timestep_cond: (`torch.Tensor`, *optional*, defaults to `None`): - Conditional embeddings for timestep. If provided, the embeddings will be summed with the samples passed - through the `self.time_embedding` layer to obtain the timestep embeddings. - attention_mask (`torch.Tensor`, *optional*, defaults to `None`): - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask - is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large - negative values to the attention scores corresponding to "discard" tokens. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - down_block_additional_residuals: (`tuple` of `torch.Tensor`, *optional*): - A tuple of tensors that if specified are added to the residuals of down unet blocks. - mid_block_additional_residual: (`torch.Tensor`, *optional*): - A tensor that if specified is added to the residual of the middle unet block. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unets.unet_3d_condition.UNet3DConditionOutput`] instead of a plain - tuple. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the [`AttnProcessor`]. - - Returns: - [`~models.unets.unet_3d_condition.UNet3DConditionOutput`] or `tuple`: - If `return_dict` is True, an [`~models.unets.unet_3d_condition.UNet3DConditionOutput`] is returned, - otherwise a `tuple` is returned where the first element is the sample tensor. - """ - # By default samples have to be AT least a multiple of the overall upsampling factor. - # The overall upsampling factor is equal to 2 ** (# num of upsampling layears). - # However, the upsampling interpolation output size can be forced to fit any upsampling size - # on the fly if necessary. - default_overall_up_factor = 2**self.num_upsamplers - - # upsample size should be forwarded when sample is not a multiple of `default_overall_up_factor` - forward_upsample_size = False - upsample_size = None - - if any(s % default_overall_up_factor != 0 for s in sample.shape[-2:]): - logger.info("Forward upsample size to force interpolation output size.") - forward_upsample_size = True - - # prepare attention_mask - if attention_mask is not None: - attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # 1. time - timesteps = timestep - if not torch.is_tensor(timesteps): - # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can - # This would be a good case for the `match` statement (Python 3.10+) - dtype = maybe_adjust_dtype_for_device( - torch.float64 if isinstance(timestep, float) else torch.int64, sample.device - ) - timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) - elif len(timesteps.shape) == 0: - timesteps = timesteps[None].to(sample.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - num_frames = sample.shape[2] - timesteps = timesteps.expand(sample.shape[0]) - - t_emb = self.time_proj(timesteps) - - # timesteps does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=self.dtype) - - emb = self.time_embedding(t_emb, timestep_cond) - emb = emb.repeat_interleave(num_frames, dim=0, output_size=emb.shape[0] * num_frames) - encoder_hidden_states = encoder_hidden_states.repeat_interleave( - num_frames, dim=0, output_size=encoder_hidden_states.shape[0] * num_frames - ) - - # 2. pre-process - sample = sample.permute(0, 2, 1, 3, 4).reshape((sample.shape[0] * num_frames, -1) + sample.shape[3:]) - sample = self.conv_in(sample) - - sample = self.transformer_in( - sample, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - return_dict=False, - )[0] - - # 3. down - down_block_res_samples = (sample,) - for downsample_block in self.down_blocks: - if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention: - sample, res_samples = downsample_block( - hidden_states=sample, - temb=emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample, res_samples = downsample_block(hidden_states=sample, temb=emb, num_frames=num_frames) - - down_block_res_samples += res_samples - - if down_block_additional_residuals is not None: - new_down_block_res_samples = () - - for down_block_res_sample, down_block_additional_residual in zip( - down_block_res_samples, down_block_additional_residuals - ): - down_block_res_sample = down_block_res_sample + down_block_additional_residual - new_down_block_res_samples += (down_block_res_sample,) - - down_block_res_samples = new_down_block_res_samples - - # 4. mid - if self.mid_block is not None: - sample = self.mid_block( - sample, - emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - ) - - if mid_block_additional_residual is not None: - sample = sample + mid_block_additional_residual - - # 5. up - for i, upsample_block in enumerate(self.up_blocks): - is_final_block = i == len(self.up_blocks) - 1 - - res_samples = down_block_res_samples[-len(upsample_block.resnets) :] - down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)] - - # if we have not reached the final block and need to forward the - # upsample size, we do it here - if not is_final_block and forward_upsample_size: - upsample_size = down_block_res_samples[-1].shape[2:] - - if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention: - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - encoder_hidden_states=encoder_hidden_states, - upsample_size=upsample_size, - attention_mask=attention_mask, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - upsample_size=upsample_size, - num_frames=num_frames, - ) - - # 6. post-process - if self.conv_norm_out: - sample = self.conv_norm_out(sample) - sample = self.conv_act(sample) - - sample = self.conv_out(sample) - - # reshape to (batch, channel, framerate, width, height) - sample = sample[None, :].reshape((-1, num_frames) + sample.shape[1:]).permute(0, 2, 1, 3, 4) - - if not return_dict: - return (sample,) - - return UNet3DConditionOutput(sample=sample) diff --git a/diffusers/models/unets/unet_dreamlite.py b/diffusers/models/unets/unet_dreamlite.py deleted file mode 100644 index e9d3397c16dd8295211f477177b1ca8d5de99495..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_dreamlite.py +++ /dev/null @@ -1,2041 +0,0 @@ -# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -""" -DreamLite UNet model and its constituent 2D blocks. - -This single file mirrors the structure used by recent diffusers transformer model files: it defines all DreamLite -building blocks (Down / Mid / Up) and the top-level :class:`DreamLiteUNetModel` together. - -Compared to the upstream ``unet_2d_blocks`` Down/Mid/Up cross-attention blocks, the DreamLite variants additionally -thread the following knobs: - -- ``use_sep_conv``: replace standard convs in :class:`ResnetBlock2DDreamLite` with depthwise-separable convs - (mobile-friendly). -- ``qk_norm``, ``num_kv_heads``, ``ff_mult``: propagated into :class:`DreamLiteTransformer2DModel` / - :class:`BasicTransformerBlockDreamLite`. - -The two "no self-attention" variants hard-code ``use_self_attention=False`` in their -:class:`DreamLiteTransformer2DModel` calls. - -The U-Net itself defaults its attention processors to :class:`DreamLiteAttnProcessor2_0` (GQA-aware SDPA), which is -required because the upstream ``AttnProcessor2_0`` does not handle ``kv_heads != heads`` correctly. -""" - -from __future__ import annotations - -from functools import partial -from typing import Any, Optional - -import torch -import torch.nn.functional as F -from torch import nn - -from ...configuration_utils import register_to_config -from ..activations import get_activation -from ..attention_dispatch import dispatch_attention_fn -from ..attention_processor import Attention -from ..downsampling import Downsample2D as _CoreDownsample2D -from ..downsampling import downsample_2d -from ..modeling_utils import ModelMixin -from ..normalization import RMSNorm -from ..transformers.dual_transformer_2d import DualTransformer2DModel -from ..transformers.transformer_2d_dreamlite import DreamLiteTransformer2DModel -from ..upsampling import Upsample2D as _CoreUpsample2D -from ..upsampling import upsample_2d -from .unet_2d_blocks import Downsample2D, Upsample2D, apply_freeu -from .unet_2d_condition import UNet2DConditionModel - - -# --------------------------------------------------------------------------- -# Building blocks (resnet + attention processor) -# --------------------------------------------------------------------------- -class DepthwiseSeparableConv(nn.Module): - """ - Depthwise separable convolution used by DreamLite mobile-friendly ResNet blocks. - - A depthwise convolution (groups == in_channels) followed by a 1x1 pointwise convolution. The pointwise output - channel count is multiplied by `expand_ratio` to support inverted-residual style expansion / contraction inside - [`ResnetBlock2DDreamLite`]. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int, - stride: int = 1, - padding: int = 0, - bias: bool = False, - expand_ratio: float = 1, - ): - super().__init__() - self.depthwise = nn.Conv2d( - in_channels, - in_channels, - kernel_size=kernel_size, - stride=stride, - padding=padding, - groups=in_channels, - bias=bias, - ) - self.pointwise = nn.Conv2d(in_channels, int(out_channels * expand_ratio), kernel_size=1, bias=bias) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = self.depthwise(hidden_states) - hidden_states = self.pointwise(hidden_states) - return hidden_states - - -class ResnetBlock2DDreamLite(nn.Module): - r""" - A ResNet block used by DreamLite. Mirrors [`diffusers.models.resnet.ResnetBlock2D`] with one extra option: - - use_sep_conv (`bool`, *optional*, defaults to `False`): - Replace the two 3x3 convolutions with [`DepthwiseSeparableConv`]. The first conv expands the channel count - by 2x; the second conv contracts it back. Used by the mobile-friendly DreamLite checkpoints. - - All other parameters behave identically to [`diffusers.models.resnet.ResnetBlock2D`]. - """ - - def __init__( - self, - *, - in_channels: int, - out_channels: Optional[int] = None, - conv_shortcut: bool = False, - dropout: float = 0.0, - temb_channels: int = 512, - groups: int = 32, - groups_out: Optional[int] = None, - pre_norm: bool = True, - eps: float = 1e-6, - non_linearity: str = "swish", - skip_time_act: bool = False, - time_embedding_norm: str = "default", - kernel: Optional[torch.Tensor] = None, - output_scale_factor: float = 1.0, - use_in_shortcut: Optional[bool] = None, - up: bool = False, - down: bool = False, - conv_shortcut_bias: bool = True, - conv_2d_out_channels: Optional[int] = None, - use_sep_conv: bool = False, - ): - super().__init__() - if time_embedding_norm in ("ada_group", "spatial"): - raise ValueError( - f"`time_embedding_norm`={time_embedding_norm!r} is not supported by `ResnetBlock2DDreamLite`. " - "Use `diffusers.models.resnet.ResnetBlockCondNorm2D` instead." - ) - - self.pre_norm = True - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - self.use_conv_shortcut = conv_shortcut - self.up = up - self.down = down - self.output_scale_factor = output_scale_factor - self.time_embedding_norm = time_embedding_norm - self.skip_time_act = skip_time_act - - if groups_out is None: - groups_out = groups - - self.norm1 = nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True) - - # Inverted-residual style expansion when `use_sep_conv=True`: conv1 expands channels by 2x, - # conv2 contracts them back. For the standard branch this is just a regular 3x3 conv. - if use_sep_conv: - expand_ratio = 2 - self.conv1 = DepthwiseSeparableConv( - in_channels, out_channels, kernel_size=3, stride=1, padding=1, expand_ratio=expand_ratio - ) - out_channels = out_channels * expand_ratio - else: - expand_ratio = 1 - self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) - - if temb_channels is not None: - if self.time_embedding_norm == "default": - self.time_emb_proj = nn.Linear(temb_channels, out_channels) - elif self.time_embedding_norm == "scale_shift": - self.time_emb_proj = nn.Linear(temb_channels, 2 * out_channels) - else: - raise ValueError(f"unknown time_embedding_norm : {self.time_embedding_norm}") - else: - self.time_emb_proj = None - - self.norm2 = nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True) - - self.dropout = nn.Dropout(dropout) - conv_2d_out_channels = conv_2d_out_channels or out_channels - if use_sep_conv: - self.conv2 = DepthwiseSeparableConv( - out_channels, - conv_2d_out_channels, - kernel_size=3, - stride=1, - padding=1, - expand_ratio=1 / expand_ratio, - ) - conv_2d_out_channels = conv_2d_out_channels // expand_ratio - else: - self.conv2 = nn.Conv2d(out_channels, conv_2d_out_channels, kernel_size=3, stride=1, padding=1) - - self.nonlinearity = get_activation(non_linearity) - - self.upsample = self.downsample = None - if self.up: - if kernel == "fir": - fir_kernel = (1, 3, 3, 1) - self.upsample = lambda x: upsample_2d(x, kernel=fir_kernel) - elif kernel == "sde_vp": - self.upsample = partial(F.interpolate, scale_factor=2.0, mode="nearest") - else: - self.upsample = _CoreUpsample2D(in_channels, use_conv=False) - elif self.down: - if kernel == "fir": - fir_kernel = (1, 3, 3, 1) - self.downsample = lambda x: downsample_2d(x, kernel=fir_kernel) - elif kernel == "sde_vp": - self.downsample = partial(F.avg_pool2d, kernel_size=2, stride=2) - else: - self.downsample = _CoreDownsample2D(in_channels, use_conv=False, padding=1, name="op") - - self.use_in_shortcut = self.in_channels != conv_2d_out_channels if use_in_shortcut is None else use_in_shortcut - - self.conv_shortcut = None - if self.use_in_shortcut: - self.conv_shortcut = nn.Conv2d( - in_channels, - conv_2d_out_channels, - kernel_size=1, - stride=1, - padding=0, - bias=conv_shortcut_bias, - ) - - def forward(self, input_tensor: torch.Tensor, temb: torch.Tensor) -> torch.Tensor: - hidden_states = input_tensor - - hidden_states = self.norm1(hidden_states) - hidden_states = self.nonlinearity(hidden_states) - - if self.upsample is not None: - # upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984 - if hidden_states.shape[0] >= 64: - input_tensor = input_tensor.contiguous() - hidden_states = hidden_states.contiguous() - input_tensor = self.upsample(input_tensor) - hidden_states = self.upsample(hidden_states) - elif self.downsample is not None: - input_tensor = self.downsample(input_tensor) - hidden_states = self.downsample(hidden_states) - - hidden_states = self.conv1(hidden_states) - - if self.time_emb_proj is not None: - if not self.skip_time_act: - temb = self.nonlinearity(temb) - temb = self.time_emb_proj(temb)[:, :, None, None] - - if self.time_embedding_norm == "default": - if temb is not None: - hidden_states = hidden_states + temb - hidden_states = self.norm2(hidden_states) - elif self.time_embedding_norm == "scale_shift": - if temb is None: - raise ValueError(f"`temb` should not be None when `time_embedding_norm` is {self.time_embedding_norm}") - time_scale, time_shift = torch.chunk(temb, 2, dim=1) - hidden_states = self.norm2(hidden_states) - hidden_states = hidden_states * (1 + time_scale) + time_shift - else: - hidden_states = self.norm2(hidden_states) - - hidden_states = self.nonlinearity(hidden_states) - - hidden_states = self.dropout(hidden_states) - hidden_states = self.conv2(hidden_states) - - if self.conv_shortcut is not None: - # Only call .contiguous() under training, to avoid DDP gradient-stride warnings while keeping - # inference fast (especially on CPU). Mirrors the upstream fix from huggingface/diffusers#12975. - if self.training: - input_tensor = input_tensor.contiguous() - input_tensor = self.conv_shortcut(input_tensor) - - output_tensor = (input_tensor + hidden_states) / self.output_scale_factor - - return output_tensor - - -class DreamLiteAttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention with Grouped Query Attention (GQA / MQA) support. - - Identical to :class:`AttnProcessor2_0` except the key/value reshape branch correctly handles ``attn.kv_heads != - attn.heads`` by reshaping K/V to ``kv_heads`` and then ``repeat_interleave``-ing them up to ``attn.heads``. This is - required by the DreamLite UNet, which combines GQA with ``qk_norm`` — a combination the default - :class:`AttnProcessor2_0` does not handle. SDPA is delegated to :func:`dispatch_attention_fn` so any of the - diffusers attention backends (native PyTorch SDPA, FlashAttention, etc.) can be used. - """ - - _attention_backend = None - _parallel_config = None - - def __call__( - self, - attn: Attention, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - temb: torch.Tensor | None = None, - ) -> torch.Tensor: - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - # --- GQA-aware reshape (the only real difference vs AttnProcessor2_0) --- - # ``dispatch_attention_fn`` expects (batch, seq, heads, head_dim) — keep Q/K/V in that layout - # and let the dispatched backend handle the transpose internally. - head_dim = query.shape[-1] // attn.heads - kv_heads = key.shape[-1] // head_dim - - query = query.view(batch_size, -1, attn.heads, head_dim) - key = key.view(batch_size, -1, kv_heads, head_dim) - value = value.view(batch_size, -1, kv_heads, head_dim) - - if attn.norm_q is not None: - query = attn.norm_q(query) - if attn.norm_k is not None: - key = attn.norm_k(key) - - if kv_heads != attn.heads: - # GQA / MQA: repeat K/V heads up to query heads for SDPA. - heads_per_kv_head = attn.heads // kv_heads - key = torch.repeat_interleave(key, heads_per_kv_head, dim=2, output_size=key.shape[2] * heads_per_kv_head) - value = torch.repeat_interleave( - value, heads_per_kv_head, dim=2, output_size=value.shape[2] * heads_per_kv_head - ) - # ------------------------------------------------------------------------ - - # the output of sdp = (batch, seq_len, num_heads, head_dim) - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=self._parallel_config, - ) - - hidden_states = hidden_states.flatten(2, 3) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -# --------------------------------------------------------------------------- -# Mid block -# --------------------------------------------------------------------------- -class DreamLiteUNetMidBlock2DCrossAttn(nn.Module): - def __init__( - self, - in_channels: int, - temb_channels: int, - out_channels: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_groups_out: int | None = None, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - output_scale_factor: float = 1.0, - cross_attention_dim: int = 1280, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - upcast_attention: bool = False, - attention_type: str = "default", - # DreamLite extras - qk_norm: str | None = None, - use_sep_conv: bool = False, - ff_mult: int = 4, - num_kv_heads: int | None = None, - num_mid_layers: int = 1, - ): - super().__init__() - - out_channels = out_channels or in_channels - self.in_channels = in_channels - self.out_channels = out_channels - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32) - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - resnet_groups_out = resnet_groups_out or resnet_groups - - resnets = [ - ResnetBlock2DDreamLite( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - groups_out=resnet_groups_out, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_sep_conv=use_sep_conv, - ) - ] - attentions = [] - - for i in range(num_layers): - if not dual_cross_attention: - attentions.append( - DreamLiteTransformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups_out, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - attention_type=attention_type, - qk_norm=qk_norm, - ff_mult=ff_mult, - num_kv_heads=num_kv_heads, - ) - ) - else: - attentions.append( - DualTransformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - resnets.append( - ResnetBlock2DDreamLite( - in_channels=out_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups_out, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_sep_conv=use_sep_conv, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - hidden_states = self.resnets[0](hidden_states, temb) - for attn, resnet in zip(self.attentions, self.resnets[1:]): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - hidden_states = resnet(hidden_states, temb) - - return hidden_states - - -# --------------------------------------------------------------------------- -# Down blocks -# --------------------------------------------------------------------------- -class DreamLiteCrossAttnDownBlock2D(nn.Module): - """DreamLite down block with both self- and cross-attention in each transformer layer.""" - - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - downsample_padding: int = 1, - add_downsample: bool = True, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - attention_type: str = "default", - # DreamLite extras - qk_norm: str | None = None, - use_sep_conv: bool = False, - ff_mult: int = 4, - num_kv_heads: int | None = None, - ): - super().__init__() - resnets = [] - attentions = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - for i in range(num_layers): - in_ch = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2DDreamLite( - in_channels=in_ch, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_sep_conv=use_sep_conv, - ) - ) - if not dual_cross_attention: - attentions.append( - DreamLiteTransformer2DModel( - num_attention_heads=num_attention_heads, - attention_head_dim=out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - attention_type=attention_type, - qk_norm=qk_norm, - ff_mult=ff_mult, - num_kv_heads=num_kv_heads, - ) - ) - else: - attentions.append( - DualTransformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, - use_conv=True, - out_channels=out_channels, - padding=downsample_padding, - name="op", - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - additional_residuals: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - output_states: tuple[torch.Tensor, ...] = () - blocks = list(zip(self.resnets, self.attentions)) - - for i, (resnet, attn) in enumerate(blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - if i == len(blocks) - 1 and additional_residuals is not None: - hidden_states = hidden_states + additional_residuals - - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - output_states = output_states + (hidden_states,) - - return hidden_states, output_states - - -class DreamLiteCrossAttnNoSelfAttnDownBlock2D(nn.Module): - """DreamLite down block with cross-attention only (self-attention is removed).""" - - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - downsample_padding: int = 1, - add_downsample: bool = True, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - attention_type: str = "default", - # DreamLite extras - qk_norm: str | None = None, - use_sep_conv: bool = False, - ff_mult: int = 4, - num_kv_heads: int | None = None, - ): - super().__init__() - resnets = [] - attentions = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - for i in range(num_layers): - in_ch = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2DDreamLite( - in_channels=in_ch, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_sep_conv=use_sep_conv, - ) - ) - if not dual_cross_attention: - attentions.append( - DreamLiteTransformer2DModel( - num_attention_heads=num_attention_heads, - attention_head_dim=out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - attention_type=attention_type, - qk_norm=qk_norm, - ff_mult=ff_mult, - num_kv_heads=num_kv_heads, - # DreamLite "remove self-attention" path: - use_self_attention=False, - ) - ) - else: - attentions.append( - DualTransformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, - use_conv=True, - out_channels=out_channels, - padding=downsample_padding, - name="op", - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - additional_residuals: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - output_states: tuple[torch.Tensor, ...] = () - blocks = list(zip(self.resnets, self.attentions)) - - for i, (resnet, attn) in enumerate(blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - if i == len(blocks) - 1 and additional_residuals is not None: - hidden_states = hidden_states + additional_residuals - - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - output_states = output_states + (hidden_states,) - - return hidden_states, output_states - - -class DreamLiteDownBlock2D(nn.Module): - """DreamLite plain resnet-only down block (no attention).""" - - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_downsample: bool = True, - downsample_padding: int = 1, - use_sep_conv: bool = False, - ): - super().__init__() - resnets = [] - for i in range(num_layers): - in_ch = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2DDreamLite( - in_channels=in_ch, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_sep_conv=use_sep_conv, - ) - ) - self.resnets = nn.ModuleList(resnets) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, - use_conv=True, - out_channels=out_channels, - padding=downsample_padding, - name="op", - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - **kwargs, - ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: - output_states: tuple[torch.Tensor, ...] = () - for resnet in self.resnets: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) - output_states = output_states + (hidden_states,) - - return hidden_states, output_states - - -# --------------------------------------------------------------------------- -# Up blocks -# --------------------------------------------------------------------------- -class DreamLiteCrossAttnUpBlock2D(nn.Module): - """DreamLite up block with both self- and cross-attention in each transformer layer.""" - - def __init__( - self, - in_channels: int, - out_channels: int, - prev_output_channel: int, - temb_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - attention_type: str = "default", - # DreamLite extras - qk_norm: str | None = None, - use_sep_conv: bool = False, - ff_mult: int = 4, - num_kv_heads: int | None = None, - ): - super().__init__() - resnets = [] - attentions = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2DDreamLite( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_sep_conv=use_sep_conv, - ) - ) - if not dual_cross_attention: - attentions.append( - DreamLiteTransformer2DModel( - num_attention_heads=num_attention_heads, - attention_head_dim=out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - attention_type=attention_type, - qk_norm=qk_norm, - ff_mult=ff_mult, - num_kv_heads=num_kv_heads, - ) - ) - else: - attentions.append( - DualTransformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - upsample_size: int | None = None, - attention_mask: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - is_freeu_enabled = ( - getattr(self, "s1", None) - and getattr(self, "s2", None) - and getattr(self, "b1", None) - and getattr(self, "b2", None) - ) - - for resnet, attn in zip(self.resnets, self.attentions): - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - - if is_freeu_enabled: - hidden_states, res_hidden_states = apply_freeu( - self.resolution_idx, - hidden_states, - res_hidden_states, - s1=self.s1, - s2=self.s2, - b1=self.b1, - b2=self.b2, - ) - - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size) - - return hidden_states - - -class DreamLiteCrossAttnNoSelfAttnUpBlock2D(nn.Module): - """DreamLite up block with cross-attention only (self-attention is removed).""" - - def __init__( - self, - in_channels: int, - out_channels: int, - prev_output_channel: int, - temb_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - attention_type: str = "default", - # DreamLite extras - qk_norm: str | None = None, - use_sep_conv: bool = False, - ff_mult: int = 4, - num_kv_heads: int | None = None, - ): - super().__init__() - resnets = [] - attentions = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * num_layers - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2DDreamLite( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_sep_conv=use_sep_conv, - ) - ) - if not dual_cross_attention: - attentions.append( - DreamLiteTransformer2DModel( - num_attention_heads=num_attention_heads, - attention_head_dim=out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - attention_type=attention_type, - qk_norm=qk_norm, - ff_mult=ff_mult, - num_kv_heads=num_kv_heads, - # DreamLite "remove self-attention" path: - use_self_attention=False, - ) - ) - else: - attentions.append( - DualTransformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - upsample_size: int | None = None, - attention_mask: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - ) -> torch.Tensor: - is_freeu_enabled = ( - getattr(self, "s1", None) - and getattr(self, "s2", None) - and getattr(self, "b1", None) - and getattr(self, "b2", None) - ) - - for resnet, attn in zip(self.resnets, self.attentions): - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - - if is_freeu_enabled: - hidden_states, res_hidden_states = apply_freeu( - self.resolution_idx, - hidden_states, - res_hidden_states, - s1=self.s1, - s2=self.s2, - b1=self.b1, - b2=self.b2, - ) - - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size) - - return hidden_states - - -class DreamLiteUpBlock2D(nn.Module): - """DreamLite plain resnet-only up block (no attention).""" - - def __init__( - self, - in_channels: int, - prev_output_channel: int, - out_channels: int, - temb_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - use_sep_conv: bool = False, - ): - super().__init__() - resnets = [] - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2DDreamLite( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - use_sep_conv=use_sep_conv, - ) - ) - self.resnets = nn.ModuleList(resnets) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - upsample_size: int | None = None, - **kwargs, - ) -> torch.Tensor: - is_freeu_enabled = ( - getattr(self, "s1", None) - and getattr(self, "s2", None) - and getattr(self, "b1", None) - and getattr(self, "b2", None) - ) - - for resnet in self.resnets: - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - - if is_freeu_enabled: - hidden_states, res_hidden_states = apply_freeu( - self.resolution_idx, - hidden_states, - res_hidden_states, - s1=self.s1, - s2=self.s2, - b1=self.b1, - b2=self.b2, - ) - - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(hidden_states, temb) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size) - - return hidden_states - - -# --------------------------------------------------------------------------- -# Local block dispatch (DreamLite-only) -# -# The string ``down_block_type`` / ``up_block_type`` / ``mid_block_type`` keys -# persisted in saved checkpoints' ``config.json`` usually mirror the Python class -# names defined above. Some configs use upstream UNet block names instead. -# --------------------------------------------------------------------------- -_DREAMLITE_DOWN_BLOCK_ALIASES = { - "CrossAttnDownRemoveSelfAttnBlock2D": "DreamLiteCrossAttnNoSelfAttnDownBlock2D", - "CrossAttnDownBlock2D": "DreamLiteCrossAttnDownBlock2D", - "DownBlock2D": "DreamLiteDownBlock2D", -} - -_DREAMLITE_MID_BLOCK_ALIASES = { - "UNetMidBlock2DCrossAttn": "DreamLiteUNetMidBlock2DCrossAttn", -} - -_DREAMLITE_UP_BLOCK_ALIASES = { - "CrossAttnUpRemoveSelfAttnBlock2D": "DreamLiteCrossAttnNoSelfAttnUpBlock2D", - "CrossAttnUpRemoveSelfAttnBlock2DV1": "DreamLiteCrossAttnNoSelfAttnUpBlock2D", - "CrossAttnUpBlock2D": "DreamLiteCrossAttnUpBlock2D", - "UpBlock2D": "DreamLiteUpBlock2D", -} - - -def _get_down_block_dreamlite( - down_block_type: str, - *, - num_layers, - transformer_layers_per_block, - in_channels, - out_channels, - temb_channels, - add_downsample, - resnet_eps, - resnet_act_fn, - resnet_groups, - cross_attention_dim, - num_attention_heads, - downsample_padding, - dual_cross_attention, - use_linear_projection, - only_cross_attention, - upcast_attention, - resnet_time_scale_shift, - attention_type, - dropout, - qk_norm, - use_sep_conv, - ff_mult, - num_kv_heads, -): - down_block_type = _DREAMLITE_DOWN_BLOCK_ALIASES.get(down_block_type, down_block_type) - - if down_block_type == "DreamLiteDownBlock2D": - return DreamLiteDownBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - downsample_padding=downsample_padding, - resnet_time_scale_shift=resnet_time_scale_shift, - use_sep_conv=use_sep_conv, - ) - if down_block_type in ("DreamLiteCrossAttnDownBlock2D", "DreamLiteCrossAttnNoSelfAttnDownBlock2D"): - if cross_attention_dim is None: - raise ValueError(f"cross_attention_dim must be specified for {down_block_type}") - cls = ( - DreamLiteCrossAttnDownBlock2D - if down_block_type == "DreamLiteCrossAttnDownBlock2D" - else DreamLiteCrossAttnNoSelfAttnDownBlock2D - ) - return cls( - num_layers=num_layers, - transformer_layers_per_block=transformer_layers_per_block, - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - dropout=dropout, - add_downsample=add_downsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - downsample_padding=downsample_padding, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads, - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - attention_type=attention_type, - qk_norm=qk_norm, - use_sep_conv=use_sep_conv, - ff_mult=ff_mult, - num_kv_heads=num_kv_heads, - ) - raise ValueError(f"DreamLite does not support down_block_type={down_block_type!r}") - - -def _get_mid_block_dreamlite( - mid_block_type, - *, - temb_channels, - in_channels, - resnet_eps, - resnet_act_fn, - resnet_groups, - output_scale_factor, - transformer_layers_per_block, - num_attention_heads, - cross_attention_dim, - dual_cross_attention, - use_linear_projection, - upcast_attention, - resnet_time_scale_shift, - attention_type, - dropout, - qk_norm, - use_sep_conv, - ff_mult, - num_kv_heads, - num_mid_layers=1, -): - if mid_block_type is None: - return None - mid_block_type = _DREAMLITE_MID_BLOCK_ALIASES.get(mid_block_type, mid_block_type) - - if mid_block_type == "DreamLiteUNetMidBlock2DCrossAttn": - return DreamLiteUNetMidBlock2DCrossAttn( - transformer_layers_per_block=transformer_layers_per_block, - in_channels=in_channels, - temb_channels=temb_channels, - dropout=dropout, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - output_scale_factor=output_scale_factor, - resnet_time_scale_shift=resnet_time_scale_shift, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads, - resnet_groups=resnet_groups, - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - attention_type=attention_type, - qk_norm=qk_norm, - use_sep_conv=use_sep_conv, - ff_mult=ff_mult, - num_kv_heads=num_kv_heads, - num_layers=num_mid_layers, - ) - raise ValueError(f"DreamLite does not support mid_block_type={mid_block_type!r}") - - -def _get_up_block_dreamlite( - up_block_type, - *, - num_layers, - transformer_layers_per_block, - in_channels, - out_channels, - prev_output_channel, - temb_channels, - add_upsample, - resnet_eps, - resnet_act_fn, - resolution_idx, - resnet_groups, - cross_attention_dim, - num_attention_heads, - dual_cross_attention, - use_linear_projection, - only_cross_attention, - upcast_attention, - resnet_time_scale_shift, - attention_type, - dropout, - qk_norm, - use_sep_conv, - ff_mult, - num_kv_heads, -): - up_block_type = _DREAMLITE_UP_BLOCK_ALIASES.get(up_block_type, up_block_type) - - if up_block_type == "DreamLiteUpBlock2D": - return DreamLiteUpBlock2D( - num_layers=num_layers, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - resnet_time_scale_shift=resnet_time_scale_shift, - use_sep_conv=use_sep_conv, - ) - if up_block_type in ("DreamLiteCrossAttnUpBlock2D", "DreamLiteCrossAttnNoSelfAttnUpBlock2D"): - if cross_attention_dim is None: - raise ValueError(f"cross_attention_dim must be specified for {up_block_type}") - cls = ( - DreamLiteCrossAttnUpBlock2D - if up_block_type == "DreamLiteCrossAttnUpBlock2D" - else DreamLiteCrossAttnNoSelfAttnUpBlock2D - ) - return cls( - num_layers=num_layers, - transformer_layers_per_block=transformer_layers_per_block, - in_channels=in_channels, - out_channels=out_channels, - prev_output_channel=prev_output_channel, - temb_channels=temb_channels, - resolution_idx=resolution_idx, - dropout=dropout, - add_upsample=add_upsample, - resnet_eps=resnet_eps, - resnet_act_fn=resnet_act_fn, - resnet_groups=resnet_groups, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads, - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - attention_type=attention_type, - qk_norm=qk_norm, - use_sep_conv=use_sep_conv, - ff_mult=ff_mult, - num_kv_heads=num_kv_heads, - ) - raise ValueError(f"DreamLite does not support up_block_type={up_block_type!r}") - - -# --------------------------------------------------------------------------- -# Model -# --------------------------------------------------------------------------- -class DreamLiteUNetModel(UNet2DConditionModel): - r""" - DreamLite variant of :class:`UNet2DConditionModel`. - - Differences vs the parent class: - - * Down / Mid / Up blocks are dispatched to the DreamLite variants defined above, which support depthwise-separable - convolutions in resnets and Grouped Query Attention with RMSNorm ``qk_norm`` in attention. - * ``default_attn_processor`` returns :class:`DreamLiteAttnProcessor2_0` so SDPA is GQA-aware out of the box. - """ - - _supports_gradient_checkpointing = True - _no_split_modules = [ - "BasicTransformerBlockDreamLite", - "ResnetBlock2DDreamLite", - "DreamLiteCrossAttnUpBlock2D", - "DreamLiteUpBlock2D", - ] - _repeated_blocks = ["BasicTransformerBlockDreamLite"] - - @register_to_config - def __init__( - self, - sample_size: int | tuple[int, int] | None = None, - in_channels: int = 4, - out_channels: int = 4, - center_input_sample: bool = False, - flip_sin_to_cos: bool = True, - freq_shift: int = 0, - down_block_types: tuple[str, ...] = ( - "DreamLiteCrossAttnNoSelfAttnDownBlock2D", - "DreamLiteCrossAttnNoSelfAttnDownBlock2D", - "DreamLiteCrossAttnDownBlock2D", - ), - mid_block_type: str | None = "DreamLiteUNetMidBlock2DCrossAttn", - up_block_types: tuple[str, ...] = ( - "DreamLiteCrossAttnUpBlock2D", - "DreamLiteCrossAttnNoSelfAttnUpBlock2D", - "DreamLiteUpBlock2D", - ), - only_cross_attention: bool | tuple[bool, ...] = False, - block_out_channels: tuple[int, ...] = (320, 640, 1280), - layers_per_block: int | tuple[int, ...] = 2, - downsample_padding: int = 1, - mid_block_scale_factor: float = 1, - dropout: float = 0.0, - act_fn: str = "silu", - norm_num_groups: int | None = 32, - norm_eps: float = 1e-5, - cross_attention_dim: int | tuple[int, ...] = 2048, - transformer_layers_per_block: int | tuple[int, ...] | tuple[tuple, ...] = 1, - reverse_transformer_layers_per_block: tuple[tuple[int, ...], ...] | None = None, - encoder_hid_dim: int | None = None, - encoder_hid_dim_type: str | None = None, - attention_head_dim: int | tuple[int, ...] = 64, - num_attention_heads: int | tuple[int, ...] | None = None, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - class_embed_type: str | None = None, - addition_embed_type: str | None = None, - addition_time_embed_dim: int | None = None, - num_class_embeds: int | None = None, - upcast_attention: bool = False, - resnet_time_scale_shift: str = "default", - resnet_skip_time_act: bool = False, - resnet_out_scale_factor: float = 1.0, - time_embedding_type: str = "positional", - time_embedding_dim: int | None = None, - time_embedding_act_fn: str | None = None, - timestep_post_act: str | None = None, - time_cond_proj_dim: int | None = None, - conv_in_kernel: int = 3, - conv_out_kernel: int = 3, - projection_class_embeddings_input_dim: int | None = None, - attention_type: str = "default", - class_embeddings_concat: bool = False, - mid_block_only_cross_attention: bool | None = None, - cross_attention_norm: str | None = None, - addition_embed_type_num_heads: int = 64, - # ---- DreamLite extras ---- - qk_norm: str | None = "rms_norm", - use_sep_conv: bool = True, - ff_mult: int = 6, - num_kv_heads: int | None = 1, - num_mid_layers: int = 1, - ): - # NOTE: deliberately skip UNet2DConditionModel.__init__ because we replicate - # the body with DreamLite block dispatch, but call ModelMixin.__init__ so that - # mixin state (e.g. _gradient_checkpointing_func) is properly initialised. - ModelMixin.__init__(self) - - self.sample_size = sample_size - - if num_attention_heads is not None: - raise ValueError( - "At the moment it is not possible to define the number of attention heads via " - "`num_attention_heads` because of a naming issue as described in " - "https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131. " - "Passing `num_attention_heads` will only be supported in diffusers v0.19." - ) - num_attention_heads = num_attention_heads or attention_head_dim - - # Reuse parent helpers (they only touch self, no super().__init__ required). - self._check_config( - down_block_types=down_block_types, - up_block_types=up_block_types, - only_cross_attention=only_cross_attention, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - cross_attention_dim=cross_attention_dim, - transformer_layers_per_block=transformer_layers_per_block, - reverse_transformer_layers_per_block=reverse_transformer_layers_per_block, - attention_head_dim=attention_head_dim, - num_attention_heads=num_attention_heads, - ) - - self.projection_class_embeddings_input_dim = projection_class_embeddings_input_dim - - # input - conv_in_padding = (conv_in_kernel - 1) // 2 - self.conv_in = nn.Conv2d( - in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding - ) - - # time - time_embed_dim, timestep_input_dim = self._set_time_proj( - time_embedding_type, - block_out_channels=block_out_channels, - flip_sin_to_cos=flip_sin_to_cos, - freq_shift=freq_shift, - time_embedding_dim=time_embedding_dim, - ) - - from ..embeddings import TimestepEmbedding # local import to avoid cycle - - self.time_embedding = TimestepEmbedding( - timestep_input_dim, - time_embed_dim, - act_fn=act_fn, - post_act_fn=timestep_post_act, - cond_proj_dim=time_cond_proj_dim, - ) - - self._set_encoder_hid_proj( - encoder_hid_dim_type, - cross_attention_dim=cross_attention_dim, - encoder_hid_dim=encoder_hid_dim, - ) - self._set_class_embedding( - class_embed_type, - act_fn=act_fn, - num_class_embeds=num_class_embeds, - projection_class_embeddings_input_dim=projection_class_embeddings_input_dim, - time_embed_dim=time_embed_dim, - timestep_input_dim=timestep_input_dim, - ) - self._set_add_embedding( - addition_embed_type, - addition_embed_type_num_heads=addition_embed_type_num_heads, - addition_time_embed_dim=addition_time_embed_dim, - cross_attention_dim=cross_attention_dim, - encoder_hid_dim=encoder_hid_dim, - flip_sin_to_cos=flip_sin_to_cos, - freq_shift=freq_shift, - projection_class_embeddings_input_dim=projection_class_embeddings_input_dim, - time_embed_dim=time_embed_dim, - ) - - self.time_embed_act = None if time_embedding_act_fn is None else get_activation(time_embedding_act_fn) - - self.down_blocks = nn.ModuleList([]) - self.up_blocks = nn.ModuleList([]) - - # Normalize per-stage args - if isinstance(only_cross_attention, bool): - if mid_block_only_cross_attention is None: - mid_block_only_cross_attention = only_cross_attention - only_cross_attention = [only_cross_attention] * len(down_block_types) - if mid_block_only_cross_attention is None: - mid_block_only_cross_attention = False - if isinstance(num_attention_heads, int): - num_attention_heads = (num_attention_heads,) * len(down_block_types) - if isinstance(attention_head_dim, int): - attention_head_dim = (attention_head_dim,) * len(down_block_types) - if isinstance(cross_attention_dim, int): - cross_attention_dim = (cross_attention_dim,) * len(down_block_types) - if isinstance(layers_per_block, int): - layers_per_block = [layers_per_block] * len(down_block_types) - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types) - - blocks_time_embed_dim = time_embed_dim * 2 if class_embeddings_concat else time_embed_dim - - # ---- Down ---- - output_channel = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - self.down_blocks.append( - _get_down_block_dreamlite( - down_block_type, - num_layers=layers_per_block[i], - transformer_layers_per_block=transformer_layers_per_block[i], - in_channels=input_channel, - out_channels=output_channel, - temb_channels=blocks_time_embed_dim, - add_downsample=not is_final_block, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - cross_attention_dim=cross_attention_dim[i], - num_attention_heads=num_attention_heads[i], - downsample_padding=downsample_padding, - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention[i], - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - attention_type=attention_type, - dropout=dropout, - qk_norm=qk_norm, - use_sep_conv=use_sep_conv, - ff_mult=ff_mult, - num_kv_heads=num_kv_heads, - ) - ) - - # ---- Mid ---- - self.mid_block = _get_mid_block_dreamlite( - mid_block_type, - temb_channels=blocks_time_embed_dim, - in_channels=block_out_channels[-1], - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - output_scale_factor=mid_block_scale_factor, - transformer_layers_per_block=transformer_layers_per_block[-1], - num_attention_heads=num_attention_heads[-1], - cross_attention_dim=cross_attention_dim[-1], - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - attention_type=attention_type, - dropout=dropout, - qk_norm=qk_norm, - use_sep_conv=use_sep_conv, - ff_mult=ff_mult, - num_kv_heads=num_kv_heads, - num_mid_layers=num_mid_layers, - ) - - # ---- Up ---- - self.num_upsamplers = 0 - reversed_block_out_channels = list(reversed(block_out_channels)) - reversed_num_attention_heads = list(reversed(num_attention_heads)) - reversed_layers_per_block = list(reversed(layers_per_block)) - reversed_cross_attention_dim = list(reversed(cross_attention_dim)) - reversed_transformer_layers_per_block = ( - list(reversed(transformer_layers_per_block)) - if reverse_transformer_layers_per_block is None - else reverse_transformer_layers_per_block - ) - only_cross_attention = list(reversed(only_cross_attention)) - - output_channel = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - is_final_block = i == len(block_out_channels) - 1 - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)] - - if not is_final_block: - add_upsample = True - self.num_upsamplers += 1 - else: - add_upsample = False - - self.up_blocks.append( - _get_up_block_dreamlite( - up_block_type, - num_layers=reversed_layers_per_block[i] + 1, - transformer_layers_per_block=reversed_transformer_layers_per_block[i], - in_channels=input_channel, - out_channels=output_channel, - prev_output_channel=prev_output_channel, - temb_channels=blocks_time_embed_dim, - add_upsample=add_upsample, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resolution_idx=i, - resnet_groups=norm_num_groups, - cross_attention_dim=reversed_cross_attention_dim[i], - num_attention_heads=reversed_num_attention_heads[i], - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention[i], - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - attention_type=attention_type, - dropout=dropout, - qk_norm=qk_norm, - use_sep_conv=use_sep_conv, - ff_mult=ff_mult, - num_kv_heads=num_kv_heads, - ) - ) - - # ---- Out ---- - if norm_num_groups is not None: - self.conv_norm_out = nn.GroupNorm( - num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=norm_eps - ) - self.conv_act = get_activation(act_fn) - else: - self.conv_norm_out = None - self.conv_act = None - - conv_out_padding = (conv_out_kernel - 1) // 2 - self.conv_out = nn.Conv2d( - block_out_channels[0], out_channels, kernel_size=conv_out_kernel, padding=conv_out_padding - ) - - self._set_pos_net_if_use_gligen(attention_type=attention_type, cross_attention_dim=cross_attention_dim) - - # ---- DreamLite: install GQA-aware processor everywhere ---- - for module in self.modules(): - if isinstance(module, Attention): - module.set_processor(DreamLiteAttnProcessor2_0()) - - # ----- override default processor so set_attn_processor("default") restores GQA ---- - @property - def default_attn_processor(self): # type: ignore[override] - return DreamLiteAttnProcessor2_0() - - def set_default_attn_processor(self): # type: ignore[override] - """Reinstall :class:`DreamLiteAttnProcessor2_0` everywhere. - - The parent implementation only knows about the diffusers stock processor sets and would raise for our GQA-aware - processor; override so utilities that round-trip through this method (CPU offload, save/load, layerwise - casting, ...) keep working unchanged. - """ - self.set_attn_processor(DreamLiteAttnProcessor2_0()) - - # ----- DreamLite extension: support `text_proj_rms` encoder_hid_proj ----- - def _set_encoder_hid_proj( # type: ignore[override] - self, - encoder_hid_dim_type, - cross_attention_dim, - encoder_hid_dim, - ): - """ - Override to support DreamLite's `text_proj_rms` variant (Linear → RMSNorm). All other variants fall back to the - parent implementation, preserving full compatibility with upstream configs (`text_proj`, `text_image_proj`, - `image_proj`, ...). - """ - if encoder_hid_dim_type == "text_proj_rms": - if encoder_hid_dim is None: - raise ValueError( - "`encoder_hid_dim` has to be defined when `encoder_hid_dim_type` is set to 'text_proj_rms'." - ) - self.encoder_hid_proj = nn.Sequential( - nn.Linear(encoder_hid_dim, cross_attention_dim), - RMSNorm(cross_attention_dim, eps=1e-5, elementwise_affine=True), - ) - return - super()._set_encoder_hid_proj( - encoder_hid_dim_type=encoder_hid_dim_type, - cross_attention_dim=cross_attention_dim, - encoder_hid_dim=encoder_hid_dim, - ) - - # ----- DreamLite extension: dispatch `text_proj_rms` like `text_proj` ----- - def process_encoder_hidden_states( # type: ignore[override] - self, encoder_hidden_states, added_cond_kwargs - ): - """ - For `text_proj_rms`, the projection is a plain `nn.Sequential` applied to `encoder_hidden_states` (same call - signature as `text_proj`). All other variants are delegated to the parent. - """ - if self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "text_proj_rms": - return self.encoder_hid_proj(encoder_hidden_states) - return super().process_encoder_hidden_states( - encoder_hidden_states=encoder_hidden_states, - added_cond_kwargs=added_cond_kwargs, - ) - - # ----- DreamLite extension: support `addition_embed_type == "time"` ----- - def _set_add_embedding( # type: ignore[override] - self, - addition_embed_type, - addition_embed_type_num_heads, - addition_time_embed_dim, - flip_sin_to_cos, - freq_shift, - cross_attention_dim, - encoder_hid_dim, - projection_class_embeddings_input_dim, - time_embed_dim, - ): - """ - Override to support DreamLite's `addition_embed_type == "time"` variant (same module layout as `text_time` but - `get_aug_embed` does not require `text_embeds`). All other variants delegate to the parent implementation. - """ - if addition_embed_type == "time": - from ..embeddings import TimestepEmbedding, Timesteps - - self.add_time_proj = Timesteps(addition_time_embed_dim, flip_sin_to_cos, freq_shift) - self.add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) - return - super()._set_add_embedding( - addition_embed_type=addition_embed_type, - addition_embed_type_num_heads=addition_embed_type_num_heads, - addition_time_embed_dim=addition_time_embed_dim, - flip_sin_to_cos=flip_sin_to_cos, - freq_shift=freq_shift, - cross_attention_dim=cross_attention_dim, - encoder_hid_dim=encoder_hid_dim, - projection_class_embeddings_input_dim=projection_class_embeddings_input_dim, - time_embed_dim=time_embed_dim, - ) - - # ----- DreamLite extension: dispatch `addition_embed_type == "time"` ----- - def get_aug_embed( # type: ignore[override] - self, emb, encoder_hidden_states, added_cond_kwargs - ): - """ - For `addition_embed_type == "time"`, build aug_emb from `time_ids` only (no `text_embeds`). All other variants - are delegated to the parent. - """ - if self.config.addition_embed_type == "time": - if "time_ids" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'time' " - "which requires the keyword argument `time_ids` to be passed in `added_cond_kwargs`" - ) - time_ids = added_cond_kwargs.get("time_ids") - time_embeds = self.add_time_proj(time_ids.flatten()) - time_embeds = time_embeds.reshape((-1, self.config.projection_class_embeddings_input_dim)) - add_embeds = time_embeds.to(emb.dtype) - return self.add_embedding(add_embeds) - return super().get_aug_embed( - emb=emb, - encoder_hidden_states=encoder_hidden_states, - added_cond_kwargs=added_cond_kwargs, - ) - - -__all__ = [ - "DreamLiteUNetModel", - "DreamLiteUNetMidBlock2DCrossAttn", - "DreamLiteCrossAttnDownBlock2D", - "DreamLiteCrossAttnNoSelfAttnDownBlock2D", - "DreamLiteCrossAttnUpBlock2D", - "DreamLiteCrossAttnNoSelfAttnUpBlock2D", - "DreamLiteDownBlock2D", - "DreamLiteUpBlock2D", -] diff --git a/diffusers/models/unets/unet_i2vgen_xl.py b/diffusers/models/unets/unet_i2vgen_xl.py deleted file mode 100644 index 9e7841f95e582f41212a8c238571513ed29df0b8..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_i2vgen_xl.py +++ /dev/null @@ -1,652 +0,0 @@ -# Copyright 2025 Alibaba DAMO-VILAB and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import UNet2DConditionLoadersMixin -from ...utils import logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..activations import get_activation -from ..attention import Attention, AttentionMixin, FeedForward -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - AttnAddedKVProcessor, - AttnProcessor, - FusedAttnProcessor2_0, -) -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from ..transformers.transformer_temporal import TransformerTemporalModel -from .unet_3d_blocks import ( - UNetMidBlock3DCrossAttn, - get_down_block, - get_up_block, -) -from .unet_3d_condition import UNet3DConditionOutput - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class I2VGenXLTransformerTemporalEncoder(nn.Module): - def __init__( - self, - dim: int, - num_attention_heads: int, - attention_head_dim: int, - activation_fn: str = "geglu", - upcast_attention: bool = False, - ff_inner_dim: int | None = None, - dropout: int = 0.0, - ): - super().__init__() - self.norm1 = nn.LayerNorm(dim, elementwise_affine=True, eps=1e-5) - self.attn1 = Attention( - query_dim=dim, - heads=num_attention_heads, - dim_head=attention_head_dim, - dropout=dropout, - bias=False, - upcast_attention=upcast_attention, - out_bias=True, - ) - self.ff = FeedForward( - dim, - dropout=dropout, - activation_fn=activation_fn, - final_dropout=False, - inner_dim=ff_inner_dim, - bias=True, - ) - - def forward( - self, - hidden_states: torch.Tensor, - ) -> torch.Tensor: - norm_hidden_states = self.norm1(hidden_states) - attn_output = self.attn1(norm_hidden_states, encoder_hidden_states=None) - hidden_states = attn_output + hidden_states - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - - ff_output = self.ff(hidden_states) - hidden_states = ff_output + hidden_states - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - - return hidden_states - - -class I2VGenXLUNet(ModelMixin, AttentionMixin, ConfigMixin, UNet2DConditionLoadersMixin): - r""" - I2VGenXL UNet. It is a conditional 3D UNet model that takes a noisy sample, conditional state, and a timestep and - returns a sample-shaped output. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - sample_size (`int` or `tuple[int, int]`, *optional*, defaults to `None`): - Height and width of input/output sample. - in_channels (`int`, *optional*, defaults to 4): The number of channels in the input sample. - out_channels (`int`, *optional*, defaults to 4): The number of channels in the output. - down_block_types (`tuple[str]`, *optional*, defaults to `("CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "DownBlock2D")`): - The tuple of downsample blocks to use. - up_block_types (`tuple[str]`, *optional*, defaults to `("UpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D")`): - The tuple of upsample blocks to use. - block_out_channels (`tuple[int]`, *optional*, defaults to `(320, 640, 1280, 1280)`): - The tuple of output channels for each block. - layers_per_block (`int`, *optional*, defaults to 2): The number of layers per block. - norm_num_groups (`int`, *optional*, defaults to 32): The number of groups to use for the normalization. - If `None`, normalization and activation layers is skipped in post-processing. - cross_attention_dim (`int`, *optional*, defaults to 1280): The dimension of the cross attention features. - attention_head_dim (`int`, *optional*, defaults to 64): Attention head dim. - num_attention_heads (`int`, *optional*): The number of attention heads. - """ - - _supports_gradient_checkpointing = False - - @register_to_config - def __init__( - self, - sample_size: int | None = None, - in_channels: int = 4, - out_channels: int = 4, - down_block_types: tuple[str, ...] = ( - "CrossAttnDownBlock3D", - "CrossAttnDownBlock3D", - "CrossAttnDownBlock3D", - "DownBlock3D", - ), - up_block_types: tuple[str, ...] = ( - "UpBlock3D", - "CrossAttnUpBlock3D", - "CrossAttnUpBlock3D", - "CrossAttnUpBlock3D", - ), - block_out_channels: tuple[int, ...] = (320, 640, 1280, 1280), - layers_per_block: int = 2, - norm_num_groups: int | None = 32, - cross_attention_dim: int = 1024, - attention_head_dim: int | tuple[int] = 64, - num_attention_heads: int | tuple[int] | None = None, - ): - super().__init__() - - # When we first integrated the UNet into the library, we didn't have `attention_head_dim`. As a consequence - # of that, we used `num_attention_heads` for arguments that actually denote attention head dimension. This - # is why we ignore `num_attention_heads` and calculate it from `attention_head_dims` below. - # This is still an incorrect way of calculating `num_attention_heads` but we need to stick to it - # without running proper deprecation cycles for the {down,mid,up} blocks which are a - # part of the public API. - num_attention_heads = attention_head_dim - - # Check inputs - if len(down_block_types) != len(up_block_types): - raise ValueError( - f"Must provide the same number of `down_block_types` as `up_block_types`. `down_block_types`: {down_block_types}. `up_block_types`: {up_block_types}." - ) - - if len(block_out_channels) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}." - ) - - # input - self.conv_in = nn.Conv2d(in_channels + in_channels, block_out_channels[0], kernel_size=3, padding=1) - - self.transformer_in = TransformerTemporalModel( - num_attention_heads=8, - attention_head_dim=num_attention_heads, - in_channels=block_out_channels[0], - num_layers=1, - norm_num_groups=norm_num_groups, - ) - - # image embedding - self.image_latents_proj_in = nn.Sequential( - nn.Conv2d(4, in_channels * 4, 3, padding=1), - nn.SiLU(), - nn.Conv2d(in_channels * 4, in_channels * 4, 3, stride=1, padding=1), - nn.SiLU(), - nn.Conv2d(in_channels * 4, in_channels, 3, stride=1, padding=1), - ) - self.image_latents_temporal_encoder = I2VGenXLTransformerTemporalEncoder( - dim=in_channels, - num_attention_heads=2, - ff_inner_dim=in_channels * 4, - attention_head_dim=in_channels, - activation_fn="gelu", - ) - self.image_latents_context_embedding = nn.Sequential( - nn.Conv2d(4, in_channels * 8, 3, padding=1), - nn.SiLU(), - nn.AdaptiveAvgPool2d((32, 32)), - nn.Conv2d(in_channels * 8, in_channels * 16, 3, stride=2, padding=1), - nn.SiLU(), - nn.Conv2d(in_channels * 16, cross_attention_dim, 3, stride=2, padding=1), - ) - - # other embeddings -- time, context, fps, etc. - time_embed_dim = block_out_channels[0] * 4 - self.time_proj = Timesteps(block_out_channels[0], True, 0) - timestep_input_dim = block_out_channels[0] - - self.time_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim, act_fn="silu") - self.context_embedding = nn.Sequential( - nn.Linear(cross_attention_dim, time_embed_dim), - nn.SiLU(), - nn.Linear(time_embed_dim, cross_attention_dim * in_channels), - ) - self.fps_embedding = nn.Sequential( - nn.Linear(timestep_input_dim, time_embed_dim), nn.SiLU(), nn.Linear(time_embed_dim, time_embed_dim) - ) - - # blocks - self.down_blocks = nn.ModuleList([]) - self.up_blocks = nn.ModuleList([]) - - if isinstance(num_attention_heads, int): - num_attention_heads = (num_attention_heads,) * len(down_block_types) - - # down - output_channel = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - down_block = get_down_block( - down_block_type, - num_layers=layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - temb_channels=time_embed_dim, - add_downsample=not is_final_block, - resnet_eps=1e-05, - resnet_act_fn="silu", - resnet_groups=norm_num_groups, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads[i], - downsample_padding=1, - dual_cross_attention=False, - ) - self.down_blocks.append(down_block) - - # mid - self.mid_block = UNetMidBlock3DCrossAttn( - in_channels=block_out_channels[-1], - temb_channels=time_embed_dim, - resnet_eps=1e-05, - resnet_act_fn="silu", - output_scale_factor=1, - cross_attention_dim=cross_attention_dim, - num_attention_heads=num_attention_heads[-1], - resnet_groups=norm_num_groups, - dual_cross_attention=False, - ) - - # count how many layers upsample the images - self.num_upsamplers = 0 - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - reversed_num_attention_heads = list(reversed(num_attention_heads)) - - output_channel = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - is_final_block = i == len(block_out_channels) - 1 - - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)] - - # add upsample block for all BUT final layer - if not is_final_block: - add_upsample = True - self.num_upsamplers += 1 - else: - add_upsample = False - - up_block = get_up_block( - up_block_type, - num_layers=layers_per_block + 1, - in_channels=input_channel, - out_channels=output_channel, - prev_output_channel=prev_output_channel, - temb_channels=time_embed_dim, - add_upsample=add_upsample, - resnet_eps=1e-05, - resnet_act_fn="silu", - resnet_groups=norm_num_groups, - cross_attention_dim=cross_attention_dim, - num_attention_heads=reversed_num_attention_heads[i], - dual_cross_attention=False, - resolution_idx=i, - ) - self.up_blocks.append(up_block) - prev_output_channel = output_channel - - # out - self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-05) - self.conv_act = get_activation("silu") - self.conv_out = nn.Conv2d(block_out_channels[0], out_channels, kernel_size=3, padding=1) - - # Copied from diffusers.models.unets.unet_3d_condition.UNet3DConditionModel.enable_forward_chunking - def enable_forward_chunking(self, chunk_size: int | None = None, dim: int = 0) -> None: - """ - Sets the attention processor to use [feed forward - chunking](https://huggingface.co/blog/reformer#2-chunked-feed-forward-layers). - - Parameters: - chunk_size (`int`, *optional*): - The chunk size of the feed-forward layers. If not specified, will run feed-forward layer individually - over each tensor of dim=`dim`. - dim (`int`, *optional*, defaults to `0`): - The dimension over which the feed-forward computation should be chunked. Choose between dim=0 (batch) - or dim=1 (sequence length). - """ - if dim not in [0, 1]: - raise ValueError(f"Make sure to set `dim` to either 0 or 1, not {dim}") - - # By default chunk size is 1 - chunk_size = chunk_size or 1 - - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, chunk_size, dim) - - # Copied from diffusers.models.unets.unet_3d_condition.UNet3DConditionModel.disable_forward_chunking - def disable_forward_chunking(self): - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, None, 0) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.enable_freeu - def enable_freeu(self, s1, s2, b1, b2): - r"""Enables the FreeU mechanism from https://huggingface.co/papers/2309.11497. - - The suffixes after the scaling factors represent the stage blocks where they are being applied. - - Please refer to the [official repository](https://github.com/ChenyangSi/FreeU) for combinations of values that - are known to work well for different pipelines such as Stable Diffusion v1, v2, and Stable Diffusion XL. - - Args: - s1 (`float`): - Scaling factor for stage 1 to attenuate the contributions of the skip features. This is done to - mitigate the "oversmoothing effect" in the enhanced denoising process. - s2 (`float`): - Scaling factor for stage 2 to attenuate the contributions of the skip features. This is done to - mitigate the "oversmoothing effect" in the enhanced denoising process. - b1 (`float`): Scaling factor for stage 1 to amplify the contributions of backbone features. - b2 (`float`): Scaling factor for stage 2 to amplify the contributions of backbone features. - """ - for i, upsample_block in enumerate(self.up_blocks): - setattr(upsample_block, "s1", s1) - setattr(upsample_block, "s2", s2) - setattr(upsample_block, "b1", b1) - setattr(upsample_block, "b2", b2) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.disable_freeu - def disable_freeu(self): - """Disables the FreeU mechanism.""" - freeu_keys = {"s1", "s2", "b1", "b2"} - for i, upsample_block in enumerate(self.up_blocks): - for k in freeu_keys: - if hasattr(upsample_block, k) or getattr(upsample_block, k, None) is not None: - setattr(upsample_block, k, None) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) - - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - fps: torch.Tensor, - image_latents: torch.Tensor, - image_embeddings: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - timestep_cond: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - return_dict: bool = True, - ) -> UNet3DConditionOutput | tuple[torch.Tensor]: - r""" - The [`I2VGenXLUNet`] forward method. - - Args: - sample (`torch.Tensor`): - The noisy input tensor with the following shape `(batch, num_frames, channel, height, width`. - timestep (`torch.Tensor` or `float` or `int`): The number of timesteps to denoise an input. - fps (`torch.Tensor`): Frames per second for the video being generated. Used as a "micro-condition". - image_latents (`torch.Tensor`): Image encodings from the VAE. - image_embeddings (`torch.Tensor`): - Projection embeddings of the conditioning image computed with a vision encoder. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states with shape `(batch, sequence_length, feature_dim)`. - timestep_cond (`torch.Tensor`, *optional*): - Additional conditional embeddings for timestep. If provided, the embeddings will be summed with the - timestep_embedding passed through the `self.time_embedding` layer to obtain the final timestep - embeddings. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unets.unet_3d_condition.UNet3DConditionOutput`] instead of a plain - tuple. - - Returns: - [`~models.unets.unet_3d_condition.UNet3DConditionOutput`] or `tuple`: - If `return_dict` is True, an [`~models.unets.unet_3d_condition.UNet3DConditionOutput`] is returned, - otherwise a `tuple` is returned where the first element is the sample tensor. - """ - batch_size, channels, num_frames, height, width = sample.shape - - # By default samples have to be AT least a multiple of the overall upsampling factor. - # The overall upsampling factor is equal to 2 ** (# num of upsampling layears). - # However, the upsampling interpolation output size can be forced to fit any upsampling size - # on the fly if necessary. - default_overall_up_factor = 2**self.num_upsamplers - - # upsample size should be forwarded when sample is not a multiple of `default_overall_up_factor` - forward_upsample_size = False - upsample_size = None - - if any(s % default_overall_up_factor != 0 for s in sample.shape[-2:]): - logger.info("Forward upsample size to force interpolation output size.") - forward_upsample_size = True - - # 1. time - timesteps = timestep - if not torch.is_tensor(timesteps): - # TODO: this requires sync between CPU and GPU. So try to pass `timesteps` as tensors if you can - # This would be a good case for the `match` statement (Python 3.10+) - dtype = maybe_adjust_dtype_for_device( - torch.float64 if isinstance(timesteps, float) else torch.int64, sample.device - ) - timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) - elif len(timesteps.shape) == 0: - timesteps = timesteps[None].to(sample.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - timesteps = timesteps.expand(sample.shape[0]) - t_emb = self.time_proj(timesteps) - - # timesteps does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=self.dtype) - t_emb = self.time_embedding(t_emb, timestep_cond) - - # 2. FPS - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - fps = fps.expand(fps.shape[0]) - fps_emb = self.fps_embedding(self.time_proj(fps).to(dtype=self.dtype)) - - # 3. time + FPS embeddings. - emb = t_emb + fps_emb - emb = emb.repeat_interleave(num_frames, dim=0, output_size=emb.shape[0] * num_frames) - - # 4. context embeddings. - # The context embeddings consist of both text embeddings from the input prompt - # AND the image embeddings from the input image. For images, both VAE encodings - # and the CLIP image embeddings are incorporated. - # So the final `context_embeddings` becomes the query for cross-attention. - context_emb = sample.new_zeros(batch_size, 0, self.config.cross_attention_dim) - context_emb = torch.cat([context_emb, encoder_hidden_states], dim=1) - - image_latents_for_context_embds = image_latents[:, :, :1, :] - image_latents_context_embs = image_latents_for_context_embds.permute(0, 2, 1, 3, 4).reshape( - image_latents_for_context_embds.shape[0] * image_latents_for_context_embds.shape[2], - image_latents_for_context_embds.shape[1], - image_latents_for_context_embds.shape[3], - image_latents_for_context_embds.shape[4], - ) - image_latents_context_embs = self.image_latents_context_embedding(image_latents_context_embs) - - _batch_size, _channels, _height, _width = image_latents_context_embs.shape - image_latents_context_embs = image_latents_context_embs.permute(0, 2, 3, 1).reshape( - _batch_size, _height * _width, _channels - ) - context_emb = torch.cat([context_emb, image_latents_context_embs], dim=1) - - image_emb = self.context_embedding(image_embeddings) - image_emb = image_emb.view(-1, self.config.in_channels, self.config.cross_attention_dim) - context_emb = torch.cat([context_emb, image_emb], dim=1) - context_emb = context_emb.repeat_interleave(num_frames, dim=0, output_size=context_emb.shape[0] * num_frames) - - image_latents = image_latents.permute(0, 2, 1, 3, 4).reshape( - image_latents.shape[0] * image_latents.shape[2], - image_latents.shape[1], - image_latents.shape[3], - image_latents.shape[4], - ) - image_latents = self.image_latents_proj_in(image_latents) - image_latents = ( - image_latents[None, :] - .reshape(batch_size, num_frames, channels, height, width) - .permute(0, 3, 4, 1, 2) - .reshape(batch_size * height * width, num_frames, channels) - ) - image_latents = self.image_latents_temporal_encoder(image_latents) - image_latents = image_latents.reshape(batch_size, height, width, num_frames, channels).permute(0, 4, 3, 1, 2) - - # 5. pre-process - sample = torch.cat([sample, image_latents], dim=1) - sample = sample.permute(0, 2, 1, 3, 4).reshape((sample.shape[0] * num_frames, -1) + sample.shape[3:]) - sample = self.conv_in(sample) - sample = self.transformer_in( - sample, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - return_dict=False, - )[0] - - # 6. down - down_block_res_samples = (sample,) - for downsample_block in self.down_blocks: - if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention: - sample, res_samples = downsample_block( - hidden_states=sample, - temb=emb, - encoder_hidden_states=context_emb, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample, res_samples = downsample_block(hidden_states=sample, temb=emb, num_frames=num_frames) - - down_block_res_samples += res_samples - - # 7. mid - if self.mid_block is not None: - sample = self.mid_block( - sample, - emb, - encoder_hidden_states=context_emb, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - ) - # 8. up - for i, upsample_block in enumerate(self.up_blocks): - is_final_block = i == len(self.up_blocks) - 1 - - res_samples = down_block_res_samples[-len(upsample_block.resnets) :] - down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)] - - # if we have not reached the final block and need to forward the - # upsample size, we do it here - if not is_final_block and forward_upsample_size: - upsample_size = down_block_res_samples[-1].shape[2:] - - if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention: - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - encoder_hidden_states=context_emb, - upsample_size=upsample_size, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - upsample_size=upsample_size, - num_frames=num_frames, - ) - - # 9. post-process - sample = self.conv_norm_out(sample) - sample = self.conv_act(sample) - - sample = self.conv_out(sample) - - # reshape to (batch, channel, framerate, width, height) - sample = sample[None, :].reshape((-1, num_frames) + sample.shape[1:]).permute(0, 2, 1, 3, 4) - - if not return_dict: - return (sample,) - - return UNet3DConditionOutput(sample=sample) diff --git a/diffusers/models/unets/unet_kandinsky3.py b/diffusers/models/unets/unet_kandinsky3.py deleted file mode 100644 index 790d255101a4ce846dbaeaf4e8af7409afac2782..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_kandinsky3.py +++ /dev/null @@ -1,485 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass - -import torch -from torch import nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...utils import BaseOutput, logging -from ..attention import AttentionMixin -from ..attention_processor import Attention, AttnProcessor -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class Kandinsky3UNetOutput(BaseOutput): - sample: torch.Tensor = None - - -class Kandinsky3EncoderProj(nn.Module): - def __init__(self, encoder_hid_dim, cross_attention_dim): - super().__init__() - self.projection_linear = nn.Linear(encoder_hid_dim, cross_attention_dim, bias=False) - self.projection_norm = nn.LayerNorm(cross_attention_dim) - - def forward(self, x): - x = self.projection_linear(x) - x = self.projection_norm(x) - return x - - -class Kandinsky3UNet(ModelMixin, AttentionMixin, ConfigMixin): - @register_to_config - def __init__( - self, - in_channels: int = 4, - time_embedding_dim: int = 1536, - groups: int = 32, - attention_head_dim: int = 64, - layers_per_block: int | tuple[int] = 3, - block_out_channels: tuple[int, ...] = (384, 768, 1536, 3072), - cross_attention_dim: int | tuple[int] = 4096, - encoder_hid_dim: int = 4096, - ): - super().__init__() - - # TODO(Yiyi): Give better name and put into config for the following 4 parameters - expansion_ratio = 4 - compression_ratio = 2 - add_cross_attention = (False, True, True, True) - add_self_attention = (False, True, True, True) - - out_channels = in_channels - init_channels = block_out_channels[0] // 2 - self.time_proj = Timesteps(init_channels, flip_sin_to_cos=False, downscale_freq_shift=1) - - self.time_embedding = TimestepEmbedding( - init_channels, - time_embedding_dim, - ) - - self.add_time_condition = Kandinsky3AttentionPooling( - time_embedding_dim, cross_attention_dim, attention_head_dim - ) - - self.conv_in = nn.Conv2d(in_channels, init_channels, kernel_size=3, padding=1) - - self.encoder_hid_proj = Kandinsky3EncoderProj(encoder_hid_dim, cross_attention_dim) - - hidden_dims = [init_channels] + list(block_out_channels) - in_out_dims = list(zip(hidden_dims[:-1], hidden_dims[1:])) - text_dims = [cross_attention_dim if is_exist else None for is_exist in add_cross_attention] - num_blocks = len(block_out_channels) * [layers_per_block] - layer_params = [num_blocks, text_dims, add_self_attention] - rev_layer_params = map(reversed, layer_params) - - cat_dims = [] - self.num_levels = len(in_out_dims) - self.down_blocks = nn.ModuleList([]) - for level, ((in_dim, out_dim), res_block_num, text_dim, self_attention) in enumerate( - zip(in_out_dims, *layer_params) - ): - down_sample = level != (self.num_levels - 1) - cat_dims.append(out_dim if level != (self.num_levels - 1) else 0) - self.down_blocks.append( - Kandinsky3DownSampleBlock( - in_dim, - out_dim, - time_embedding_dim, - text_dim, - res_block_num, - groups, - attention_head_dim, - expansion_ratio, - compression_ratio, - down_sample, - self_attention, - ) - ) - - self.up_blocks = nn.ModuleList([]) - for level, ((out_dim, in_dim), res_block_num, text_dim, self_attention) in enumerate( - zip(reversed(in_out_dims), *rev_layer_params) - ): - up_sample = level != 0 - self.up_blocks.append( - Kandinsky3UpSampleBlock( - in_dim, - cat_dims.pop(), - out_dim, - time_embedding_dim, - text_dim, - res_block_num, - groups, - attention_head_dim, - expansion_ratio, - compression_ratio, - up_sample, - self_attention, - ) - ) - - self.conv_norm_out = nn.GroupNorm(groups, init_channels) - self.conv_act_out = nn.SiLU() - self.conv_out = nn.Conv2d(init_channels, out_channels, kernel_size=3, padding=1) - - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - self.set_attn_processor(AttnProcessor()) - - def forward(self, sample, timestep, encoder_hidden_states=None, encoder_attention_mask=None, return_dict=True): - r""" - Args: - sample (`torch.Tensor`): Input sample. - timestep (`torch.Tensor`, `float`, or `int`): - The number of timesteps to denoise an input. - encoder_hidden_states (`torch.Tensor`, *optional*): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - encoder_attention_mask (`torch.Tensor`, *optional*): - Attention mask applied to `encoder_hidden_states`. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unets.unet_2d_condition.UNet2DConditionOutput`] instead of a plain - tuple. - """ - if encoder_attention_mask is not None: - encoder_attention_mask = (1 - encoder_attention_mask.to(sample.dtype)) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - if not torch.is_tensor(timestep): - dtype = torch.float32 if isinstance(timestep, float) else torch.int32 - timestep = torch.tensor([timestep], dtype=dtype, device=sample.device) - elif len(timestep.shape) == 0: - timestep = timestep[None].to(sample.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - timestep = timestep.expand(sample.shape[0]) - time_embed_input = self.time_proj(timestep).to(sample.dtype) - time_embed = self.time_embedding(time_embed_input) - - encoder_hidden_states = self.encoder_hid_proj(encoder_hidden_states) - - if encoder_hidden_states is not None: - time_embed = self.add_time_condition(time_embed, encoder_hidden_states, encoder_attention_mask) - - hidden_states = [] - sample = self.conv_in(sample) - for level, down_sample in enumerate(self.down_blocks): - sample = down_sample(sample, time_embed, encoder_hidden_states, encoder_attention_mask) - if level != self.num_levels - 1: - hidden_states.append(sample) - - for level, up_sample in enumerate(self.up_blocks): - if level != 0: - sample = torch.cat([sample, hidden_states.pop()], dim=1) - sample = up_sample(sample, time_embed, encoder_hidden_states, encoder_attention_mask) - - sample = self.conv_norm_out(sample) - sample = self.conv_act_out(sample) - sample = self.conv_out(sample) - - if not return_dict: - return (sample,) - return Kandinsky3UNetOutput(sample=sample) - - -class Kandinsky3UpSampleBlock(nn.Module): - def __init__( - self, - in_channels, - cat_dim, - out_channels, - time_embed_dim, - context_dim=None, - num_blocks=3, - groups=32, - head_dim=64, - expansion_ratio=4, - compression_ratio=2, - up_sample=True, - self_attention=True, - ): - super().__init__() - up_resolutions = [[None, True if up_sample else None, None, None]] + [[None] * 4] * (num_blocks - 1) - hidden_channels = ( - [(in_channels + cat_dim, in_channels)] - + [(in_channels, in_channels)] * (num_blocks - 2) - + [(in_channels, out_channels)] - ) - attentions = [] - resnets_in = [] - resnets_out = [] - - self.self_attention = self_attention - self.context_dim = context_dim - - if self_attention: - attentions.append( - Kandinsky3AttentionBlock(out_channels, time_embed_dim, None, groups, head_dim, expansion_ratio) - ) - else: - attentions.append(nn.Identity()) - - for (in_channel, out_channel), up_resolution in zip(hidden_channels, up_resolutions): - resnets_in.append( - Kandinsky3ResNetBlock(in_channel, in_channel, time_embed_dim, groups, compression_ratio, up_resolution) - ) - - if context_dim is not None: - attentions.append( - Kandinsky3AttentionBlock( - in_channel, time_embed_dim, context_dim, groups, head_dim, expansion_ratio - ) - ) - else: - attentions.append(nn.Identity()) - - resnets_out.append( - Kandinsky3ResNetBlock(in_channel, out_channel, time_embed_dim, groups, compression_ratio) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets_in = nn.ModuleList(resnets_in) - self.resnets_out = nn.ModuleList(resnets_out) - - def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None): - for attention, resnet_in, resnet_out in zip(self.attentions[1:], self.resnets_in, self.resnets_out): - x = resnet_in(x, time_embed) - if self.context_dim is not None: - x = attention(x, time_embed, context, context_mask, image_mask) - x = resnet_out(x, time_embed) - - if self.self_attention: - x = self.attentions[0](x, time_embed, image_mask=image_mask) - return x - - -class Kandinsky3DownSampleBlock(nn.Module): - def __init__( - self, - in_channels, - out_channels, - time_embed_dim, - context_dim=None, - num_blocks=3, - groups=32, - head_dim=64, - expansion_ratio=4, - compression_ratio=2, - down_sample=True, - self_attention=True, - ): - super().__init__() - attentions = [] - resnets_in = [] - resnets_out = [] - - self.self_attention = self_attention - self.context_dim = context_dim - - if self_attention: - attentions.append( - Kandinsky3AttentionBlock(in_channels, time_embed_dim, None, groups, head_dim, expansion_ratio) - ) - else: - attentions.append(nn.Identity()) - - up_resolutions = [[None] * 4] * (num_blocks - 1) + [[None, None, False if down_sample else None, None]] - hidden_channels = [(in_channels, out_channels)] + [(out_channels, out_channels)] * (num_blocks - 1) - for (in_channel, out_channel), up_resolution in zip(hidden_channels, up_resolutions): - resnets_in.append( - Kandinsky3ResNetBlock(in_channel, out_channel, time_embed_dim, groups, compression_ratio) - ) - - if context_dim is not None: - attentions.append( - Kandinsky3AttentionBlock( - out_channel, time_embed_dim, context_dim, groups, head_dim, expansion_ratio - ) - ) - else: - attentions.append(nn.Identity()) - - resnets_out.append( - Kandinsky3ResNetBlock( - out_channel, out_channel, time_embed_dim, groups, compression_ratio, up_resolution - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets_in = nn.ModuleList(resnets_in) - self.resnets_out = nn.ModuleList(resnets_out) - - def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None): - if self.self_attention: - x = self.attentions[0](x, time_embed, image_mask=image_mask) - - for attention, resnet_in, resnet_out in zip(self.attentions[1:], self.resnets_in, self.resnets_out): - x = resnet_in(x, time_embed) - if self.context_dim is not None: - x = attention(x, time_embed, context, context_mask, image_mask) - x = resnet_out(x, time_embed) - return x - - -class Kandinsky3ConditionalGroupNorm(nn.Module): - def __init__(self, groups, normalized_shape, context_dim): - super().__init__() - self.norm = nn.GroupNorm(groups, normalized_shape, affine=False) - self.context_mlp = nn.Sequential(nn.SiLU(), nn.Linear(context_dim, 2 * normalized_shape)) - self.context_mlp[1].weight.data.zero_() - self.context_mlp[1].bias.data.zero_() - - def forward(self, x, context): - context = self.context_mlp(context) - - for _ in range(len(x.shape[2:])): - context = context.unsqueeze(-1) - - scale, shift = context.chunk(2, dim=1) - x = self.norm(x) * (scale + 1.0) + shift - return x - - -class Kandinsky3Block(nn.Module): - def __init__(self, in_channels, out_channels, time_embed_dim, kernel_size=3, norm_groups=32, up_resolution=None): - super().__init__() - self.group_norm = Kandinsky3ConditionalGroupNorm(norm_groups, in_channels, time_embed_dim) - self.activation = nn.SiLU() - if up_resolution is not None and up_resolution: - self.up_sample = nn.ConvTranspose2d(in_channels, in_channels, kernel_size=2, stride=2) - else: - self.up_sample = nn.Identity() - - padding = int(kernel_size > 1) - self.projection = nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, padding=padding) - - if up_resolution is not None and not up_resolution: - self.down_sample = nn.Conv2d(out_channels, out_channels, kernel_size=2, stride=2) - else: - self.down_sample = nn.Identity() - - def forward(self, x, time_embed): - x = self.group_norm(x, time_embed) - x = self.activation(x) - x = self.up_sample(x) - x = self.projection(x) - x = self.down_sample(x) - return x - - -class Kandinsky3ResNetBlock(nn.Module): - def __init__( - self, in_channels, out_channels, time_embed_dim, norm_groups=32, compression_ratio=2, up_resolutions=4 * [None] - ): - super().__init__() - kernel_sizes = [1, 3, 3, 1] - hidden_channel = max(in_channels, out_channels) // compression_ratio - hidden_channels = ( - [(in_channels, hidden_channel)] + [(hidden_channel, hidden_channel)] * 2 + [(hidden_channel, out_channels)] - ) - self.resnet_blocks = nn.ModuleList( - [ - Kandinsky3Block(in_channel, out_channel, time_embed_dim, kernel_size, norm_groups, up_resolution) - for (in_channel, out_channel), kernel_size, up_resolution in zip( - hidden_channels, kernel_sizes, up_resolutions - ) - ] - ) - self.shortcut_up_sample = ( - nn.ConvTranspose2d(in_channels, in_channels, kernel_size=2, stride=2) - if True in up_resolutions - else nn.Identity() - ) - self.shortcut_projection = ( - nn.Conv2d(in_channels, out_channels, kernel_size=1) if in_channels != out_channels else nn.Identity() - ) - self.shortcut_down_sample = ( - nn.Conv2d(out_channels, out_channels, kernel_size=2, stride=2) - if False in up_resolutions - else nn.Identity() - ) - - def forward(self, x, time_embed): - out = x - for resnet_block in self.resnet_blocks: - out = resnet_block(out, time_embed) - - x = self.shortcut_up_sample(x) - x = self.shortcut_projection(x) - x = self.shortcut_down_sample(x) - x = x + out - return x - - -class Kandinsky3AttentionPooling(nn.Module): - def __init__(self, num_channels, context_dim, head_dim=64): - super().__init__() - self.attention = Attention( - context_dim, - context_dim, - dim_head=head_dim, - out_dim=num_channels, - out_bias=False, - ) - - def forward(self, x, context, context_mask=None): - context_mask = context_mask.to(dtype=context.dtype) - context = self.attention(context.mean(dim=1, keepdim=True), context, context_mask) - return x + context.squeeze(1) - - -class Kandinsky3AttentionBlock(nn.Module): - def __init__(self, num_channels, time_embed_dim, context_dim=None, norm_groups=32, head_dim=64, expansion_ratio=4): - super().__init__() - self.in_norm = Kandinsky3ConditionalGroupNorm(norm_groups, num_channels, time_embed_dim) - self.attention = Attention( - num_channels, - context_dim or num_channels, - dim_head=head_dim, - out_dim=num_channels, - out_bias=False, - ) - - hidden_channels = expansion_ratio * num_channels - self.out_norm = Kandinsky3ConditionalGroupNorm(norm_groups, num_channels, time_embed_dim) - self.feed_forward = nn.Sequential( - nn.Conv2d(num_channels, hidden_channels, kernel_size=1, bias=False), - nn.SiLU(), - nn.Conv2d(hidden_channels, num_channels, kernel_size=1, bias=False), - ) - - def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None): - height, width = x.shape[-2:] - out = self.in_norm(x, time_embed) - out = out.reshape(x.shape[0], -1, height * width).permute(0, 2, 1) - context = context if context is not None else out - if context_mask is not None: - context_mask = context_mask.to(dtype=context.dtype) - - out = self.attention(out, context, context_mask) - out = out.permute(0, 2, 1).unsqueeze(-1).reshape(out.shape[0], -1, height, width) - x = x + out - - out = self.out_norm(x, time_embed) - out = self.feed_forward(out) - x = x + out - return x diff --git a/diffusers/models/unets/unet_motion_model.py b/diffusers/models/unets/unet_motion_model.py deleted file mode 100644 index faa181d9bfd5c85f5b952ffccb1f06c07a6aadc9..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_motion_model.py +++ /dev/null @@ -1,2112 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from dataclasses import dataclass -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ...configuration_utils import ConfigMixin, FrozenDict, register_to_config -from ...loaders import FromOriginalModelMixin, PeftAdapterMixin, UNet2DConditionLoadersMixin -from ...utils import BaseOutput, apply_lora_scale, deprecate, logging -from ...utils.torch_utils import apply_freeu, maybe_adjust_dtype_for_device -from ..attention import AttentionMixin, BasicTransformerBlock -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - Attention, - AttnAddedKVProcessor, - AttnProcessor, - AttnProcessor2_0, - FusedAttnProcessor2_0, - IPAdapterAttnProcessor, - IPAdapterAttnProcessor2_0, -) -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from ..resnet import Downsample2D, ResnetBlock2D, Upsample2D -from ..transformers.dual_transformer_2d import DualTransformer2DModel -from ..transformers.transformer_2d import Transformer2DModel -from .unet_2d_blocks import UNetMidBlock2DCrossAttn -from .unet_2d_condition import UNet2DConditionModel - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class UNetMotionOutput(BaseOutput): - """ - The output of [`UNetMotionOutput`]. - - Args: - sample (`torch.Tensor` of shape `(batch_size, num_channels, num_frames, height, width)`): - The hidden states output conditioned on `encoder_hidden_states` input. Output of last layer of model. - """ - - sample: torch.Tensor - - -class AnimateDiffTransformer3D(nn.Module): - """ - A Transformer model for video-like data. - - Parameters: - num_attention_heads (`int`, *optional*, defaults to 16): The number of heads to use for multi-head attention. - attention_head_dim (`int`, *optional*, defaults to 88): The number of channels in each head. - in_channels (`int`, *optional*): - The number of channels in the input and output (specify if the input is **continuous**). - num_layers (`int`, *optional*, defaults to 1): The number of layers of Transformer blocks to use. - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - cross_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use. - attention_bias (`bool`, *optional*): - Configure if the `TransformerBlock` attention should contain a bias parameter. - sample_size (`int`, *optional*): The width of the latent images (specify if the input is **discrete**). - This is fixed during training since it is used to learn a number of position embeddings. - activation_fn (`str`, *optional*, defaults to `"geglu"`): - Activation function to use in feed-forward. See `diffusers.models.activations.get_activation` for supported - activation functions. - norm_elementwise_affine (`bool`, *optional*): - Configure if the `TransformerBlock` should use learnable elementwise affine parameters for normalization. - double_self_attention (`bool`, *optional*): - Configure if each `TransformerBlock` should contain two self-attention layers. - positional_embeddings: (`str`, *optional*): - The type of positional embeddings to apply to the sequence input before passing use. - num_positional_embeddings: (`int`, *optional*): - The maximum length of the sequence over which to apply positional embeddings. - """ - - def __init__( - self, - num_attention_heads: int = 16, - attention_head_dim: int = 88, - in_channels: int | None = None, - out_channels: int | None = None, - num_layers: int = 1, - dropout: float = 0.0, - norm_num_groups: int = 32, - cross_attention_dim: int | None = None, - attention_bias: bool = False, - sample_size: int | None = None, - activation_fn: str = "geglu", - norm_elementwise_affine: bool = True, - double_self_attention: bool = True, - positional_embeddings: str | None = None, - num_positional_embeddings: int | None = None, - ): - super().__init__() - self.num_attention_heads = num_attention_heads - self.attention_head_dim = attention_head_dim - inner_dim = num_attention_heads * attention_head_dim - - self.in_channels = in_channels - - self.norm = nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True) - self.proj_in = nn.Linear(in_channels, inner_dim) - - # 3. Define transformers blocks - self.transformer_blocks = nn.ModuleList( - [ - BasicTransformerBlock( - inner_dim, - num_attention_heads, - attention_head_dim, - dropout=dropout, - cross_attention_dim=cross_attention_dim, - activation_fn=activation_fn, - attention_bias=attention_bias, - double_self_attention=double_self_attention, - norm_elementwise_affine=norm_elementwise_affine, - positional_embeddings=positional_embeddings, - num_positional_embeddings=num_positional_embeddings, - ) - for _ in range(num_layers) - ] - ) - - self.proj_out = nn.Linear(inner_dim, in_channels) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.LongTensor | None = None, - timestep: torch.LongTensor | None = None, - class_labels: torch.LongTensor | None = None, - num_frames: int = 1, - cross_attention_kwargs: dict[str, Any] | None = None, - ) -> torch.Tensor: - """ - The [`AnimateDiffTransformer3D`] forward method. - - Args: - hidden_states (`torch.LongTensor` of shape `(batch size, num latent pixels)` if discrete, `torch.Tensor` of shape `(batch size, channel, height, width)` if continuous): - Input hidden_states. - encoder_hidden_states ( `torch.LongTensor` of shape `(batch size, encoder_hidden_states dim)`, *optional*): - Conditional embeddings for cross attention layer. If not given, cross-attention defaults to - self-attention. - timestep ( `torch.LongTensor`, *optional*): - Used to indicate denoising step. Optional timestep to be applied as an embedding in `AdaLayerNorm`. - class_labels ( `torch.LongTensor` of shape `(batch size, num classes)`, *optional*): - Used to indicate class labels conditioning. Optional class labels to be applied as an embedding in - `AdaLayerZeroNorm`. - num_frames (`int`, *optional*, defaults to 1): - The number of frames to be processed per batch. This is used to reshape the hidden states. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - - Returns: - torch.Tensor: - The output tensor. - """ - # 1. Input - batch_frames, channel, height, width = hidden_states.shape - batch_size = batch_frames // num_frames - - residual = hidden_states - - hidden_states = hidden_states[None, :].reshape(batch_size, num_frames, channel, height, width) - hidden_states = hidden_states.permute(0, 2, 1, 3, 4) - - hidden_states = self.norm(hidden_states) - hidden_states = hidden_states.permute(0, 3, 4, 2, 1).reshape(batch_size * height * width, num_frames, channel) - - hidden_states = self.proj_in(input=hidden_states) - - # 2. Blocks - for block in self.transformer_blocks: - hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - timestep=timestep, - cross_attention_kwargs=cross_attention_kwargs, - class_labels=class_labels, - ) - - # 3. Output - hidden_states = self.proj_out(input=hidden_states) - hidden_states = ( - hidden_states[None, None, :] - .reshape(batch_size, height, width, num_frames, channel) - .permute(0, 3, 4, 1, 2) - .contiguous() - ) - hidden_states = hidden_states.reshape(batch_frames, channel, height, width) - - output = hidden_states + residual - return output - - -class DownBlockMotion(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_downsample: bool = True, - downsample_padding: int = 1, - temporal_num_attention_heads: int | tuple[int] = 1, - temporal_cross_attention_dim: int | None = None, - temporal_max_seq_length: int = 32, - temporal_transformer_layers_per_block: int | tuple[int] = 1, - temporal_double_self_attention: bool = True, - ): - super().__init__() - resnets = [] - motion_modules = [] - - # support for variable transformer layers per temporal block - if isinstance(temporal_transformer_layers_per_block, int): - temporal_transformer_layers_per_block = (temporal_transformer_layers_per_block,) * num_layers - elif len(temporal_transformer_layers_per_block) != num_layers: - raise ValueError( - f"`temporal_transformer_layers_per_block` must be an integer or a tuple of integers of length {num_layers}" - ) - - # support for variable number of attention head per temporal layers - if isinstance(temporal_num_attention_heads, int): - temporal_num_attention_heads = (temporal_num_attention_heads,) * num_layers - elif len(temporal_num_attention_heads) != num_layers: - raise ValueError( - f"`temporal_num_attention_heads` must be an integer or a tuple of integers of length {num_layers}" - ) - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - motion_modules.append( - AnimateDiffTransformer3D( - num_attention_heads=temporal_num_attention_heads[i], - in_channels=out_channels, - num_layers=temporal_transformer_layers_per_block[i], - norm_num_groups=resnet_groups, - cross_attention_dim=temporal_cross_attention_dim, - attention_bias=False, - activation_fn="geglu", - positional_embeddings="sinusoidal", - num_positional_embeddings=temporal_max_seq_length, - attention_head_dim=out_channels // temporal_num_attention_heads[i], - double_self_attention=temporal_double_self_attention, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.motion_modules = nn.ModuleList(motion_modules) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, - use_conv=True, - out_channels=out_channels, - padding=downsample_padding, - name="op", - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - num_frames: int = 1, - *args, - **kwargs, - ) -> torch.Tensor | tuple[torch.Tensor, ...]: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - output_states = () - - blocks = zip(self.resnets, self.motion_modules) - for resnet, motion_module in blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(input_tensor=hidden_states, temb=temb) - - hidden_states = motion_module(hidden_states, num_frames=num_frames) - - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states=hidden_states) - - output_states = output_states + (hidden_states,) - - return hidden_states, output_states - - -class CrossAttnDownBlockMotion(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - downsample_padding: int = 1, - add_downsample: bool = True, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - attention_type: str = "default", - temporal_cross_attention_dim: int | None = None, - temporal_num_attention_heads: int = 8, - temporal_max_seq_length: int = 32, - temporal_transformer_layers_per_block: int | tuple[int] = 1, - temporal_double_self_attention: bool = True, - ): - super().__init__() - resnets = [] - attentions = [] - motion_modules = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - - # support for variable transformer layers per block - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = (transformer_layers_per_block,) * num_layers - elif len(transformer_layers_per_block) != num_layers: - raise ValueError( - f"transformer_layers_per_block must be an integer or a list of integers of length {num_layers}" - ) - - # support for variable transformer layers per temporal block - if isinstance(temporal_transformer_layers_per_block, int): - temporal_transformer_layers_per_block = (temporal_transformer_layers_per_block,) * num_layers - elif len(temporal_transformer_layers_per_block) != num_layers: - raise ValueError( - f"temporal_transformer_layers_per_block must be an integer or a list of integers of length {num_layers}" - ) - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - if not dual_cross_attention: - attentions.append( - Transformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - attention_type=attention_type, - ) - ) - else: - attentions.append( - DualTransformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - - motion_modules.append( - AnimateDiffTransformer3D( - num_attention_heads=temporal_num_attention_heads, - in_channels=out_channels, - num_layers=temporal_transformer_layers_per_block[i], - norm_num_groups=resnet_groups, - cross_attention_dim=temporal_cross_attention_dim, - attention_bias=False, - activation_fn="geglu", - positional_embeddings="sinusoidal", - num_positional_embeddings=temporal_max_seq_length, - attention_head_dim=out_channels // temporal_num_attention_heads, - double_self_attention=temporal_double_self_attention, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - self.motion_modules = nn.ModuleList(motion_modules) - - if add_downsample: - self.downsamplers = nn.ModuleList( - [ - Downsample2D( - out_channels, - use_conv=True, - out_channels=out_channels, - padding=downsample_padding, - name="op", - ) - ] - ) - else: - self.downsamplers = None - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - num_frames: int = 1, - encoder_attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - additional_residuals: torch.Tensor | None = None, - ): - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - output_states = () - - blocks = list(zip(self.resnets, self.attentions, self.motion_modules)) - for i, (resnet, attn, motion_module) in enumerate(blocks): - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(input_tensor=hidden_states, temb=temb) - - hidden_states = attn( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - hidden_states = motion_module(hidden_states, num_frames=num_frames) - - # apply additional residuals to the output of the last pair of resnet and attention blocks - if i == len(blocks) - 1 and additional_residuals is not None: - hidden_states = hidden_states + additional_residuals - - output_states = output_states + (hidden_states,) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states=hidden_states) - - output_states = output_states + (hidden_states,) - - return hidden_states, output_states - - -class CrossAttnUpBlockMotion(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - prev_output_channel: int, - temb_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - cross_attention_dim: int = 1280, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - only_cross_attention: bool = False, - upcast_attention: bool = False, - attention_type: str = "default", - temporal_cross_attention_dim: int | None = None, - temporal_num_attention_heads: int = 8, - temporal_max_seq_length: int = 32, - temporal_transformer_layers_per_block: int | tuple[int] = 1, - ): - super().__init__() - resnets = [] - attentions = [] - motion_modules = [] - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - - # support for variable transformer layers per block - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = (transformer_layers_per_block,) * num_layers - elif len(transformer_layers_per_block) != num_layers: - raise ValueError( - f"transformer_layers_per_block must be an integer or a list of integers of length {num_layers}, got {len(transformer_layers_per_block)}" - ) - - # support for variable transformer layers per temporal block - if isinstance(temporal_transformer_layers_per_block, int): - temporal_transformer_layers_per_block = (temporal_transformer_layers_per_block,) * num_layers - elif len(temporal_transformer_layers_per_block) != num_layers: - raise ValueError( - f"temporal_transformer_layers_per_block must be an integer or a list of integers of length {num_layers}, got {len(temporal_transformer_layers_per_block)}" - ) - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - if not dual_cross_attention: - attentions.append( - Transformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - only_cross_attention=only_cross_attention, - upcast_attention=upcast_attention, - attention_type=attention_type, - ) - ) - else: - attentions.append( - DualTransformer2DModel( - num_attention_heads, - out_channels // num_attention_heads, - in_channels=out_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - motion_modules.append( - AnimateDiffTransformer3D( - num_attention_heads=temporal_num_attention_heads, - in_channels=out_channels, - num_layers=temporal_transformer_layers_per_block[i], - norm_num_groups=resnet_groups, - cross_attention_dim=temporal_cross_attention_dim, - attention_bias=False, - activation_fn="geglu", - positional_embeddings="sinusoidal", - num_positional_embeddings=temporal_max_seq_length, - attention_head_dim=out_channels // temporal_num_attention_heads, - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - self.motion_modules = nn.ModuleList(motion_modules) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - upsample_size: int | None = None, - attention_mask: torch.Tensor | None = None, - encoder_attention_mask: torch.Tensor | None = None, - num_frames: int = 1, - ) -> torch.Tensor: - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - is_freeu_enabled = ( - getattr(self, "s1", None) - and getattr(self, "s2", None) - and getattr(self, "b1", None) - and getattr(self, "b2", None) - ) - - blocks = zip(self.resnets, self.attentions, self.motion_modules) - for resnet, attn, motion_module in blocks: - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - - # FreeU: Only operate on the first two stages - if is_freeu_enabled: - hidden_states, res_hidden_states = apply_freeu( - self.resolution_idx, - hidden_states, - res_hidden_states, - s1=self.s1, - s2=self.s2, - b1=self.b1, - b2=self.b2, - ) - - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(input_tensor=hidden_states, temb=temb) - - hidden_states = attn( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - hidden_states = motion_module(hidden_states, num_frames=num_frames) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states=hidden_states, output_size=upsample_size) - - return hidden_states - - -class UpBlockMotion(nn.Module): - def __init__( - self, - in_channels: int, - prev_output_channel: int, - out_channels: int, - temb_channels: int, - resolution_idx: int | None = None, - dropout: float = 0.0, - num_layers: int = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - output_scale_factor: float = 1.0, - add_upsample: bool = True, - temporal_cross_attention_dim: int | None = None, - temporal_num_attention_heads: int = 8, - temporal_max_seq_length: int = 32, - temporal_transformer_layers_per_block: int | tuple[int] = 1, - ): - super().__init__() - resnets = [] - motion_modules = [] - - # support for variable transformer layers per temporal block - if isinstance(temporal_transformer_layers_per_block, int): - temporal_transformer_layers_per_block = (temporal_transformer_layers_per_block,) * num_layers - elif len(temporal_transformer_layers_per_block) != num_layers: - raise ValueError( - f"temporal_transformer_layers_per_block must be an integer or a list of integers of length {num_layers}" - ) - - for i in range(num_layers): - res_skip_channels = in_channels if (i == num_layers - 1) else out_channels - resnet_in_channels = prev_output_channel if i == 0 else out_channels - - resnets.append( - ResnetBlock2D( - in_channels=resnet_in_channels + res_skip_channels, - out_channels=out_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - - motion_modules.append( - AnimateDiffTransformer3D( - num_attention_heads=temporal_num_attention_heads, - in_channels=out_channels, - num_layers=temporal_transformer_layers_per_block[i], - norm_num_groups=resnet_groups, - cross_attention_dim=temporal_cross_attention_dim, - attention_bias=False, - activation_fn="geglu", - positional_embeddings="sinusoidal", - num_positional_embeddings=temporal_max_seq_length, - attention_head_dim=out_channels // temporal_num_attention_heads, - ) - ) - - self.resnets = nn.ModuleList(resnets) - self.motion_modules = nn.ModuleList(motion_modules) - - if add_upsample: - self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) - else: - self.upsamplers = None - - self.gradient_checkpointing = False - self.resolution_idx = resolution_idx - - def forward( - self, - hidden_states: torch.Tensor, - res_hidden_states_tuple: tuple[torch.Tensor, ...], - temb: torch.Tensor | None = None, - upsample_size=None, - num_frames: int = 1, - *args, - **kwargs, - ) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - is_freeu_enabled = ( - getattr(self, "s1", None) - and getattr(self, "s2", None) - and getattr(self, "b1", None) - and getattr(self, "b2", None) - ) - - blocks = zip(self.resnets, self.motion_modules) - - for resnet, motion_module in blocks: - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - - # FreeU: Only operate on the first two stages - if is_freeu_enabled: - hidden_states, res_hidden_states = apply_freeu( - self.resolution_idx, - hidden_states, - res_hidden_states, - s1=self.s1, - s2=self.s2, - b1=self.b1, - b2=self.b2, - ) - - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = resnet(input_tensor=hidden_states, temb=temb) - - hidden_states = motion_module(hidden_states, num_frames=num_frames) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states=hidden_states, output_size=upsample_size) - - return hidden_states - - -class UNetMidBlockCrossAttnMotion(nn.Module): - def __init__( - self, - in_channels: int, - temb_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - transformer_layers_per_block: int | tuple[int] = 1, - resnet_eps: float = 1e-6, - resnet_time_scale_shift: str = "default", - resnet_act_fn: str = "swish", - resnet_groups: int = 32, - resnet_pre_norm: bool = True, - num_attention_heads: int = 1, - output_scale_factor: float = 1.0, - cross_attention_dim: int = 1280, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - upcast_attention: bool = False, - attention_type: str = "default", - temporal_num_attention_heads: int = 1, - temporal_cross_attention_dim: int | None = None, - temporal_max_seq_length: int = 32, - temporal_transformer_layers_per_block: int | tuple[int] = 1, - ): - super().__init__() - - self.has_cross_attention = True - self.num_attention_heads = num_attention_heads - resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32) - - # support for variable transformer layers per block - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = (transformer_layers_per_block,) * num_layers - elif len(transformer_layers_per_block) != num_layers: - raise ValueError( - f"`transformer_layers_per_block` should be an integer or a list of integers of length {num_layers}." - ) - - # support for variable transformer layers per temporal block - if isinstance(temporal_transformer_layers_per_block, int): - temporal_transformer_layers_per_block = (temporal_transformer_layers_per_block,) * num_layers - elif len(temporal_transformer_layers_per_block) != num_layers: - raise ValueError( - f"`temporal_transformer_layers_per_block` should be an integer or a list of integers of length {num_layers}." - ) - - # there is always at least one resnet - resnets = [ - ResnetBlock2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ] - attentions = [] - motion_modules = [] - - for i in range(num_layers): - if not dual_cross_attention: - attentions.append( - Transformer2DModel( - num_attention_heads, - in_channels // num_attention_heads, - in_channels=in_channels, - num_layers=transformer_layers_per_block[i], - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - use_linear_projection=use_linear_projection, - upcast_attention=upcast_attention, - attention_type=attention_type, - ) - ) - else: - attentions.append( - DualTransformer2DModel( - num_attention_heads, - in_channels // num_attention_heads, - in_channels=in_channels, - num_layers=1, - cross_attention_dim=cross_attention_dim, - norm_num_groups=resnet_groups, - ) - ) - resnets.append( - ResnetBlock2D( - in_channels=in_channels, - out_channels=in_channels, - temb_channels=temb_channels, - eps=resnet_eps, - groups=resnet_groups, - dropout=dropout, - time_embedding_norm=resnet_time_scale_shift, - non_linearity=resnet_act_fn, - output_scale_factor=output_scale_factor, - pre_norm=resnet_pre_norm, - ) - ) - motion_modules.append( - AnimateDiffTransformer3D( - num_attention_heads=temporal_num_attention_heads, - attention_head_dim=in_channels // temporal_num_attention_heads, - in_channels=in_channels, - num_layers=temporal_transformer_layers_per_block[i], - norm_num_groups=resnet_groups, - cross_attention_dim=temporal_cross_attention_dim, - attention_bias=False, - positional_embeddings="sinusoidal", - num_positional_embeddings=temporal_max_seq_length, - activation_fn="geglu", - ) - ) - - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - self.motion_modules = nn.ModuleList(motion_modules) - - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.Tensor, - temb: torch.Tensor | None = None, - encoder_hidden_states: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - encoder_attention_mask: torch.Tensor | None = None, - num_frames: int = 1, - ) -> torch.Tensor: - if cross_attention_kwargs is not None: - if cross_attention_kwargs.get("scale", None) is not None: - logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - - hidden_states = self.resnets[0](input_tensor=hidden_states, temb=temb) - - blocks = zip(self.attentions, self.resnets[1:], self.motion_modules) - for attn, resnet, motion_module in blocks: - hidden_states = attn( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] - - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - motion_module, hidden_states, None, None, None, num_frames, None - ) - hidden_states = self._gradient_checkpointing_func(resnet, hidden_states, temb) - else: - hidden_states = motion_module(hidden_states, None, None, None, num_frames, None) - hidden_states = resnet(input_tensor=hidden_states, temb=temb) - - return hidden_states - - -class MotionModules(nn.Module): - def __init__( - self, - in_channels: int, - layers_per_block: int = 2, - transformer_layers_per_block: int | tuple[int] = 8, - num_attention_heads: int | tuple[int] = 8, - attention_bias: bool = False, - cross_attention_dim: int | None = None, - activation_fn: str = "geglu", - norm_num_groups: int = 32, - max_seq_length: int = 32, - ): - super().__init__() - self.motion_modules = nn.ModuleList([]) - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = (transformer_layers_per_block,) * layers_per_block - elif len(transformer_layers_per_block) != layers_per_block: - raise ValueError( - f"The number of transformer layers per block must match the number of layers per block, " - f"got {layers_per_block} and {len(transformer_layers_per_block)}" - ) - - for i in range(layers_per_block): - self.motion_modules.append( - AnimateDiffTransformer3D( - in_channels=in_channels, - num_layers=transformer_layers_per_block[i], - norm_num_groups=norm_num_groups, - cross_attention_dim=cross_attention_dim, - activation_fn=activation_fn, - attention_bias=attention_bias, - num_attention_heads=num_attention_heads, - attention_head_dim=in_channels // num_attention_heads, - positional_embeddings="sinusoidal", - num_positional_embeddings=max_seq_length, - ) - ) - - -class MotionAdapter(ModelMixin, ConfigMixin, FromOriginalModelMixin): - @register_to_config - def __init__( - self, - block_out_channels: tuple[int, ...] = (320, 640, 1280, 1280), - motion_layers_per_block: int | tuple[int] = 2, - motion_transformer_layers_per_block: int | tuple[int] | tuple[tuple[int]] = 1, - motion_mid_block_layers_per_block: int = 1, - motion_transformer_layers_per_mid_block: int | tuple[int] = 1, - motion_num_attention_heads: int | tuple[int] = 8, - motion_norm_num_groups: int = 32, - motion_max_seq_length: int = 32, - use_motion_mid_block: bool = True, - conv_in_channels: int | None = None, - ): - """Container to store AnimateDiff Motion Modules - - Args: - block_out_channels (`tuple[int]`, *optional*, defaults to `(320, 640, 1280, 1280)`): - The tuple of output channels for each UNet block. - motion_layers_per_block (`int` or `tuple[int]`, *optional*, defaults to 2): - The number of motion layers per UNet block. - motion_transformer_layers_per_block (`int`, `tuple[int]`, or `tuple[tuple[int]]`, *optional*, defaults to 1): - The number of transformer layers to use in each motion layer in each block. - motion_mid_block_layers_per_block (`int`, *optional*, defaults to 1): - The number of motion layers in the middle UNet block. - motion_transformer_layers_per_mid_block (`int` or `tuple[int]`, *optional*, defaults to 1): - The number of transformer layers to use in each motion layer in the middle block. - motion_num_attention_heads (`int` or `tuple[int]`, *optional*, defaults to 8): - The number of heads to use in each attention layer of the motion module. - motion_norm_num_groups (`int`, *optional*, defaults to 32): - The number of groups to use in each group normalization layer of the motion module. - motion_max_seq_length (`int`, *optional*, defaults to 32): - The maximum sequence length to use in the motion module. - use_motion_mid_block (`bool`, *optional*, defaults to True): - Whether to use a motion module in the middle of the UNet. - """ - - super().__init__() - down_blocks = [] - up_blocks = [] - - if isinstance(motion_layers_per_block, int): - motion_layers_per_block = (motion_layers_per_block,) * len(block_out_channels) - elif len(motion_layers_per_block) != len(block_out_channels): - raise ValueError( - f"The number of motion layers per block must match the number of blocks, " - f"got {len(block_out_channels)} and {len(motion_layers_per_block)}" - ) - - if isinstance(motion_transformer_layers_per_block, int): - motion_transformer_layers_per_block = (motion_transformer_layers_per_block,) * len(block_out_channels) - - if isinstance(motion_transformer_layers_per_mid_block, int): - motion_transformer_layers_per_mid_block = ( - motion_transformer_layers_per_mid_block, - ) * motion_mid_block_layers_per_block - elif len(motion_transformer_layers_per_mid_block) != motion_mid_block_layers_per_block: - raise ValueError( - f"The number of layers per mid block ({motion_mid_block_layers_per_block}) " - f"must match the length of motion_transformer_layers_per_mid_block ({len(motion_transformer_layers_per_mid_block)})" - ) - - if isinstance(motion_num_attention_heads, int): - motion_num_attention_heads = (motion_num_attention_heads,) * len(block_out_channels) - elif len(motion_num_attention_heads) != len(block_out_channels): - raise ValueError( - f"The length of the attention head number tuple in the motion module must match the " - f"number of block, got {len(motion_num_attention_heads)} and {len(block_out_channels)}" - ) - - if conv_in_channels: - # input - self.conv_in = nn.Conv2d(conv_in_channels, block_out_channels[0], kernel_size=3, padding=1) - else: - self.conv_in = None - - for i, channel in enumerate(block_out_channels): - output_channel = block_out_channels[i] - down_blocks.append( - MotionModules( - in_channels=output_channel, - norm_num_groups=motion_norm_num_groups, - cross_attention_dim=None, - activation_fn="geglu", - attention_bias=False, - num_attention_heads=motion_num_attention_heads[i], - max_seq_length=motion_max_seq_length, - layers_per_block=motion_layers_per_block[i], - transformer_layers_per_block=motion_transformer_layers_per_block[i], - ) - ) - - if use_motion_mid_block: - self.mid_block = MotionModules( - in_channels=block_out_channels[-1], - norm_num_groups=motion_norm_num_groups, - cross_attention_dim=None, - activation_fn="geglu", - attention_bias=False, - num_attention_heads=motion_num_attention_heads[-1], - max_seq_length=motion_max_seq_length, - layers_per_block=motion_mid_block_layers_per_block, - transformer_layers_per_block=motion_transformer_layers_per_mid_block, - ) - else: - self.mid_block = None - - reversed_block_out_channels = list(reversed(block_out_channels)) - output_channel = reversed_block_out_channels[0] - - reversed_motion_layers_per_block = list(reversed(motion_layers_per_block)) - reversed_motion_transformer_layers_per_block = list(reversed(motion_transformer_layers_per_block)) - reversed_motion_num_attention_heads = list(reversed(motion_num_attention_heads)) - for i, channel in enumerate(reversed_block_out_channels): - output_channel = reversed_block_out_channels[i] - up_blocks.append( - MotionModules( - in_channels=output_channel, - norm_num_groups=motion_norm_num_groups, - cross_attention_dim=None, - activation_fn="geglu", - attention_bias=False, - num_attention_heads=reversed_motion_num_attention_heads[i], - max_seq_length=motion_max_seq_length, - layers_per_block=reversed_motion_layers_per_block[i] + 1, - transformer_layers_per_block=reversed_motion_transformer_layers_per_block[i], - ) - ) - - self.down_blocks = nn.ModuleList(down_blocks) - self.up_blocks = nn.ModuleList(up_blocks) - - def forward(self, sample): - r""" - Args: - sample (`torch.Tensor`): Input sample. - """ - pass - - -class UNetMotionModel(ModelMixin, AttentionMixin, ConfigMixin, UNet2DConditionLoadersMixin, PeftAdapterMixin): - r""" - A modified conditional 2D UNet model that takes a noisy sample, conditional state, and a timestep and returns a - sample shaped output. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - """ - - _supports_gradient_checkpointing = True - _skip_layerwise_casting_patterns = ["norm"] - - @register_to_config - def __init__( - self, - sample_size: int | None = None, - in_channels: int = 4, - out_channels: int = 4, - down_block_types: tuple[str, ...] = ( - "CrossAttnDownBlockMotion", - "CrossAttnDownBlockMotion", - "CrossAttnDownBlockMotion", - "DownBlockMotion", - ), - up_block_types: tuple[str, ...] = ( - "UpBlockMotion", - "CrossAttnUpBlockMotion", - "CrossAttnUpBlockMotion", - "CrossAttnUpBlockMotion", - ), - block_out_channels: tuple[int, ...] = (320, 640, 1280, 1280), - layers_per_block: int | tuple[int] = 2, - downsample_padding: int = 1, - mid_block_scale_factor: float = 1, - act_fn: str = "silu", - norm_num_groups: int = 32, - norm_eps: float = 1e-5, - cross_attention_dim: int = 1280, - transformer_layers_per_block: int | tuple[int] | tuple[tuple] = 1, - reverse_transformer_layers_per_block: int | tuple[int] | tuple[tuple] | None = None, - temporal_transformer_layers_per_block: int | tuple[int] | tuple[tuple] = 1, - reverse_temporal_transformer_layers_per_block: int | tuple[int] | tuple[tuple] | None = None, - transformer_layers_per_mid_block: int | tuple[int] | None = None, - temporal_transformer_layers_per_mid_block: int | tuple[int] | None = 1, - use_linear_projection: bool = False, - num_attention_heads: int | tuple[int, ...] = 8, - motion_max_seq_length: int = 32, - motion_num_attention_heads: int | tuple[int, ...] = 8, - reverse_motion_num_attention_heads: int | tuple[int, ...] | tuple[tuple[int, ...], ...] | None = None, - use_motion_mid_block: bool = True, - mid_block_layers: int = 1, - encoder_hid_dim: int | None = None, - encoder_hid_dim_type: str | None = None, - addition_embed_type: str | None = None, - addition_time_embed_dim: int | None = None, - projection_class_embeddings_input_dim: int | None = None, - time_cond_proj_dim: int | None = None, - ): - super().__init__() - - self.sample_size = sample_size - - # Check inputs - if len(down_block_types) != len(up_block_types): - raise ValueError( - f"Must provide the same number of `down_block_types` as `up_block_types`. `down_block_types`: {down_block_types}. `up_block_types`: {up_block_types}." - ) - - if len(block_out_channels) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}." - ) - - if isinstance(cross_attention_dim, list) and len(cross_attention_dim) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `cross_attention_dim` as `down_block_types`. `cross_attention_dim`: {cross_attention_dim}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(layers_per_block, int) and len(layers_per_block) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `layers_per_block` as `down_block_types`. `layers_per_block`: {layers_per_block}. `down_block_types`: {down_block_types}." - ) - - if isinstance(transformer_layers_per_block, list) and reverse_transformer_layers_per_block is None: - for layer_number_per_block in transformer_layers_per_block: - if isinstance(layer_number_per_block, list): - raise ValueError("Must provide 'reverse_transformer_layers_per_block` if using asymmetrical UNet.") - - if ( - isinstance(temporal_transformer_layers_per_block, list) - and reverse_temporal_transformer_layers_per_block is None - ): - for layer_number_per_block in temporal_transformer_layers_per_block: - if isinstance(layer_number_per_block, list): - raise ValueError( - "Must provide 'reverse_temporal_transformer_layers_per_block` if using asymmetrical motion module in UNet." - ) - - # input - conv_in_kernel = 3 - conv_out_kernel = 3 - conv_in_padding = (conv_in_kernel - 1) // 2 - self.conv_in = nn.Conv2d( - in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding - ) - - # time - time_embed_dim = block_out_channels[0] * 4 - self.time_proj = Timesteps(block_out_channels[0], True, 0) - timestep_input_dim = block_out_channels[0] - - self.time_embedding = TimestepEmbedding( - timestep_input_dim, time_embed_dim, act_fn=act_fn, cond_proj_dim=time_cond_proj_dim - ) - - if encoder_hid_dim_type is None: - self.encoder_hid_proj = None - - if addition_embed_type == "text_time": - self.add_time_proj = Timesteps(addition_time_embed_dim, True, 0) - self.add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) - - # class embedding - self.down_blocks = nn.ModuleList([]) - self.up_blocks = nn.ModuleList([]) - - if isinstance(num_attention_heads, int): - num_attention_heads = (num_attention_heads,) * len(down_block_types) - - if isinstance(cross_attention_dim, int): - cross_attention_dim = (cross_attention_dim,) * len(down_block_types) - - if isinstance(layers_per_block, int): - layers_per_block = [layers_per_block] * len(down_block_types) - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types) - - if isinstance(reverse_transformer_layers_per_block, int): - reverse_transformer_layers_per_block = [reverse_transformer_layers_per_block] * len(down_block_types) - - if isinstance(temporal_transformer_layers_per_block, int): - temporal_transformer_layers_per_block = [temporal_transformer_layers_per_block] * len(down_block_types) - - if isinstance(reverse_temporal_transformer_layers_per_block, int): - reverse_temporal_transformer_layers_per_block = [reverse_temporal_transformer_layers_per_block] * len( - down_block_types - ) - - if isinstance(motion_num_attention_heads, int): - motion_num_attention_heads = (motion_num_attention_heads,) * len(down_block_types) - - # down - output_channel = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - if down_block_type == "CrossAttnDownBlockMotion": - down_block = CrossAttnDownBlockMotion( - in_channels=input_channel, - out_channels=output_channel, - temb_channels=time_embed_dim, - num_layers=layers_per_block[i], - transformer_layers_per_block=transformer_layers_per_block[i], - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - num_attention_heads=num_attention_heads[i], - cross_attention_dim=cross_attention_dim[i], - downsample_padding=downsample_padding, - add_downsample=not is_final_block, - use_linear_projection=use_linear_projection, - temporal_num_attention_heads=motion_num_attention_heads[i], - temporal_max_seq_length=motion_max_seq_length, - temporal_transformer_layers_per_block=temporal_transformer_layers_per_block[i], - ) - elif down_block_type == "DownBlockMotion": - down_block = DownBlockMotion( - in_channels=input_channel, - out_channels=output_channel, - temb_channels=time_embed_dim, - num_layers=layers_per_block[i], - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - add_downsample=not is_final_block, - downsample_padding=downsample_padding, - temporal_num_attention_heads=motion_num_attention_heads[i], - temporal_max_seq_length=motion_max_seq_length, - temporal_transformer_layers_per_block=temporal_transformer_layers_per_block[i], - ) - else: - raise ValueError( - "Invalid `down_block_type` encountered. Must be one of `CrossAttnDownBlockMotion` or `DownBlockMotion`" - ) - - self.down_blocks.append(down_block) - - # mid - if transformer_layers_per_mid_block is None: - transformer_layers_per_mid_block = ( - transformer_layers_per_block[-1] if isinstance(transformer_layers_per_block[-1], int) else 1 - ) - - if use_motion_mid_block: - self.mid_block = UNetMidBlockCrossAttnMotion( - in_channels=block_out_channels[-1], - temb_channels=time_embed_dim, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - output_scale_factor=mid_block_scale_factor, - cross_attention_dim=cross_attention_dim[-1], - num_attention_heads=num_attention_heads[-1], - resnet_groups=norm_num_groups, - dual_cross_attention=False, - use_linear_projection=use_linear_projection, - num_layers=mid_block_layers, - temporal_num_attention_heads=motion_num_attention_heads[-1], - temporal_max_seq_length=motion_max_seq_length, - transformer_layers_per_block=transformer_layers_per_mid_block, - temporal_transformer_layers_per_block=temporal_transformer_layers_per_mid_block, - ) - - else: - self.mid_block = UNetMidBlock2DCrossAttn( - in_channels=block_out_channels[-1], - temb_channels=time_embed_dim, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - output_scale_factor=mid_block_scale_factor, - cross_attention_dim=cross_attention_dim[-1], - num_attention_heads=num_attention_heads[-1], - resnet_groups=norm_num_groups, - dual_cross_attention=False, - use_linear_projection=use_linear_projection, - num_layers=mid_block_layers, - transformer_layers_per_block=transformer_layers_per_mid_block, - ) - - # count how many layers upsample the images - self.num_upsamplers = 0 - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - reversed_num_attention_heads = list(reversed(num_attention_heads)) - reversed_layers_per_block = list(reversed(layers_per_block)) - reversed_cross_attention_dim = list(reversed(cross_attention_dim)) - reversed_motion_num_attention_heads = list(reversed(motion_num_attention_heads)) - - if reverse_transformer_layers_per_block is None: - reverse_transformer_layers_per_block = list(reversed(transformer_layers_per_block)) - - if reverse_temporal_transformer_layers_per_block is None: - reverse_temporal_transformer_layers_per_block = list(reversed(temporal_transformer_layers_per_block)) - - output_channel = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - is_final_block = i == len(block_out_channels) - 1 - - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)] - - # add upsample block for all BUT final layer - if not is_final_block: - add_upsample = True - self.num_upsamplers += 1 - else: - add_upsample = False - - if up_block_type == "CrossAttnUpBlockMotion": - up_block = CrossAttnUpBlockMotion( - in_channels=input_channel, - out_channels=output_channel, - prev_output_channel=prev_output_channel, - temb_channels=time_embed_dim, - resolution_idx=i, - num_layers=reversed_layers_per_block[i] + 1, - transformer_layers_per_block=reverse_transformer_layers_per_block[i], - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - num_attention_heads=reversed_num_attention_heads[i], - cross_attention_dim=reversed_cross_attention_dim[i], - add_upsample=add_upsample, - use_linear_projection=use_linear_projection, - temporal_num_attention_heads=reversed_motion_num_attention_heads[i], - temporal_max_seq_length=motion_max_seq_length, - temporal_transformer_layers_per_block=reverse_temporal_transformer_layers_per_block[i], - ) - elif up_block_type == "UpBlockMotion": - up_block = UpBlockMotion( - in_channels=input_channel, - prev_output_channel=prev_output_channel, - out_channels=output_channel, - temb_channels=time_embed_dim, - resolution_idx=i, - num_layers=reversed_layers_per_block[i] + 1, - resnet_eps=norm_eps, - resnet_act_fn=act_fn, - resnet_groups=norm_num_groups, - add_upsample=add_upsample, - temporal_num_attention_heads=reversed_motion_num_attention_heads[i], - temporal_max_seq_length=motion_max_seq_length, - temporal_transformer_layers_per_block=reverse_temporal_transformer_layers_per_block[i], - ) - else: - raise ValueError( - "Invalid `up_block_type` encountered. Must be one of `CrossAttnUpBlockMotion` or `UpBlockMotion`" - ) - - self.up_blocks.append(up_block) - prev_output_channel = output_channel - - # out - if norm_num_groups is not None: - self.conv_norm_out = nn.GroupNorm( - num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=norm_eps - ) - self.conv_act = nn.SiLU() - else: - self.conv_norm_out = None - self.conv_act = None - - conv_out_padding = (conv_out_kernel - 1) // 2 - self.conv_out = nn.Conv2d( - block_out_channels[0], out_channels, kernel_size=conv_out_kernel, padding=conv_out_padding - ) - - @classmethod - def from_unet2d( - cls, - unet: UNet2DConditionModel, - motion_adapter: MotionAdapter | None = None, - load_weights: bool = True, - ): - has_motion_adapter = motion_adapter is not None - - if has_motion_adapter: - motion_adapter.to(device=unet.device) - - # check compatibility of number of blocks - if len(unet.config["down_block_types"]) != len(motion_adapter.config["block_out_channels"]): - raise ValueError("Incompatible Motion Adapter, got different number of blocks") - - # check layers compatibility for each block - if isinstance(unet.config["layers_per_block"], int): - expanded_layers_per_block = [unet.config["layers_per_block"]] * len(unet.config["down_block_types"]) - else: - expanded_layers_per_block = list(unet.config["layers_per_block"]) - if isinstance(motion_adapter.config["motion_layers_per_block"], int): - expanded_adapter_layers_per_block = [motion_adapter.config["motion_layers_per_block"]] * len( - motion_adapter.config["block_out_channels"] - ) - else: - expanded_adapter_layers_per_block = list(motion_adapter.config["motion_layers_per_block"]) - if expanded_layers_per_block != expanded_adapter_layers_per_block: - raise ValueError("Incompatible Motion Adapter, got different number of layers per block") - - # based on https://github.com/guoyww/AnimateDiff/blob/895f3220c06318ea0760131ec70408b466c49333/animatediff/models/unet.py#L459 - config = dict(unet.config) - config["_class_name"] = cls.__name__ - - down_blocks = [] - for down_blocks_type in config["down_block_types"]: - if "CrossAttn" in down_blocks_type: - down_blocks.append("CrossAttnDownBlockMotion") - else: - down_blocks.append("DownBlockMotion") - config["down_block_types"] = down_blocks - - up_blocks = [] - for down_blocks_type in config["up_block_types"]: - if "CrossAttn" in down_blocks_type: - up_blocks.append("CrossAttnUpBlockMotion") - else: - up_blocks.append("UpBlockMotion") - config["up_block_types"] = up_blocks - - if has_motion_adapter: - config["motion_num_attention_heads"] = motion_adapter.config["motion_num_attention_heads"] - config["motion_max_seq_length"] = motion_adapter.config["motion_max_seq_length"] - config["use_motion_mid_block"] = motion_adapter.config["use_motion_mid_block"] - config["layers_per_block"] = motion_adapter.config["motion_layers_per_block"] - config["temporal_transformer_layers_per_mid_block"] = motion_adapter.config[ - "motion_transformer_layers_per_mid_block" - ] - config["temporal_transformer_layers_per_block"] = motion_adapter.config[ - "motion_transformer_layers_per_block" - ] - config["motion_num_attention_heads"] = motion_adapter.config["motion_num_attention_heads"] - - # For PIA UNets we need to set the number input channels to 9 - if motion_adapter.config["conv_in_channels"]: - config["in_channels"] = motion_adapter.config["conv_in_channels"] - - # Need this for backwards compatibility with UNet2DConditionModel checkpoints - if not config.get("num_attention_heads"): - config["num_attention_heads"] = config["attention_head_dim"] - - expected_kwargs, optional_kwargs = cls._get_signature_keys(cls) - config = FrozenDict({k: config.get(k) for k in config if k in expected_kwargs or k in optional_kwargs}) - config["_class_name"] = cls.__name__ - model = cls.from_config(config) - - if not load_weights: - return model - - # Logic for loading PIA UNets which allow the first 4 channels to be any UNet2DConditionModel conv_in weight - # while the last 5 channels must be PIA conv_in weights. - if has_motion_adapter and motion_adapter.config["conv_in_channels"]: - model.conv_in = motion_adapter.conv_in - updated_conv_in_weight = torch.cat( - [unet.conv_in.weight, motion_adapter.conv_in.weight[:, 4:, :, :]], dim=1 - ) - model.conv_in.load_state_dict({"weight": updated_conv_in_weight, "bias": unet.conv_in.bias}) - else: - model.conv_in.load_state_dict(unet.conv_in.state_dict()) - - model.time_proj.load_state_dict(unet.time_proj.state_dict()) - model.time_embedding.load_state_dict(unet.time_embedding.state_dict()) - - if any( - isinstance(proc, (IPAdapterAttnProcessor, IPAdapterAttnProcessor2_0)) - for proc in unet.attn_processors.values() - ): - attn_procs = {} - for name, processor in unet.attn_processors.items(): - if name.endswith("attn1.processor"): - attn_processor_class = ( - AttnProcessor2_0 if hasattr(F, "scaled_dot_product_attention") else AttnProcessor - ) - attn_procs[name] = attn_processor_class() - else: - attn_processor_class = ( - IPAdapterAttnProcessor2_0 - if hasattr(F, "scaled_dot_product_attention") - else IPAdapterAttnProcessor - ) - attn_procs[name] = attn_processor_class( - hidden_size=processor.hidden_size, - cross_attention_dim=processor.cross_attention_dim, - scale=processor.scale, - num_tokens=processor.num_tokens, - ) - for name, processor in model.attn_processors.items(): - if name not in attn_procs: - attn_procs[name] = processor.__class__() - model.set_attn_processor(attn_procs) - model.config.encoder_hid_dim_type = "ip_image_proj" - model.encoder_hid_proj = unet.encoder_hid_proj - - for i, down_block in enumerate(unet.down_blocks): - model.down_blocks[i].resnets.load_state_dict(down_block.resnets.state_dict()) - if hasattr(model.down_blocks[i], "attentions"): - model.down_blocks[i].attentions.load_state_dict(down_block.attentions.state_dict()) - if model.down_blocks[i].downsamplers: - model.down_blocks[i].downsamplers.load_state_dict(down_block.downsamplers.state_dict()) - - for i, up_block in enumerate(unet.up_blocks): - model.up_blocks[i].resnets.load_state_dict(up_block.resnets.state_dict()) - if hasattr(model.up_blocks[i], "attentions"): - model.up_blocks[i].attentions.load_state_dict(up_block.attentions.state_dict()) - if model.up_blocks[i].upsamplers: - model.up_blocks[i].upsamplers.load_state_dict(up_block.upsamplers.state_dict()) - - model.mid_block.resnets.load_state_dict(unet.mid_block.resnets.state_dict()) - model.mid_block.attentions.load_state_dict(unet.mid_block.attentions.state_dict()) - - if unet.conv_norm_out is not None: - model.conv_norm_out.load_state_dict(unet.conv_norm_out.state_dict()) - if unet.conv_act is not None: - model.conv_act.load_state_dict(unet.conv_act.state_dict()) - model.conv_out.load_state_dict(unet.conv_out.state_dict()) - - if has_motion_adapter: - model.load_motion_modules(motion_adapter) - - # ensure that the Motion UNet is the same dtype as the UNet2DConditionModel - model.to(unet.dtype) - - return model - - def freeze_unet2d_params(self) -> None: - """Freeze the weights of just the UNet2DConditionModel, and leave the motion modules - unfrozen for fine tuning. - """ - # Freeze everything - for param in self.parameters(): - param.requires_grad = False - - # Unfreeze Motion Modules - for down_block in self.down_blocks: - motion_modules = down_block.motion_modules - for param in motion_modules.parameters(): - param.requires_grad = True - - for up_block in self.up_blocks: - motion_modules = up_block.motion_modules - for param in motion_modules.parameters(): - param.requires_grad = True - - if hasattr(self.mid_block, "motion_modules"): - motion_modules = self.mid_block.motion_modules - for param in motion_modules.parameters(): - param.requires_grad = True - - def load_motion_modules(self, motion_adapter: MotionAdapter | None) -> None: - for i, down_block in enumerate(motion_adapter.down_blocks): - self.down_blocks[i].motion_modules.load_state_dict(down_block.motion_modules.state_dict()) - for i, up_block in enumerate(motion_adapter.up_blocks): - self.up_blocks[i].motion_modules.load_state_dict(up_block.motion_modules.state_dict()) - - # to support older motion modules that don't have a mid_block - if hasattr(self.mid_block, "motion_modules"): - self.mid_block.motion_modules.load_state_dict(motion_adapter.mid_block.motion_modules.state_dict()) - - def save_motion_modules( - self, - save_directory: str, - is_main_process: bool = True, - safe_serialization: bool = True, - variant: str | None = None, - push_to_hub: bool = False, - **kwargs, - ) -> None: - state_dict = self.state_dict() - - # Extract all motion modules - motion_state_dict = {} - for k, v in state_dict.items(): - if "motion_modules" in k: - motion_state_dict[k] = v - - adapter = MotionAdapter( - block_out_channels=self.config["block_out_channels"], - motion_layers_per_block=self.config["layers_per_block"], - motion_norm_num_groups=self.config["norm_num_groups"], - motion_num_attention_heads=self.config["motion_num_attention_heads"], - motion_max_seq_length=self.config["motion_max_seq_length"], - use_motion_mid_block=self.config["use_motion_mid_block"], - ) - adapter.load_state_dict(motion_state_dict) - adapter.save_pretrained( - save_directory=save_directory, - is_main_process=is_main_process, - safe_serialization=safe_serialization, - variant=variant, - push_to_hub=push_to_hub, - **kwargs, - ) - - def enable_forward_chunking(self, chunk_size: int | None = None, dim: int = 0) -> None: - """ - Sets the attention processor to use [feed forward - chunking](https://huggingface.co/blog/reformer#2-chunked-feed-forward-layers). - - Parameters: - chunk_size (`int`, *optional*): - The chunk size of the feed-forward layers. If not specified, will run feed-forward layer individually - over each tensor of dim=`dim`. - dim (`int`, *optional*, defaults to `0`): - The dimension over which the feed-forward computation should be chunked. Choose between dim=0 (batch) - or dim=1 (sequence length). - """ - if dim not in [0, 1]: - raise ValueError(f"Make sure to set `dim` to either 0 or 1, not {dim}") - - # By default chunk size is 1 - chunk_size = chunk_size or 1 - - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, chunk_size, dim) - - def disable_forward_chunking(self) -> None: - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, None, 0) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self) -> None: - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.enable_freeu - def enable_freeu(self, s1: float, s2: float, b1: float, b2: float) -> None: - r"""Enables the FreeU mechanism from https://huggingface.co/papers/2309.11497. - - The suffixes after the scaling factors represent the stage blocks where they are being applied. - - Please refer to the [official repository](https://github.com/ChenyangSi/FreeU) for combinations of values that - are known to work well for different pipelines such as Stable Diffusion v1, v2, and Stable Diffusion XL. - - Args: - s1 (`float`): - Scaling factor for stage 1 to attenuate the contributions of the skip features. This is done to - mitigate the "oversmoothing effect" in the enhanced denoising process. - s2 (`float`): - Scaling factor for stage 2 to attenuate the contributions of the skip features. This is done to - mitigate the "oversmoothing effect" in the enhanced denoising process. - b1 (`float`): Scaling factor for stage 1 to amplify the contributions of backbone features. - b2 (`float`): Scaling factor for stage 2 to amplify the contributions of backbone features. - """ - for i, upsample_block in enumerate(self.up_blocks): - setattr(upsample_block, "s1", s1) - setattr(upsample_block, "s2", s2) - setattr(upsample_block, "b1", b1) - setattr(upsample_block, "b2", b2) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.disable_freeu - def disable_freeu(self) -> None: - """Disables the FreeU mechanism.""" - freeu_keys = {"s1", "s2", "b1", "b2"} - for i, upsample_block in enumerate(self.up_blocks): - for k in freeu_keys: - if hasattr(upsample_block, k) or getattr(upsample_block, k, None) is not None: - setattr(upsample_block, k, None) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections - def fuse_qkv_projections(self): - """ - Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) - are fused. For cross-attention modules, key and value projection matrices are fused. - - > [!WARNING] > This API is 🧪 experimental. - """ - self.original_attn_processors = None - - for _, attn_processor in self.attn_processors.items(): - if "Added" in str(attn_processor.__class__.__name__): - raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") - - self.original_attn_processors = self.attn_processors - - for module in self.modules(): - if isinstance(module, Attention): - module.fuse_projections(fuse=True) - - self.set_attn_processor(FusedAttnProcessor2_0()) - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections - def unfuse_qkv_projections(self): - """Disables the fused QKV projection if enabled. - - > [!WARNING] > This API is 🧪 experimental. - - """ - if self.original_attn_processors is not None: - self.set_attn_processor(self.original_attn_processors) - - @apply_lora_scale("cross_attention_kwargs") - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - encoder_hidden_states: torch.Tensor, - timestep_cond: torch.Tensor | None = None, - attention_mask: torch.Tensor | None = None, - cross_attention_kwargs: dict[str, Any] | None = None, - added_cond_kwargs: dict[str, torch.Tensor] | None = None, - down_block_additional_residuals: tuple[torch.Tensor] | None = None, - mid_block_additional_residual: torch.Tensor | None = None, - return_dict: bool = True, - ) -> UNetMotionOutput | tuple[torch.Tensor]: - r""" - The [`UNetMotionModel`] forward method. - - Args: - sample (`torch.Tensor`): - The noisy input tensor with the following shape `(batch, num_frames, channel, height, width`. - timestep (`torch.Tensor` or `float` or `int`): The number of timesteps to denoise an input. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states with shape `(batch, sequence_length, feature_dim)`. - timestep_cond: (`torch.Tensor`, *optional*, defaults to `None`): - Conditional embeddings for timestep. If provided, the embeddings will be summed with the samples passed - through the `self.time_embedding` layer to obtain the timestep embeddings. - attention_mask (`torch.Tensor`, *optional*, defaults to `None`): - An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask - is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large - negative values to the attention scores corresponding to "discard" tokens. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - added_cond_kwargs (`dict`, *optional*): - A dictionary of additional embeddings (e.g. text and time embeddings) used to condition the model. - down_block_additional_residuals: (`tuple` of `torch.Tensor`, *optional*): - A tuple of tensors that if specified are added to the residuals of down unet blocks. - mid_block_additional_residual: (`torch.Tensor`, *optional*): - A tensor that if specified is added to the residual of the middle unet block. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unets.unet_motion_model.UNetMotionOutput`] instead of a plain - tuple. - - Returns: - [`~models.unets.unet_motion_model.UNetMotionOutput`] or `tuple`: - If `return_dict` is True, an [`~models.unets.unet_motion_model.UNetMotionOutput`] is returned, - otherwise a `tuple` is returned where the first element is the sample tensor. - """ - # By default samples have to be AT least a multiple of the overall upsampling factor. - # The overall upsampling factor is equal to 2 ** (# num of upsampling layears). - # However, the upsampling interpolation output size can be forced to fit any upsampling size - # on the fly if necessary. - default_overall_up_factor = 2**self.num_upsamplers - - # upsample size should be forwarded when sample is not a multiple of `default_overall_up_factor` - forward_upsample_size = False - upsample_size = None - - if any(s % default_overall_up_factor != 0 for s in sample.shape[-2:]): - logger.info("Forward upsample size to force interpolation output size.") - forward_upsample_size = True - - # prepare attention_mask - if attention_mask is not None: - attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # 1. time - timesteps = timestep - if not torch.is_tensor(timesteps): - # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can - # This would be a good case for the `match` statement (Python 3.10+) - dtype = maybe_adjust_dtype_for_device( - torch.float64 if isinstance(timestep, float) else torch.int64, sample.device - ) - timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) - elif len(timesteps.shape) == 0: - timesteps = timesteps[None].to(sample.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - num_frames = sample.shape[2] - timesteps = timesteps.expand(sample.shape[0]) - - t_emb = self.time_proj(timesteps) - - # timesteps does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=self.dtype) - - emb = self.time_embedding(t_emb, timestep_cond) - aug_emb = None - - if self.config.addition_embed_type == "text_time": - if "text_embeds" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `text_embeds` to be passed in `added_cond_kwargs`" - ) - - text_embeds = added_cond_kwargs.get("text_embeds") - if "time_ids" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `time_ids` to be passed in `added_cond_kwargs`" - ) - time_ids = added_cond_kwargs.get("time_ids") - time_embeds = self.add_time_proj(time_ids.flatten()) - time_embeds = time_embeds.reshape((text_embeds.shape[0], -1)) - - add_embeds = torch.concat([text_embeds, time_embeds], dim=-1) - add_embeds = add_embeds.to(emb.dtype) - aug_emb = self.add_embedding(add_embeds) - - emb = emb if aug_emb is None else emb + aug_emb - emb = emb.repeat_interleave(num_frames, dim=0, output_size=emb.shape[0] * num_frames) - - if self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "ip_image_proj": - if "image_embeds" not in added_cond_kwargs: - raise ValueError( - f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'ip_image_proj' which requires the keyword argument `image_embeds` to be passed in `added_conditions`" - ) - image_embeds = added_cond_kwargs.get("image_embeds") - image_embeds = self.encoder_hid_proj(image_embeds) - image_embeds = [ - image_embed.repeat_interleave(num_frames, dim=0, output_size=image_embed.shape[0] * num_frames) - for image_embed in image_embeds - ] - encoder_hidden_states = (encoder_hidden_states, image_embeds) - - # 2. pre-process - sample = sample.permute(0, 2, 1, 3, 4).reshape((sample.shape[0] * num_frames, -1) + sample.shape[3:]) - sample = self.conv_in(sample) - - # 3. down - down_block_res_samples = (sample,) - for downsample_block in self.down_blocks: - if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention: - sample, res_samples = downsample_block( - hidden_states=sample, - temb=emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample, res_samples = downsample_block(hidden_states=sample, temb=emb, num_frames=num_frames) - - down_block_res_samples += res_samples - - if down_block_additional_residuals is not None: - new_down_block_res_samples = () - - for down_block_res_sample, down_block_additional_residual in zip( - down_block_res_samples, down_block_additional_residuals - ): - down_block_res_sample = down_block_res_sample + down_block_additional_residual - new_down_block_res_samples += (down_block_res_sample,) - - down_block_res_samples = new_down_block_res_samples - - # 4. mid - if self.mid_block is not None: - # To support older versions of motion modules that don't have a mid_block - if hasattr(self.mid_block, "motion_modules"): - sample = self.mid_block( - sample, - emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample = self.mid_block( - sample, - emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - ) - - if mid_block_additional_residual is not None: - sample = sample + mid_block_additional_residual - - # 5. up - for i, upsample_block in enumerate(self.up_blocks): - is_final_block = i == len(self.up_blocks) - 1 - - res_samples = down_block_res_samples[-len(upsample_block.resnets) :] - down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)] - - # if we have not reached the final block and need to forward the - # upsample size, we do it here - if not is_final_block and forward_upsample_size: - upsample_size = down_block_res_samples[-1].shape[2:] - - if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention: - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - encoder_hidden_states=encoder_hidden_states, - upsample_size=upsample_size, - attention_mask=attention_mask, - num_frames=num_frames, - cross_attention_kwargs=cross_attention_kwargs, - ) - else: - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - upsample_size=upsample_size, - num_frames=num_frames, - ) - - # 6. post-process - if self.conv_norm_out: - sample = self.conv_norm_out(sample) - sample = self.conv_act(sample) - - sample = self.conv_out(sample) - - # reshape to (batch, channel, framerate, width, height) - sample = sample[None, :].reshape((-1, num_frames) + sample.shape[1:]).permute(0, 2, 1, 3, 4) - - if not return_dict: - return (sample,) - - return UNetMotionOutput(sample=sample) diff --git a/diffusers/models/unets/unet_spatio_temporal_condition.py b/diffusers/models/unets/unet_spatio_temporal_condition.py deleted file mode 100644 index d38be0b0675fbf628d66eb93ea46a36b7e4202cc..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_spatio_temporal_condition.py +++ /dev/null @@ -1,448 +0,0 @@ -from dataclasses import dataclass - -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import UNet2DConditionLoadersMixin -from ...utils import BaseOutput, logging -from ...utils.torch_utils import maybe_adjust_dtype_for_device -from ..attention import AttentionMixin -from ..attention_processor import CROSS_ATTENTION_PROCESSORS, AttnProcessor -from ..embeddings import TimestepEmbedding, Timesteps -from ..modeling_utils import ModelMixin -from .unet_3d_blocks import UNetMidBlockSpatioTemporal, get_down_block, get_up_block - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -@dataclass -class UNetSpatioTemporalConditionOutput(BaseOutput): - """ - The output of [`UNetSpatioTemporalConditionModel`]. - - Args: - sample (`torch.Tensor` of shape `(batch_size, num_frames, num_channels, height, width)`): - The hidden states output conditioned on `encoder_hidden_states` input. Output of last layer of model. - """ - - sample: torch.Tensor = None - - -class UNetSpatioTemporalConditionModel(ModelMixin, AttentionMixin, ConfigMixin, UNet2DConditionLoadersMixin): - r""" - A conditional Spatio-Temporal UNet model that takes a noisy video frames, conditional state, and a timestep and - returns a sample shaped output. - - This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented - for all models (such as downloading or saving). - - Parameters: - sample_size (`int` or `tuple[int, int]`, *optional*, defaults to `None`): - Height and width of input/output sample. - in_channels (`int`, *optional*, defaults to 8): Number of channels in the input sample. - out_channels (`int`, *optional*, defaults to 4): Number of channels in the output. - down_block_types (`tuple[str]`, *optional*, defaults to `("CrossAttnDownBlockSpatioTemporal", "CrossAttnDownBlockSpatioTemporal", "CrossAttnDownBlockSpatioTemporal", "DownBlockSpatioTemporal")`): - The tuple of downsample blocks to use. - up_block_types (`tuple[str]`, *optional*, defaults to `("UpBlockSpatioTemporal", "CrossAttnUpBlockSpatioTemporal", "CrossAttnUpBlockSpatioTemporal", "CrossAttnUpBlockSpatioTemporal")`): - The tuple of upsample blocks to use. - block_out_channels (`tuple[int]`, *optional*, defaults to `(320, 640, 1280, 1280)`): - The tuple of output channels for each block. - addition_time_embed_dim: (`int`, defaults to 256): - Dimension to to encode the additional time ids. - projection_class_embeddings_input_dim (`int`, defaults to 768): - The dimension of the projection of encoded `added_time_ids`. - layers_per_block (`int`, *optional*, defaults to 2): The number of layers per block. - cross_attention_dim (`int` or `tuple[int]`, *optional*, defaults to 1280): - The dimension of the cross attention features. - transformer_layers_per_block (`int`, `tuple[int]`, or `tuple[tuple]` , *optional*, defaults to 1): - The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`]. Only relevant for - [`~models.unets.unet_3d_blocks.CrossAttnDownBlockSpatioTemporal`], - [`~models.unets.unet_3d_blocks.CrossAttnUpBlockSpatioTemporal`], - [`~models.unets.unet_3d_blocks.UNetMidBlockSpatioTemporal`]. - num_attention_heads (`int`, `tuple[int]`, defaults to `(5, 10, 10, 20)`): - The number of attention heads. - dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. - """ - - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - sample_size: int | None = None, - in_channels: int = 8, - out_channels: int = 4, - down_block_types: tuple[str, ...] = ( - "CrossAttnDownBlockSpatioTemporal", - "CrossAttnDownBlockSpatioTemporal", - "CrossAttnDownBlockSpatioTemporal", - "DownBlockSpatioTemporal", - ), - up_block_types: tuple[str, ...] = ( - "UpBlockSpatioTemporal", - "CrossAttnUpBlockSpatioTemporal", - "CrossAttnUpBlockSpatioTemporal", - "CrossAttnUpBlockSpatioTemporal", - ), - block_out_channels: tuple[int, ...] = (320, 640, 1280, 1280), - addition_time_embed_dim: int = 256, - projection_class_embeddings_input_dim: int = 768, - layers_per_block: int | tuple[int] = 2, - cross_attention_dim: int | tuple[int] = 1024, - transformer_layers_per_block: int | tuple[int, tuple[tuple]] = 1, - num_attention_heads: int | tuple[int, ...] = (5, 10, 20, 20), - num_frames: int = 25, - ): - super().__init__() - - self.sample_size = sample_size - - # Check inputs - if len(down_block_types) != len(up_block_types): - raise ValueError( - f"Must provide the same number of `down_block_types` as `up_block_types`. `down_block_types`: {down_block_types}. `up_block_types`: {up_block_types}." - ) - - if len(block_out_channels) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}." - ) - - if isinstance(cross_attention_dim, list) and len(cross_attention_dim) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `cross_attention_dim` as `down_block_types`. `cross_attention_dim`: {cross_attention_dim}. `down_block_types`: {down_block_types}." - ) - - if not isinstance(layers_per_block, int) and len(layers_per_block) != len(down_block_types): - raise ValueError( - f"Must provide the same number of `layers_per_block` as `down_block_types`. `layers_per_block`: {layers_per_block}. `down_block_types`: {down_block_types}." - ) - - # input - self.conv_in = nn.Conv2d( - in_channels, - block_out_channels[0], - kernel_size=3, - padding=1, - ) - - # time - time_embed_dim = block_out_channels[0] * 4 - - self.time_proj = Timesteps(block_out_channels[0], True, downscale_freq_shift=0) - timestep_input_dim = block_out_channels[0] - - self.time_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim) - - self.add_time_proj = Timesteps(addition_time_embed_dim, True, downscale_freq_shift=0) - self.add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) - - self.down_blocks = nn.ModuleList([]) - self.up_blocks = nn.ModuleList([]) - - if isinstance(num_attention_heads, int): - num_attention_heads = (num_attention_heads,) * len(down_block_types) - - if isinstance(cross_attention_dim, int): - cross_attention_dim = (cross_attention_dim,) * len(down_block_types) - - if isinstance(layers_per_block, int): - layers_per_block = [layers_per_block] * len(down_block_types) - - if isinstance(transformer_layers_per_block, int): - transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types) - - blocks_time_embed_dim = time_embed_dim - - # down - output_channel = block_out_channels[0] - for i, down_block_type in enumerate(down_block_types): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - - down_block = get_down_block( - down_block_type, - num_layers=layers_per_block[i], - transformer_layers_per_block=transformer_layers_per_block[i], - in_channels=input_channel, - out_channels=output_channel, - temb_channels=blocks_time_embed_dim, - add_downsample=not is_final_block, - resnet_eps=1e-5, - cross_attention_dim=cross_attention_dim[i], - num_attention_heads=num_attention_heads[i], - resnet_act_fn="silu", - ) - self.down_blocks.append(down_block) - - # mid - self.mid_block = UNetMidBlockSpatioTemporal( - block_out_channels[-1], - temb_channels=blocks_time_embed_dim, - transformer_layers_per_block=transformer_layers_per_block[-1], - cross_attention_dim=cross_attention_dim[-1], - num_attention_heads=num_attention_heads[-1], - ) - - # count how many layers upsample the images - self.num_upsamplers = 0 - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - reversed_num_attention_heads = list(reversed(num_attention_heads)) - reversed_layers_per_block = list(reversed(layers_per_block)) - reversed_cross_attention_dim = list(reversed(cross_attention_dim)) - reversed_transformer_layers_per_block = list(reversed(transformer_layers_per_block)) - - output_channel = reversed_block_out_channels[0] - for i, up_block_type in enumerate(up_block_types): - is_final_block = i == len(block_out_channels) - 1 - - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)] - - # add upsample block for all BUT final layer - if not is_final_block: - add_upsample = True - self.num_upsamplers += 1 - else: - add_upsample = False - - up_block = get_up_block( - up_block_type, - num_layers=reversed_layers_per_block[i] + 1, - transformer_layers_per_block=reversed_transformer_layers_per_block[i], - in_channels=input_channel, - out_channels=output_channel, - prev_output_channel=prev_output_channel, - temb_channels=blocks_time_embed_dim, - add_upsample=add_upsample, - resnet_eps=1e-5, - resolution_idx=i, - cross_attention_dim=reversed_cross_attention_dim[i], - num_attention_heads=reversed_num_attention_heads[i], - resnet_act_fn="silu", - ) - self.up_blocks.append(up_block) - prev_output_channel = output_channel - - # out - self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=32, eps=1e-5) - self.conv_act = nn.SiLU() - - self.conv_out = nn.Conv2d( - block_out_channels[0], - out_channels, - kernel_size=3, - padding=1, - ) - - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - # Copied from diffusers.models.unets.unet_3d_condition.UNet3DConditionModel.enable_forward_chunking - def enable_forward_chunking(self, chunk_size: int | None = None, dim: int = 0) -> None: - """ - Sets the attention processor to use [feed forward - chunking](https://huggingface.co/blog/reformer#2-chunked-feed-forward-layers). - - Parameters: - chunk_size (`int`, *optional*): - The chunk size of the feed-forward layers. If not specified, will run feed-forward layer individually - over each tensor of dim=`dim`. - dim (`int`, *optional*, defaults to `0`): - The dimension over which the feed-forward computation should be chunked. Choose between dim=0 (batch) - or dim=1 (sequence length). - """ - if dim not in [0, 1]: - raise ValueError(f"Make sure to set `dim` to either 0 or 1, not {dim}") - - # By default chunk size is 1 - chunk_size = chunk_size or 1 - - def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int): - if hasattr(module, "set_chunk_feed_forward"): - module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim) - - for child in module.children(): - fn_recursive_feed_forward(child, chunk_size, dim) - - for module in self.children(): - fn_recursive_feed_forward(module, chunk_size, dim) - - def forward( - self, - sample: torch.Tensor, - timestep: torch.Tensor | float | int, - encoder_hidden_states: torch.Tensor, - added_time_ids: torch.Tensor, - return_dict: bool = True, - ) -> UNetSpatioTemporalConditionOutput | tuple: - r""" - The [`UNetSpatioTemporalConditionModel`] forward method. - - Args: - sample (`torch.Tensor`): - The noisy input tensor with the following shape `(batch, num_frames, channel, height, width)`. - timestep (`torch.Tensor` or `float` or `int`): The number of timesteps to denoise an input. - encoder_hidden_states (`torch.Tensor`): - The encoder hidden states with shape `(batch, sequence_length, cross_attention_dim)`. - added_time_ids: (`torch.Tensor`): - The additional time ids with shape `(batch, num_additional_ids)`. These are encoded with sinusoidal - embeddings and added to the time embeddings. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`~models.unet_slatio_temporal.UNetSpatioTemporalConditionOutput`] instead - of a plain tuple. - Returns: - [`~models.unet_slatio_temporal.UNetSpatioTemporalConditionOutput`] or `tuple`: - If `return_dict` is True, an [`~models.unet_slatio_temporal.UNetSpatioTemporalConditionOutput`] is - returned, otherwise a `tuple` is returned where the first element is the sample tensor. - """ - # By default samples have to be AT least a multiple of the overall upsampling factor. - # The overall upsampling factor is equal to 2 ** (# num of upsampling layears). - # However, the upsampling interpolation output size can be forced to fit any upsampling size - # on the fly if necessary. - default_overall_up_factor = 2**self.num_upsamplers - - # upsample size should be forwarded when sample is not a multiple of `default_overall_up_factor` - forward_upsample_size = False - upsample_size = None - - if any(s % default_overall_up_factor != 0 for s in sample.shape[-2:]): - logger.info("Forward upsample size to force interpolation output size.") - forward_upsample_size = True - - # 1. time - timesteps = timestep - if not torch.is_tensor(timesteps): - # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can - # This would be a good case for the `match` statement (Python 3.10+) - dtype = maybe_adjust_dtype_for_device( - torch.float64 if isinstance(timestep, float) else torch.int64, sample.device - ) - timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) - elif len(timesteps.shape) == 0: - timesteps = timesteps[None].to(sample.device) - - # broadcast to batch dimension in a way that's compatible with ONNX/Core ML - batch_size, num_frames = sample.shape[:2] - timesteps = timesteps.expand(batch_size) - - t_emb = self.time_proj(timesteps) - - # `Timesteps` does not contain any weights and will always return f32 tensors - # but time_embedding might actually be running in fp16. so we need to cast here. - # there might be better ways to encapsulate this. - t_emb = t_emb.to(dtype=sample.dtype) - - emb = self.time_embedding(t_emb) - - time_embeds = self.add_time_proj(added_time_ids.flatten()) - time_embeds = time_embeds.reshape((batch_size, -1)) - time_embeds = time_embeds.to(emb.dtype) - aug_emb = self.add_embedding(time_embeds) - emb = emb + aug_emb - - # Flatten the batch and frames dimensions - # sample: [batch, frames, channels, height, width] -> [batch * frames, channels, height, width] - sample = sample.flatten(0, 1) - # Repeat the embeddings num_video_frames times - # emb: [batch, channels] -> [batch * frames, channels] - emb = emb.repeat_interleave(num_frames, dim=0, output_size=emb.shape[0] * num_frames) - # encoder_hidden_states: [batch, 1, channels] -> [batch * frames, 1, channels] - encoder_hidden_states = encoder_hidden_states.repeat_interleave( - num_frames, dim=0, output_size=encoder_hidden_states.shape[0] * num_frames - ) - - # 2. pre-process - sample = self.conv_in(sample) - - image_only_indicator = torch.zeros(batch_size, num_frames, dtype=sample.dtype, device=sample.device) - - down_block_res_samples = (sample,) - for downsample_block in self.down_blocks: - if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention: - sample, res_samples = downsample_block( - hidden_states=sample, - temb=emb, - encoder_hidden_states=encoder_hidden_states, - image_only_indicator=image_only_indicator, - ) - else: - sample, res_samples = downsample_block( - hidden_states=sample, - temb=emb, - image_only_indicator=image_only_indicator, - ) - - down_block_res_samples += res_samples - - # 4. mid - sample = self.mid_block( - hidden_states=sample, - temb=emb, - encoder_hidden_states=encoder_hidden_states, - image_only_indicator=image_only_indicator, - ) - - # 5. up - for i, upsample_block in enumerate(self.up_blocks): - is_final_block = i == len(self.up_blocks) - 1 - - res_samples = down_block_res_samples[-len(upsample_block.resnets) :] - down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)] - - # if we have not reached the final block and need to forward the - # upsample size, we do it here - if not is_final_block and forward_upsample_size: - upsample_size = down_block_res_samples[-1].shape[2:] - - if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention: - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - encoder_hidden_states=encoder_hidden_states, - upsample_size=upsample_size, - image_only_indicator=image_only_indicator, - ) - else: - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - upsample_size=upsample_size, - image_only_indicator=image_only_indicator, - ) - - # 6. post-process - sample = self.conv_norm_out(sample) - sample = self.conv_act(sample) - sample = self.conv_out(sample) - - # 7. Reshape back to original shape - sample = sample.reshape(batch_size, num_frames, *sample.shape[1:]) - - if not return_dict: - return (sample,) - - return UNetSpatioTemporalConditionOutput(sample=sample) diff --git a/diffusers/models/unets/unet_stable_cascade.py b/diffusers/models/unets/unet_stable_cascade.py deleted file mode 100644 index e000fdc51e06fbd2d74a42c29f9ed6920dfde84e..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/unet_stable_cascade.py +++ /dev/null @@ -1,605 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import math -from dataclasses import dataclass - -import numpy as np -import torch -import torch.nn as nn - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import FromOriginalModelMixin -from ...utils import BaseOutput -from ..attention_processor import Attention -from ..modeling_utils import ModelMixin - - -# Copied from diffusers.pipelines.deprecated.wuerstchen.modeling_wuerstchen_common.WuerstchenLayerNorm with WuerstchenLayerNorm -> SDCascadeLayerNorm -class SDCascadeLayerNorm(nn.LayerNorm): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - - def forward(self, x): - x = x.permute(0, 2, 3, 1) - x = super().forward(x) - return x.permute(0, 3, 1, 2) - - -class SDCascadeTimestepBlock(nn.Module): - def __init__(self, c, c_timestep, conds=[]): - super().__init__() - - self.mapper = nn.Linear(c_timestep, c * 2) - self.conds = conds - for cname in conds: - setattr(self, f"mapper_{cname}", nn.Linear(c_timestep, c * 2)) - - def forward(self, x, t): - t = t.chunk(len(self.conds) + 1, dim=1) - a, b = self.mapper(t[0])[:, :, None, None].chunk(2, dim=1) - for i, c in enumerate(self.conds): - ac, bc = getattr(self, f"mapper_{c}")(t[i + 1])[:, :, None, None].chunk(2, dim=1) - a, b = a + ac, b + bc - return x * (1 + a) + b - - -class SDCascadeResBlock(nn.Module): - def __init__(self, c, c_skip=0, kernel_size=3, dropout=0.0): - super().__init__() - self.depthwise = nn.Conv2d(c, c, kernel_size=kernel_size, padding=kernel_size // 2, groups=c) - self.norm = SDCascadeLayerNorm(c, elementwise_affine=False, eps=1e-6) - self.channelwise = nn.Sequential( - nn.Linear(c + c_skip, c * 4), - nn.GELU(), - GlobalResponseNorm(c * 4), - nn.Dropout(dropout), - nn.Linear(c * 4, c), - ) - - def forward(self, x, x_skip=None): - x_res = x - x = self.norm(self.depthwise(x)) - if x_skip is not None: - x = torch.cat([x, x_skip], dim=1) - x = self.channelwise(x.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) - return x + x_res - - -# from https://github.com/facebookresearch/ConvNeXt-V2/blob/3608f67cc1dae164790c5d0aead7bf2d73d9719b/models/utils.py#L105 -class GlobalResponseNorm(nn.Module): - def __init__(self, dim): - super().__init__() - self.gamma = nn.Parameter(torch.zeros(1, 1, 1, dim)) - self.beta = nn.Parameter(torch.zeros(1, 1, 1, dim)) - - def forward(self, x): - agg_norm = torch.norm(x, p=2, dim=(1, 2), keepdim=True) - stand_div_norm = agg_norm / (agg_norm.mean(dim=-1, keepdim=True) + 1e-6) - return self.gamma * (x * stand_div_norm) + self.beta + x - - -class SDCascadeAttnBlock(nn.Module): - def __init__(self, c, c_cond, nhead, self_attn=True, dropout=0.0): - super().__init__() - - self.self_attn = self_attn - self.norm = SDCascadeLayerNorm(c, elementwise_affine=False, eps=1e-6) - self.attention = Attention(query_dim=c, heads=nhead, dim_head=c // nhead, dropout=dropout, bias=True) - self.kv_mapper = nn.Sequential(nn.SiLU(), nn.Linear(c_cond, c)) - - def forward(self, x, kv): - kv = self.kv_mapper(kv) - norm_x = self.norm(x) - if self.self_attn: - batch_size, channel, _, _ = x.shape - kv = torch.cat([norm_x.view(batch_size, channel, -1).transpose(1, 2), kv], dim=1) - x = x + self.attention(norm_x, encoder_hidden_states=kv) - return x - - -class UpDownBlock2d(nn.Module): - def __init__(self, in_channels, out_channels, mode, enabled=True): - super().__init__() - if mode not in ["up", "down"]: - raise ValueError(f"{mode} not supported") - interpolation = ( - nn.Upsample(scale_factor=2 if mode == "up" else 0.5, mode="bilinear", align_corners=True) - if enabled - else nn.Identity() - ) - mapping = nn.Conv2d(in_channels, out_channels, kernel_size=1) - self.blocks = nn.ModuleList([interpolation, mapping] if mode == "up" else [mapping, interpolation]) - - def forward(self, x): - for block in self.blocks: - x = block(x) - return x - - -@dataclass -class StableCascadeUNetOutput(BaseOutput): - sample: torch.Tensor = None - - -class StableCascadeUNet(ModelMixin, ConfigMixin, FromOriginalModelMixin): - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - in_channels: int = 16, - out_channels: int = 16, - timestep_ratio_embedding_dim: int = 64, - patch_size: int = 1, - conditioning_dim: int = 2048, - block_out_channels: tuple[int, ...] = (2048, 2048), - num_attention_heads: tuple[int, ...] = (32, 32), - down_num_layers_per_block: tuple[int, ...] = (8, 24), - up_num_layers_per_block: tuple[int, ...] = (24, 8), - down_blocks_repeat_mappers: tuple[int] | None = ( - 1, - 1, - ), - up_blocks_repeat_mappers: tuple[int] | None = (1, 1), - block_types_per_layer: tuple[tuple[str]] = ( - ("SDCascadeResBlock", "SDCascadeTimestepBlock", "SDCascadeAttnBlock"), - ("SDCascadeResBlock", "SDCascadeTimestepBlock", "SDCascadeAttnBlock"), - ), - clip_text_in_channels: int | None = None, - clip_text_pooled_in_channels=1280, - clip_image_in_channels: int | None = None, - clip_seq=4, - effnet_in_channels: int | None = None, - pixel_mapper_in_channels: int | None = None, - kernel_size=3, - dropout: float | tuple[float] = (0.1, 0.1), - self_attn: bool | tuple[bool] = True, - timestep_conditioning_type: tuple[str, ...] = ("sca", "crp"), - switch_level: tuple[bool] | None = None, - ): - """ - - Parameters: - in_channels (`int`, defaults to 16): - Number of channels in the input sample. - out_channels (`int`, defaults to 16): - Number of channels in the output sample. - timestep_ratio_embedding_dim (`int`, defaults to 64): - Dimension of the projected time embedding. - patch_size (`int`, defaults to 1): - Patch size to use for pixel unshuffling layer - conditioning_dim (`int`, defaults to 2048): - Dimension of the image and text conditional embedding. - block_out_channels (tuple[int], defaults to (2048, 2048)): - tuple of output channels for each block. - num_attention_heads (tuple[int], defaults to (32, 32)): - Number of attention heads in each attention block. Set to -1 to if block types in a layer do not have - attention. - down_num_layers_per_block (tuple[int], defaults to [8, 24]): - Number of layers in each down block. - up_num_layers_per_block (tuple[int], defaults to [24, 8]): - Number of layers in each up block. - down_blocks_repeat_mappers (tuple[int], optional, defaults to [1, 1]): - Number of 1x1 Convolutional layers to repeat in each down block. - up_blocks_repeat_mappers (tuple[int], optional, defaults to [1, 1]): - Number of 1x1 Convolutional layers to repeat in each up block. - block_types_per_layer (tuple[tuple[str]], optional, - defaults to ( - ("SDCascadeResBlock", "SDCascadeTimestepBlock", "SDCascadeAttnBlock"), ("SDCascadeResBlock", - "SDCascadeTimestepBlock", "SDCascadeAttnBlock") - ): Block types used in each layer of the up/down blocks. - clip_text_in_channels (`int`, *optional*, defaults to `None`): - Number of input channels for CLIP based text conditioning. - clip_text_pooled_in_channels (`int`, *optional*, defaults to 1280): - Number of input channels for pooled CLIP text embeddings. - clip_image_in_channels (`int`, *optional*): - Number of input channels for CLIP based image conditioning. - clip_seq (`int`, *optional*, defaults to 4): - effnet_in_channels (`int`, *optional*, defaults to `None`): - Number of input channels for effnet conditioning. - pixel_mapper_in_channels (`int`, defaults to `None`): - Number of input channels for pixel mapper conditioning. - kernel_size (`int`, *optional*, defaults to 3): - Kernel size to use in the block convolutional layers. - dropout (tuple[float], *optional*, defaults to (0.1, 0.1)): - Dropout to use per block. - self_attn (bool | tuple[bool]): - tuple of booleans that determine whether to use self attention in a block or not. - timestep_conditioning_type (tuple[str], defaults to ("sca", "crp")): - Timestep conditioning type. - switch_level (tuple[bool] | None, *optional*, defaults to `None`): - tuple that indicates whether upsampling or downsampling should be applied in a block - """ - - super().__init__() - - if len(block_out_channels) != len(down_num_layers_per_block): - raise ValueError( - f"Number of elements in `down_num_layers_per_block` must match the length of `block_out_channels`: {len(block_out_channels)}" - ) - - elif len(block_out_channels) != len(up_num_layers_per_block): - raise ValueError( - f"Number of elements in `up_num_layers_per_block` must match the length of `block_out_channels`: {len(block_out_channels)}" - ) - - elif len(block_out_channels) != len(down_blocks_repeat_mappers): - raise ValueError( - f"Number of elements in `down_blocks_repeat_mappers` must match the length of `block_out_channels`: {len(block_out_channels)}" - ) - - elif len(block_out_channels) != len(up_blocks_repeat_mappers): - raise ValueError( - f"Number of elements in `up_blocks_repeat_mappers` must match the length of `block_out_channels`: {len(block_out_channels)}" - ) - - elif len(block_out_channels) != len(block_types_per_layer): - raise ValueError( - f"Number of elements in `block_types_per_layer` must match the length of `block_out_channels`: {len(block_out_channels)}" - ) - - if isinstance(dropout, float): - dropout = (dropout,) * len(block_out_channels) - if isinstance(self_attn, bool): - self_attn = (self_attn,) * len(block_out_channels) - - # CONDITIONING - if effnet_in_channels is not None: - self.effnet_mapper = nn.Sequential( - nn.Conv2d(effnet_in_channels, block_out_channels[0] * 4, kernel_size=1), - nn.GELU(), - nn.Conv2d(block_out_channels[0] * 4, block_out_channels[0], kernel_size=1), - SDCascadeLayerNorm(block_out_channels[0], elementwise_affine=False, eps=1e-6), - ) - if pixel_mapper_in_channels is not None: - self.pixels_mapper = nn.Sequential( - nn.Conv2d(pixel_mapper_in_channels, block_out_channels[0] * 4, kernel_size=1), - nn.GELU(), - nn.Conv2d(block_out_channels[0] * 4, block_out_channels[0], kernel_size=1), - SDCascadeLayerNorm(block_out_channels[0], elementwise_affine=False, eps=1e-6), - ) - - self.clip_txt_pooled_mapper = nn.Linear(clip_text_pooled_in_channels, conditioning_dim * clip_seq) - if clip_text_in_channels is not None: - self.clip_txt_mapper = nn.Linear(clip_text_in_channels, conditioning_dim) - if clip_image_in_channels is not None: - self.clip_img_mapper = nn.Linear(clip_image_in_channels, conditioning_dim * clip_seq) - self.clip_norm = nn.LayerNorm(conditioning_dim, elementwise_affine=False, eps=1e-6) - - self.embedding = nn.Sequential( - nn.PixelUnshuffle(patch_size), - nn.Conv2d(in_channels * (patch_size**2), block_out_channels[0], kernel_size=1), - SDCascadeLayerNorm(block_out_channels[0], elementwise_affine=False, eps=1e-6), - ) - - def get_block(block_type, in_channels, nhead, c_skip=0, dropout=0, self_attn=True): - if block_type == "SDCascadeResBlock": - return SDCascadeResBlock(in_channels, c_skip, kernel_size=kernel_size, dropout=dropout) - elif block_type == "SDCascadeAttnBlock": - return SDCascadeAttnBlock(in_channels, conditioning_dim, nhead, self_attn=self_attn, dropout=dropout) - elif block_type == "SDCascadeTimestepBlock": - return SDCascadeTimestepBlock( - in_channels, timestep_ratio_embedding_dim, conds=timestep_conditioning_type - ) - else: - raise ValueError(f"Block type {block_type} not supported") - - # BLOCKS - # -- down blocks - self.down_blocks = nn.ModuleList() - self.down_downscalers = nn.ModuleList() - self.down_repeat_mappers = nn.ModuleList() - for i in range(len(block_out_channels)): - if i > 0: - self.down_downscalers.append( - nn.Sequential( - SDCascadeLayerNorm(block_out_channels[i - 1], elementwise_affine=False, eps=1e-6), - UpDownBlock2d( - block_out_channels[i - 1], block_out_channels[i], mode="down", enabled=switch_level[i - 1] - ) - if switch_level is not None - else nn.Conv2d(block_out_channels[i - 1], block_out_channels[i], kernel_size=2, stride=2), - ) - ) - else: - self.down_downscalers.append(nn.Identity()) - - down_block = nn.ModuleList() - for _ in range(down_num_layers_per_block[i]): - for block_type in block_types_per_layer[i]: - block = get_block( - block_type, - block_out_channels[i], - num_attention_heads[i], - dropout=dropout[i], - self_attn=self_attn[i], - ) - down_block.append(block) - self.down_blocks.append(down_block) - - if down_blocks_repeat_mappers is not None: - block_repeat_mappers = nn.ModuleList() - for _ in range(down_blocks_repeat_mappers[i] - 1): - block_repeat_mappers.append(nn.Conv2d(block_out_channels[i], block_out_channels[i], kernel_size=1)) - self.down_repeat_mappers.append(block_repeat_mappers) - - # -- up blocks - self.up_blocks = nn.ModuleList() - self.up_upscalers = nn.ModuleList() - self.up_repeat_mappers = nn.ModuleList() - for i in reversed(range(len(block_out_channels))): - if i > 0: - self.up_upscalers.append( - nn.Sequential( - SDCascadeLayerNorm(block_out_channels[i], elementwise_affine=False, eps=1e-6), - UpDownBlock2d( - block_out_channels[i], block_out_channels[i - 1], mode="up", enabled=switch_level[i - 1] - ) - if switch_level is not None - else nn.ConvTranspose2d( - block_out_channels[i], block_out_channels[i - 1], kernel_size=2, stride=2 - ), - ) - ) - else: - self.up_upscalers.append(nn.Identity()) - - up_block = nn.ModuleList() - for j in range(up_num_layers_per_block[::-1][i]): - for k, block_type in enumerate(block_types_per_layer[i]): - c_skip = block_out_channels[i] if i < len(block_out_channels) - 1 and j == k == 0 else 0 - block = get_block( - block_type, - block_out_channels[i], - num_attention_heads[i], - c_skip=c_skip, - dropout=dropout[i], - self_attn=self_attn[i], - ) - up_block.append(block) - self.up_blocks.append(up_block) - - if up_blocks_repeat_mappers is not None: - block_repeat_mappers = nn.ModuleList() - for _ in range(up_blocks_repeat_mappers[::-1][i] - 1): - block_repeat_mappers.append(nn.Conv2d(block_out_channels[i], block_out_channels[i], kernel_size=1)) - self.up_repeat_mappers.append(block_repeat_mappers) - - # OUTPUT - self.clf = nn.Sequential( - SDCascadeLayerNorm(block_out_channels[0], elementwise_affine=False, eps=1e-6), - nn.Conv2d(block_out_channels[0], out_channels * (patch_size**2), kernel_size=1), - nn.PixelShuffle(patch_size), - ) - - self.gradient_checkpointing = False - - def _init_weights(self, m): - if isinstance(m, (nn.Conv2d, nn.Linear)): - torch.nn.init.xavier_uniform_(m.weight) - if m.bias is not None: - nn.init.constant_(m.bias, 0) - - nn.init.normal_(self.clip_txt_pooled_mapper.weight, std=0.02) - nn.init.normal_(self.clip_txt_mapper.weight, std=0.02) if hasattr(self, "clip_txt_mapper") else None - nn.init.normal_(self.clip_img_mapper.weight, std=0.02) if hasattr(self, "clip_img_mapper") else None - - if hasattr(self, "effnet_mapper"): - nn.init.normal_(self.effnet_mapper[0].weight, std=0.02) # conditionings - nn.init.normal_(self.effnet_mapper[2].weight, std=0.02) # conditionings - - if hasattr(self, "pixels_mapper"): - nn.init.normal_(self.pixels_mapper[0].weight, std=0.02) # conditionings - nn.init.normal_(self.pixels_mapper[2].weight, std=0.02) # conditionings - - torch.nn.init.xavier_uniform_(self.embedding[1].weight, 0.02) # inputs - nn.init.constant_(self.clf[1].weight, 0) # outputs - - # blocks - for level_block in self.down_blocks + self.up_blocks: - for block in level_block: - if isinstance(block, SDCascadeResBlock): - block.channelwise[-1].weight.data *= np.sqrt(1 / sum(self.config.blocks[0])) - elif isinstance(block, SDCascadeTimestepBlock): - nn.init.constant_(block.mapper.weight, 0) - - def get_timestep_ratio_embedding(self, timestep_ratio, max_positions=10000): - r = timestep_ratio * max_positions - half_dim = self.config.timestep_ratio_embedding_dim // 2 - - emb = math.log(max_positions) / (half_dim - 1) - emb = torch.arange(half_dim, device=r.device).float().mul(-emb).exp() - emb = r[:, None] * emb[None, :] - emb = torch.cat([emb.sin(), emb.cos()], dim=1) - - if self.config.timestep_ratio_embedding_dim % 2 == 1: # zero pad - emb = nn.functional.pad(emb, (0, 1), mode="constant") - - return emb.to(dtype=r.dtype) - - def get_clip_embeddings(self, clip_txt_pooled, clip_txt=None, clip_img=None): - if len(clip_txt_pooled.shape) == 2: - clip_txt_pool = clip_txt_pooled.unsqueeze(1) - clip_txt_pool = self.clip_txt_pooled_mapper(clip_txt_pooled).view( - clip_txt_pooled.size(0), clip_txt_pooled.size(1) * self.config.clip_seq, -1 - ) - if clip_txt is not None and clip_img is not None: - clip_txt = self.clip_txt_mapper(clip_txt) - if len(clip_img.shape) == 2: - clip_img = clip_img.unsqueeze(1) - clip_img = self.clip_img_mapper(clip_img).view( - clip_img.size(0), clip_img.size(1) * self.config.clip_seq, -1 - ) - clip = torch.cat([clip_txt, clip_txt_pool, clip_img], dim=1) - else: - clip = clip_txt_pool - return self.clip_norm(clip) - - def _down_encode(self, x, r_embed, clip): - level_outputs = [] - block_group = zip(self.down_blocks, self.down_downscalers, self.down_repeat_mappers) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for down_block, downscaler, repmap in block_group: - x = downscaler(x) - for i in range(len(repmap) + 1): - for block in down_block: - if isinstance(block, SDCascadeResBlock): - x = self._gradient_checkpointing_func(block, x) - elif isinstance(block, SDCascadeAttnBlock): - x = self._gradient_checkpointing_func(block, x, clip) - elif isinstance(block, SDCascadeTimestepBlock): - x = self._gradient_checkpointing_func(block, x, r_embed) - else: - x = self._gradient_checkpointing_func(block) - if i < len(repmap): - x = repmap[i](x) - level_outputs.insert(0, x) - else: - for down_block, downscaler, repmap in block_group: - x = downscaler(x) - for i in range(len(repmap) + 1): - for block in down_block: - if isinstance(block, SDCascadeResBlock): - x = block(x) - elif isinstance(block, SDCascadeAttnBlock): - x = block(x, clip) - elif isinstance(block, SDCascadeTimestepBlock): - x = block(x, r_embed) - else: - x = block(x) - if i < len(repmap): - x = repmap[i](x) - level_outputs.insert(0, x) - return level_outputs - - def _up_decode(self, level_outputs, r_embed, clip): - x = level_outputs[0] - block_group = zip(self.up_blocks, self.up_upscalers, self.up_repeat_mappers) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for i, (up_block, upscaler, repmap) in enumerate(block_group): - for j in range(len(repmap) + 1): - for k, block in enumerate(up_block): - if isinstance(block, SDCascadeResBlock): - skip = level_outputs[i] if k == 0 and i > 0 else None - if skip is not None and (x.size(-1) != skip.size(-1) or x.size(-2) != skip.size(-2)): - orig_type = x.dtype - x = torch.nn.functional.interpolate( - x.float(), skip.shape[-2:], mode="bilinear", align_corners=True - ) - x = x.to(orig_type) - x = self._gradient_checkpointing_func(block, x, skip) - elif isinstance(block, SDCascadeAttnBlock): - x = self._gradient_checkpointing_func(block, x, clip) - elif isinstance(block, SDCascadeTimestepBlock): - x = self._gradient_checkpointing_func(block, x, r_embed) - else: - x = self._gradient_checkpointing_func(block, x) - if j < len(repmap): - x = repmap[j](x) - x = upscaler(x) - else: - for i, (up_block, upscaler, repmap) in enumerate(block_group): - for j in range(len(repmap) + 1): - for k, block in enumerate(up_block): - if isinstance(block, SDCascadeResBlock): - skip = level_outputs[i] if k == 0 and i > 0 else None - if skip is not None and (x.size(-1) != skip.size(-1) or x.size(-2) != skip.size(-2)): - orig_type = x.dtype - x = torch.nn.functional.interpolate( - x.float(), skip.shape[-2:], mode="bilinear", align_corners=True - ) - x = x.to(orig_type) - x = block(x, skip) - elif isinstance(block, SDCascadeAttnBlock): - x = block(x, clip) - elif isinstance(block, SDCascadeTimestepBlock): - x = block(x, r_embed) - else: - x = block(x) - if j < len(repmap): - x = repmap[j](x) - x = upscaler(x) - return x - - def forward( - self, - sample, - timestep_ratio, - clip_text_pooled, - clip_text=None, - clip_img=None, - effnet=None, - pixels=None, - sca=None, - crp=None, - return_dict=True, - ): - r""" - Args: - sample (`torch.Tensor`): The noisy input sample. - timestep_ratio (`torch.Tensor`): - Timestep ratio used to compute the timestep embedding. - clip_text_pooled (`torch.Tensor`): - Pooled CLIP text embeddings. - clip_text (`torch.Tensor`, *optional*): - Sequence-level CLIP text embeddings. - clip_img (`torch.Tensor`, *optional*): - CLIP image embeddings. - effnet (`torch.Tensor`, *optional*): - EfficientNet feature map used as additional conditioning. - pixels (`torch.Tensor`, *optional*): - Pixel-level conditioning tensor. If `None`, a tensor of zeros is used. - sca (`torch.Tensor`, *optional*): - Optional `sca` conditioning value used to build the timestep embedding. - crp (`torch.Tensor`, *optional*): - Optional `crp` conditioning value used to build the timestep embedding. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`StableCascadeUNetOutput`] instead of a plain tuple. - """ - if pixels is None: - pixels = sample.new_zeros(sample.size(0), 3, 8, 8) - - # Process the conditioning embeddings - timestep_ratio_embed = self.get_timestep_ratio_embedding(timestep_ratio) - for c in self.config.timestep_conditioning_type: - if c == "sca": - cond = sca - elif c == "crp": - cond = crp - else: - cond = None - t_cond = cond or torch.zeros_like(timestep_ratio) - timestep_ratio_embed = torch.cat([timestep_ratio_embed, self.get_timestep_ratio_embedding(t_cond)], dim=1) - clip = self.get_clip_embeddings(clip_txt_pooled=clip_text_pooled, clip_txt=clip_text, clip_img=clip_img) - - # Model Blocks - x = self.embedding(sample) - if hasattr(self, "effnet_mapper") and effnet is not None: - x = x + self.effnet_mapper( - nn.functional.interpolate(effnet, size=x.shape[-2:], mode="bilinear", align_corners=True) - ) - if hasattr(self, "pixels_mapper"): - x = x + nn.functional.interpolate( - self.pixels_mapper(pixels), size=x.shape[-2:], mode="bilinear", align_corners=True - ) - level_outputs = self._down_encode(x, timestep_ratio_embed, clip) - x = self._up_decode(level_outputs, timestep_ratio_embed, clip) - sample = self.clf(x) - - if not return_dict: - return (sample,) - return StableCascadeUNetOutput(sample=sample) diff --git a/diffusers/models/unets/uvit_2d.py b/diffusers/models/unets/uvit_2d.py deleted file mode 100644 index 317abe80b1ebe58a6b5f55ed6728b7e5e5bdd97f..0000000000000000000000000000000000000000 --- a/diffusers/models/unets/uvit_2d.py +++ /dev/null @@ -1,420 +0,0 @@ -# coding=utf-8 -# Copyright 2025 The HuggingFace Inc. team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch -import torch.nn.functional as F -from torch import nn -from torch.utils.checkpoint import checkpoint - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin -from ...utils import apply_lora_scale -from ..attention import AttentionMixin, BasicTransformerBlock, SkipFFTransformerBlock -from ..attention_processor import ( - ADDED_KV_ATTENTION_PROCESSORS, - CROSS_ATTENTION_PROCESSORS, - AttnAddedKVProcessor, - AttnProcessor, -) -from ..embeddings import TimestepEmbedding, get_timestep_embedding -from ..modeling_utils import ModelMixin -from ..normalization import GlobalResponseNorm, RMSNorm -from ..resnet import Downsample2D, Upsample2D - - -class UVit2DModel(ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin): - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - # global config - hidden_size: int = 1024, - use_bias: bool = False, - hidden_dropout: float = 0.0, - # conditioning dimensions - cond_embed_dim: int = 768, - micro_cond_encode_dim: int = 256, - micro_cond_embed_dim: int = 1280, - encoder_hidden_size: int = 768, - # num tokens - vocab_size: int = 8256, # codebook_size + 1 (for the mask token) rounded - codebook_size: int = 8192, - # `UVit2DConvEmbed` - in_channels: int = 768, - block_out_channels: int = 768, - num_res_blocks: int = 3, - downsample: bool = False, - upsample: bool = False, - block_num_heads: int = 12, - # `TransformerLayer` - num_hidden_layers: int = 22, - num_attention_heads: int = 16, - # `Attention` - attention_dropout: float = 0.0, - # `FeedForward` - intermediate_size: int = 2816, - # `Norm` - layer_norm_eps: float = 1e-6, - ln_elementwise_affine: bool = True, - sample_size: int = 64, - ): - super().__init__() - - self.encoder_proj = nn.Linear(encoder_hidden_size, hidden_size, bias=use_bias) - self.encoder_proj_layer_norm = RMSNorm(hidden_size, layer_norm_eps, ln_elementwise_affine) - - self.embed = UVit2DConvEmbed( - in_channels, block_out_channels, vocab_size, ln_elementwise_affine, layer_norm_eps, use_bias - ) - - self.cond_embed = TimestepEmbedding( - micro_cond_embed_dim + cond_embed_dim, hidden_size, sample_proj_bias=use_bias - ) - - self.down_block = UVitBlock( - block_out_channels, - num_res_blocks, - hidden_size, - hidden_dropout, - ln_elementwise_affine, - layer_norm_eps, - use_bias, - block_num_heads, - attention_dropout, - downsample, - False, - ) - - self.project_to_hidden_norm = RMSNorm(block_out_channels, layer_norm_eps, ln_elementwise_affine) - self.project_to_hidden = nn.Linear(block_out_channels, hidden_size, bias=use_bias) - - self.transformer_layers = nn.ModuleList( - [ - BasicTransformerBlock( - dim=hidden_size, - num_attention_heads=num_attention_heads, - attention_head_dim=hidden_size // num_attention_heads, - dropout=hidden_dropout, - cross_attention_dim=hidden_size, - attention_bias=use_bias, - norm_type="ada_norm_continuous", - ada_norm_continous_conditioning_embedding_dim=hidden_size, - norm_elementwise_affine=ln_elementwise_affine, - norm_eps=layer_norm_eps, - ada_norm_bias=use_bias, - ff_inner_dim=intermediate_size, - ff_bias=use_bias, - attention_out_bias=use_bias, - ) - for _ in range(num_hidden_layers) - ] - ) - - self.project_from_hidden_norm = RMSNorm(hidden_size, layer_norm_eps, ln_elementwise_affine) - self.project_from_hidden = nn.Linear(hidden_size, block_out_channels, bias=use_bias) - - self.up_block = UVitBlock( - block_out_channels, - num_res_blocks, - hidden_size, - hidden_dropout, - ln_elementwise_affine, - layer_norm_eps, - use_bias, - block_num_heads, - attention_dropout, - downsample=False, - upsample=upsample, - ) - - self.mlm_layer = ConvMlmLayer( - block_out_channels, in_channels, use_bias, ln_elementwise_affine, layer_norm_eps, codebook_size - ) - - self.gradient_checkpointing = False - - @apply_lora_scale("cross_attention_kwargs") - def forward(self, input_ids, encoder_hidden_states, pooled_text_emb, micro_conds, cross_attention_kwargs=None): - r""" - Args: - input_ids (`torch.LongTensor`): - Token ids of the masked latent image tokens, with shape `(batch_size, height, width)`. - encoder_hidden_states (`torch.Tensor`): - Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. - pooled_text_emb (`torch.Tensor`): - Pooled text embeddings used for additional conditioning. - micro_conds (`torch.Tensor`): - Micro-conditioning values that are embedded and combined with `pooled_text_emb`. - cross_attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor`. - """ - encoder_hidden_states = self.encoder_proj(encoder_hidden_states) - encoder_hidden_states = self.encoder_proj_layer_norm(encoder_hidden_states) - - micro_cond_embeds = get_timestep_embedding( - micro_conds.flatten(), self.config.micro_cond_encode_dim, flip_sin_to_cos=True, downscale_freq_shift=0 - ) - - micro_cond_embeds = micro_cond_embeds.reshape((input_ids.shape[0], -1)) - - pooled_text_emb = torch.cat([pooled_text_emb, micro_cond_embeds], dim=1) - pooled_text_emb = pooled_text_emb.to(dtype=self.dtype) - pooled_text_emb = self.cond_embed(pooled_text_emb).to(encoder_hidden_states.dtype) - - hidden_states = self.embed(input_ids) - - hidden_states = self.down_block( - hidden_states, - pooled_text_emb=pooled_text_emb, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - ) - - batch_size, channels, height, width = hidden_states.shape - hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch_size, height * width, channels) - - hidden_states = self.project_to_hidden_norm(hidden_states) - hidden_states = self.project_to_hidden(hidden_states) - - for layer in self.transformer_layers: - if torch.is_grad_enabled() and self.gradient_checkpointing: - - def layer_(*args): - return checkpoint(layer, *args) - - else: - layer_ = layer - - hidden_states = layer_( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - added_cond_kwargs={"pooled_text_emb": pooled_text_emb}, - ) - - hidden_states = self.project_from_hidden_norm(hidden_states) - hidden_states = self.project_from_hidden(hidden_states) - - hidden_states = hidden_states.reshape(batch_size, height, width, channels).permute(0, 3, 1, 2) - - hidden_states = self.up_block( - hidden_states, - pooled_text_emb=pooled_text_emb, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - ) - - logits = self.mlm_layer(hidden_states) - - return logits - - # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor - def set_default_attn_processor(self): - """ - Disables custom attention processors and sets the default attention implementation. - """ - if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnAddedKVProcessor() - elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): - processor = AttnProcessor() - else: - raise ValueError( - f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" - ) - - self.set_attn_processor(processor) - - -class UVit2DConvEmbed(nn.Module): - def __init__(self, in_channels, block_out_channels, vocab_size, elementwise_affine, eps, bias): - super().__init__() - self.embeddings = nn.Embedding(vocab_size, in_channels) - self.layer_norm = RMSNorm(in_channels, eps, elementwise_affine) - self.conv = nn.Conv2d(in_channels, block_out_channels, kernel_size=1, bias=bias) - - def forward(self, input_ids): - embeddings = self.embeddings(input_ids) - embeddings = self.layer_norm(embeddings) - embeddings = embeddings.permute(0, 3, 1, 2) - embeddings = self.conv(embeddings) - return embeddings - - -class UVitBlock(nn.Module): - def __init__( - self, - channels, - num_res_blocks: int, - hidden_size, - hidden_dropout, - ln_elementwise_affine, - layer_norm_eps, - use_bias, - block_num_heads, - attention_dropout, - downsample: bool, - upsample: bool, - ): - super().__init__() - - if downsample: - self.downsample = Downsample2D( - channels, - use_conv=True, - padding=0, - name="Conv2d_0", - kernel_size=2, - norm_type="rms_norm", - eps=layer_norm_eps, - elementwise_affine=ln_elementwise_affine, - bias=use_bias, - ) - else: - self.downsample = None - - self.res_blocks = nn.ModuleList( - [ - ConvNextBlock( - channels, - layer_norm_eps, - ln_elementwise_affine, - use_bias, - hidden_dropout, - hidden_size, - ) - for i in range(num_res_blocks) - ] - ) - - self.attention_blocks = nn.ModuleList( - [ - SkipFFTransformerBlock( - channels, - block_num_heads, - channels // block_num_heads, - hidden_size, - use_bias, - attention_dropout, - channels, - attention_bias=use_bias, - attention_out_bias=use_bias, - ) - for _ in range(num_res_blocks) - ] - ) - - if upsample: - self.upsample = Upsample2D( - channels, - use_conv_transpose=True, - kernel_size=2, - padding=0, - name="conv", - norm_type="rms_norm", - eps=layer_norm_eps, - elementwise_affine=ln_elementwise_affine, - bias=use_bias, - interpolate=False, - ) - else: - self.upsample = None - - def forward(self, x, pooled_text_emb, encoder_hidden_states, cross_attention_kwargs): - if self.downsample is not None: - x = self.downsample(x) - - for res_block, attention_block in zip(self.res_blocks, self.attention_blocks): - x = res_block(x, pooled_text_emb) - - batch_size, channels, height, width = x.shape - x = x.view(batch_size, channels, height * width).permute(0, 2, 1) - x = attention_block( - x, encoder_hidden_states=encoder_hidden_states, cross_attention_kwargs=cross_attention_kwargs - ) - x = x.permute(0, 2, 1).view(batch_size, channels, height, width) - - if self.upsample is not None: - x = self.upsample(x) - - return x - - -class ConvNextBlock(nn.Module): - def __init__( - self, channels, layer_norm_eps, ln_elementwise_affine, use_bias, hidden_dropout, hidden_size, res_ffn_factor=4 - ): - super().__init__() - self.depthwise = nn.Conv2d( - channels, - channels, - kernel_size=3, - padding=1, - groups=channels, - bias=use_bias, - ) - self.norm = RMSNorm(channels, layer_norm_eps, ln_elementwise_affine) - self.channelwise_linear_1 = nn.Linear(channels, int(channels * res_ffn_factor), bias=use_bias) - self.channelwise_act = nn.GELU() - self.channelwise_norm = GlobalResponseNorm(int(channels * res_ffn_factor)) - self.channelwise_linear_2 = nn.Linear(int(channels * res_ffn_factor), channels, bias=use_bias) - self.channelwise_dropout = nn.Dropout(hidden_dropout) - self.cond_embeds_mapper = nn.Linear(hidden_size, channels * 2, use_bias) - - def forward(self, x, cond_embeds): - x_res = x - - x = self.depthwise(x) - - x = x.permute(0, 2, 3, 1) - x = self.norm(x) - - x = self.channelwise_linear_1(x) - x = self.channelwise_act(x) - x = self.channelwise_norm(x) - x = self.channelwise_linear_2(x) - x = self.channelwise_dropout(x) - - x = x.permute(0, 3, 1, 2) - - x = x + x_res - - scale, shift = self.cond_embeds_mapper(F.silu(cond_embeds)).chunk(2, dim=1) - x = x * (1 + scale[:, :, None, None]) + shift[:, :, None, None] - - return x - - -class ConvMlmLayer(nn.Module): - def __init__( - self, - block_out_channels: int, - in_channels: int, - use_bias: bool, - ln_elementwise_affine: bool, - layer_norm_eps: float, - codebook_size: int, - ): - super().__init__() - self.conv1 = nn.Conv2d(block_out_channels, in_channels, kernel_size=1, bias=use_bias) - self.layer_norm = RMSNorm(in_channels, layer_norm_eps, ln_elementwise_affine) - self.conv2 = nn.Conv2d(in_channels, codebook_size, kernel_size=1, bias=use_bias) - - def forward(self, hidden_states): - hidden_states = self.conv1(hidden_states) - hidden_states = self.layer_norm(hidden_states.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) - logits = self.conv2(hidden_states) - return logits diff --git a/diffusers/models/upsampling.py b/diffusers/models/upsampling.py deleted file mode 100644 index 36f22250a873634025f617bb86c1a956beb40995..0000000000000000000000000000000000000000 --- a/diffusers/models/upsampling.py +++ /dev/null @@ -1,515 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from ..utils import deprecate -from ..utils.import_utils import is_torch_version -from .normalization import RMSNorm - - -class Upsample1D(nn.Module): - """A 1D upsampling layer with an optional convolution. - - Parameters: - channels (`int`): - number of channels in the inputs and outputs. - use_conv (`bool`, default `False`): - option to use a convolution. - use_conv_transpose (`bool`, default `False`): - option to use a convolution transpose. - out_channels (`int`, optional): - number of output channels. Defaults to `channels`. - name (`str`, default `conv`): - name of the upsampling 1D layer. - """ - - def __init__( - self, - channels: int, - use_conv: bool = False, - use_conv_transpose: bool = False, - out_channels: int | None = None, - name: str = "conv", - ): - super().__init__() - self.channels = channels - self.out_channels = out_channels or channels - self.use_conv = use_conv - self.use_conv_transpose = use_conv_transpose - self.name = name - - self.conv = None - if use_conv_transpose: - self.conv = nn.ConvTranspose1d(channels, self.out_channels, 4, 2, 1) - elif use_conv: - self.conv = nn.Conv1d(self.channels, self.out_channels, 3, padding=1) - - def forward(self, inputs: torch.Tensor) -> torch.Tensor: - assert inputs.shape[1] == self.channels - if self.use_conv_transpose: - return self.conv(inputs) - - outputs = F.interpolate(inputs, scale_factor=2.0, mode="nearest") - - if self.use_conv: - outputs = self.conv(outputs) - - return outputs - - -class Upsample2D(nn.Module): - """A 2D upsampling layer with an optional convolution. - - Parameters: - channels (`int`): - number of channels in the inputs and outputs. - use_conv (`bool`, default `False`): - option to use a convolution. - use_conv_transpose (`bool`, default `False`): - option to use a convolution transpose. - out_channels (`int`, optional): - number of output channels. Defaults to `channels`. - name (`str`, default `conv`): - name of the upsampling 2D layer. - """ - - def __init__( - self, - channels: int, - use_conv: bool = False, - use_conv_transpose: bool = False, - out_channels: int | None = None, - name: str = "conv", - kernel_size: int | None = None, - padding=1, - norm_type=None, - eps=None, - elementwise_affine=None, - bias=True, - interpolate=True, - ): - super().__init__() - self.channels = channels - self.out_channels = out_channels or channels - self.use_conv = use_conv - self.use_conv_transpose = use_conv_transpose - self.name = name - self.interpolate = interpolate - - if norm_type == "ln_norm": - self.norm = nn.LayerNorm(channels, eps, elementwise_affine) - elif norm_type == "rms_norm": - self.norm = RMSNorm(channels, eps, elementwise_affine) - elif norm_type is None: - self.norm = None - else: - raise ValueError(f"unknown norm_type: {norm_type}") - - conv = None - if use_conv_transpose: - if kernel_size is None: - kernel_size = 4 - conv = nn.ConvTranspose2d( - channels, self.out_channels, kernel_size=kernel_size, stride=2, padding=padding, bias=bias - ) - elif use_conv: - if kernel_size is None: - kernel_size = 3 - conv = nn.Conv2d(self.channels, self.out_channels, kernel_size=kernel_size, padding=padding, bias=bias) - - # TODO(Suraj, Patrick) - clean up after weight dicts are correctly renamed - if name == "conv": - self.conv = conv - else: - self.Conv2d_0 = conv - - def forward(self, hidden_states: torch.Tensor, output_size: int | None = None, *args, **kwargs) -> torch.Tensor: - if len(args) > 0 or kwargs.get("scale", None) is not None: - deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." - deprecate("scale", "1.0.0", deprecation_message) - - assert hidden_states.shape[1] == self.channels - - if self.norm is not None: - hidden_states = self.norm(hidden_states.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) - - if self.use_conv_transpose: - return self.conv(hidden_states) - - # Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16 until PyTorch 2.1 - # https://github.com/pytorch/pytorch/issues/86679#issuecomment-1783978767 - dtype = hidden_states.dtype - if dtype == torch.bfloat16 and is_torch_version("<", "2.1"): - hidden_states = hidden_states.to(torch.float32) - - # upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984 - if hidden_states.shape[0] >= 64: - hidden_states = hidden_states.contiguous() - - # if `output_size` is passed we force the interpolation output - # size and do not make use of `scale_factor=2` - if self.interpolate: - # upsample_nearest_nhwc also fails when the number of output elements is large - # https://github.com/pytorch/pytorch/issues/141831 - scale_factor = ( - 2 if output_size is None else max([f / s for f, s in zip(output_size, hidden_states.shape[-2:])]) - ) - if hidden_states.numel() * scale_factor > pow(2, 31): - hidden_states = hidden_states.contiguous() - - if output_size is None: - hidden_states = F.interpolate(hidden_states, scale_factor=2.0, mode="nearest") - else: - hidden_states = F.interpolate(hidden_states, size=output_size, mode="nearest") - - # Cast back to original dtype - if dtype == torch.bfloat16 and is_torch_version("<", "2.1"): - hidden_states = hidden_states.to(dtype) - - # TODO(Suraj, Patrick) - clean up after weight dicts are correctly renamed - if self.use_conv: - if self.name == "conv": - hidden_states = self.conv(hidden_states) - else: - hidden_states = self.Conv2d_0(hidden_states) - - return hidden_states - - -class FirUpsample2D(nn.Module): - """A 2D FIR upsampling layer with an optional convolution. - - Parameters: - channels (`int`, optional): - number of channels in the inputs and outputs. - use_conv (`bool`, default `False`): - option to use a convolution. - out_channels (`int`, optional): - number of output channels. Defaults to `channels`. - fir_kernel (`tuple`, default `(1, 3, 3, 1)`): - kernel for the FIR filter. - """ - - def __init__( - self, - channels: int | None = None, - out_channels: int | None = None, - use_conv: bool = False, - fir_kernel: tuple[int, int, int, int] = (1, 3, 3, 1), - ): - super().__init__() - out_channels = out_channels if out_channels else channels - if use_conv: - self.Conv2d_0 = nn.Conv2d(channels, out_channels, kernel_size=3, stride=1, padding=1) - self.use_conv = use_conv - self.fir_kernel = fir_kernel - self.out_channels = out_channels - - def _upsample_2d( - self, - hidden_states: torch.Tensor, - weight: torch.Tensor | None = None, - kernel: torch.Tensor | None = None, - factor: int = 2, - gain: float = 1, - ) -> torch.Tensor: - """Fused `upsample_2d()` followed by `Conv2d()`. - - Padding is performed only once at the beginning, not between the operations. The fused op is considerably more - efficient than performing the same calculation using standard TensorFlow ops. It supports gradients of - arbitrary order. - - Args: - hidden_states (`torch.Tensor`): - Input tensor of the shape `[N, C, H, W]` or `[N, H, W, C]`. - weight (`torch.Tensor`, *optional*): - Weight tensor of the shape `[filterH, filterW, inChannels, outChannels]`. Grouped convolution can be - performed by `inChannels = x.shape[0] // numGroups`. - kernel (`torch.Tensor`, *optional*): - FIR filter of the shape `[firH, firW]` or `[firN]` (separable). The default is `[1] * factor`, which - corresponds to nearest-neighbor upsampling. - factor (`int`, *optional*): Integer upsampling factor (default: 2). - gain (`float`, *optional*): Scaling factor for signal magnitude (default: 1.0). - - Returns: - output (`torch.Tensor`): - Tensor of the shape `[N, C, H * factor, W * factor]` or `[N, H * factor, W * factor, C]`, and same - datatype as `hidden_states`. - """ - - assert isinstance(factor, int) and factor >= 1 - - # Setup filter kernel. - if kernel is None: - kernel = [1] * factor - - # setup kernel - kernel = torch.tensor(kernel, dtype=torch.float32) - if kernel.ndim == 1: - kernel = torch.outer(kernel, kernel) - kernel /= torch.sum(kernel) - - kernel = kernel * (gain * (factor**2)) - - if self.use_conv: - convH = weight.shape[2] - convW = weight.shape[3] - inC = weight.shape[1] - - pad_value = (kernel.shape[0] - factor) - (convW - 1) - - stride = (factor, factor) - # Determine data dimensions. - output_shape = ( - (hidden_states.shape[2] - 1) * factor + convH, - (hidden_states.shape[3] - 1) * factor + convW, - ) - output_padding = ( - output_shape[0] - (hidden_states.shape[2] - 1) * stride[0] - convH, - output_shape[1] - (hidden_states.shape[3] - 1) * stride[1] - convW, - ) - assert output_padding[0] >= 0 and output_padding[1] >= 0 - num_groups = hidden_states.shape[1] // inC - - # Transpose weights. - weight = torch.reshape(weight, (num_groups, -1, inC, convH, convW)) - weight = torch.flip(weight, dims=[3, 4]).permute(0, 2, 1, 3, 4) - weight = torch.reshape(weight, (num_groups * inC, -1, convH, convW)) - - inverse_conv = F.conv_transpose2d( - hidden_states, - weight, - stride=stride, - output_padding=output_padding, - padding=0, - ) - - output = upfirdn2d_native( - inverse_conv, - kernel.to(device=inverse_conv.device, dtype=inverse_conv.dtype), - pad=((pad_value + 1) // 2 + factor - 1, pad_value // 2 + 1), - ) - else: - pad_value = kernel.shape[0] - factor - output = upfirdn2d_native( - hidden_states, - kernel.to(device=hidden_states.device, dtype=hidden_states.dtype), - up=factor, - pad=((pad_value + 1) // 2 + factor - 1, pad_value // 2), - ) - - return output - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.use_conv: - height = self._upsample_2d(hidden_states, self.Conv2d_0.weight, kernel=self.fir_kernel) - height = height + self.Conv2d_0.bias.reshape(1, -1, 1, 1) - else: - height = self._upsample_2d(hidden_states, kernel=self.fir_kernel, factor=2) - - return height - - -class KUpsample2D(nn.Module): - r"""A 2D K-upsampling layer. - - Parameters: - pad_mode (`str`, *optional*, default to `"reflect"`): the padding mode to use. - """ - - def __init__(self, pad_mode: str = "reflect"): - super().__init__() - self.pad_mode = pad_mode - kernel_1d = torch.tensor([[1 / 8, 3 / 8, 3 / 8, 1 / 8]]) * 2 - self.pad = kernel_1d.shape[1] // 2 - 1 - self.register_buffer("kernel", kernel_1d.T @ kernel_1d, persistent=False) - - def forward(self, inputs: torch.Tensor) -> torch.Tensor: - inputs = F.pad(inputs, ((self.pad + 1) // 2,) * 4, self.pad_mode) - weight = inputs.new_zeros( - [ - inputs.shape[1], - inputs.shape[1], - self.kernel.shape[0], - self.kernel.shape[1], - ] - ) - indices = torch.arange(inputs.shape[1], device=inputs.device) - kernel = self.kernel.to(weight)[None, :].expand(inputs.shape[1], -1, -1) - weight[indices, indices] = kernel - return F.conv_transpose2d(inputs, weight, stride=2, padding=self.pad * 2 + 1) - - -class CogVideoXUpsample3D(nn.Module): - r""" - A 3D Upsample layer using in CogVideoX by Tsinghua University & ZhipuAI # Todo: Wait for paper release. - - Args: - in_channels (`int`): - Number of channels in the input image. - out_channels (`int`): - Number of channels produced by the convolution. - kernel_size (`int`, defaults to `3`): - Size of the convolving kernel. - stride (`int`, defaults to `1`): - Stride of the convolution. - padding (`int`, defaults to `1`): - Padding added to all four sides of the input. - compress_time (`bool`, defaults to `False`): - Whether or not to compress the time dimension. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int = 3, - stride: int = 1, - padding: int = 1, - compress_time: bool = False, - ) -> None: - super().__init__() - - self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding) - self.compress_time = compress_time - - def forward(self, inputs: torch.Tensor) -> torch.Tensor: - if self.compress_time: - if inputs.shape[2] > 1 and inputs.shape[2] % 2 == 1: - # split first frame - x_first, x_rest = inputs[:, :, 0], inputs[:, :, 1:] - - x_first = F.interpolate(x_first, scale_factor=2.0) - x_rest = F.interpolate(x_rest, scale_factor=2.0) - x_first = x_first[:, :, None, :, :] - inputs = torch.cat([x_first, x_rest], dim=2) - elif inputs.shape[2] > 1: - inputs = F.interpolate(inputs, scale_factor=2.0) - else: - inputs = inputs.squeeze(2) - inputs = F.interpolate(inputs, scale_factor=2.0) - inputs = inputs[:, :, None, :, :] - else: - # only interpolate 2D - b, c, t, h, w = inputs.shape - inputs = inputs.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w) - inputs = F.interpolate(inputs, scale_factor=2.0) - inputs = inputs.reshape(b, t, c, *inputs.shape[2:]).permute(0, 2, 1, 3, 4) - - b, c, t, h, w = inputs.shape - inputs = inputs.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w) - inputs = self.conv(inputs) - inputs = inputs.reshape(b, t, *inputs.shape[1:]).permute(0, 2, 1, 3, 4) - - return inputs - - -def upfirdn2d_native( - tensor: torch.Tensor, - kernel: torch.Tensor, - up: int = 1, - down: int = 1, - pad: tuple[int, int] = (0, 0), -) -> torch.Tensor: - up_x = up_y = up - down_x = down_y = down - pad_x0 = pad_y0 = pad[0] - pad_x1 = pad_y1 = pad[1] - - _, channel, in_h, in_w = tensor.shape - tensor = tensor.reshape(-1, in_h, in_w, 1) - - _, in_h, in_w, minor = tensor.shape - kernel_h, kernel_w = kernel.shape - - out = tensor.view(-1, in_h, 1, in_w, 1, minor) - out = F.pad(out, [0, 0, 0, up_x - 1, 0, 0, 0, up_y - 1]) - out = out.view(-1, in_h * up_y, in_w * up_x, minor) - - out = F.pad(out, [0, 0, max(pad_x0, 0), max(pad_x1, 0), max(pad_y0, 0), max(pad_y1, 0)]) - out = out.to(tensor.device) # Move back to mps if necessary - out = out[ - :, - max(-pad_y0, 0) : out.shape[1] - max(-pad_y1, 0), - max(-pad_x0, 0) : out.shape[2] - max(-pad_x1, 0), - :, - ] - - out = out.permute(0, 3, 1, 2) - out = out.reshape([-1, 1, in_h * up_y + pad_y0 + pad_y1, in_w * up_x + pad_x0 + pad_x1]) - w = torch.flip(kernel, [0, 1]).view(1, 1, kernel_h, kernel_w) - out = F.conv2d(out, w) - out = out.reshape( - -1, - minor, - in_h * up_y + pad_y0 + pad_y1 - kernel_h + 1, - in_w * up_x + pad_x0 + pad_x1 - kernel_w + 1, - ) - out = out.permute(0, 2, 3, 1) - out = out[:, ::down_y, ::down_x, :] - - out_h = (in_h * up_y + pad_y0 + pad_y1 - kernel_h) // down_y + 1 - out_w = (in_w * up_x + pad_x0 + pad_x1 - kernel_w) // down_x + 1 - - return out.view(-1, channel, out_h, out_w) - - -def upsample_2d( - hidden_states: torch.Tensor, - kernel: torch.Tensor | None = None, - factor: int = 2, - gain: float = 1, -) -> torch.Tensor: - r"""Upsample2D a batch of 2D images with the given filter. - Accepts a batch of 2D images of the shape `[N, C, H, W]` or `[N, H, W, C]` and upsamples each image with the given - filter. The filter is normalized so that if the input pixels are constant, they will be scaled by the specified - `gain`. Pixels outside the image are assumed to be zero, and the filter is padded with zeros so that its shape is - a: multiple of the upsampling factor. - - Args: - hidden_states (`torch.Tensor`): - Input tensor of the shape `[N, C, H, W]` or `[N, H, W, C]`. - kernel (`torch.Tensor`, *optional*): - FIR filter of the shape `[firH, firW]` or `[firN]` (separable). The default is `[1] * factor`, which - corresponds to nearest-neighbor upsampling. - factor (`int`, *optional*, default to `2`): - Integer upsampling factor. - gain (`float`, *optional*, default to `1.0`): - Scaling factor for signal magnitude (default: 1.0). - - Returns: - output (`torch.Tensor`): - Tensor of the shape `[N, C, H * factor, W * factor]` - """ - assert isinstance(factor, int) and factor >= 1 - if kernel is None: - kernel = [1] * factor - - kernel = torch.tensor(kernel, dtype=torch.float32) - if kernel.ndim == 1: - kernel = torch.outer(kernel, kernel) - kernel /= torch.sum(kernel) - - kernel = kernel * (gain * (factor**2)) - pad_value = kernel.shape[0] - factor - output = upfirdn2d_native( - hidden_states, - kernel.to(device=hidden_states.device, dtype=hidden_states.dtype), - up=factor, - pad=((pad_value + 1) // 2 + factor - 1, pad_value // 2), - ) - return output diff --git a/diffusers/models/vq_model.py b/diffusers/models/vq_model.py deleted file mode 100644 index 635db53102588cfa266d4a8c539c72a45394bcd8..0000000000000000000000000000000000000000 --- a/diffusers/models/vq_model.py +++ /dev/null @@ -1,29 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from ..utils import deprecate -from .autoencoders.vq_model import VQEncoderOutput, VQModel - - -class VQEncoderOutput(VQEncoderOutput): - def __init__(self, *args, **kwargs): - deprecation_message = "Importing `VQEncoderOutput` from `diffusers.models.vq_model` is deprecated and this will be removed in a future version. Please use `from diffusers.models.autoencoders.vq_model import VQEncoderOutput`, instead." - deprecate("VQEncoderOutput", "0.31", deprecation_message) - super().__init__(*args, **kwargs) - - -class VQModel(VQModel): - def __init__(self, *args, **kwargs): - deprecation_message = "Importing `VQModel` from `diffusers.models.vq_model` is deprecated and this will be removed in a future version. Please use `from diffusers.models.autoencoders.vq_model import VQModel`, instead." - deprecate("VQModel", "0.31", deprecation_message) - super().__init__(*args, **kwargs) diff --git a/diffusers/modular_pipelines/__init__.py b/diffusers/modular_pipelines/__init__.py deleted file mode 100644 index a107a004b7f29cf97f07357748f45f4a28458d96..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/__init__.py +++ /dev/null @@ -1,234 +0,0 @@ -from typing import TYPE_CHECKING - -from ..utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, - logging, -) - - -logger = logging.get_logger(__name__) -logger.warning( - "Modular Diffusers is currently an experimental feature under active development. The API is subject to breaking changes in future releases." -) - -# These modules contain pipelines from multiple libraries/frameworks -_dummy_objects = {} -_import_structure = {} - -try: - if not is_torch_available(): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ..utils import dummy_pt_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_pt_objects)) -else: - _import_structure["modular_pipeline"] = [ - "ModularPipelineBlocks", - "ModularPipeline", - "AutoPipelineBlocks", - "SequentialPipelineBlocks", - "ConditionalPipelineBlocks", - "LoopSequentialPipelineBlocks", - "PipelineState", - "BlockState", - ] - _import_structure["modular_pipeline_utils"] = [ - "ComponentSpec", - "ConfigSpec", - "InputParam", - "OutputParam", - "InsertableDict", - ] - _import_structure["stable_diffusion_xl"] = ["StableDiffusionXLAutoBlocks", "StableDiffusionXLModularPipeline"] - _import_structure["stable_diffusion_3"] = ["StableDiffusion3AutoBlocks", "StableDiffusion3ModularPipeline"] - _import_structure["wan"] = [ - "WanBlocks", - "Wan22Blocks", - "WanImage2VideoAutoBlocks", - "Wan22Image2VideoBlocks", - "WanModularPipeline", - "Wan22ModularPipeline", - "WanImage2VideoModularPipeline", - "Wan22Image2VideoModularPipeline", - ] - _import_structure["helios"] = [ - "HeliosAutoBlocks", - "HeliosModularPipeline", - "HeliosPyramidAutoBlocks", - "HeliosPyramidDistilledAutoBlocks", - "HeliosPyramidDistilledModularPipeline", - "HeliosPyramidModularPipeline", - ] - _import_structure["flux"] = [ - "FluxAutoBlocks", - "FluxModularPipeline", - "FluxKontextAutoBlocks", - "FluxKontextModularPipeline", - ] - _import_structure["flux2"] = [ - "Flux2AutoBlocks", - "Flux2KleinAutoBlocks", - "Flux2KleinBaseAutoBlocks", - "Flux2ModularPipeline", - "Flux2KleinModularPipeline", - "Flux2KleinBaseModularPipeline", - ] - _import_structure["ideogram4"] = [ - "Ideogram4AutoBlocks", - "Ideogram4ModularPipeline", - ] - _import_structure["krea2"] = [ - "Krea2AutoBlocks", - "Krea2ModularPipeline", - "Krea2TurboAutoBlocks", - "Krea2TurboModularPipeline", - ] - _import_structure["qwenimage"] = [ - "QwenImageAutoBlocks", - "QwenImageModularPipeline", - "QwenImageEditModularPipeline", - "QwenImageEditAutoBlocks", - "QwenImageEditPlusModularPipeline", - "QwenImageEditPlusAutoBlocks", - "QwenImageLayeredModularPipeline", - "QwenImageLayeredAutoBlocks", - ] - _import_structure["anima"] = [ - "AnimaAutoBlocks", - "AnimaModularPipeline", - ] - _import_structure["cosmos"] = [ - "Cosmos3DistilledBlocks", - "Cosmos3DistilledModularPipeline", - "Cosmos3OmniBlocks", - "Cosmos3OmniModularPipeline", - ] - _import_structure["ernie_image"] = [ - "ErnieImageAutoBlocks", - "ErnieImageModularPipeline", - ] - _import_structure["hunyuan_video1_5"] = [ - "HunyuanVideo15AutoBlocks", - "HunyuanVideo15ModularPipeline", - ] - _import_structure["ltx"] = [ - "LTXAutoBlocks", - "LTXModularPipeline", - ] - _import_structure["minimax_h3"] = [ - "MiniMaxH3Blocks", - "MiniMaxH3ModularPipeline", - "MiniMaxH3Ref2VABlocks", - "MiniMaxH3Ref2VAModularPipeline", - ] - _import_structure["z_image"] = [ - "ZImageAutoBlocks", - "ZImageModularPipeline", - ] - _import_structure["components_manager"] = ["ComponentsManager"] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not is_torch_available(): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ..utils.dummy_pt_objects import * # noqa F403 - else: - from .anima import AnimaAutoBlocks, AnimaModularPipeline - from .components_manager import ComponentsManager - from .cosmos import ( - Cosmos3DistilledBlocks, - Cosmos3DistilledModularPipeline, - Cosmos3OmniBlocks, - Cosmos3OmniModularPipeline, - ) - from .ernie_image import ErnieImageAutoBlocks, ErnieImageModularPipeline - from .flux import FluxAutoBlocks, FluxKontextAutoBlocks, FluxKontextModularPipeline, FluxModularPipeline - from .flux2 import ( - Flux2AutoBlocks, - Flux2KleinAutoBlocks, - Flux2KleinBaseAutoBlocks, - Flux2KleinBaseModularPipeline, - Flux2KleinModularPipeline, - Flux2ModularPipeline, - ) - from .helios import ( - HeliosAutoBlocks, - HeliosModularPipeline, - HeliosPyramidAutoBlocks, - HeliosPyramidDistilledAutoBlocks, - HeliosPyramidDistilledModularPipeline, - HeliosPyramidModularPipeline, - ) - from .hunyuan_video1_5 import ( - HunyuanVideo15AutoBlocks, - HunyuanVideo15ModularPipeline, - ) - from .ideogram4 import ( - Ideogram4AutoBlocks, - Ideogram4ModularPipeline, - ) - from .krea2 import ( - Krea2AutoBlocks, - Krea2ModularPipeline, - Krea2TurboAutoBlocks, - Krea2TurboModularPipeline, - ) - from .ltx import LTXAutoBlocks, LTXModularPipeline - from .minimax_h3 import ( - MiniMaxH3Blocks, - MiniMaxH3ModularPipeline, - MiniMaxH3Ref2VABlocks, - MiniMaxH3Ref2VAModularPipeline, - ) - from .modular_pipeline import ( - AutoPipelineBlocks, - BlockState, - ConditionalPipelineBlocks, - LoopSequentialPipelineBlocks, - ModularPipeline, - ModularPipelineBlocks, - PipelineState, - SequentialPipelineBlocks, - ) - from .modular_pipeline_utils import ComponentSpec, ConfigSpec, InputParam, InsertableDict, OutputParam - from .qwenimage import ( - QwenImageAutoBlocks, - QwenImageEditAutoBlocks, - QwenImageEditModularPipeline, - QwenImageEditPlusAutoBlocks, - QwenImageEditPlusModularPipeline, - QwenImageLayeredAutoBlocks, - QwenImageLayeredModularPipeline, - QwenImageModularPipeline, - ) - from .stable_diffusion_3 import StableDiffusion3AutoBlocks, StableDiffusion3ModularPipeline - from .stable_diffusion_xl import StableDiffusionXLAutoBlocks, StableDiffusionXLModularPipeline - from .wan import ( - Wan22Blocks, - Wan22Image2VideoBlocks, - Wan22Image2VideoModularPipeline, - Wan22ModularPipeline, - WanBlocks, - WanImage2VideoAutoBlocks, - WanImage2VideoModularPipeline, - WanModularPipeline, - ) - from .z_image import ZImageAutoBlocks, ZImageModularPipeline -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/anima/__init__.py b/diffusers/modular_pipelines/anima/__init__.py deleted file mode 100644 index 4772d906e03b74a73634c9db88497e2b63463abe..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/anima/__init__.py +++ /dev/null @@ -1,47 +0,0 @@ -from typing import TYPE_CHECKING - -from ...utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, -) - - -_dummy_objects = {} -_import_structure = {} - -try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ...utils import dummy_torch_and_transformers_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) -else: - _import_structure["modular_blocks_anima"] = ["AnimaAutoBlocks"] - _import_structure["modular_pipeline"] = ["AnimaModularPipeline"] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .modular_blocks_anima import AnimaAutoBlocks - from .modular_pipeline import AnimaModularPipeline -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/anima/before_denoise.py b/diffusers/modular_pipelines/anima/before_denoise.py deleted file mode 100644 index dbfe82d7f35d3a179951afead27fa7c72d5fb5c3..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/anima/before_denoise.py +++ /dev/null @@ -1,714 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect - -import numpy as np -import torch - -from ...models import AnimaTextConditioner, CosmosTransformer3DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import AnimaModularPipeline - - -def retrieve_timesteps( - scheduler, - num_inference_steps: int | None = None, - device: str | torch.device | None = None, - timesteps: list[int] | None = None, - sigmas: list[float] | None = None, - **kwargs, -): - if timesteps is not None and sigmas is not None: - raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values") - if timesteps is not None: - accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) - if not accepts_timesteps: - raise ValueError( - f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" - f" timestep schedules. Please check whether you are using the correct scheduler." - ) - scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) - timesteps = scheduler.timesteps - num_inference_steps = len(timesteps) - elif sigmas is not None: - accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) - if not accept_sigmas: - raise ValueError( - f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" - f" sigmas schedules. Please check whether you are using the correct scheduler." - ) - scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs) - timesteps = scheduler.timesteps - num_inference_steps = len(timesteps) - else: - scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) - timesteps = scheduler.timesteps - return timesteps, num_inference_steps - - -# Copied from diffusers.modular_pipelines.z_image.before_denoise.repeat_tensor_to_batch_size -def repeat_tensor_to_batch_size( - input_name: str, - input_tensor: torch.Tensor, - batch_size: int, - num_images_per_prompt: int = 1, -) -> torch.Tensor: - """Repeat tensor elements to match the final batch size. - - This function expands a tensor's batch dimension to match the final batch size (batch_size * num_images_per_prompt) - by repeating each element along dimension 0. - - The input tensor must have batch size 1 or batch_size. The function will: - - If batch size is 1: repeat each element (batch_size * num_images_per_prompt) times - - If batch size equals batch_size: repeat each element num_images_per_prompt times - - Args: - input_name (str): Name of the input tensor (used for error messages) - input_tensor (torch.Tensor): The tensor to repeat. Must have batch size 1 or batch_size. - batch_size (int): The base batch size (number of prompts) - num_images_per_prompt (int, optional): Number of images to generate per prompt. Defaults to 1. - - Returns: - torch.Tensor: The repeated tensor with final batch size (batch_size * num_images_per_prompt) - - Raises: - ValueError: If input_tensor is not a torch.Tensor or has invalid batch size - - Examples: - tensor = torch.tensor([[1, 2, 3]]) # shape: [1, 3] repeated = repeat_tensor_to_batch_size("image", tensor, - batch_size=2, num_images_per_prompt=2) repeated # tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]]) - shape: - [4, 3] - - tensor = torch.tensor([[1, 2, 3], [4, 5, 6]]) # shape: [2, 3] repeated = repeat_tensor_to_batch_size("image", - tensor, batch_size=2, num_images_per_prompt=2) repeated # tensor([[1, 2, 3], [1, 2, 3], [4, 5, 6], [4, 5, 6]]) - - shape: [4, 3] - """ - # make sure input is a tensor - if not isinstance(input_tensor, torch.Tensor): - raise ValueError(f"`{input_name}` must be a tensor") - - # make sure input tensor e.g. image_latents has batch size 1 or batch_size same as prompts - if input_tensor.shape[0] == 1: - repeat_by = batch_size * num_images_per_prompt - elif input_tensor.shape[0] == batch_size: - repeat_by = num_images_per_prompt - else: - raise ValueError( - f"`{input_name}` must have have batch size 1 or {batch_size}, but got {input_tensor.shape[0]}" - ) - - # expand the tensor to match the batch_size * num_images_per_prompt - input_tensor = input_tensor.repeat_interleave(repeat_by, dim=0) - - return input_tensor - - -class AnimaTextConditioningStep(ModularPipelineBlocks): - model_name = "anima" - - @property - def description(self) -> str: - return "Map Qwen text encoder states and T5 token ids to Cosmos text conditioning for Anima." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_conditioner", AnimaTextConditioner), - ComponentSpec("transformer", CosmosTransformer3DModel), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "qwen_prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Qwen prompt embeddings generated by the text encoder step.", - ), - InputParam( - "qwen_attention_mask", - required=True, - type_hint=torch.Tensor, - description="Qwen prompt attention mask generated by the text encoder step.", - ), - InputParam( - "t5_input_ids", - required=True, - type_hint=torch.Tensor, - description="T5 prompt token ids generated by the text encoder step.", - ), - InputParam( - "t5_attention_mask", - required=True, - type_hint=torch.Tensor, - description="T5 prompt attention mask generated by the text encoder step.", - ), - InputParam( - "negative_qwen_prompt_embeds", - type_hint=torch.Tensor, - description="Negative Qwen prompt embeddings generated by the text encoder step.", - ), - InputParam( - "negative_qwen_attention_mask", - type_hint=torch.Tensor, - description="Negative Qwen prompt attention mask generated by the text encoder step.", - ), - InputParam( - "negative_t5_input_ids", - type_hint=torch.Tensor, - description="Negative T5 prompt token ids generated by the text encoder step.", - ), - InputParam( - "negative_t5_attention_mask", - type_hint=torch.Tensor, - description="Negative T5 prompt attention mask generated by the text encoder step.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "prompt_embeds", - type_hint=torch.Tensor, - description="Conditioned prompt embeddings generated by the Anima text conditioner.", - ), - OutputParam( - "negative_prompt_embeds", - type_hint=torch.Tensor, - description="Conditioned negative prompt embeddings generated by the Anima text conditioner.", - ), - ] - - @staticmethod - def _condition_prompt_embeds( - components: AnimaModularPipeline, - qwen_prompt_embeds: torch.Tensor, - qwen_attention_mask: torch.Tensor, - t5_input_ids: torch.Tensor, - t5_attention_mask: torch.Tensor, - device: torch.device, - conditioning_dtype: torch.dtype, - output_dtype: torch.dtype, - ) -> torch.Tensor: - prompt_embeds = components.text_conditioner( - source_hidden_states=qwen_prompt_embeds.to(device=device, dtype=conditioning_dtype), - target_input_ids=t5_input_ids.to(device), - target_attention_mask=t5_attention_mask.to(device), - source_attention_mask=qwen_attention_mask.to(device), - ) - return prompt_embeds.to(dtype=output_dtype, device=device) - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - conditioning_dtype = components.text_conditioner.dtype - output_dtype = components.transformer.dtype - - block_state.prompt_embeds = self._condition_prompt_embeds( - components, - qwen_prompt_embeds=block_state.qwen_prompt_embeds, - qwen_attention_mask=block_state.qwen_attention_mask, - t5_input_ids=block_state.t5_input_ids, - t5_attention_mask=block_state.t5_attention_mask, - device=device, - conditioning_dtype=conditioning_dtype, - output_dtype=output_dtype, - ) - - block_state.negative_prompt_embeds = None - if block_state.negative_qwen_prompt_embeds is not None: - block_state.negative_prompt_embeds = self._condition_prompt_embeds( - components, - qwen_prompt_embeds=block_state.negative_qwen_prompt_embeds, - qwen_attention_mask=block_state.negative_qwen_attention_mask, - t5_input_ids=block_state.negative_t5_input_ids, - t5_attention_mask=block_state.negative_t5_attention_mask, - device=device, - conditioning_dtype=conditioning_dtype, - output_dtype=output_dtype, - ) - - self.set_block_state(state, block_state) - return components, state - - -class AnimaTextInputStep(ModularPipelineBlocks): - model_name = "anima" - - @property - def description(self) -> str: - return "Input processing step that expands Anima prompt embeddings for the requested image batch." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", CosmosTransformer3DModel)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_images_per_prompt"), - InputParam( - "prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Conditioned prompt embeddings generated by the Anima text conditioner.", - ), - InputParam( - "negative_prompt_embeds", - type_hint=torch.Tensor, - description="Conditioned negative prompt embeddings generated by the Anima text conditioner.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "prompt_embeds", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Prompt embeddings expanded to the final denoising batch.", - ), - OutputParam( - "negative_prompt_embeds", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Negative prompt embeddings expanded to the final denoising batch.", - ), - OutputParam( - "batch_size", - type_hint=int, - description="Number of input prompts before `num_images_per_prompt` expansion.", - ), - OutputParam("dtype", type_hint=torch.dtype, description="Dtype used by the Anima denoiser."), - ] - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - block_state.batch_size = block_state.prompt_embeds.shape[0] - block_state.dtype = components.transformer.dtype - - _, seq_len, _ = block_state.prompt_embeds.shape - block_state.prompt_embeds = block_state.prompt_embeds.repeat(1, block_state.num_images_per_prompt, 1) - block_state.prompt_embeds = block_state.prompt_embeds.view( - block_state.batch_size * block_state.num_images_per_prompt, seq_len, -1 - ) - - if block_state.negative_prompt_embeds is not None: - _, seq_len, _ = block_state.negative_prompt_embeds.shape - block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.repeat( - 1, block_state.num_images_per_prompt, 1 - ) - block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.view( - block_state.batch_size * block_state.num_images_per_prompt, seq_len, -1 - ) - - self.set_block_state(state, block_state) - return components, state - - -class AnimaImageInputStep(ModularPipelineBlocks): - model_name = "anima" - - @property - def description(self) -> str: - return ( - "Input processing step that expands Anima image latents to the final denoising batch " - "and derives height/width from the latents when not provided." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("image_latents"), - InputParam( - "batch_size", - required=True, - type_hint=int, - description="Number of input prompts before `num_images_per_prompt` expansion.", - ), - InputParam.template("num_images_per_prompt"), - InputParam.template("height"), - InputParam.template("width"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "image_latents", - type_hint=torch.Tensor, - description="Image latents expanded to the final denoising batch.", - ), - OutputParam("height", type_hint=int, description="Image height used for generation."), - OutputParam("width", type_hint=int, description="Image width used for generation."), - ] - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - latent_height, latent_width = block_state.image_latents.shape[-2:] - block_state.height = block_state.height or latent_height * components.vae_scale_factor - block_state.width = block_state.width or latent_width * components.vae_scale_factor - - block_state.image_latents = repeat_tensor_to_batch_size( - input_name="image_latents", - input_tensor=block_state.image_latents, - batch_size=block_state.batch_size, - num_images_per_prompt=block_state.num_images_per_prompt, - ) - - self.set_block_state(state, block_state) - return components, state - - -class AnimaPrepareLatentsStep(ModularPipelineBlocks): - model_name = "anima" - - @property - def description(self) -> str: - return "Prepare noisy image latents and padding mask for Anima denoising." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", CosmosTransformer3DModel)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("height"), - InputParam.template("width"), - InputParam.template("latents"), - InputParam.template("num_images_per_prompt"), - InputParam.template("generator"), - InputParam( - "batch_size", - required=True, - type_hint=int, - description="Number of input prompts before `num_images_per_prompt` expansion.", - ), - InputParam("dtype", type_hint=torch.dtype, description="Dtype used by the Anima denoiser."), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("height", type_hint=int, description="Image height used for generation."), - OutputParam("width", type_hint=int, description="Image width used for generation."), - OutputParam("latents", type_hint=torch.Tensor, description="Noisy latents for the denoising process."), - OutputParam("padding_mask", type_hint=torch.Tensor, description="Cosmos padding mask for image latents."), - ] - - def check_inputs(self, components: AnimaModularPipeline, block_state): - divisor = components.vae_scale_factor * 2 - if block_state.height % divisor != 0 or block_state.width % divisor != 0: - raise ValueError( - f"`height` and `width` have to be divisible by {divisor} but are {block_state.height} and" - f" {block_state.width}." - ) - - @staticmethod - def prepare_latents( - batch_size: int, - num_channels_latents: int, - height: int, - width: int, - vae_scale_factor: int, - dtype: torch.dtype, - device: torch.device, - generator: torch.Generator | list[torch.Generator] | None, - latents: torch.Tensor | None = None, - ) -> torch.Tensor: - if latents is not None: - return latents.to(device=device, dtype=dtype) - - latent_height = height // vae_scale_factor - latent_width = width // vae_scale_factor - shape = (batch_size, num_channels_latents, 1, latent_height, latent_width) - - if isinstance(generator, list) and len(generator) != batch_size: - raise ValueError( - f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" - f" size of {batch_size}. Make sure the batch size matches the length of the generators." - ) - - return randn_tensor(shape, generator=generator, device=device, dtype=dtype) - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - block_state.height = block_state.height or components.default_height - block_state.width = block_state.width or components.default_width - self.check_inputs(components, block_state) - - device = components._execution_device - block_state.latents = self.prepare_latents( - batch_size=block_state.batch_size * block_state.num_images_per_prompt, - num_channels_latents=components.num_channels_latents, - height=block_state.height, - width=block_state.width, - vae_scale_factor=components.vae_scale_factor, - dtype=torch.float32, - device=device, - generator=block_state.generator, - latents=block_state.latents, - ) - block_state.padding_mask = block_state.latents.new_zeros( - 1, 1, block_state.height, block_state.width, dtype=block_state.dtype - ) - - self.set_block_state(state, block_state) - return components, state - - -# Copied from diffusers.modular_pipelines.qwenimage.before_denoise.get_timesteps -def get_timesteps(scheduler, num_inference_steps, strength): - # get the original timestep using init_timestep - init_timestep = min(num_inference_steps * strength, num_inference_steps) - - t_start = int(max(num_inference_steps - init_timestep, 0)) - timesteps = scheduler.timesteps[t_start * scheduler.order :] - if hasattr(scheduler, "set_begin_index"): - scheduler.set_begin_index(t_start * scheduler.order) - - return timesteps, num_inference_steps - t_start - - -class AnimaSetTimestepsStep(ModularPipelineBlocks): - model_name = "anima" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def description(self) -> str: - return "Set the scheduler timesteps for Anima inference." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_inference_steps"), - InputParam.template("sigmas"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("timesteps", type_hint=torch.Tensor, description="Timesteps for the denoising loop."), - OutputParam("num_inference_steps", type_hint=int, description="Number of denoising steps."), - ] - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - sigmas = ( - np.linspace(1.0, 1 / block_state.num_inference_steps, block_state.num_inference_steps) - if block_state.sigmas is None - else block_state.sigmas - ) - block_state.timesteps, block_state.num_inference_steps = retrieve_timesteps( - components.scheduler, - device=device, - sigmas=sigmas, - ) - components.scheduler.set_begin_index(0) - - self.set_block_state(state, block_state) - return components, state - - -class AnimaImg2ImgSetTimestepsStep(ModularPipelineBlocks): - """Set the scheduler timesteps for Anima image-to-image inference. - - This step computes the full timestep schedule, then slices it based on ``strength`` via ``get_timesteps()``, which - also sets the scheduler's begin index. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - strength (`float`, *optional*, defaults to 0.9): - How much to transform the reference image. - - Outputs: - timesteps (`Tensor`): - Timestep schedule sliced by ``strength``. - num_inference_steps (`int`): - Number of denoising steps after strength-based slicing. - """ - - model_name = "anima" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def description(self) -> str: - return "Set the scheduler timesteps for Anima image-to-image inference, sliced by strength." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_inference_steps"), - InputParam.template("sigmas"), - InputParam.template("strength"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "timesteps", - type_hint=torch.Tensor, - description="Timestep schedule sliced by strength.", - ), - OutputParam( - "num_inference_steps", - type_hint=int, - description="Number of denoising steps after strength-based slicing.", - ), - ] - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - sigmas = ( - np.linspace(1.0, 1 / block_state.num_inference_steps, block_state.num_inference_steps) - if block_state.sigmas is None - else block_state.sigmas - ) - block_state.timesteps, block_state.num_inference_steps = retrieve_timesteps( - components.scheduler, - device=device, - sigmas=sigmas, - ) - block_state.timesteps, block_state.num_inference_steps = get_timesteps( - components.scheduler, block_state.num_inference_steps, block_state.strength - ) - - self.set_block_state(state, block_state) - return components, state - - -class AnimaImg2ImgPrepareLatentsStep(ModularPipelineBlocks): - """Prepares noisy latents for Anima image-to-image generation. - - Generates noise and mixes it with the image latents via ``scheduler.scale_noise()`` at the first sliced timestep. - The image latents are expected to already be expanded to the final batch size by ``AnimaImageInputStep``. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - image_latents (`Tensor`): - Encoded image latents, expanded to the final denoising batch. - timesteps (`Tensor`): - Timestep schedule sliced by ``strength`` from ``AnimaImg2ImgSetTimestepsStep``. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - latents (`Tensor`, *optional*): - Pre-computed noise tensor. Generated randomly if ``None``. - dtype (`torch.dtype`): - Dtype used by the Anima denoiser. - height (`int`): - Image height. - width (`int`): - Image width. - - Outputs: - latents (`Tensor`): - Noisy image latents for the denoising loop. - padding_mask (`Tensor`): - Cosmos padding mask for the image latents. - """ - - model_name = "anima" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def description(self) -> str: - return ( - "Prepares noisy image-to-image latents for Anima by adding noise to the encoded " - "image latents via scheduler.scale_noise()." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("image_latents"), - InputParam.template("timesteps", required=True), - InputParam.template("generator"), - InputParam.template("latents"), - InputParam("dtype", type_hint=torch.dtype, description="Dtype used by the Anima denoiser."), - InputParam.template("height"), - InputParam.template("width"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("latents", type_hint=torch.Tensor, description="Noisy latents for the denoising loop."), - OutputParam("padding_mask", type_hint=torch.Tensor, description="Cosmos padding mask for image latents."), - ] - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - image_latents = block_state.image_latents.to(device=device, dtype=torch.float32) - - if block_state.latents is None: - noise = randn_tensor( - image_latents.shape, - generator=block_state.generator, - device=device, - dtype=torch.float32, - ) - else: - noise = block_state.latents.to(device=device, dtype=torch.float32) - - latent_timestep = block_state.timesteps[:1].repeat(image_latents.shape[0]) - block_state.latents = components.scheduler.scale_noise(image_latents, latent_timestep, noise) - - block_state.padding_mask = block_state.latents.new_zeros( - 1, 1, block_state.height, block_state.width, dtype=block_state.dtype - ) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/anima/decoders.py b/diffusers/modular_pipelines/anima/decoders.py deleted file mode 100644 index f1f4b475a4b88f210164d898d6c373b5798300c6..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/anima/decoders.py +++ /dev/null @@ -1,120 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import numpy as np -import PIL -import torch - -from ...configuration_utils import FrozenDict -from ...image_processor import VaeImageProcessor -from ...models import AutoencoderKLQwenImage -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import AnimaModularPipeline - - -class AnimaVaeDecoderStep(ModularPipelineBlocks): - model_name = "anima" - - @property - def description(self) -> str: - return "Step that decodes Anima latents into image tensors." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("vae", AutoencoderKLQwenImage)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("latents", required=True, type_hint=torch.Tensor, description="Denoised Anima latents."), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam.template("images", note="tensor output of the VAE decoder")] - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - latents = block_state.latents.to(components.vae.dtype) - latents_mean = ( - torch.tensor(components.vae.config.latents_mean) - .view(1, components.vae.config.z_dim, 1, 1, 1) - .to(latents.device, latents.dtype) - ) - latents_std = 1.0 / torch.tensor(components.vae.config.latents_std).view( - 1, components.vae.config.z_dim, 1, 1, 1 - ).to(latents.device, latents.dtype) - latents = latents / latents_std + latents_mean - - block_state.images = components.vae.decode(latents, return_dict=False)[0][:, :, 0] - - self.set_block_state(state, block_state) - return components, state - - -class AnimaProcessImagesOutputStep(ModularPipelineBlocks): - model_name = "anima" - - @property - def description(self) -> str: - return "Postprocess decoded Anima image tensors." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "image_processor", - VaeImageProcessor, - config=FrozenDict({"vae_scale_factor": 8}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("images", required=True, type_hint=torch.Tensor, description="Decoded Anima image tensors."), - InputParam.template("output_type"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "images", - type_hint=list[PIL.Image.Image] | np.ndarray | torch.Tensor, - description="Generated images.", - ) - ] - - @staticmethod - def check_inputs(output_type): - if output_type not in ["pil", "np", "pt"]: - raise ValueError(f"Invalid output_type: {output_type}") - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - self.check_inputs(block_state.output_type) - - block_state.images = components.image_processor.postprocess( - image=block_state.images, - output_type=block_state.output_type, - ) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/anima/denoise.py b/diffusers/modular_pipelines/anima/denoise.py deleted file mode 100644 index d8146beefe72443bedcd6685000457e8c30011d6..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/anima/denoise.py +++ /dev/null @@ -1,211 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...models import CosmosTransformer3DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ..modular_pipeline import BlockState, LoopSequentialPipelineBlocks, ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam -from .modular_pipeline import AnimaModularPipeline - - -class AnimaLoopBeforeDenoiser(ModularPipelineBlocks): - model_name = "anima" - - @property - def description(self) -> str: - return "Step within the denoising loop that prepares Anima latent and timestep inputs." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("latents", required=True, type_hint=torch.Tensor, description="Current Anima latents."), - InputParam("dtype", required=True, type_hint=torch.dtype, description="Dtype used by the Anima denoiser."), - ] - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - block_state.latent_model_input = block_state.latents.to(block_state.dtype) - - timestep = t.expand(block_state.latents.shape[0]).to(block_state.dtype) - block_state.timestep = timestep / components.scheduler.config.num_train_timesteps - return components, block_state - - -class AnimaLoopDenoiser(ModularPipelineBlocks): - model_name = "anima" - - def __init__( - self, - guider_input_fields: dict[str, Any] | None = None, - ): - if guider_input_fields is None: - guider_input_fields = {"encoder_hidden_states": ("prompt_embeds", "negative_prompt_embeds")} - if not isinstance(guider_input_fields, dict): - raise ValueError(f"`guider_input_fields` must be a dictionary but is {type(guider_input_fields)}") - self._guider_input_fields = guider_input_fields - super().__init__() - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 4.0}), - default_creation_method="from_config", - ), - ComponentSpec("transformer", CosmosTransformer3DModel), - ] - - @property - def description(self) -> str: - return "Step within the denoising loop that predicts Anima noise with guidance." - - @property - def inputs(self) -> list[InputParam]: - inputs = [ - InputParam( - "num_inference_steps", - required=True, - type_hint=int, - description="Number of denoising steps.", - ), - InputParam( - "padding_mask", - required=True, - type_hint=torch.Tensor, - description="Cosmos padding mask for image latents.", - ), - InputParam( - kwargs_type="denoiser_input_fields", - description="The conditional model inputs for the Anima denoiser.", - ), - ] - - guider_input_names = [] - uncond_guider_input_names = [] - for value in self._guider_input_fields.values(): - if isinstance(value, tuple): - guider_input_names.append(value[0]) - uncond_guider_input_names.append(value[1]) - else: - guider_input_names.append(value) - - for name in guider_input_names: - inputs.append(InputParam(name=name, required=True)) - for name in uncond_guider_input_names: - inputs.append(InputParam(name=name)) - return inputs - - @torch.no_grad() - def __call__( - self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: - components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) - guider_state = components.guider.prepare_inputs_from_block_state(block_state, self._guider_input_fields) - - for guider_state_batch in guider_state: - components.guider.prepare_models(components.transformer) - cond_kwargs = { - key: getattr(guider_state_batch, key).to(block_state.dtype) for key in self._guider_input_fields.keys() - } - guider_state_batch.noise_pred = components.transformer( - hidden_states=block_state.latent_model_input, - timestep=block_state.timestep, - padding_mask=block_state.padding_mask, - return_dict=False, - **cond_kwargs, - )[0] - components.guider.cleanup_models(components.transformer) - - block_state.noise_pred = components.guider(guider_state)[0] - return components, block_state - - -class AnimaLoopAfterDenoiser(ModularPipelineBlocks): - model_name = "anima" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def description(self) -> str: - return "Step within the denoising loop that updates Anima latents." - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - latents_dtype = block_state.latents.dtype - block_state.latents = components.scheduler.step( - block_state.noise_pred, t, block_state.latents, return_dict=False - )[0] - if block_state.latents.dtype != latents_dtype and torch.backends.mps.is_available(): - block_state.latents = block_state.latents.to(latents_dtype) - - return components, block_state - - -class AnimaDenoiseLoopWrapper(LoopSequentialPipelineBlocks): - model_name = "anima" - - @property - def description(self) -> str: - return "Pipeline block that iteratively denoises Anima latents over scheduler timesteps." - - @property - def loop_expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def loop_inputs(self) -> list[InputParam]: - return [ - InputParam("timesteps", required=True, type_hint=torch.Tensor, description="Timesteps to denoise over."), - InputParam("num_inference_steps", required=True, type_hint=int, description="Number of denoising steps."), - ] - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - num_warmup_steps = len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order - - with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) - if i == len(block_state.timesteps) - 1 or ( - (i + 1) > num_warmup_steps and (i + 1) % components.scheduler.order == 0 - ): - progress_bar.update() - - self.set_block_state(state, block_state) - return components, state - - -class AnimaDenoiseStep(AnimaDenoiseLoopWrapper): - block_classes = [ - AnimaLoopBeforeDenoiser, - AnimaLoopDenoiser(guider_input_fields={"encoder_hidden_states": ("prompt_embeds", "negative_prompt_embeds")}), - AnimaLoopAfterDenoiser, - ] - block_names = ["before_denoiser", "denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return "Denoise step that iteratively denoises image latents for Anima." diff --git a/diffusers/modular_pipelines/anima/encoders.py b/diffusers/modular_pipelines/anima/encoders.py deleted file mode 100644 index 68950f97be83dcbb23efa2070c82931ceecd3e97..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/anima/encoders.py +++ /dev/null @@ -1,404 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -from transformers import Qwen2Tokenizer, Qwen3Model, T5TokenizerFast - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...image_processor import VaeImageProcessor -from ...models import AutoencoderKLQwenImage -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import AnimaModularPipeline - - -class AnimaTextEncoderStep(ModularPipelineBlocks): - model_name = "anima" - - @property - def description(self) -> str: - return "Text encoder step that encodes Anima prompts into Qwen states and T5 token ids." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", Qwen3Model), - ComponentSpec("tokenizer", Qwen2Tokenizer), - ComponentSpec("t5_tokenizer", T5TokenizerFast), - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 4.0}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt"), - InputParam.template("negative_prompt"), - InputParam.template("max_sequence_length"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "qwen_prompt_embeds", - type_hint=torch.Tensor, - description="Qwen prompt embeddings to be consumed by the Anima text conditioner.", - ), - OutputParam( - "qwen_attention_mask", - type_hint=torch.Tensor, - description="Qwen prompt attention mask to be consumed by the Anima text conditioner.", - ), - OutputParam( - "t5_input_ids", - type_hint=torch.Tensor, - description="T5 prompt token ids to be consumed by the Anima text conditioner.", - ), - OutputParam( - "t5_attention_mask", - type_hint=torch.Tensor, - description="T5 prompt attention mask to be consumed by the Anima text conditioner.", - ), - OutputParam( - "negative_qwen_prompt_embeds", - type_hint=torch.Tensor, - description="Negative Qwen prompt embeddings to be consumed by the Anima text conditioner.", - ), - OutputParam( - "negative_qwen_attention_mask", - type_hint=torch.Tensor, - description="Negative Qwen prompt attention mask to be consumed by the Anima text conditioner.", - ), - OutputParam( - "negative_t5_input_ids", - type_hint=torch.Tensor, - description="Negative T5 prompt token ids to be consumed by the Anima text conditioner.", - ), - OutputParam( - "negative_t5_attention_mask", - type_hint=torch.Tensor, - description="Negative T5 prompt attention mask to be consumed by the Anima text conditioner.", - ), - ] - - @staticmethod - def check_inputs(block_state): - if not isinstance(block_state.prompt, str) and not isinstance(block_state.prompt, list): - raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(block_state.prompt)}") - if block_state.max_sequence_length is not None and block_state.max_sequence_length > 4096: - raise ValueError( - f"`max_sequence_length` cannot be greater than 4096 but is {block_state.max_sequence_length}" - ) - - @staticmethod - def _get_qwen_prompt_embeds( - components: AnimaModularPipeline, - prompt: str | list[str], - max_sequence_length: int, - device: torch.device, - dtype: torch.dtype, - ) -> tuple[torch.Tensor, torch.Tensor]: - prompt = [prompt] if isinstance(prompt, str) else prompt - - text_inputs = components.tokenizer( - prompt, - padding="longest", - max_length=max_sequence_length, - truncation=True, - return_tensors="pt", - ) - text_input_ids = text_inputs.input_ids.to(device) - prompt_attention_mask = text_inputs.attention_mask.to(device) - if text_input_ids.shape[-1] == 0: - text_input_ids = text_input_ids.new_zeros((text_input_ids.shape[0], 1)) - prompt_attention_mask = prompt_attention_mask.new_zeros((prompt_attention_mask.shape[0], 1)) - - prompt_embeds = components.text_encoder( - input_ids=text_input_ids, - attention_mask=prompt_attention_mask, - output_hidden_states=False, - ).last_hidden_state - prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) - prompt_embeds = prompt_embeds * prompt_attention_mask.to(prompt_embeds).unsqueeze(-1) - - return prompt_embeds, prompt_attention_mask - - @staticmethod - def _get_t5_prompt_ids( - components: AnimaModularPipeline, - prompt: str | list[str], - max_sequence_length: int, - device: torch.device, - ) -> tuple[torch.Tensor, torch.Tensor]: - prompt = [prompt] if isinstance(prompt, str) else prompt - - text_inputs = components.t5_tokenizer( - prompt, - padding="longest", - max_length=max_sequence_length, - truncation=True, - return_tensors="pt", - ) - return text_inputs.input_ids.to(device), text_inputs.attention_mask.to(device) - - @classmethod - def encode_prompt( - cls, - components: AnimaModularPipeline, - prompt: str | list[str], - negative_prompt: str | list[str] | None = None, - prepare_unconditional_embeds: bool = True, - max_sequence_length: int = 512, - device: torch.device | None = None, - dtype: torch.dtype | None = None, - ) -> dict[str, torch.Tensor | None]: - device = device or components._execution_device - dtype = dtype or components.text_encoder.dtype - - prompt = [prompt] if isinstance(prompt, str) else prompt - batch_size = len(prompt) - - prompt_embeds, prompt_attention_mask = cls._get_qwen_prompt_embeds( - components=components, - prompt=prompt, - max_sequence_length=max_sequence_length, - device=device, - dtype=dtype, - ) - t5_input_ids, t5_attention_mask = cls._get_t5_prompt_ids( - components=components, - prompt=prompt, - max_sequence_length=max_sequence_length, - device=device, - ) - - negative_prompt_embeds = None - negative_prompt_attention_mask = None - negative_t5_input_ids = None - negative_t5_attention_mask = None - if prepare_unconditional_embeds: - negative_prompt = negative_prompt if negative_prompt is not None else "" - negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt - - if prompt is not None and type(prompt) is not type(negative_prompt): - raise TypeError( - f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" - f" {type(prompt)}." - ) - if batch_size != len(negative_prompt): - raise ValueError( - f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" - f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches" - " the batch size of `prompt`." - ) - - negative_prompt_embeds, negative_prompt_attention_mask = cls._get_qwen_prompt_embeds( - components=components, - prompt=negative_prompt, - max_sequence_length=max_sequence_length, - device=device, - dtype=dtype, - ) - negative_t5_input_ids, negative_t5_attention_mask = cls._get_t5_prompt_ids( - components=components, - prompt=negative_prompt, - max_sequence_length=max_sequence_length, - device=device, - ) - - return { - "qwen_prompt_embeds": prompt_embeds, - "qwen_attention_mask": prompt_attention_mask, - "t5_input_ids": t5_input_ids, - "t5_attention_mask": t5_attention_mask, - "negative_qwen_prompt_embeds": negative_prompt_embeds, - "negative_qwen_attention_mask": negative_prompt_attention_mask, - "negative_t5_input_ids": negative_t5_input_ids, - "negative_t5_attention_mask": negative_t5_attention_mask, - } - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - self.check_inputs(block_state) - - prompt_outputs = self.encode_prompt( - components=components, - prompt=block_state.prompt, - negative_prompt=block_state.negative_prompt, - prepare_unconditional_embeds=components.guider.num_conditions > 1, - max_sequence_length=block_state.max_sequence_length, - device=components._execution_device, - dtype=components.text_encoder.dtype, - ) - for name, value in prompt_outputs.items(): - setattr(block_state, name, value) - - self.set_block_state(state, block_state) - return components, state - - -# Copied from diffusers.modular_pipelines.qwenimage.encoders.retrieve_latents -def retrieve_latents( - encoder_output: torch.Tensor, generator: torch.Generator | None = None, sample_mode: str = "sample" -): - if hasattr(encoder_output, "latent_dist") and sample_mode == "sample": - return encoder_output.latent_dist.sample(generator) - elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax": - return encoder_output.latent_dist.mode() - elif hasattr(encoder_output, "latents"): - return encoder_output.latents - else: - raise AttributeError("Could not access latents of provided encoder_output") - - -# Copied from diffusers.modular_pipelines.qwenimage.encoders.encode_vae_image -def encode_vae_image( - image: torch.Tensor, - vae: AutoencoderKLQwenImage, - generator: torch.Generator, - device: torch.device, - dtype: torch.dtype, - latent_channels: int = 16, - sample_mode: str = "argmax", -): - if not isinstance(image, torch.Tensor): - raise ValueError(f"Expected image to be a tensor, got {type(image)}.") - - # preprocessed image should be a 4D tensor: batch_size, num_channels, height, width - if image.dim() == 4: - image = image.unsqueeze(2) - elif image.dim() != 5: - raise ValueError(f"Expected image dims 4 or 5, got {image.dim()}.") - - image = image.to(device=device, dtype=dtype) - - if isinstance(generator, list): - image_latents = [ - retrieve_latents(vae.encode(image[i : i + 1]), generator=generator[i], sample_mode=sample_mode) - for i in range(image.shape[0]) - ] - image_latents = torch.cat(image_latents, dim=0) - else: - image_latents = retrieve_latents(vae.encode(image), generator=generator, sample_mode=sample_mode) - latents_mean = ( - torch.tensor(vae.config.latents_mean) - .view(1, latent_channels, 1, 1, 1) - .to(image_latents.device, image_latents.dtype) - ) - latents_std = ( - torch.tensor(vae.config.latents_std) - .view(1, latent_channels, 1, 1, 1) - .to(image_latents.device, image_latents.dtype) - ) - image_latents = (image_latents - latents_mean) / latents_std - - return image_latents - - -class AnimaImg2ImgVaeEncoderStep(ModularPipelineBlocks): - """VAE Encoder step for Anima image-to-image generation. - - Preprocesses the input image and encodes it with the VAE, producing ``image_latents``. Timestep slicing is handled - downstream by ``AnimaImg2ImgSetTimestepsStep`` and noise addition by ``AnimaImg2ImgPrepareLatentsStep``. - - Components: - vae (`AutoencoderKLQwenImage`) image_processor (`VaeImageProcessor`) - - Inputs: - image (`PIL.Image.Image`): - Input image to encode. - height (`int`, *optional*): - Height of the output image. Defaults to pipeline default. - width (`int`, *optional*): - Width of the output image. Defaults to pipeline default. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - - Outputs: - image_latents (`Tensor`): - Encoded image latents. - height (`int`): - Output image height. - width (`int`): - Output image width. - """ - - model_name = "anima" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLQwenImage), - ComponentSpec( - "image_processor", - VaeImageProcessor, - config=FrozenDict({"vae_scale_factor": 8}), - default_creation_method="from_config", - ), - ] - - @property - def description(self) -> str: - return ( - "VAE Encoder step for Anima image-to-image generation. Encodes the input image to produce image_latents." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("image"), - InputParam.template("height"), - InputParam.template("width"), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("image_latents", type_hint=torch.Tensor, description="Encoded image latents."), - OutputParam("height", type_hint=int, description="Image height used for generation."), - OutputParam("width", type_hint=int, description="Image width used for generation."), - ] - - @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - - block_state.height = block_state.height or components.default_height - block_state.width = block_state.width or components.default_width - - processed_image = components.image_processor.preprocess( - image=block_state.image, height=block_state.height, width=block_state.width - ) - - block_state.image_latents = encode_vae_image( - image=processed_image, - vae=components.vae, - generator=block_state.generator, - device=device, - dtype=components.vae.dtype, - latent_channels=components.num_channels_latents, - ) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/anima/modular_blocks_anima.py b/diffusers/modular_pipelines/anima/modular_blocks_anima.py deleted file mode 100644 index f17538fd258ff902334b86d918345a558f4f5829..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/anima/modular_blocks_anima.py +++ /dev/null @@ -1,381 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from ..modular_pipeline import AutoPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import OutputParam -from .before_denoise import ( - AnimaImageInputStep, - AnimaImg2ImgPrepareLatentsStep, - AnimaImg2ImgSetTimestepsStep, - AnimaPrepareLatentsStep, - AnimaSetTimestepsStep, - AnimaTextConditioningStep, - AnimaTextInputStep, -) -from .decoders import AnimaProcessImagesOutputStep, AnimaVaeDecoderStep -from .denoise import AnimaDenoiseStep -from .encoders import AnimaImg2ImgVaeEncoderStep, AnimaTextEncoderStep - - -# auto_docstring -class AnimaCoreDenoiseStep(SequentialPipelineBlocks): - """ - Denoise block that takes encoded Anima text inputs and runs the denoising process. - - Components: - text_conditioner (`AnimaTextConditioner`) transformer (`CosmosTransformer3DModel`) scheduler - (`FlowMatchEulerDiscreteScheduler`) guider (`ClassifierFreeGuidance`) - - Inputs: - qwen_prompt_embeds (`Tensor`): - Qwen prompt embeddings generated by the text encoder step. - qwen_attention_mask (`Tensor`): - Qwen prompt attention mask generated by the text encoder step. - t5_input_ids (`Tensor`): - T5 prompt token ids generated by the text encoder step. - t5_attention_mask (`Tensor`): - T5 prompt attention mask generated by the text encoder step. - negative_qwen_prompt_embeds (`Tensor`, *optional*): - Negative Qwen prompt embeddings generated by the text encoder step. - negative_qwen_attention_mask (`Tensor`, *optional*): - Negative Qwen prompt attention mask generated by the text encoder step. - negative_t5_input_ids (`Tensor`, *optional*): - Negative T5 prompt token ids generated by the text encoder step. - negative_t5_attention_mask (`Tensor`, *optional*): - Negative T5 prompt attention mask generated by the text encoder step. - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - **denoiser_input_fields (`None`, *optional*): - The conditional model inputs for the Anima denoiser. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - block_classes = [ - AnimaTextConditioningStep, - AnimaTextInputStep, - AnimaPrepareLatentsStep, - AnimaSetTimestepsStep, - AnimaDenoiseStep, - ] - block_names = ["text_conditioning", "input", "prepare_latents", "set_timesteps", "denoise"] - - @property - def description(self) -> str: - return "Denoise block that takes encoded Anima text inputs and runs the denoising process." - - @property - def outputs(self): - return [OutputParam.template("latents")] - - -# auto_docstring -class AnimaDecodeStep(SequentialPipelineBlocks): - """ - Decode Anima latents into generated images. - - Components: - vae (`AutoencoderKLQwenImage`) image_processor (`VaeImageProcessor`) - - Inputs: - latents (`Tensor`): - Denoised Anima latents. - output_type (`str`, *optional*, defaults to pil): - Output format: 'pil', 'np', 'pt'. - - Outputs: - images (`list`): - Generated images. - """ - - block_classes = [AnimaVaeDecoderStep, AnimaProcessImagesOutputStep] - block_names = ["decode", "postprocess"] - - @property - def description(self) -> str: - return "Decode Anima latents into generated images." - - @property - def outputs(self): - return [OutputParam.template("images")] - - -# auto_docstring -class AnimaImg2ImgCoreDenoiseStep(SequentialPipelineBlocks): - """ - Denoise block for Anima image-to-image generation. Uses image_latents already in state from - AnimaImg2ImgVaeEncoderStep. - - Components: - text_conditioner (`AnimaTextConditioner`) transformer (`CosmosTransformer3DModel`) scheduler - (`FlowMatchEulerDiscreteScheduler`) guider (`ClassifierFreeGuidance`) - - Inputs: - qwen_prompt_embeds (`Tensor`): - Qwen prompt embeddings generated by the text encoder step. - qwen_attention_mask (`Tensor`): - Qwen prompt attention mask generated by the text encoder step. - t5_input_ids (`Tensor`): - T5 prompt token ids generated by the text encoder step. - t5_attention_mask (`Tensor`): - T5 prompt attention mask generated by the text encoder step. - negative_qwen_prompt_embeds (`Tensor`, *optional*): - Negative Qwen prompt embeddings generated by the text encoder step. - negative_qwen_attention_mask (`Tensor`, *optional*): - Negative Qwen prompt attention mask generated by the text encoder step. - negative_t5_input_ids (`Tensor`, *optional*): - Negative T5 prompt token ids generated by the text encoder step. - negative_t5_attention_mask (`Tensor`, *optional*): - Negative T5 prompt attention mask generated by the text encoder step. - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - image_latents (`Tensor`): - image latents used to guide the image generation. Can be generated from vae_encoder step. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - strength (`float`, *optional*, defaults to 0.9): - Strength for img2img/inpainting. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - The conditional model inputs for the Anima denoiser. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - block_classes = [ - AnimaTextConditioningStep, - AnimaTextInputStep, - AnimaImageInputStep, - AnimaImg2ImgSetTimestepsStep, - AnimaImg2ImgPrepareLatentsStep, - AnimaDenoiseStep, - ] - block_names = ["text_conditioning", "input", "image_input", "set_timesteps", "prepare_latents", "denoise"] - - @property - def description(self) -> str: - return ( - "Denoise block for Anima image-to-image generation. " - "Uses image_latents already in state from AnimaImg2ImgVaeEncoderStep." - ) - - @property - def outputs(self): - return [OutputParam.template("latents")] - - -# auto_docstring -class AnimaAutoCoreDenoiseStep(AutoPipelineBlocks): - """ - Denoise step that selects between text-to-image and image-to-image denoising based on whether image_latents is - present in state. - `AnimaCoreDenoiseStep` (text2image) is used when no image_latents are present. - - `AnimaImg2ImgCoreDenoiseStep` (img2img) is used when image_latents are present. - - Components: - text_conditioner (`AnimaTextConditioner`) transformer (`CosmosTransformer3DModel`) scheduler - (`FlowMatchEulerDiscreteScheduler`) guider (`ClassifierFreeGuidance`) - - Inputs: - qwen_prompt_embeds (`Tensor`): - Qwen prompt embeddings generated by the text encoder step. - qwen_attention_mask (`Tensor`): - Qwen prompt attention mask generated by the text encoder step. - t5_input_ids (`Tensor`): - T5 prompt token ids generated by the text encoder step. - t5_attention_mask (`Tensor`): - T5 prompt attention mask generated by the text encoder step. - negative_qwen_prompt_embeds (`Tensor`, *optional*): - Negative Qwen prompt embeddings generated by the text encoder step. - negative_qwen_attention_mask (`Tensor`, *optional*): - Negative Qwen prompt attention mask generated by the text encoder step. - negative_t5_input_ids (`Tensor`, *optional*): - Negative T5 prompt token ids generated by the text encoder step. - negative_t5_attention_mask (`Tensor`, *optional*): - Negative T5 prompt attention mask generated by the text encoder step. - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - image_latents (`Tensor`, *optional*): - image latents used to guide the image generation. Can be generated from vae_encoder step. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - num_inference_steps (`int`): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - strength (`float`, *optional*, defaults to 0.9): - Strength for img2img/inpainting. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - latents (`Tensor`): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - The conditional model inputs for the Anima denoiser. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - block_classes = [AnimaImg2ImgCoreDenoiseStep, AnimaCoreDenoiseStep] - block_names = ["img2img", "text2image"] - block_trigger_inputs = ["image_latents", None] - - @property - def description(self) -> str: - return ( - "Denoise step that selects between text-to-image and image-to-image denoising based on whether " - "image_latents is present in state." - " - `AnimaCoreDenoiseStep` (text2image) is used when no image_latents are present." - " - `AnimaImg2ImgCoreDenoiseStep` (img2img) is used when image_latents are present." - ) - - -# auto_docstring -class AnimaAutoVaeImageEncoderStep(AutoPipelineBlocks): - """ - VAE Image Encoder step that encodes the input image to produce image_latents. Skipped when no image is provided - (text-to-image workflow). - - Components: - vae (`AutoencoderKLQwenImage`) image_processor (`VaeImageProcessor`) - - Inputs: - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - - Outputs: - image_latents (`Tensor`): - Encoded image latents. - height (`int`): - Image height used for generation. - width (`int`): - Image width used for generation. - """ - - block_classes = [AnimaImg2ImgVaeEncoderStep] - block_names = ["vae_encoder"] - block_trigger_inputs = ["image"] - - @property - def description(self) -> str: - return ( - "VAE Image Encoder step that encodes the input image to produce image_latents. " - "Skipped when no image is provided (text-to-image workflow)." - ) - - -# auto_docstring -class AnimaAutoBlocks(SequentialPipelineBlocks): - """ - Auto Modular pipeline for text-to-image and image-to-image generation using Anima. - - Supported workflows: - - `text2image`: requires `prompt` - - `img2img`: requires `image`, `prompt` - - Components: - text_encoder (`Qwen3Model`) tokenizer (`Qwen2Tokenizer`) t5_tokenizer (`T5Tokenizer`) guider - (`ClassifierFreeGuidance`) vae (`AutoencoderKLQwenImage`) image_processor (`VaeImageProcessor`) - text_conditioner (`AnimaTextConditioner`) transformer (`CosmosTransformer3DModel`) scheduler - (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The prompt or prompts not to guide the image generation. - max_sequence_length (`int`, *optional*, defaults to 512): - Maximum sequence length for prompt encoding. - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - image_latents (`Tensor`, *optional*): - image latents used to guide the image generation. Can be generated from vae_encoder step. - num_inference_steps (`int`): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - strength (`float`, *optional*, defaults to 0.9): - Strength for img2img/inpainting. - latents (`Tensor`): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - The conditional model inputs for the Anima denoiser. - output_type (`str`, *optional*, defaults to pil): - Output format: 'pil', 'np', 'pt'. - - Outputs: - images (`list`): - Generated images. - """ - - block_classes = [ - AnimaTextEncoderStep, - AnimaAutoVaeImageEncoderStep, - AnimaAutoCoreDenoiseStep, - AnimaDecodeStep, - ] - block_names = ["text_encoder", "vae_encoder", "denoise", "decode"] - _workflow_map = { - "text2image": {"prompt": True}, - "img2img": {"image": True, "prompt": True}, - } - - @property - def description(self) -> str: - return "Auto Modular pipeline for text-to-image and image-to-image generation using Anima." - - @property - def outputs(self): - return [OutputParam.template("images")] diff --git a/diffusers/modular_pipelines/anima/modular_pipeline.py b/diffusers/modular_pipelines/anima/modular_pipeline.py deleted file mode 100644 index 44fce4657c6f3b7358fa6124f718dc0d3750e8db..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/anima/modular_pipeline.py +++ /dev/null @@ -1,52 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from ...loaders import AnimaLoraLoaderMixin -from ..modular_pipeline import ModularPipeline - - -class AnimaModularPipeline(ModularPipeline, AnimaLoraLoaderMixin): - """ - A ModularPipeline for Anima. - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "AnimaAutoBlocks" - - @property - def default_height(self): - return self.default_sample_size * self.vae_scale_factor - - @property - def default_width(self): - return self.default_sample_size * self.vae_scale_factor - - @property - def default_sample_size(self): - return 128 - - @property - def vae_scale_factor(self): - vae_scale_factor = 8 - if self.vae is not None: - vae_scale_factor = 2 ** len(self.vae.temperal_downsample) - return vae_scale_factor - - @property - def num_channels_latents(self): - num_channels_latents = 16 - if self.transformer is not None: - num_channels_latents = self.transformer.config.in_channels - return num_channels_latents diff --git a/diffusers/modular_pipelines/components_manager.py b/diffusers/modular_pipelines/components_manager.py deleted file mode 100644 index 31ba2c9422032369cc2847137d8f8de3b147baf1..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/components_manager.py +++ /dev/null @@ -1,1109 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -import copy -import time -from collections import OrderedDict -from itertools import combinations -from typing import Any - -import torch - -from ..hooks import ModelHook -from ..utils import ( - is_accelerate_available, - logging, -) -from ..utils.torch_utils import get_device - - -if is_accelerate_available(): - from accelerate.hooks import add_hook_to_module, remove_hook_from_module - from accelerate.state import PartialState - from accelerate.utils import send_to_device - from accelerate.utils.memory import clear_device_cache - from accelerate.utils.modeling import convert_file_size_to_int - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class CustomOffloadHook(ModelHook): - """ - A hook that offloads a model on the CPU until its forward pass is called. It ensures the model and its inputs are - on the given device. Optionally offloads other models to the CPU before the forward pass is called. - - Args: - execution_device(`str`, `int` or `torch.device`, *optional*): - The device on which the model should be executed. Will default to the MPS device if it's available, then - GPU 0 if there is a GPU, and finally to the CPU. - """ - - no_grad = False - - def __init__( - self, - execution_device: str | int | torch.device | None = None, - other_hooks: list["UserCustomOffloadHook"] | None = None, - offload_strategy: "AutoOffloadStrategy" | None = None, - ): - self.execution_device = execution_device if execution_device is not None else PartialState().default_device - self.other_hooks = other_hooks - self.offload_strategy = offload_strategy - self.model_id = None - - def set_strategy(self, offload_strategy: "AutoOffloadStrategy"): - self.offload_strategy = offload_strategy - - def add_other_hook(self, hook: "UserCustomOffloadHook"): - """ - Add a hook to the list of hooks to consider for offloading. - """ - if self.other_hooks is None: - self.other_hooks = [] - self.other_hooks.append(hook) - - def init_hook(self, module): - return module.to("cpu") - - def pre_forward(self, module, *args, **kwargs): - if module.device != self.execution_device: - if self.other_hooks is not None: - hooks_to_offload = [hook for hook in self.other_hooks if hook.model.device == self.execution_device] - # offload all other hooks - start_time = time.perf_counter() - if self.offload_strategy is not None: - hooks_to_offload = self.offload_strategy( - hooks=hooks_to_offload, - model_id=self.model_id, - model=module, - execution_device=self.execution_device, - ) - end_time = time.perf_counter() - logger.info( - f" time taken to apply offload strategy for {self.model_id}: {(end_time - start_time):.2f} seconds" - ) - - for hook in hooks_to_offload: - logger.info( - f"moving {self.model_id} to {self.execution_device}, offloading {hook.model_id} to cpu" - ) - hook.offload() - - if hooks_to_offload: - clear_device_cache() - module.to(self.execution_device) - return send_to_device(args, self.execution_device), send_to_device(kwargs, self.execution_device) - - -class UserCustomOffloadHook: - """ - A simple hook grouping a model and a `CustomOffloadHook`, which provides easy APIs for to call the init method of - the hook or remove it entirely. - """ - - def __init__(self, model_id, model, hook): - self.model_id = model_id - self.model = model - self.hook = hook - - def offload(self): - self.hook.init_hook(self.model) - - def attach(self): - add_hook_to_module(self.model, self.hook) - self.hook.model_id = self.model_id - - def remove(self): - remove_hook_from_module(self.model) - self.hook.model_id = None - - def add_other_hook(self, hook: "UserCustomOffloadHook"): - self.hook.add_other_hook(hook) - - -def custom_offload_with_hook( - model_id: str, - model: torch.nn.Module, - execution_device: str | int | torch.device = None, - offload_strategy: "AutoOffloadStrategy" | None = None, -): - hook = CustomOffloadHook(execution_device=execution_device, offload_strategy=offload_strategy) - user_hook = UserCustomOffloadHook(model_id=model_id, model=model, hook=hook) - user_hook.attach() - return user_hook - - -# this is the class that user can customize to implement their own offload strategy -class AutoOffloadStrategy: - """ - Offload strategy that should be used with `CustomOffloadHook` to automatically offload models to the CPU based on - the available memory on the device. - """ - - # YiYi TODO: instead of memory_reserve_margin, we should let user set the maximum_total_models_size to keep on device - # the actual memory usage would be higher. But it's simpler this way, and can be tested - def __init__(self, memory_reserve_margin="3GB"): - self.memory_reserve_margin = convert_file_size_to_int(memory_reserve_margin) - - def __call__(self, hooks, model_id, model, execution_device): - if len(hooks) == 0: - return [] - - try: - current_module_size = model.get_memory_footprint() - except AttributeError: - raise AttributeError(f"Do not know how to compute memory footprint of `{model.__class__.__name__}.") - - device_type = execution_device.type - device_module = getattr(torch, device_type, torch.cuda) - try: - mem_on_device = device_module.mem_get_info(execution_device.index)[0] - except AttributeError: - raise AttributeError(f"Do not know how to obtain obtain memory info for {str(device_module)}.") - - mem_on_device = mem_on_device - self.memory_reserve_margin - if current_module_size < mem_on_device: - return [] - - min_memory_offload = current_module_size - mem_on_device - logger.info(f" search for models to offload in order to free up {min_memory_offload / 1024**3:.2f} GB memory") - - # exlucde models that's not currently loaded on the device - module_sizes = dict( - sorted( - {hook.model_id: hook.model.get_memory_footprint() for hook in hooks}.items(), - key=lambda x: x[1], - reverse=True, - ) - ) - - # YiYi/Dhruv TODO: sort smallest to largest, and offload in that order we would tend to keep the larger models on GPU more often - def search_best_candidate(module_sizes, min_memory_offload): - """ - search the optimal combination of models to offload to cpu, given a dictionary of module sizes and a - minimum memory offload size. the combination of models should add up to the smallest modulesize that is - larger than `min_memory_offload` - """ - model_ids = list(module_sizes.keys()) - best_candidate = None - best_size = float("inf") - for r in range(1, len(model_ids) + 1): - for candidate_model_ids in combinations(model_ids, r): - candidate_size = sum( - module_sizes[candidate_model_id] for candidate_model_id in candidate_model_ids - ) - if candidate_size < min_memory_offload: - continue - else: - if best_candidate is None or candidate_size < best_size: - best_candidate = candidate_model_ids - best_size = candidate_size - - return best_candidate - - best_offload_model_ids = search_best_candidate(module_sizes, min_memory_offload) - - if best_offload_model_ids is None: - # if no combination is found, meaning that we cannot meet the memory requirement, offload all models - logger.warning("no combination of models to offload to cpu is found, offloading all models") - hooks_to_offload = hooks - else: - hooks_to_offload = [hook for hook in hooks if hook.model_id in best_offload_model_ids] - - return hooks_to_offload - - -# utils for display component info in a readable format -# TODO: move to a different file -def summarize_dict_by_value_and_parts(d: dict[str, Any]) -> dict[str, Any]: - """Summarizes a dictionary by finding common prefixes that share the same value. - - For a dictionary with dot-separated keys like: { - 'down_blocks.1.attentions.1.transformer_blocks.0.attn2.processor': [0.6], - 'down_blocks.1.attentions.1.transformer_blocks.1.attn2.processor': [0.6], - 'up_blocks.1.attentions.0.transformer_blocks.0.attn2.processor': [0.3], - } - - Returns a dictionary where keys are the shortest common prefixes and values are their shared values: { - 'down_blocks': [0.6], 'up_blocks': [0.3] - } - """ - # First group by values - convert lists to tuples to make them hashable - value_to_keys = {} - for key, value in d.items(): - value_tuple = tuple(value) if isinstance(value, list) else value - if value_tuple not in value_to_keys: - value_to_keys[value_tuple] = [] - value_to_keys[value_tuple].append(key) - - def find_common_prefix(keys: list[str]) -> str: - """Find the shortest common prefix among a list of dot-separated keys.""" - if not keys: - return "" - if len(keys) == 1: - return keys[0] - - # Split all keys into parts - key_parts = [k.split(".") for k in keys] - - # Find how many initial parts are common - common_length = 0 - for parts in zip(*key_parts): - if len(set(parts)) == 1: # All parts at this position are the same - common_length += 1 - else: - break - - if common_length == 0: - return "" - - # Return the common prefix - return ".".join(key_parts[0][:common_length]) - - # Create summary by finding common prefixes for each value group - summary = {} - for value_tuple, keys in value_to_keys.items(): - prefix = find_common_prefix(keys) - if prefix: # Only add if we found a common prefix - # Convert tuple back to list if it was originally a list - value = list(value_tuple) if isinstance(d[keys[0]], list) else value_tuple - summary[prefix] = value - else: - summary[""] = value # Use empty string if no common prefix - - return summary - - -class ComponentsManager: - """ - A central registry and management system for model components across multiple pipelines. - - [`ComponentsManager`] provides a unified way to register, track, and reuse model components (like UNet, VAE, text - encoders, etc.) across different modular pipelines. It includes features for duplicate detection, memory - management, and component organization. - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - - Example: - ```python - from diffusers import ComponentsManager - - # Create a components manager - cm = ComponentsManager() - - # Add components - cm.add("unet", unet_model, collection="sdxl") - cm.add("vae", vae_model, collection="sdxl") - - # Enable auto offloading - cm.enable_auto_cpu_offload() - - # Retrieve components - unet = cm.get_one(name="unet", collection="sdxl") - ``` - """ - - _available_info_fields = [ - "model_id", - "added_time", - "collection", - "class_name", - "size_gb", - "adapters", - "has_hook", - "execution_device", - "ip_adapter", - "quantization", - ] - - def __init__(self): - self.components = OrderedDict() - # YiYi TODO: can remove once confirm we don't need this in mellon - self.added_time = OrderedDict() # Store when components were added - self.collections = OrderedDict() # collection_name -> set of component_names - self.model_hooks = None - self._auto_offload_enabled = False - - def _lookup_ids( - self, - name: str | None = None, - collection: str | None = None, - load_id: str | None = None, - components: OrderedDict | None = None, - ): - """ - Lookup component_ids by name, collection, or load_id. Does not support pattern matching. Returns a set of - component_ids - """ - if components is None: - components = self.components - - if name: - ids_by_name = set() - for component_id, component in components.items(): - comp_name = self._id_to_name(component_id) - if comp_name == name: - ids_by_name.add(component_id) - else: - ids_by_name = set(components.keys()) - if collection and collection not in self.collections: - return set() - elif collection and collection in self.collections: - ids_by_collection = set() - for component_id, component in components.items(): - if component_id in self.collections[collection]: - ids_by_collection.add(component_id) - else: - ids_by_collection = set(components.keys()) - if load_id: - ids_by_load_id = set() - for name, component in components.items(): - if hasattr(component, "_diffusers_load_id") and component._diffusers_load_id == load_id: - ids_by_load_id.add(name) - else: - ids_by_load_id = set(components.keys()) - - ids = ids_by_name.intersection(ids_by_collection).intersection(ids_by_load_id) - return ids - - @staticmethod - def _id_to_name(component_id: str): - return "_".join(component_id.split("_")[:-1]) - - def add(self, name: str, component: Any, collection: str | None = None): - """ - Add a component to the ComponentsManager. - - Args: - name (str): The name of the component - component (Any): The component to add - collection (str | None): The collection to add the component to - - Returns: - str: The unique component ID, which is generated as "{name}_{id(component)}" where - id(component) is Python's built-in unique identifier for the object - """ - component_id = f"{name}_{id(component)}" - is_new_component = True - - # check for duplicated components - for comp_id, comp in self.components.items(): - if comp == component: - comp_name = self._id_to_name(comp_id) - if comp_name == name: - logger.warning(f"ComponentsManager: component '{name}' already exists as '{comp_id}'") - component_id = comp_id - is_new_component = False - break - else: - logger.warning( - f"ComponentsManager: adding component '{name}' as '{component_id}', but it is duplicate of '{comp_id}'" - f"To remove a duplicate, call `components_manager.remove('')`." - ) - - # check for duplicated load_id and warn (we do not delete for you) - if hasattr(component, "_diffusers_load_id") and component._diffusers_load_id != "null": - components_with_same_load_id = self._lookup_ids(load_id=component._diffusers_load_id) - components_with_same_load_id = [id for id in components_with_same_load_id if id != component_id] - - if components_with_same_load_id: - existing = ", ".join(components_with_same_load_id) - logger.warning( - f"ComponentsManager: adding component '{component_id}', but it has duplicate load_id '{component._diffusers_load_id}' with existing components: {existing}. " - f"To remove a duplicate, call `components_manager.remove('')`." - ) - - # add component to components manager - self.components[component_id] = component - if is_new_component: - self.added_time[component_id] = time.time() - - if collection: - if collection not in self.collections: - self.collections[collection] = set() - if component_id not in self.collections[collection]: - comp_ids_in_collection = self._lookup_ids(name=name, collection=collection) - for comp_id in comp_ids_in_collection: - logger.warning( - f"ComponentsManager: removing existing {name} from collection '{collection}': {comp_id}" - ) - # remove existing component from this collection (if it is not in any other collection, will be removed from ComponentsManager) - self.remove_from_collection(comp_id, collection) - - self.collections[collection].add(component_id) - logger.info( - f"ComponentsManager: added component '{name}' in collection '{collection}': {component_id}" - ) - else: - logger.info(f"ComponentsManager: added component '{name}' as '{component_id}'") - - if self._auto_offload_enabled and is_new_component: - self.enable_auto_cpu_offload(self._auto_offload_device) - - return component_id - - def remove_from_collection(self, component_id: str, collection: str): - """ - Remove a component from a collection. - """ - if collection not in self.collections: - logger.warning(f"Collection '{collection}' not found in ComponentsManager") - return - if component_id not in self.collections[collection]: - logger.warning(f"Component '{component_id}' not found in collection '{collection}'") - return - # remove from the collection - self.collections[collection].remove(component_id) - # check if this component is in any other collection - comp_colls = [coll for coll, comps in self.collections.items() if component_id in comps] - if not comp_colls: # only if no other collection contains this component, remove it - logger.warning(f"ComponentsManager: removing component '{component_id}' from ComponentsManager") - self.remove(component_id) - - def remove(self, component_id: str = None): - """ - Remove a component from the ComponentsManager. - - Args: - component_id (str): The ID of the component to remove - """ - if component_id not in self.components: - logger.warning(f"Component '{component_id}' not found in ComponentsManager") - return - - component = self.components.pop(component_id) - self.added_time.pop(component_id) - - for collection in self.collections: - if component_id in self.collections[collection]: - self.collections[collection].remove(component_id) - - if self._auto_offload_enabled: - self.enable_auto_cpu_offload(self._auto_offload_device) - else: - if isinstance(component, torch.nn.Module): - component.to("cpu") - del component - import gc - - gc.collect() - if torch.cuda.is_available(): - torch.cuda.empty_cache() - if torch.xpu.is_available(): - torch.xpu.empty_cache() - - # YiYi TODO: rename to search_components for now, may remove this method - def search_components( - self, - names: str | None = None, - collection: str | None = None, - load_id: str | None = None, - return_dict_with_names: bool = True, - ): - """ - Search components by name with simple pattern matching. Optionally filter by collection or load_id. - - Args: - names: Component name(s) or pattern(s) - Patterns: - - "unet" : match any component with base name "unet" (e.g., unet_123abc) - - "!unet" : everything except components with base name "unet" - - "unet*" : anything with base name starting with "unet" - - "!unet*" : anything with base name NOT starting with "unet" - - "*unet*" : anything with base name containing "unet" - - "!*unet*" : anything with base name NOT containing "unet" - - "refiner|vae|unet" : anything with base name exactly matching "refiner", "vae", or "unet" - - "!refiner|vae|unet" : anything with base name NOT exactly matching "refiner", "vae", or "unet" - - "unet*|vae*" : anything with base name starting with "unet" OR starting with "vae" - collection: Optional collection to filter by - load_id: Optional load_id to filter by - return_dict_with_names: - If True, returns a dictionary with component names as keys, throw an error if - multiple components with the same name are found If False, returns a dictionary - with component IDs as keys - - Returns: - Dictionary mapping component names to components if return_dict_with_names=True, or a dictionary mapping - component IDs to components if return_dict_with_names=False - """ - - # select components based on collection and load_id filters - selected_ids = self._lookup_ids(collection=collection, load_id=load_id) - components = {k: self.components[k] for k in selected_ids} - - def get_return_dict(components, return_dict_with_names): - """ - Create a dictionary mapping component names to components if return_dict_with_names=True, or a dictionary - mapping component IDs to components if return_dict_with_names=False, throw an error if duplicate component - names are found when return_dict_with_names=True - """ - if return_dict_with_names: - dict_to_return = {} - for comp_id, comp in components.items(): - comp_name = self._id_to_name(comp_id) - if comp_name in dict_to_return: - raise ValueError( - f"Duplicate component names found in the search results: {comp_name}, please set `return_dict_with_names=False` to return a dictionary with component IDs as keys" - ) - dict_to_return[comp_name] = comp - return dict_to_return - else: - return components - - # if no names are provided, return the filtered components as it is - if names is None: - return get_return_dict(components, return_dict_with_names) - - # if names is not a string, raise an error - elif not isinstance(names, str): - raise ValueError(f"Invalid type for `names: {type(names)}, only support string") - - # Create mapping from component_id to base_name for components to be used for pattern matching - base_names = {comp_id: self._id_to_name(comp_id) for comp_id in components.keys()} - - # Helper function to check if a component matches a pattern based on its base name - def matches_pattern(component_id, pattern, exact_match=False): - """ - Helper function to check if a component matches a pattern based on its base name. - - Args: - component_id: The component ID to check - pattern: The pattern to match against - exact_match: If True, only exact matches to base_name are considered - """ - base_name = base_names[component_id] - - # Exact match with base name - if exact_match: - return pattern == base_name - - # Prefix match (ends with *) - elif pattern.endswith("*"): - prefix = pattern[:-1] - return base_name.startswith(prefix) - - # Contains match (starts with *) - elif pattern.startswith("*"): - search = pattern[1:-1] if pattern.endswith("*") else pattern[1:] - return search in base_name - - # Exact match (no wildcards) - else: - return pattern == base_name - - # Check if this is a "not" pattern - is_not_pattern = names.startswith("!") - if is_not_pattern: - names = names[1:] # Remove the ! prefix - - # Handle OR patterns (containing |) - if "|" in names: - terms = names.split("|") - matches = {} - - for comp_id, comp in components.items(): - # For OR patterns with exact names (no wildcards), we do exact matching on base names - exact_match = all(not (term.startswith("*") or term.endswith("*")) for term in terms) - - # Check if any of the terms match this component - should_include = any(matches_pattern(comp_id, term, exact_match) for term in terms) - - # Flip the decision if this is a NOT pattern - if is_not_pattern: - should_include = not should_include - - if should_include: - matches[comp_id] = comp - - log_msg = "NOT " if is_not_pattern else "" - match_type = "exactly matching" if exact_match else "matching any of patterns" - logger.info(f"Getting components {log_msg}{match_type} {terms}: {list(matches.keys())}") - - # Try exact match with a base name - elif any(names == base_name for base_name in base_names.values()): - # Find all components with this base name - matches = { - comp_id: comp - for comp_id, comp in components.items() - if (base_names[comp_id] == names) != is_not_pattern - } - - if is_not_pattern: - logger.info(f"Getting all components except those with base name '{names}': {list(matches.keys())}") - else: - logger.info(f"Getting components with base name '{names}': {list(matches.keys())}") - - # Prefix match (ends with *) - elif names.endswith("*"): - prefix = names[:-1] - matches = { - comp_id: comp - for comp_id, comp in components.items() - if base_names[comp_id].startswith(prefix) != is_not_pattern - } - if is_not_pattern: - logger.info(f"Getting components NOT starting with '{prefix}': {list(matches.keys())}") - else: - logger.info(f"Getting components starting with '{prefix}': {list(matches.keys())}") - - # Contains match (starts with *) - elif names.startswith("*"): - search = names[1:-1] if names.endswith("*") else names[1:] - matches = { - comp_id: comp - for comp_id, comp in components.items() - if (search in base_names[comp_id]) != is_not_pattern - } - if is_not_pattern: - logger.info(f"Getting components NOT containing '{search}': {list(matches.keys())}") - else: - logger.info(f"Getting components containing '{search}': {list(matches.keys())}") - - # Substring match (no wildcards, but not an exact component name) - elif any(names in base_name for base_name in base_names.values()): - matches = { - comp_id: comp - for comp_id, comp in components.items() - if (names in base_names[comp_id]) != is_not_pattern - } - if is_not_pattern: - logger.info(f"Getting components NOT containing '{names}': {list(matches.keys())}") - else: - logger.info(f"Getting components containing '{names}': {list(matches.keys())}") - - else: - raise ValueError(f"Component or pattern '{names}' not found in ComponentsManager") - - if not matches: - raise ValueError(f"No components found matching pattern '{names}'") - - return get_return_dict(matches, return_dict_with_names) - - def enable_auto_cpu_offload(self, device: str | int | torch.device = None, memory_reserve_margin="3GB"): - """ - Enable automatic CPU offloading for all components. - - The algorithm works as follows: - 1. All models start on CPU by default - 2. When a model's forward pass is called, it's moved to the execution device - 3. If there's insufficient memory, other models on the device are moved back to CPU - 4. The system tries to offload the smallest combination of models that frees enough memory - 5. Models stay on the execution device until another model needs memory and forces them off - - Args: - device (str | int | torch.device): The execution device where models are moved for forward passes - memory_reserve_margin (str): The memory reserve margin to use, default is 3GB. This is the amount of - memory to keep free on the device to avoid running out of memory during model - execution (e.g., for intermediate activations, gradients, etc.) - """ - if not is_accelerate_available(): - raise ImportError("Make sure to install accelerate to use auto_cpu_offload") - - if device is None: - device = get_device() - if not isinstance(device, torch.device): - device = torch.device(device) - - device_type = device.type - device_module = getattr(torch, device_type, torch.cuda) - if not hasattr(device_module, "mem_get_info"): - raise NotImplementedError( - f"`enable_auto_cpu_offload() relies on the `mem_get_info()` method. It's not implemented for {str(device.type)}." - ) - - if device.index is None: - device = torch.device(f"{device.type}:{0}") - - for name, component in self.components.items(): - if isinstance(component, torch.nn.Module) and hasattr(component, "_hf_hook"): - remove_hook_from_module(component, recurse=True) - - self.disable_auto_cpu_offload() - offload_strategy = AutoOffloadStrategy(memory_reserve_margin=memory_reserve_margin) - - all_hooks = [] - for name, component in self.components.items(): - if isinstance(component, torch.nn.Module): - hook = custom_offload_with_hook(name, component, device, offload_strategy=offload_strategy) - all_hooks.append(hook) - - for hook in all_hooks: - other_hooks = [h for h in all_hooks if h is not hook] - for other_hook in other_hooks: - if other_hook.hook.execution_device == hook.hook.execution_device: - hook.add_other_hook(other_hook) - - self.model_hooks = all_hooks - self._auto_offload_enabled = True - self._auto_offload_device = device - - def disable_auto_cpu_offload(self): - """ - Disable automatic CPU offloading for all components. - """ - if self.model_hooks is None: - self._auto_offload_enabled = False - return - - for hook in self.model_hooks: - hook.offload() - hook.remove() - if self.model_hooks: - clear_device_cache() - self.model_hooks = None - self._auto_offload_enabled = False - - def get_model_info( - self, - component_id: str, - fields: str | list[str] | None = None, - ) -> dict[str, Any] | None: - """Get comprehensive information about a component. - - Args: - component_id (str): Name of the component to get info for - fields (str | list[str] | None): - Field(s) to return. Can be a string for single field or list of fields. If None, uses the - available_info_fields setting. - - Returns: - Dictionary containing requested component metadata. If fields is specified, returns only those fields. - Otherwise, returns all fields. - """ - if component_id not in self.components: - raise ValueError(f"Component '{component_id}' not found in ComponentsManager") - - component = self.components[component_id] - - # Validate fields if specified - if fields is not None: - if isinstance(fields, str): - fields = [fields] - for field in fields: - if field not in self._available_info_fields: - raise ValueError(f"Field '{field}' not found in available_info_fields") - - # Build complete info dict first - info = { - "model_id": component_id, - "added_time": self.added_time[component_id], - "collection": ", ".join([coll for coll, comps in self.collections.items() if component_id in comps]) - or None, - } - - # Additional info for torch.nn.Module components - if isinstance(component, torch.nn.Module): - # Check for hook information - has_hook = hasattr(component, "_hf_hook") - execution_device = None - if has_hook and hasattr(component._hf_hook, "execution_device"): - execution_device = component._hf_hook.execution_device - - info.update( - { - "class_name": component.__class__.__name__, - "size_gb": component.get_memory_footprint() / (1024**3), - "adapters": None, # Default to None - "has_hook": has_hook, - "execution_device": execution_device, - } - ) - - # Get adapters if applicable - if hasattr(component, "peft_config"): - info["adapters"] = list(component.peft_config.keys()) - - # Check for IP-Adapter scales - if hasattr(component, "_load_ip_adapter_weights") and hasattr(component, "attn_processors"): - processors = copy.deepcopy(component.attn_processors) - # First check if any processor is an IP-Adapter - processor_types = [v.__class__.__name__ for v in processors.values()] - if any("IPAdapter" in ptype for ptype in processor_types): - # Then get scales only from IP-Adapter processors - scales = { - k: v.scale - for k, v in processors.items() - if hasattr(v, "scale") and "IPAdapter" in v.__class__.__name__ - } - if scales: - info["ip_adapter"] = summarize_dict_by_value_and_parts(scales) - - # Check for quantization - hf_quantizer = getattr(component, "hf_quantizer", None) - if hf_quantizer is not None: - quant_config = hf_quantizer.quantization_config - if hasattr(quant_config, "to_diff_dict"): - info["quantization"] = quant_config.to_diff_dict() - else: - info["quantization"] = quant_config.to_dict() - else: - info["quantization"] = None - - # If fields specified, filter info - if fields is not None: - return {k: v for k, v in info.items() if k in fields} - else: - return info - - # YiYi TODO: (1) add display fields, allow user to set which fields to display in the comnponents table - def __repr__(self): - # Handle empty components case - if not self.components: - return "Components:\n" + "=" * 50 + "\nNo components registered.\n" + "=" * 50 - - # Extract load_id if available - def get_load_id(component): - if hasattr(component, "_diffusers_load_id"): - return component._diffusers_load_id - return "N/A" - - # Format device info compactly - def format_device(component, info): - if not info["has_hook"]: - return str(getattr(component, "device", "N/A")) - else: - device = str(getattr(component, "device", "N/A")) - exec_device = str(info["execution_device"] or "N/A") - return f"{device}({exec_device})" - - # Get max length of load_ids for models - load_ids = [ - get_load_id(component) - for component in self.components.values() - if isinstance(component, torch.nn.Module) and hasattr(component, "_diffusers_load_id") - ] - max_load_id_len = max([15] + [len(str(lid)) for lid in load_ids]) if load_ids else 15 - - # Get all collections for each component - component_collections = {} - for name in self.components.keys(): - component_collections[name] = [] - for coll, comps in self.collections.items(): - if name in comps: - component_collections[name].append(coll) - if not component_collections[name]: - component_collections[name] = ["N/A"] - - # Find the maximum collection name length - all_collections = [coll for colls in component_collections.values() for coll in colls] - max_collection_len = max(10, max(len(str(c)) for c in all_collections)) if all_collections else 10 - - col_widths = { - "id": max(15, max(len(name) for name in self.components.keys())), - "class": max(25, max(len(component.__class__.__name__) for component in self.components.values())), - "device": 20, - "dtype": 15, - "size": 10, - "load_id": max_load_id_len, - "collection": max_collection_len, - } - - # Create the header lines - sep_line = "=" * (sum(col_widths.values()) + len(col_widths) * 3 - 1) + "\n" - dash_line = "-" * (sum(col_widths.values()) + len(col_widths) * 3 - 1) + "\n" - - output = "Components:\n" + sep_line - - # Separate components into models and others - models = {k: v for k, v in self.components.items() if isinstance(v, torch.nn.Module)} - others = {k: v for k, v in self.components.items() if not isinstance(v, torch.nn.Module)} - - # Models section - if models: - output += "Models:\n" + dash_line - # Column headers - output += f"{'Name_ID':<{col_widths['id']}} | {'Class':<{col_widths['class']}} | " - output += f"{'Device: act(exec)':<{col_widths['device']}} | {'Dtype':<{col_widths['dtype']}} | " - output += f"{'Size (GB)':<{col_widths['size']}} | {'Load ID':<{col_widths['load_id']}} | Collection\n" - output += dash_line - - # Model entries - for name, component in models.items(): - info = self.get_model_info(name) - device_str = format_device(component, info) - dtype = str(component.dtype) if hasattr(component, "dtype") else "N/A" - load_id = get_load_id(component) - - # Print first collection on the main line - first_collection = component_collections[name][0] if component_collections[name] else "N/A" - - output += f"{name:<{col_widths['id']}} | {info['class_name']:<{col_widths['class']}} | " - output += f"{device_str:<{col_widths['device']}} | {dtype:<{col_widths['dtype']}} | " - output += f"{info['size_gb']:<{col_widths['size']}.2f} | {load_id:<{col_widths['load_id']}} | {first_collection}\n" - - # Print additional collections on separate lines if they exist - for i in range(1, len(component_collections[name])): - collection = component_collections[name][i] - output += f"{'':<{col_widths['id']}} | {'':<{col_widths['class']}} | " - output += f"{'':<{col_widths['device']}} | {'':<{col_widths['dtype']}} | " - output += f"{'':<{col_widths['size']}} | {'':<{col_widths['load_id']}} | {collection}\n" - - output += dash_line - - # Other components section - if others: - if models: # Add extra newline if we had models section - output += "\n" - output += "Other Components:\n" + dash_line - # Column headers for other components - output += f"{'ID':<{col_widths['id']}} | {'Class':<{col_widths['class']}} | Collection\n" - output += dash_line - - # Other component entries - for name, component in others.items(): - info = self.get_model_info(name) - - # Print first collection on the main line - first_collection = component_collections[name][0] if component_collections[name] else "N/A" - - output += f"{name:<{col_widths['id']}} | {component.__class__.__name__:<{col_widths['class']}} | {first_collection}\n" - - # Print additional collections on separate lines if they exist - for i in range(1, len(component_collections[name])): - collection = component_collections[name][i] - output += f"{'':<{col_widths['id']}} | {'':<{col_widths['class']}} | {collection}\n" - - output += dash_line - - # Add additional component info - output += "\nAdditional Component Info:\n" + "=" * 50 + "\n" - for name in self.components: - info = self.get_model_info(name) - if info is not None and ( - info.get("adapters") is not None or info.get("ip_adapter") or info.get("quantization") - ): - output += f"\n{name}:\n" - if info.get("adapters") is not None: - output += f" Adapters: {info['adapters']}\n" - if info.get("ip_adapter"): - output += " IP-Adapter: Enabled\n" - if info.get("quantization"): - output += f" Quantization: {info['quantization']}\n" - - return output - - def get_one( - self, - component_id: str | None = None, - name: str | None = None, - collection: str | None = None, - load_id: str | None = None, - ) -> Any: - """ - Get a single component by either: - - searching name (pattern matching), collection, or load_id. - - passing in a component_id - Raises an error if multiple components match or none are found. - - Args: - component_id (str | None): Optional component ID to get - name (str | None): Component name or pattern - collection (str | None): Optional collection to filter by - load_id (str | None): Optional load_id to filter by - - Returns: - A single component - - Raises: - ValueError: If no components match or multiple components match - """ - - if component_id is not None and (name is not None or collection is not None or load_id is not None): - raise ValueError("If searching by component_id, do not pass name, collection, or load_id") - - # search by component_id - if component_id is not None: - if component_id not in self.components: - raise ValueError(f"Component '{component_id}' not found in ComponentsManager") - return self.components[component_id] - # search with name/collection/load_id - results = self.search_components(name, collection, load_id) - - if not results: - raise ValueError(f"No components found matching '{name}'") - - if len(results) > 1: - raise ValueError(f"Multiple components found matching '{name}': {list(results.keys())}") - - return next(iter(results.values())) - - def get_ids(self, names: str | list[str] = None, collection: str | None = None): - """ - Get component IDs by a list of names, optionally filtered by collection. - - Args: - names (str | list[str]): list of component names - collection (str | None): Optional collection to filter by - - Returns: - list[str]: list of component IDs - """ - ids = set() - if not isinstance(names, list): - names = [names] - for name in names: - ids.update(self._lookup_ids(name=name, collection=collection)) - return list(ids) - - def get_components_by_ids(self, ids: list[str], return_dict_with_names: bool | None = True): - """ - Get components by a list of IDs. - - Args: - ids (list[str]): - list of component IDs - return_dict_with_names (bool | None): - Whether to return a dictionary with component names as keys: - - Returns: - dict[str, Any]: Dictionary of components. - - If return_dict_with_names=True, keys are component names. - - If return_dict_with_names=False, keys are component IDs. - - Raises: - ValueError: If duplicate component names are found in the search results when return_dict_with_names=True - """ - components = {id: self.components[id] for id in ids} - - if return_dict_with_names: - dict_to_return = {} - for comp_id, comp in components.items(): - comp_name = self._id_to_name(comp_id) - if comp_name in dict_to_return: - raise ValueError( - f"Duplicate component names found in the search results: {comp_name}, please set `return_dict_with_names=False` to return a dictionary with component IDs as keys" - ) - dict_to_return[comp_name] = comp - return dict_to_return - else: - return components - - def get_components_by_names(self, names: list[str], collection: str | None = None): - """ - Get components by a list of names, optionally filtered by collection. - - Args: - names (list[str]): list of component names - collection (str | None): Optional collection to filter by - - Returns: - dict[str, Any]: Dictionary of components with component names as keys - - Raises: - ValueError: If duplicate component names are found in the search results - """ - ids = self.get_ids(names, collection) - return self.get_components_by_ids(ids) diff --git a/diffusers/modular_pipelines/cosmos/__init__.py b/diffusers/modular_pipelines/cosmos/__init__.py deleted file mode 100644 index 38a1b30a421ea3d2baabab5334f96428fd5ff4a3..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/cosmos/__init__.py +++ /dev/null @@ -1,49 +0,0 @@ -from typing import TYPE_CHECKING - -from ...utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, -) - - -_dummy_objects = {} -_import_structure = {} - -try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ...utils import dummy_torch_and_transformers_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) -else: - _import_structure["modular_blocks_cosmos3"] = ["Cosmos3OmniBlocks"] - _import_structure["modular_blocks_cosmos3_distilled"] = ["Cosmos3DistilledBlocks"] - _import_structure["modular_pipeline"] = ["Cosmos3DistilledModularPipeline", "Cosmos3OmniModularPipeline"] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .modular_blocks_cosmos3 import Cosmos3OmniBlocks - from .modular_blocks_cosmos3_distilled import Cosmos3DistilledBlocks - from .modular_pipeline import Cosmos3DistilledModularPipeline, Cosmos3OmniModularPipeline -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/cosmos/after_decode.py b/diffusers/modular_pipelines/cosmos/after_decode.py deleted file mode 100644 index 7f8dd903d6153bef45205b0fd93c220d2151799c..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/cosmos/after_decode.py +++ /dev/null @@ -1,113 +0,0 @@ -import torch - -from ...utils import encode_video, export_to_video -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import InputParam, OutputParam -from .modular_pipeline import Cosmos3OmniModularPipeline - - -class Cosmos3ActionOutputStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Post-processes action latents into action outputs." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="action_latents", - type_hint=torch.Tensor, - default=None, - description="Denoised action latents.", - ), - InputParam( - name="action_mode", type_hint=str, default=None, description="Requested action-generation mode." - ), - InputParam( - name="raw_action_dim_resolved", - type_hint=int, - default=None, - description="Unpadded action-vector dimension.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam("action", type_hint=list[torch.Tensor], description="Generated action vectors.")] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - action_output = None - if block_state.action_mode in {"inverse_dynamics", "policy"} and block_state.action_latents is not None: - action_output = block_state.action_latents - if block_state.raw_action_dim_resolved is not None: - action_output = action_output[:, : block_state.raw_action_dim_resolved] - action_output = [action_output.detach().cpu()] - block_state.action = action_output - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3ExportStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Optional export block that writes decoded outputs to disk. Writes `videos` to `output_path` via " - "`export_to_video`, or muxes `videos` with `sound` via `encode_video` when a waveform is present. " - "Not wired into the default blocks; add it explicitly when you want the pipeline to produce a file." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="videos", required=True, description="Generated video frames to export."), - InputParam( - name="output_path", - type_hint=str, - required=True, - description="Destination path for the exported video.", - ), - InputParam(name="fps", type_hint=float, default=24.0, description="Frame rate of the exported video."), - InputParam( - name="sound", - type_hint=torch.Tensor, - default=None, - description="Generated waveform to mux into the video.", - ), - InputParam( - name="sampling_rate", - type_hint=int, - default=None, - description="Sample rate of the generated waveform in Hz.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam("output_path", type_hint=str, description="Path of the exported video file.")] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - output_path = str(block_state.output_path) - fps = int(round(block_state.fps)) - if block_state.sound is not None: - if block_state.sampling_rate is None: - raise ValueError("`sampling_rate` is required to export a video with sound.") - encode_video( - block_state.videos, - fps=fps, - audio=block_state.sound, - audio_sample_rate=int(block_state.sampling_rate), - output_path=output_path, - ) - else: - export_to_video(block_state.videos, output_path, fps=fps, macro_block_size=1) - block_state.output_path = output_path - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/cosmos/before_denoise.py b/diffusers/modular_pipelines/cosmos/before_denoise.py deleted file mode 100644 index 7bf431aa855bdeadf40369e052fed92908f38c19..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/cosmos/before_denoise.py +++ /dev/null @@ -1,1331 +0,0 @@ -import copy - -import numpy as np -import torch - -from ...models.transformers.transformer_cosmos3 import Cosmos3OmniTransformer -from ...pipelines.cosmos.pipeline_cosmos3_omni import _EMBODIMENT_TO_DOMAIN_ID, CosmosActionCondition -from ...schedulers import FlowMatchEulerDiscreteScheduler, UniPCMultistepScheduler -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, ConfigSpec, InputParam, OutputParam -from .modular_pipeline import Cosmos3OmniModularPipeline - - -class Cosmos3PrepareTextSegmentsStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Builds cond/uncond text segments before denoising." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="cond_input_ids", required=True, description="Token IDs for the conditional prompt."), - InputParam(name="uncond_input_ids", required=True, description="Token IDs for the unconditional prompt."), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "cond_text_segment", - type_hint=dict, - kwargs_type="denoiser_input_fields", - description="Conditional text segment for the denoiser.", - ), - OutputParam( - "uncond_text_segment", - type_hint=dict, - kwargs_type="denoiser_input_fields", - description="Unconditional text segment for the denoiser.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - block_state.cond_text_segment = components._prepare_text_segment(block_state.cond_input_ids, device=device) - block_state.uncond_text_segment = components._prepare_text_segment(block_state.uncond_input_ids, device=device) - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3VisionPrepareLatentsStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Prepares noisy vision latents and the vision conditioning mask." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="x0_tokens_vision", - type_hint=torch.Tensor, - default=None, - description="Vision latents encoded from the conditioning image or video.", - ), - InputParam( - name="vision_condition_frames", - type_hint=list[int], - default=None, - description="Latent-frame indexes fixed by visual conditioning.", - ), - InputParam(name="num_frames", type_hint=int, required=True, description="Number of frames to generate."), - InputParam( - name="height", type_hint=int, required=True, description="Height of the generated video in pixels." - ), - InputParam( - name="width", type_hint=int, required=True, description="Width of the generated video in pixels." - ), - InputParam(name="fps", type_hint=float, default=24.0, description="Frame rate of the generated video."), - InputParam( - name="latents", - type_hint=torch.Tensor, - default=None, - description="Pre-generated noisy vision latents.", - ), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("latents", type_hint=torch.Tensor, description="Noisy vision latents for denoising."), - OutputParam("fps_vision", type_hint=float, description="Frame rate used to pack vision latents."), - OutputParam( - "vision_condition_mask", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Mask marking conditioned vision latent frames.", - ), - OutputParam( - "vision_condition_indexes_for_pack", - type_hint=list[int], - description="Indexes of conditioned vision latent frames.", - ), - OutputParam( - "vision_conditioning_latents", - type_hint=torch.Tensor, - description="Clean encoded vision latents used to re-anchor image conditioning each step.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - dtype = components.transformer.dtype - - x0_tokens_vision = block_state.x0_tokens_vision - if x0_tokens_vision is None: - if block_state.num_frames < 1: - raise ValueError(f"num_frames must be >= 1, got {block_state.num_frames}.") - sf_spatial = components.vae_scale_factor_spatial - if block_state.height % sf_spatial != 0 or block_state.width % sf_spatial != 0: - raise ValueError( - f"height and width must be multiples of {sf_spatial}, got ({block_state.height}, {block_state.width})." - ) - latent_shape = ( - 1, - components.num_channels_latents, - (block_state.num_frames - 1) // components.vae_scale_factor_temporal + 1, - block_state.height // sf_spatial, - block_state.width // sf_spatial, - ) - x0_tokens_vision = torch.zeros(latent_shape, device=device, dtype=torch.float32) - else: - x0_tokens_vision = x0_tokens_vision.to(device=device, dtype=torch.float32) - - block_state.fps_vision = float(block_state.fps) - condition_frames = block_state.vision_condition_frames or [] - block_state.vision_condition_mask = torch.zeros((x0_tokens_vision.shape[2], 1, 1), device=device, dtype=dtype) - for frame_idx in condition_frames: - if 0 <= frame_idx < block_state.vision_condition_mask.shape[0]: - block_state.vision_condition_mask[frame_idx, 0, 0] = 1.0 - - if block_state.latents is None: - pure_noise = randn_tensor( - tuple(x0_tokens_vision.shape), generator=block_state.generator, device=device, dtype=dtype - ) - block_state.latents = ( - block_state.vision_condition_mask * x0_tokens_vision.to(device=device, dtype=dtype) - + (1.0 - block_state.vision_condition_mask) * pure_noise - ) - else: - block_state.latents = block_state.latents.to(device=device, dtype=dtype) - - vision_condition_indexes = torch.nonzero( - block_state.vision_condition_mask[:, 0, 0] > 0, as_tuple=False - ).flatten() - block_state.vision_condition_indexes_for_pack = [int(idx.item()) for idx in vision_condition_indexes] - block_state.vision_conditioning_latents = x0_tokens_vision - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3SoundPrepareLatentsStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Prepares noisy sound latents and the sound conditioning mask." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("transformer", Cosmos3OmniTransformer), - ComponentSpec("scheduler", UniPCMultistepScheduler), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="num_frames", type_hint=int, required=True, description="Number of frames to generate."), - InputParam(name="fps", type_hint=float, default=24.0, description="Frame rate of the generated video."), - InputParam( - name="sound_latents", - type_hint=torch.Tensor, - default=None, - description="Pre-generated noisy sound latents.", - ), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("sound_latents", type_hint=torch.Tensor, description="Noisy sound latents for denoising."), - OutputParam("fps_sound", type_hint=float, description="Frame rate of the sound latent sequence."), - OutputParam( - "sound_condition_mask", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Mask marking conditioned sound latent frames.", - ), - OutputParam("sound_scheduler", description="Scheduler used to update sound latents."), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - dtype = components.transformer.dtype - - if not components.transformer.config.sound_gen: - raise ValueError("Sound generation requires a transformer trained with sound_gen=True.") - - sound_dim = components.transformer.config.sound_dim - block_state.fps_sound = float(components.transformer.config.sound_latent_fps) - n_audio_samples = int(block_state.num_frames / block_state.fps * components.sound_sampling_rate) - hop_size = components.sound_hop_size - t_sound = (n_audio_samples + hop_size - 1) // hop_size - x0_tokens_sound = torch.zeros(sound_dim, t_sound, device=device, dtype=dtype) - block_state.sound_condition_mask = torch.zeros((x0_tokens_sound.shape[1], 1), device=device, dtype=dtype) - - if block_state.sound_latents is None: - pure_noise = randn_tensor( - tuple(x0_tokens_sound.shape), generator=block_state.generator, device=device, dtype=dtype - ) - block_state.sound_latents = ( - block_state.sound_condition_mask.T * x0_tokens_sound - + (1.0 - block_state.sound_condition_mask.T) * pure_noise - ) - else: - block_state.sound_latents = block_state.sound_latents.to(device=device, dtype=dtype) - - block_state.sound_scheduler = copy.deepcopy(components.scheduler) - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3ActionPrepareLatentsStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Prepares noisy action latents and the action conditioning mask." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("transformer", Cosmos3OmniTransformer), - ComponentSpec("scheduler", UniPCMultistepScheduler), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="action", - type_hint=CosmosActionCondition, - required=True, - description="Action-conditioning metadata.", - ), - InputParam( - name="action_condition_frame_indexes", - type_hint=list[int], - default=None, - description="Action-frame indexes fixed by action conditioning.", - ), - InputParam( - name="action_latents", - type_hint=torch.Tensor, - default=None, - description="Pre-generated noisy action latents.", - ), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("action_latents", type_hint=torch.Tensor, description="Noisy action latents for denoising."), - OutputParam( - "action_condition_mask", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Mask marking conditioned action latent frames.", - ), - OutputParam( - "action_domain_ids", - type_hint=list[torch.Tensor], - kwargs_type="denoiser_input_fields", - description="Embodiment domain IDs for action conditioning.", - ), - OutputParam( - "raw_action_dim_resolved", - type_hint=int, - kwargs_type="denoiser_input_fields", - description="Unpadded action-vector dimension.", - ), - OutputParam("action_scheduler", description="Scheduler used to update action latents."), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - dtype = components.transformer.dtype - action = block_state.action - - if not components.transformer.config.action_gen: - raise ValueError("action requires a transformer trained with action_gen=True.") - - block_state.raw_action_dim_resolved = int(action.raw_action_dim) if action.raw_action_dim is not None else None - if ( - block_state.raw_action_dim_resolved is not None - and block_state.raw_action_dim_resolved > components.transformer.config.action_dim - ): - raise ValueError( - f"raw_action_dim={block_state.raw_action_dim_resolved} exceeds the model action_dim=" - f"{components.transformer.config.action_dim}." - ) - - action_chunk_size = action.chunk_size - action_dim = components.transformer.action_dim - if action.mode == "forward_dynamics": - raw_actions = action.raw_actions - if raw_actions is None: - raise ValueError("action_mode='forward_dynamics' requires an action tensor.") - raw_actions = raw_actions.to(device=device, dtype=dtype) - if raw_actions.shape[-1] > action_dim: - raise ValueError( - f"Cosmos3 action dimension {raw_actions.shape[-1]} exceeds model action_dim={action_dim}." - ) - if raw_actions.shape[0] < action_chunk_size: - raw_actions = torch.cat( - [raw_actions, raw_actions[-1:].expand(action_chunk_size - raw_actions.shape[0], -1)], - dim=0, - ) - raw_actions = raw_actions[:action_chunk_size] - if raw_actions.shape[-1] < action_dim: - action_padding = torch.zeros( - raw_actions.shape[0], - action_dim - raw_actions.shape[-1], - dtype=raw_actions.dtype, - device=raw_actions.device, - ) - raw_actions = torch.cat([raw_actions, action_padding], dim=-1) - x0_tokens_action = raw_actions - else: - x0_tokens_action = torch.zeros(action_chunk_size, action_dim, device=device, dtype=dtype) - - if action.domain_name not in _EMBODIMENT_TO_DOMAIN_ID: - raise ValueError( - f"Unknown Cosmos3 action domain_name={action.domain_name!r}; expected one of {sorted(_EMBODIMENT_TO_DOMAIN_ID)}." - ) - block_state.action_domain_ids = [ - torch.tensor([_EMBODIMENT_TO_DOMAIN_ID[action.domain_name]], dtype=torch.long, device=device) - ] - condition_frames = block_state.action_condition_frame_indexes or [] - block_state.action_condition_mask = torch.zeros((x0_tokens_action.shape[0], 1), device=device, dtype=dtype) - for frame_idx in condition_frames: - if 0 <= frame_idx < block_state.action_condition_mask.shape[0]: - block_state.action_condition_mask[frame_idx, 0] = 1.0 - - if block_state.action_latents is None: - pure_noise = randn_tensor( - tuple(x0_tokens_action.shape), generator=block_state.generator, device=device, dtype=dtype - ) - block_state.action_latents = ( - block_state.action_condition_mask * x0_tokens_action - + (1.0 - block_state.action_condition_mask) * pure_noise - ) - if block_state.raw_action_dim_resolved is not None: - block_state.action_latents[:, block_state.raw_action_dim_resolved :] = 0 - else: - block_state.action_latents = block_state.action_latents.to(device=device, dtype=dtype) - - block_state.action_scheduler = copy.deepcopy(components.scheduler) - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3VisionPackSequenceStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Builds separate cond/uncond vision sequence segments." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="cond_text_segment", type_hint=dict, required=True, description="Conditional text segment." - ), - InputParam( - name="uncond_text_segment", - type_hint=dict, - required=True, - description="Unconditional text segment.", - ), - InputParam( - name="latents", type_hint=torch.Tensor, required=True, description="Noisy vision latents to pack." - ), - InputParam( - name="fps_vision", - type_hint=float, - required=True, - description="Frame rate used to pack vision latents.", - ), - InputParam( - name="vision_condition_indexes_for_pack", - type_hint=list[int], - required=True, - description="Indexes of conditioned vision latent frames.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "cond_vision_segment", - type_hint=dict, - kwargs_type="denoiser_input_fields", - description="Conditional vision segment for the denoiser.", - ), - OutputParam( - "uncond_vision_segment", - type_hint=dict, - kwargs_type="denoiser_input_fields", - description="Unconditional vision segment for the denoiser.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - has_image_condition = bool(block_state.vision_condition_indexes_for_pack) - - block_state.cond_vision_segment = components._prepare_vision_segment( - input_vision_tokens=block_state.latents, - has_image_condition=has_image_condition, - mrope_offset=block_state.cond_text_segment["vision_start_temporal_offset"], - vision_fps=block_state.fps_vision, - curr=block_state.cond_text_segment["und_len"], - device=device, - condition_frame_indexes=block_state.vision_condition_indexes_for_pack, - ) - block_state.uncond_vision_segment = components._prepare_vision_segment( - input_vision_tokens=block_state.latents, - has_image_condition=has_image_condition, - mrope_offset=block_state.uncond_text_segment["vision_start_temporal_offset"], - vision_fps=block_state.fps_vision, - curr=block_state.uncond_text_segment["und_len"], - device=device, - condition_frame_indexes=block_state.vision_condition_indexes_for_pack, - ) - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3SoundPackSequenceStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Builds separate cond/uncond sound sequence segments." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="cond_text_segment", type_hint=dict, required=True, description="Conditional text segment." - ), - InputParam( - name="uncond_text_segment", - type_hint=dict, - required=True, - description="Unconditional text segment.", - ), - InputParam( - name="cond_sequence_length", - type_hint=int, - required=True, - description="Conditional multimodal sequence length.", - ), - InputParam( - name="uncond_sequence_length", - type_hint=int, - required=True, - description="Unconditional multimodal sequence length.", - ), - InputParam( - name="sound_latents", type_hint=torch.Tensor, required=True, description="Noisy sound latents to pack." - ), - InputParam( - name="fps_sound", - type_hint=float, - required=True, - description="Frame rate of the sound latent sequence.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "cond_sound_segment", - type_hint=dict, - kwargs_type="denoiser_input_fields", - description="Conditional sound segment for the denoiser.", - ), - OutputParam( - "uncond_sound_segment", - type_hint=dict, - kwargs_type="denoiser_input_fields", - description="Unconditional sound segment for the denoiser.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - block_state.cond_sound_segment = components._prepare_sound_segment( - input_sound_tokens=block_state.sound_latents, - mrope_offset=block_state.cond_text_segment["vision_start_temporal_offset"], - sound_fps=block_state.fps_sound, - curr=block_state.cond_sequence_length, - device=device, - ) - block_state.uncond_sound_segment = components._prepare_sound_segment( - input_sound_tokens=block_state.sound_latents, - mrope_offset=block_state.uncond_text_segment["vision_start_temporal_offset"], - sound_fps=block_state.fps_sound, - curr=block_state.uncond_sequence_length, - device=device, - ) - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3ActionPackSequenceStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Builds separate cond/uncond action sequence segments." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="cond_text_segment", type_hint=dict, required=True, description="Conditional text segment." - ), - InputParam( - name="uncond_text_segment", - type_hint=dict, - required=True, - description="Unconditional text segment.", - ), - InputParam( - name="cond_sequence_length", - type_hint=int, - required=True, - description="Conditional multimodal sequence length.", - ), - InputParam( - name="uncond_sequence_length", - type_hint=int, - required=True, - description="Unconditional multimodal sequence length.", - ), - InputParam( - name="action_latents", - type_hint=torch.Tensor, - required=True, - description="Noisy action latents to pack.", - ), - InputParam( - name="action_condition_frame_indexes", - type_hint=list[int], - default=None, - description="Action-frame indexes fixed by action conditioning.", - ), - InputParam( - name="fps_vision", - type_hint=float, - required=True, - description="Frame rate used to pack vision latents.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "cond_action_segment", - type_hint=dict, - kwargs_type="denoiser_input_fields", - description="Conditional action segment for the denoiser.", - ), - OutputParam( - "uncond_action_segment", - type_hint=dict, - kwargs_type="denoiser_input_fields", - description="Unconditional action segment for the denoiser.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - block_state.cond_action_segment = components._prepare_action_segment( - input_action_tokens=block_state.action_latents, - condition_frame_indexes=block_state.action_condition_frame_indexes, - mrope_offset=block_state.cond_text_segment["vision_start_temporal_offset"], - action_fps=block_state.fps_vision, - curr=block_state.cond_sequence_length, - device=device, - ) - block_state.uncond_action_segment = components._prepare_action_segment( - input_action_tokens=block_state.action_latents, - condition_frame_indexes=block_state.action_condition_frame_indexes, - mrope_offset=block_state.uncond_text_segment["vision_start_temporal_offset"], - action_fps=block_state.fps_vision, - curr=block_state.uncond_sequence_length, - device=device, - ) - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3VisionDenoiseInputStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Assembles text and vision sequence metadata for the denoising loop." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="cond_text_segment", type_hint=dict, required=True, description="Conditional text segment." - ), - InputParam( - name="uncond_text_segment", - type_hint=dict, - required=True, - description="Unconditional text segment.", - ), - InputParam( - name="cond_vision_segment", type_hint=dict, required=True, description="Conditional vision segment." - ), - InputParam( - name="uncond_vision_segment", - type_hint=dict, - required=True, - description="Unconditional vision segment.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "cond_position_ids", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Conditional multimodal RoPE position IDs.", - ), - OutputParam( - "uncond_position_ids", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Unconditional multimodal RoPE position IDs.", - ), - OutputParam( - "cond_sequence_length", - type_hint=int, - kwargs_type="denoiser_input_fields", - description="Conditional multimodal sequence length.", - ), - OutputParam( - "uncond_sequence_length", - type_hint=int, - kwargs_type="denoiser_input_fields", - description="Unconditional multimodal sequence length.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - block_state.cond_position_ids = torch.cat( - [ - block_state.cond_text_segment["text_mrope_ids"], - block_state.cond_vision_segment["vision_mrope_ids"], - ], - dim=1, - ) - block_state.uncond_position_ids = torch.cat( - [ - block_state.uncond_text_segment["text_mrope_ids"], - block_state.uncond_vision_segment["vision_mrope_ids"], - ], - dim=1, - ) - block_state.cond_sequence_length = ( - block_state.cond_text_segment["und_len"] + block_state.cond_vision_segment["num_vision_tokens"] - ) - block_state.uncond_sequence_length = ( - block_state.uncond_text_segment["und_len"] + block_state.uncond_vision_segment["num_vision_tokens"] - ) - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3SoundDenoiseInputStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Appends sound sequence metadata to the denoising-loop inputs." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="cond_position_ids", - type_hint=torch.Tensor, - required=True, - description="Conditional multimodal RoPE position IDs.", - ), - InputParam( - name="uncond_position_ids", - type_hint=torch.Tensor, - required=True, - description="Unconditional multimodal RoPE position IDs.", - ), - InputParam( - name="cond_sequence_length", - type_hint=int, - required=True, - description="Conditional multimodal sequence length.", - ), - InputParam( - name="uncond_sequence_length", - type_hint=int, - required=True, - description="Unconditional multimodal sequence length.", - ), - InputParam( - name="cond_sound_segment", type_hint=dict, required=True, description="Conditional sound segment." - ), - InputParam( - name="uncond_sound_segment", - type_hint=dict, - required=True, - description="Unconditional sound segment.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "cond_position_ids", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Conditional multimodal RoPE position IDs.", - ), - OutputParam( - "uncond_position_ids", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Unconditional multimodal RoPE position IDs.", - ), - OutputParam( - "cond_sequence_length", - type_hint=int, - kwargs_type="denoiser_input_fields", - description="Conditional multimodal sequence length.", - ), - OutputParam( - "uncond_sequence_length", - type_hint=int, - kwargs_type="denoiser_input_fields", - description="Unconditional multimodal sequence length.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - block_state.cond_position_ids = torch.cat( - [block_state.cond_position_ids, block_state.cond_sound_segment["sound_mrope_ids"]], dim=1 - ) - block_state.uncond_position_ids = torch.cat( - [block_state.uncond_position_ids, block_state.uncond_sound_segment["sound_mrope_ids"]], dim=1 - ) - block_state.cond_sequence_length += block_state.cond_sound_segment["sound_len"] - block_state.uncond_sequence_length += block_state.uncond_sound_segment["sound_len"] - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3ActionDenoiseInputStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Appends action sequence metadata to the denoising-loop inputs." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="cond_position_ids", - type_hint=torch.Tensor, - required=True, - description="Conditional multimodal RoPE position IDs.", - ), - InputParam( - name="uncond_position_ids", - type_hint=torch.Tensor, - required=True, - description="Unconditional multimodal RoPE position IDs.", - ), - InputParam( - name="cond_sequence_length", - type_hint=int, - required=True, - description="Conditional multimodal sequence length.", - ), - InputParam( - name="uncond_sequence_length", - type_hint=int, - required=True, - description="Unconditional multimodal sequence length.", - ), - InputParam( - name="cond_action_segment", type_hint=dict, required=True, description="Conditional action segment." - ), - InputParam( - name="uncond_action_segment", - type_hint=dict, - required=True, - description="Unconditional action segment.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "cond_position_ids", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Conditional multimodal RoPE position IDs.", - ), - OutputParam( - "uncond_position_ids", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Unconditional multimodal RoPE position IDs.", - ), - OutputParam( - "cond_sequence_length", - type_hint=int, - kwargs_type="denoiser_input_fields", - description="Conditional multimodal sequence length.", - ), - OutputParam( - "uncond_sequence_length", - type_hint=int, - kwargs_type="denoiser_input_fields", - description="Unconditional multimodal sequence length.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - block_state.cond_position_ids = torch.cat( - [block_state.cond_position_ids, block_state.cond_action_segment["action_mrope_ids"]], dim=1 - ) - block_state.uncond_position_ids = torch.cat( - [block_state.uncond_position_ids, block_state.uncond_action_segment["action_mrope_ids"]], dim=1 - ) - block_state.cond_sequence_length += block_state.cond_action_segment["action_len"] - block_state.uncond_sequence_length += block_state.uncond_action_segment["action_len"] - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3SetTimestepsStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Initializes scheduler timesteps." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", UniPCMultistepScheduler)] - - @property - def expected_configs(self) -> list[ConfigSpec]: - return [ConfigSpec(name="use_native_flow_schedule", default=False)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_inference_steps", required=True), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("timesteps", type_hint=torch.Tensor, description="Scheduler timesteps for denoising."), - OutputParam("num_warmup_steps", type_hint=int, description="Number of scheduler warmup steps."), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - if components.config.use_native_flow_schedule: - sigmas = np.linspace( - 1.0 - 1.0 / components.scheduler.config.num_train_timesteps, - 0.0, - block_state.num_inference_steps + 1, - )[:-1] - components.scheduler.set_timesteps(block_state.num_inference_steps, device=device, sigmas=sigmas) - else: - components.scheduler.set_timesteps(block_state.num_inference_steps, device=device) - block_state.timesteps = components.scheduler.timesteps - block_state.num_warmup_steps = ( - len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order - ) - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3TransferPrepareLatentsStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Per-chunk transfer latent prep: takes the clean target latents encoded by " - "Cosmos3TransferChunkVaeEncoderStep and builds the noisy target latents, velocity mask, condition latents " - "and conditioned-frame indexes for this chunk." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="x0_tokens_vision", - type_hint=torch.Tensor, - required=True, - description="Clean target vision latents encoded from the seeded target frames.", - ), - InputParam( - name="current_conditional_frames", - type_hint=int, - required=True, - description="Number of pixel frames used to seed this chunk's target.", - ), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("latents", type_hint=torch.Tensor, description="Noisy target latents for this chunk."), - OutputParam( - "velocity_mask", - type_hint=torch.Tensor, - description="Mask that zeroes the velocity on conditioned (clean) latent frames.", - ), - OutputParam( - "condition_latents", - type_hint=torch.Tensor, - description="Clean target latents on the conditioned frames (the autoregressive seed).", - ), - OutputParam( - "target_condition_indexes", - type_hint=list[int], - description="Latent-frame indexes fixed by the chunk's conditioning.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - dtype = components.transformer.dtype - tcf = components.vae_scale_factor_temporal - - target_x0 = block_state.x0_tokens_vision.to(device=device) - current_conditional_frames = block_state.current_conditional_frames - - # Build the noisy target latents + conditioning mask from the clean target latents. - latent_t = target_x0.shape[2] - condition_mask = torch.zeros((latent_t, 1, 1), device=device, dtype=dtype) - latent_condition_frames = 0 - if current_conditional_frames > 0: - latent_condition_frames = (current_conditional_frames - 1) // tcf + 1 - condition_mask[:latent_condition_frames] = 1.0 - noise = randn_tensor(tuple(target_x0.shape), generator=block_state.generator, device=device, dtype=dtype) - block_state.latents = condition_mask * target_x0 + (1.0 - condition_mask) * noise - block_state.velocity_mask = 1.0 - condition_mask - block_state.condition_latents = condition_mask * target_x0 - block_state.target_condition_indexes = list(range(latent_condition_frames)) - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3TransferPackSequenceStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Pre-packs the three transfer CFG sequence variants: cond_full / uncond_full carry every control item, " - "the no-control branch drops them (only [text, target]) so the control axis can be amplified." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="cond_text_segment", type_hint=dict, required=True, description="Conditional text segment." - ), - InputParam( - name="uncond_text_segment", type_hint=dict, required=True, description="Unconditional text segment." - ), - InputParam( - name="control_latents", - type_hint=list[torch.Tensor], - required=True, - description="Clean control latents for this chunk, one per hint in canonical order.", - ), - InputParam( - name="latents", - type_hint=torch.Tensor, - required=True, - description="Noisy target latents for this chunk.", - ), - InputParam( - name="target_condition_indexes", - type_hint=list[int], - required=True, - description="Latent-frame indexes fixed by the chunk's conditioning.", - ), - InputParam(name="fps", type_hint=float, default=24.0, description="Frame rate of the generated video."), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "cond_full_static", - type_hint=dict, - kwargs_type="denoiser_input_fields", - description="Conditional [control..., target] transfer sequence carrying every control item.", - ), - OutputParam( - "cond_no_control_static", - type_hint=dict, - kwargs_type="denoiser_input_fields", - description="Conditional [target] transfer sequence with the control items dropped.", - ), - OutputParam( - "uncond_full_static", - type_hint=dict, - kwargs_type="denoiser_input_fields", - description="Unconditional [control..., target] transfer sequence for text CFG.", - ), - OutputParam( - "num_noisy_vision_tokens", - type_hint=int, - description="Number of noisy target vision tokens denoised each step.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - num_hints = len(block_state.control_latents) - - def _vision_pack(text_segment: dict, include_controls: bool) -> dict: - if include_controls: - vision_items = [*block_state.control_latents, block_state.latents] - condition_indexes = [None] * num_hints + [block_state.target_condition_indexes] - clean_flags = [True] * num_hints + [False] - else: - vision_items = [block_state.latents] - condition_indexes = [block_state.target_condition_indexes] - clean_flags = [False] - - # Transfer packs [ctrl_1, ..., ctrl_N, target] into one vision segment - mrope_offset = text_segment["vision_start_temporal_offset"] - item_curr = text_segment["und_len"] - token_shapes = [] - sequence_index_parts = [] - mse_loss_index_parts = [] - noisy_frame_indexes_per_item = [] - mrope_id_parts = [] - num_vision_tokens = 0 - num_noisy_vision_tokens = 0 - for item, item_condition, is_clean in zip(vision_items, condition_indexes, clean_flags): - latent_t = item.shape[2] - if is_clean: - frame_condition = list(range(latent_t)) - else: - frame_condition = item_condition if item_condition is not None else [] - item_segment = components._prepare_vision_segment( - input_vision_tokens=item, - has_image_condition=False, - mrope_offset=mrope_offset, - vision_fps=block_state.fps, - curr=item_curr, - device=device, - condition_frame_indexes=frame_condition, - ) - token_shapes.extend(item_segment["vision_token_shapes"]) - sequence_index_parts.append(item_segment["vision_sequence_indexes"]) - mse_loss_index_parts.append(item_segment["vision_mse_loss_indexes"]) - noisy_frame_indexes_per_item.extend(item_segment["vision_noisy_frame_indexes"]) - mrope_id_parts.append(item_segment["vision_mrope_ids"]) - num_vision_tokens += item_segment["num_vision_tokens"] - num_noisy_vision_tokens += item_segment["num_noisy_vision_tokens"] - item_curr += item_segment["num_vision_tokens"] - - vision_segment = { - "vision_token_shapes": token_shapes, - "vision_sequence_indexes": torch.cat(sequence_index_parts, dim=0), - "vision_mse_loss_indexes": torch.cat(mse_loss_index_parts, dim=0), - "vision_noisy_frame_indexes": noisy_frame_indexes_per_item, - "vision_mrope_ids": torch.cat(mrope_id_parts, dim=1), - "num_vision_tokens": num_vision_tokens, - "num_noisy_vision_tokens": num_noisy_vision_tokens, - } - return { - **text_segment, - **vision_segment, - "position_ids": torch.cat([text_segment["text_mrope_ids"], vision_segment["vision_mrope_ids"]], dim=1), - "sequence_length": text_segment["und_len"] + vision_segment["num_vision_tokens"], - } - - block_state.cond_full_static = _vision_pack(block_state.cond_text_segment, include_controls=True) - block_state.cond_no_control_static = _vision_pack(block_state.cond_text_segment, include_controls=False) - block_state.uncond_full_static = _vision_pack(block_state.uncond_text_segment, include_controls=True) - block_state.num_noisy_vision_tokens = block_state.cond_full_static["num_noisy_vision_tokens"] - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3TransferSetTimestepsStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Resets the scheduler and computes timesteps for a single transfer chunk. UniPCMultistepScheduler keeps " - "per-step state on the instance, so it is reset per chunk (each autoregressive chunk is a full denoise)." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", UniPCMultistepScheduler)] - - @property - def inputs(self) -> list[InputParam]: - return [InputParam.template("num_inference_steps", required=True)] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("timesteps", type_hint=torch.Tensor, description="Scheduler timesteps for this chunk."), - OutputParam( - "num_warmup_steps", type_hint=int, description="Number of scheduler warmup steps for this chunk." - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - components.scheduler.set_timesteps(block_state.num_inference_steps, device=device) - block_state.timesteps = components.scheduler.timesteps - block_state.num_warmup_steps = ( - len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order - ) - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3DistilledSetTimestepsStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Initializes the fixed distilled sampling schedule from the pipeline's `distilled_sigmas` config." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def expected_configs(self) -> list[ConfigSpec]: - return [ - ConfigSpec(name="is_distilled", default=True), - ConfigSpec(name="distilled_sigmas", default=None), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_inference_steps", required=False, default=None), - InputParam( - name="guidance_scale", - type_hint=float, - default=None, - description=( - "Unused for distilled checkpoints; classifier-free guidance is baked into the weights and the " - "scale is forced to 1.0. Passing a value other than 1.0 raises an error." - ), - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("timesteps", type_hint=torch.Tensor, description="Scheduler timesteps for denoising."), - OutputParam("num_warmup_steps", type_hint=int, description="Number of scheduler warmup steps."), - OutputParam( - "num_inference_steps", - type_hint=int, - description="Resolved number of denoising steps (fixed by the distilled schedule).", - ), - OutputParam( - name="guidance_scale", - type_hint=float, - description="Resolved classifier-free guidance scale (always 1.0 for distilled checkpoints).", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - sigmas = components.config.distilled_sigmas - if not sigmas: - raise ValueError( - "Cosmos3DistilledSetTimestepsStep requires the pipeline config `distilled_sigmas` to be set " - "(populated from the distilled checkpoint's `modular_model_index.json`). Load a distilled Cosmos3 " - "checkpoint or use `Cosmos3OmniModularPipeline` for base checkpoints." - ) - sigmas = [float(s) for s in sigmas] - distilled_steps = len(sigmas) - - if block_state.num_inference_steps is not None and block_state.num_inference_steps != distilled_steps: - raise ValueError( - "This is a distilled checkpoint; the step count is fixed by the pipeline's " - f"`distilled_sigmas` config ({distilled_steps} steps). " - f"`num_inference_steps` must be {distilled_steps} or left unset (got {block_state.num_inference_steps})." - ) - if block_state.guidance_scale is not None and block_state.guidance_scale != 1.0: - raise ValueError( - "This is a distilled checkpoint; classifier-free guidance is baked into the weights. " - f"`guidance_scale` must be 1.0 or left unset (got {block_state.guidance_scale})." - ) - - components.scheduler.set_timesteps(sigmas=sigmas, device=device) - block_state.num_inference_steps = distilled_steps - block_state.guidance_scale = 1.0 - block_state.timesteps = components.scheduler.timesteps - block_state.num_warmup_steps = len(block_state.timesteps) - distilled_steps * components.scheduler.order - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/cosmos/before_encoder.py b/diffusers/modular_pipelines/cosmos/before_encoder.py deleted file mode 100644 index 2cdf68712cdfb753d6a6a0a5a065d0b0f02d6437..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/cosmos/before_encoder.py +++ /dev/null @@ -1,166 +0,0 @@ -import math - -import torch - -from ...configuration_utils import FrozenDict -from ...video_processor import VideoProcessor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Cosmos3OmniModularPipeline - - -class Cosmos3TransferSetupStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Preprocesses the transfer control videos and resolves the autoregressive chunk geometry " - "(total_frames / chunk_frames / num_chunks / stride). Chunk-invariant, so it runs once before the loop." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 16, "resample": "bilinear"}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="control_videos", - type_hint=dict, - required=True, - description="Mapping of hint name (edge/blur/depth/seg/wsm) to the control video for that modality.", - ), - InputParam( - name="height", type_hint=int, default=None, description="Height of the generated video in pixels." - ), - InputParam( - name="width", type_hint=int, default=None, description="Width of the generated video in pixels." - ), - InputParam( - name="num_frames", - type_hint=int, - default=None, - description="Optional cap on the number of output frames (defaults to the control video length).", - ), - InputParam( - name="num_video_frames_per_chunk", - type_hint=int, - default=None, - description="Number of pixel frames generated per autoregressive chunk.", - ), - InputParam( - name="num_conditional_frames", - type_hint=int, - default=1, - description="Number of frames each chunk reuses from the previous chunk's tail.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("height", type_hint=int, description="Resolved output height in pixels."), - OutputParam("width", type_hint=int, description="Resolved output width in pixels."), - OutputParam( - "control_frames", - type_hint=dict, - description="Preprocessed, time-padded control maps in canonical hint order.", - ), - OutputParam("total_frames", type_hint=int, description="Total number of output frames to generate."), - OutputParam("chunk_frames", type_hint=int, description="Number of pixel frames per autoregressive chunk."), - OutputParam("num_chunks", type_hint=int, description="Number of autoregressive chunks."), - OutputParam("stride", type_hint=int, description="Frame stride between consecutive chunks."), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - dtype = components.transformer.dtype - - if block_state.height is None: - block_state.height = 720 - if block_state.width is None: - block_state.width = 1280 - - # Canonical hint order used both to validate and to order the preprocessed control maps. - hint_order = ["edge", "blur", "depth", "seg", "wsm"] - control_videos = block_state.control_videos - if not isinstance(control_videos, dict) or not control_videos: - raise ValueError("`control_videos` must be a non-empty dict mapping hint name -> control video.") - unknown = [k for k in control_videos if k not in hint_order] - if unknown: - raise ValueError(f"`control_videos` has unknown hint(s) {unknown}; expected keys from {hint_order}.") - if any(v is None for v in control_videos.values()): - raise ValueError("`control_videos` entries must be loaded videos, not None.") - - tcf = components.vae_scale_factor_temporal - sf = components.vae_scale_factor_spatial - if block_state.height % sf != 0 or block_state.width % sf != 0: - raise ValueError( - f"`height` and `width` must be multiples of {sf}, got ({block_state.height}, {block_state.width})." - ) - - # Preprocess every control map to [1, 3, T, H, W] in [-1, 1] at target geometry, in canonical hint order. - # The dict preserves this order, so downstream blocks just iterate control_frames (no separate hint_keys). - hint_keys = [k for k in hint_order if k in control_videos] - control_frames = { - key: components.video_processor.preprocess_video( - control_videos[key], height=block_state.height, width=block_state.width - ).to(device=device, dtype=dtype) - for key in hint_keys - } - - # Output frame count / chunking come from the (first) control video, optionally capped by num_frames. - total_frames = next(iter(control_frames.values())).shape[2] - if block_state.num_frames is not None: - total_frames = min(total_frames, block_state.num_frames) - total_frames = max(1, total_frames) - - per_chunk = ( - block_state.num_video_frames_per_chunk - if block_state.num_video_frames_per_chunk is not None - else total_frames - ) - chunk_frames = 1 if total_frames == 1 else per_chunk - chunk_frames = math.ceil((chunk_frames - 1) / tcf) * tcf + 1 - - if total_frames <= chunk_frames: - num_chunks, stride = 1, chunk_frames - else: - stride = chunk_frames - block_state.num_conditional_frames - if stride <= 0: - raise ValueError("`num_conditional_frames` must be smaller than `num_video_frames_per_chunk`.") - remaining = total_frames - chunk_frames - num_chunks = 1 + (remaining // stride + (1 if remaining % stride else 0)) - - # Reflect-pad each control map along time up to `padded` (repeat the last frame once the clip is too short to - # keep reflecting). No truncation here; per-chunk slicing happens later. - padded = max(total_frames, chunk_frames) - control_frames_padded = {} - for key, frames in control_frames.items(): - while frames.shape[2] < padded: - pad_len = min(frames.shape[2] - 1, padded - frames.shape[2]) - if pad_len <= 0: - pad_frame = frames[:, :, -1:].repeat(1, 1, padded - frames.shape[2], 1, 1) - frames = torch.cat([frames, pad_frame], dim=2) - break - frames = torch.cat([frames, frames.flip(dims=[2])[:, :, :pad_len]], dim=2) - control_frames_padded[key] = frames - block_state.control_frames = control_frames_padded - block_state.total_frames = total_frames - block_state.chunk_frames = chunk_frames - block_state.num_chunks = num_chunks - block_state.stride = stride - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/cosmos/decoders.py b/diffusers/modular_pipelines/cosmos/decoders.py deleted file mode 100644 index a76e48501d85cdc7f88e4cf1c64c4a600500543a..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/cosmos/decoders.py +++ /dev/null @@ -1,259 +0,0 @@ -import torch - -from ...configuration_utils import FrozenDict -from ...models.autoencoders.autoencoder_cosmos3_audio import Cosmos3AVAEAudioTokenizer -from ...models.autoencoders.autoencoder_kl_wan import AutoencoderKLWan -from ...utils import logging -from ...video_processor import VideoProcessor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Cosmos3OmniModularPipeline - - -logger = logging.get_logger(__name__) - - -class Cosmos3VideoDecodeStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Decodes denoised vision latents into video outputs." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLWan), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 16, "resample": "bilinear"}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents", required=True, description="Denoised vision latents to decode."), - InputParam.template("output_type", default="pil"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam.template("videos")] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - - if block_state.output_type == "latent": - block_state.videos = block_state.latents - else: - in_dtype = block_state.latents.dtype - vae_dtype = components.vae.dtype - mean = components._vae_latents_mean.to(device=block_state.latents.device, dtype=vae_dtype) - inv_std = components._vae_latents_inv_std.to(device=block_state.latents.device, dtype=vae_dtype) - z_raw = block_state.latents.to(vae_dtype) / inv_std.view(1, -1, 1, 1, 1) + mean.view(1, -1, 1, 1, 1) - decoded = components.vae.decode(z_raw).sample.to(in_dtype) - block_state.videos = components.video_processor.postprocess_video( - decoded, output_type=block_state.output_type - )[0] - - if components.requires_safety_checker and block_state.output_type != "latent": - if getattr(components, "safety_checker", None) is None: - raise ValueError( - "Cosmos3 requires a safety checker by default. Call `pipe.enable_safety_checker()` to load it " - "(or pass your own), or opt out explicitly with `pipe.disable_safety_checker()`." - ) - block_state.videos = components._apply_video_safety_check( - block_state.videos, output_type=block_state.output_type, device=device - ) - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3SoundDecodeStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Decodes sound latents into waveform output." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("sound_tokenizer", Cosmos3AVAEAudioTokenizer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="sound_latents", - type_hint=torch.Tensor, - required=True, - description="Denoised sound latents to decode.", - ) - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("sound", type_hint=torch.Tensor, description="Generated waveform."), - OutputParam("sampling_rate", type_hint=int, description="Sample rate of the generated waveform in Hz."), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - if components.sound_tokenizer is None: - raise ValueError("Sound decoding requires a sound-capable checkpoint with a sound_tokenizer.") - block_state.sound = components.decode_sound(block_state.sound_latents) - block_state.sampling_rate = int(components.sound_tokenizer.config.sampling_rate) - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3TransferDecodeChunkStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Decodes one transfer chunk's latents to pixels (float32, clamped to [-1, 1]), records it as the " - "autoregressive seed for the next chunk, and appends it to output_chunks (dropping the overlap that " - "later chunks share with the previous chunk's conditioning frames)." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("vae", AutoencoderKLWan)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="latents", - type_hint=torch.Tensor, - required=True, - description="Denoised target latents for this chunk.", - ), - InputParam(name="chunk_id", type_hint=int, default=0, description="Index of the current chunk."), - InputParam( - name="current_conditional_frames", - type_hint=int, - required=True, - description="Number of pixel frames this chunk reused from the previous chunk.", - ), - InputParam( - name="output_chunks", - type_hint=list[torch.Tensor], - required=True, - description="Decoded pixel chunks accumulated so far.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "previous_output", - type_hint=torch.Tensor, - description="Decoded pixels of this chunk, used to seed the next chunk.", - ), - OutputParam( - "output_chunks", - type_hint=list[torch.Tensor], - description="Decoded pixel chunks accumulated so far (with this chunk appended).", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - latents = block_state.latents - vae_dtype = components.vae.dtype - mean = components._vae_latents_mean.to(device=latents.device, dtype=vae_dtype) - inv_std = components._vae_latents_inv_std.to(device=latents.device, dtype=vae_dtype) - z_raw = latents.to(vae_dtype) / inv_std.view(1, -1, 1, 1, 1) + mean.view(1, -1, 1, 1, 1) - output_video = components.vae.decode(z_raw).sample.to(torch.float32).clamp(-1, 1) - block_state.previous_output = output_video - chunk = ( - output_video if block_state.chunk_id == 0 else output_video[:, :, block_state.current_conditional_frames :] - ) - block_state.output_chunks = [*block_state.output_chunks, chunk] - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3TransferStitchStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Concatenates the decoded transfer chunks along time, truncates to total_frames, and post-processes to " - "the requested output type. Transfer produces no audio, so sound / sampling_rate are None." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLWan), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 16, "resample": "bilinear"}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="output_chunks", - type_hint=list[torch.Tensor], - required=True, - description="Decoded pixel chunks to stitch together.", - ), - InputParam( - name="total_frames", type_hint=int, required=True, description="Total number of output frames to keep." - ), - InputParam.template("output_type", default="pil"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("videos", description="The generated transfer video."), - OutputParam("sound", description="Always None for transfer (no audio)."), - OutputParam("sampling_rate", description="Always None for transfer (no audio)."), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - decoded = torch.cat(block_state.output_chunks, dim=2)[:, :, : block_state.total_frames] - block_state.videos = components.video_processor.postprocess_video( - decoded, output_type=block_state.output_type - )[0] - - if components.requires_safety_checker and block_state.output_type != "latent": - if getattr(components, "safety_checker", None) is None: - raise ValueError( - "Cosmos3 requires a safety checker by default. Call `pipe.enable_safety_checker()` to load it " - "(or pass your own), or opt out explicitly with `pipe.disable_safety_checker()`." - ) - block_state.videos = components._apply_video_safety_check( - block_state.videos, output_type=block_state.output_type, device=device - ) - - block_state.sound = None - block_state.sampling_rate = None - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/cosmos/denoise.py b/diffusers/modular_pipelines/cosmos/denoise.py deleted file mode 100644 index eda37c8e99cf0826d1e0ed0d9c744c94c163fc76..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/cosmos/denoise.py +++ /dev/null @@ -1,889 +0,0 @@ -import inspect - -import torch - -from ...models.transformers.transformer_cosmos3 import Cosmos3OmniTransformer -from ...schedulers import FlowMatchEulerDiscreteScheduler, UniPCMultistepScheduler -from ..modular_pipeline import ( - BlockState, - LoopSequentialPipelineBlocks, - ModularPipelineBlocks, - PipelineState, -) -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Cosmos3OmniModularPipeline - - -class Cosmos3VisionLoopPrepareStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Prepares vision tokens and timesteps for one denoising iteration." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents", required=True, description="Noisy vision latents to denoise."), - InputParam( - name="cond_vision_segment", type_hint=dict, required=True, description="Conditional vision segment." - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "vision_tokens", - type_hint=list[torch.Tensor], - description="Vision tokens for the transformer denoiser.", - ), - OutputParam("vision_timesteps", type_hint=torch.Tensor, description="Timesteps for the vision tokens."), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - device = components._execution_device - block_state.vision_tokens = [block_state.latents.to(device=device, dtype=components.transformer.dtype)] - block_state.vision_timesteps = torch.full( - (block_state.cond_vision_segment["num_noisy_vision_tokens"],), t.item(), device=device - ) - return components, block_state - - -class Cosmos3SoundLoopPrepareStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Prepares sound tokens and timesteps for one denoising iteration." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="sound_latents", - type_hint=torch.Tensor, - required=True, - description="Noisy sound latents to denoise.", - ), - InputParam( - name="cond_sound_segment", type_hint=dict, required=True, description="Conditional sound segment." - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "sound_tokens", type_hint=list[torch.Tensor], description="Sound tokens for the transformer denoiser." - ), - OutputParam("sound_timesteps", type_hint=torch.Tensor, description="Timesteps for the sound tokens."), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - device = components._execution_device - block_state.sound_tokens = [block_state.sound_latents.to(device=device, dtype=components.transformer.dtype)] - block_state.sound_timesteps = torch.full( - (block_state.cond_sound_segment["sound_len"],), t.item(), device=device - ) - return components, block_state - - -class Cosmos3ActionLoopPrepareStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Prepares action tokens and timesteps for one denoising iteration." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="action_latents", - type_hint=torch.Tensor, - required=True, - description="Noisy action latents to denoise.", - ), - InputParam( - name="cond_action_segment", type_hint=dict, required=True, description="Conditional action segment." - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "action_tokens", - type_hint=list[torch.Tensor], - description="Action tokens for the transformer denoiser.", - ), - OutputParam("action_timesteps", type_hint=torch.Tensor, description="Timesteps for the action tokens."), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - device = components._execution_device - block_state.action_tokens = [block_state.action_latents.to(device=device, dtype=components.transformer.dtype)] - block_state.action_timesteps = torch.full( - (block_state.cond_action_segment["num_noisy_action_tokens"],), t.item(), device=device - ) - return components, block_state - - -class Cosmos3LoopDenoiser(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Predicts available Cosmos3 modality velocities for one denoising iteration." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("denoiser_input_fields"), - InputParam( - name="guidance_scale", - type_hint=float, - default=6.0, - description="Scale for classifier-free guidance.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "velocity_vision", type_hint=torch.Tensor, description="Predicted velocity for vision latents." - ), - OutputParam("velocity_sound", type_hint=torch.Tensor, description="Predicted velocity for sound latents."), - OutputParam( - "velocity_action", type_hint=torch.Tensor, description="Predicted velocity for action latents." - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - denoiser_input_fields = block_state.denoiser_input_fields - loop_input_fields = block_state.as_dict() - has_sound = "sound_tokens" in loop_input_fields - has_action = "action_tokens" in loop_input_fields - do_cfg = block_state.guidance_scale != 1.0 - transformer_args = set(inspect.signature(components.transformer.forward).parameters) - - prediction_passes = ["cond"] - if do_cfg: - prediction_passes.append("uncond") - - velocities = {} - for pass_name in prediction_passes: - transformer_kwargs = {} - for field_name, field_value in denoiser_input_fields.items(): - if field_name.startswith(f"{pass_name}_"): - transformer_field_name = field_name.removeprefix(f"{pass_name}_") - if transformer_field_name.endswith("_segment"): - transformer_kwargs.update(field_value) - else: - transformer_kwargs[transformer_field_name] = field_value - elif field_name in transformer_args: - transformer_kwargs[field_name] = field_value - transformer_kwargs.update( - { - field_name: field_value - for field_name, field_value in loop_input_fields.items() - if field_name in transformer_args - } - ) - transformer_kwargs = { - name: value for name, value in transformer_kwargs.items() if name in transformer_args - } - preds_vision, preds_sound, preds_action = components.transformer(**transformer_kwargs, return_dict=False) - velocities[pass_name] = components._mask_velocity_predictions( - preds_vision, - preds_sound, - vision_condition_mask=[loop_input_fields["vision_condition_mask"]], - sound_condition_mask=[loop_input_fields["sound_condition_mask"]] if has_sound else None, - preds_action=preds_action, - action_condition_mask=[loop_input_fields["action_condition_mask"]] if has_action else None, - raw_action_dim=loop_input_fields.get("raw_action_dim_resolved"), - ) - - cond_velocity_vision, cond_velocity_sound, cond_velocity_action = velocities["cond"] - if do_cfg: - uncond_velocity_vision, uncond_velocity_sound, uncond_velocity_action = velocities["uncond"] - block_state.velocity_vision = uncond_velocity_vision + block_state.guidance_scale * ( - cond_velocity_vision - uncond_velocity_vision - ) - block_state.velocity_sound = ( - uncond_velocity_sound + block_state.guidance_scale * (cond_velocity_sound - uncond_velocity_sound) - if has_sound - else None - ) - block_state.velocity_action = ( - uncond_velocity_action + block_state.guidance_scale * (cond_velocity_action - uncond_velocity_action) - if has_action - else None - ) - else: - block_state.velocity_vision = cond_velocity_vision - block_state.velocity_sound = cond_velocity_sound if has_sound else None - block_state.velocity_action = cond_velocity_action if has_action else None - - return components, block_state - - -class Cosmos3VisionLoopSchedulerStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Updates vision latents after one denoising iteration." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", UniPCMultistepScheduler)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents", required=True, description="Noisy vision latents to update."), - InputParam( - name="velocity_vision", type_hint=torch.Tensor, required=True, description="Predicted vision velocity." - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam.template("latents")] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - block_state.latents = components.scheduler.step( - block_state.velocity_vision.unsqueeze(0), t, block_state.latents.unsqueeze(0), return_dict=False - )[0].squeeze(0) - return components, block_state - - -class Cosmos3DistilledVisionLoopSchedulerStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Updates vision latents after one distilled denoising iteration, re-anchoring conditioned frames." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents", required=True, description="Noisy vision latents to update."), - InputParam( - name="velocity_vision", type_hint=torch.Tensor, required=True, description="Predicted vision velocity." - ), - InputParam( - name="vision_condition_mask", - type_hint=torch.Tensor, - required=True, - description="Mask marking conditioned vision latent frames.", - ), - InputParam( - name="vision_conditioning_latents", - type_hint=torch.Tensor, - default=None, - description="Clean encoded vision latents for re-anchoring conditioned frames.", - ), - InputParam( - name="vision_condition_indexes_for_pack", - type_hint=list, - default=None, - description="Indexes of conditioned vision latent frames; non-empty for image-to-video.", - ), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam.template("latents")] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - # Pass the generator so the scheduler's stochastic (SDE) re-noising is seedable/reproducible. - block_state.latents = components.scheduler.step( - block_state.velocity_vision.unsqueeze(0), - t, - block_state.latents.unsqueeze(0), - generator=block_state.generator, - return_dict=False, - )[0].squeeze(0) - - # Distilled checkpoints use stochastic (SDE) scheduler steps that re-noise every position. - # Re-anchor conditioned frames to the clean encoded reference after each step. - has_image_condition = bool(block_state.vision_condition_indexes_for_pack) - if has_image_condition and block_state.vision_conditioning_latents is not None: - mask = block_state.vision_condition_mask - reference = block_state.vision_conditioning_latents.to(block_state.latents.dtype) - block_state.latents = mask * reference + (1.0 - mask) * block_state.latents - - return components, block_state - - -class Cosmos3SoundLoopSchedulerStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Updates sound latents after one denoising iteration." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="sound_latents", - type_hint=torch.Tensor, - required=True, - description="Noisy sound latents to update.", - ), - InputParam( - name="sound_scheduler", - type_hint=UniPCMultistepScheduler, - required=True, - description="Scheduler used to update sound latents.", - ), - InputParam( - name="velocity_sound", type_hint=torch.Tensor, required=True, description="Predicted sound velocity." - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam("sound_latents", type_hint=torch.Tensor, description="Updated sound latents.")] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - block_state.sound_latents = block_state.sound_scheduler.step( - block_state.velocity_sound.unsqueeze(0), t, block_state.sound_latents.unsqueeze(0), return_dict=False - )[0].squeeze(0) - return components, block_state - - -class Cosmos3ActionLoopSchedulerStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Updates action latents after one denoising iteration." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="action_latents", - type_hint=torch.Tensor, - required=True, - description="Noisy action latents to update.", - ), - InputParam( - name="action_scheduler", - type_hint=UniPCMultistepScheduler, - required=True, - description="Scheduler used to update action latents.", - ), - InputParam( - name="velocity_action", type_hint=torch.Tensor, required=True, description="Predicted action velocity." - ), - InputParam( - name="action_condition_mask", - type_hint=torch.Tensor, - required=True, - description="Mask marking conditioned action latent frames.", - ), - InputParam( - name="raw_action_dim_resolved", - type_hint=int, - default=None, - description="Unpadded action-vector dimension.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam("action_latents", type_hint=torch.Tensor, description="Updated action latents.")] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - has_noisy_action = block_state.action_condition_mask.sum() < block_state.action_condition_mask.numel() - if has_noisy_action: - block_state.action_latents = block_state.action_scheduler.step( - block_state.velocity_action.unsqueeze(0), t, block_state.action_latents.unsqueeze(0), return_dict=False - )[0].squeeze(0) - if block_state.raw_action_dim_resolved is not None: - block_state.action_latents[:, block_state.raw_action_dim_resolved :] = 0 - return components, block_state - - -class Cosmos3DenoiseLoopWrapper(LoopSequentialPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Iteratively denoises Cosmos3 latents over scheduler timesteps." - - @property - def loop_expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", UniPCMultistepScheduler), - ComponentSpec("transformer", Cosmos3OmniTransformer), - ] - - @property - def loop_inputs(self) -> list[InputParam]: - return [ - InputParam.template("timesteps", required=True), - InputParam.template("num_inference_steps", required=True), - InputParam( - name="num_warmup_steps", type_hint=int, required=True, description="Number of scheduler warmup steps." - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) - if i == len(block_state.timesteps) - 1 or ( - (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 - ): - progress_bar.update() - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3VisionDenoiseStep(Cosmos3DenoiseLoopWrapper): - block_classes = [ - Cosmos3VisionLoopPrepareStep, - Cosmos3LoopDenoiser, - Cosmos3VisionLoopSchedulerStep, - ] - block_names = ["prepare_vision", "denoiser", "update_vision"] - - @property - def description(self) -> str: - return "Runs the vision-only Cosmos3 denoising loop." - - -class Cosmos3DistilledVisionDenoiseStep(Cosmos3DenoiseLoopWrapper): - model_name = "cosmos3-omni" - block_classes = [ - Cosmos3VisionLoopPrepareStep, - Cosmos3LoopDenoiser, - Cosmos3DistilledVisionLoopSchedulerStep, - ] - block_names = ["prepare_vision", "denoiser", "update_vision"] - - @property - def description(self) -> str: - return "Runs the vision-only distilled Cosmos3 denoising loop." - - @property - def loop_expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), - ComponentSpec("transformer", Cosmos3OmniTransformer), - ] - - -class Cosmos3VisionSoundDenoiseStep(Cosmos3DenoiseLoopWrapper): - block_classes = [ - Cosmos3VisionLoopPrepareStep, - Cosmos3SoundLoopPrepareStep, - Cosmos3LoopDenoiser, - Cosmos3VisionLoopSchedulerStep, - Cosmos3SoundLoopSchedulerStep, - ] - block_names = ["prepare_vision", "prepare_sound", "denoiser", "update_vision", "update_sound"] - - @property - def description(self) -> str: - return "Runs the vision-and-sound Cosmos3 denoising loop." - - -class Cosmos3VisionActionDenoiseStep(Cosmos3DenoiseLoopWrapper): - block_classes = [ - Cosmos3VisionLoopPrepareStep, - Cosmos3ActionLoopPrepareStep, - Cosmos3LoopDenoiser, - Cosmos3VisionLoopSchedulerStep, - Cosmos3ActionLoopSchedulerStep, - ] - block_names = ["prepare_vision", "prepare_action", "denoiser", "update_vision", "update_action"] - - @property - def description(self) -> str: - return "Runs the vision-and-action Cosmos3 denoising loop." - - -class Cosmos3VisionSoundActionDenoiseStep(Cosmos3DenoiseLoopWrapper): - block_classes = [ - Cosmos3VisionLoopPrepareStep, - Cosmos3SoundLoopPrepareStep, - Cosmos3ActionLoopPrepareStep, - Cosmos3LoopDenoiser, - Cosmos3VisionLoopSchedulerStep, - Cosmos3SoundLoopSchedulerStep, - Cosmos3ActionLoopSchedulerStep, - ] - block_names = [ - "prepare_vision", - "prepare_sound", - "prepare_action", - "denoiser", - "update_vision", - "update_sound", - "update_action", - ] - - @property - def description(self) -> str: - return "Runs the vision, sound, and action Cosmos3 denoising loop." - - -class Cosmos3TransferLoopPrepareStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Prepares the full [control..., target] and target-only vision token lists plus timesteps for one transfer iteration." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="control_latents", - type_hint=list[torch.Tensor], - required=True, - description="Clean control latents for this chunk, one per hint in canonical order.", - ), - InputParam( - name="latents", type_hint=torch.Tensor, required=True, description="Noisy target latents to denoise." - ), - InputParam( - name="num_noisy_vision_tokens", - type_hint=int, - required=True, - description="Number of noisy target vision tokens denoised each step.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "vision_tokens_full", - type_hint=list[torch.Tensor], - description="Token list for the [control..., target] forward passes.", - ), - OutputParam( - "vision_tokens_target", - type_hint=list[torch.Tensor], - description="Token list for the target-only (no-control) forward pass.", - ), - OutputParam( - "vision_timesteps", type_hint=torch.Tensor, description="Timesteps for the noisy target tokens." - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - device = components._execution_device - dtype = components.transformer.dtype - block_state.vision_tokens_full = [c.to(device=device, dtype=dtype) for c in block_state.control_latents] + [ - block_state.latents.to(device=device, dtype=dtype) - ] - block_state.vision_tokens_target = [block_state.latents.to(device=device, dtype=dtype)] - block_state.vision_timesteps = torch.full((block_state.num_noisy_vision_tokens,), t.item(), device=device) - return components, block_state - - -class Cosmos3TransferLoopDenoiser(ModularPipelineBlocks): - # Dedicated (not Cosmos3LoopDenoiser): transfer runs up to 3 passes over different token sequences with nested - # control/text CFG and interval gating, which the generic cond/uncond denoiser cannot express. - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Predicts the transfer velocity with nested control/text CFG over [control..., target]. Each branch is " - "gated by its guidance interval, and the result is masked so conditioned frames get zero velocity." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Cosmos3OmniTransformer)] - - @property - def inputs(self) -> list[InputParam]: - return [ - # The three pre-packed CFG sequence variants (cond_full / cond_no_control / uncond_full) flow in as - # denoiser_input_fields, gathered generically like the other Cosmos3 denoisers. - InputParam.template("denoiser_input_fields"), - InputParam( - name="vision_tokens_full", - type_hint=list[torch.Tensor], - required=True, - description="Token list for the [control..., target] forward passes.", - ), - InputParam( - name="vision_tokens_target", - type_hint=list[torch.Tensor], - required=True, - description="Token list for the target-only (no-control) forward pass.", - ), - InputParam( - name="vision_timesteps", - type_hint=torch.Tensor, - required=True, - description="Timesteps for the noisy target tokens.", - ), - InputParam( - name="velocity_mask", - type_hint=torch.Tensor, - required=True, - description="Mask that zeroes the velocity on conditioned (clean) latent frames.", - ), - InputParam( - name="guidance_scale", - type_hint=float, - default=6.0, - description="Scale for text classifier-free guidance.", - ), - InputParam( - name="control_guidance", - type_hint=float, - default=1.0, - description="Scale for the control (structural) guidance axis.", - ), - InputParam( - name="guidance_interval", - type_hint=tuple, - default=None, - description="Timestep interval [lo, hi] over which text guidance is active (None = always).", - ), - InputParam( - name="control_guidance_interval", - type_hint=tuple, - default=None, - description="Timestep interval [lo, hi] over which control guidance is active (None = always).", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam("velocity", type_hint=torch.Tensor, description="Predicted (masked) transfer velocity.")] - - @staticmethod - def _forward(components, static, vision_tokens, vision_timesteps): - preds_vision, _, _ = components.transformer( - input_ids=static["input_ids"], - text_indexes=static["text_indexes"], - position_ids=static["position_ids"], - und_len=static["und_len"], - sequence_length=static["sequence_length"], - vision_tokens=vision_tokens, - vision_token_shapes=static["vision_token_shapes"], - vision_sequence_indexes=static["vision_sequence_indexes"], - vision_mse_loss_indexes=static["vision_mse_loss_indexes"], - vision_timesteps=vision_timesteps, - vision_noisy_frame_indexes=static["vision_noisy_frame_indexes"], - return_dict=False, - ) - return preds_vision[-1] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - # active-at: a None interval is always active; otherwise the timestep must fall within [lo, hi]. - guidance_interval = block_state.guidance_interval - guidance_active = guidance_interval is None or ( - float(guidance_interval[0]) <= float(t.item()) <= float(guidance_interval[1]) - ) - control_interval = block_state.control_guidance_interval - control_active = control_interval is None or ( - float(control_interval[0]) <= float(t.item()) <= float(control_interval[1]) - ) - step_guidance = block_state.guidance_scale if guidance_active else 1.0 - step_control = block_state.control_guidance if control_active else 1.0 - needs_text_cfg = step_guidance > 1.0 - needs_control_cfg = step_control != 1.0 - - denoiser_input_fields = block_state.denoiser_input_fields - cond_full_static = denoiser_input_fields["cond_full_static"] - cond_no_control_static = denoiser_input_fields["cond_no_control_static"] - uncond_full_static = denoiser_input_fields["uncond_full_static"] - - cond_full = self._forward( - components, cond_full_static, block_state.vision_tokens_full, block_state.vision_timesteps - ) - - cond_no_control = None - if needs_control_cfg: - cond_no_control = self._forward( - components, - cond_no_control_static, - block_state.vision_tokens_target, - block_state.vision_timesteps, - ) - - uncond_full = None - if needs_text_cfg: - uncond_full = self._forward( - components, - uncond_full_static, - block_state.vision_tokens_full, - block_state.vision_timesteps, - ) - - if needs_control_cfg and needs_text_cfg: - control_cond = cond_no_control + step_control * (cond_full - cond_no_control) - velocity = uncond_full + step_guidance * (control_cond - uncond_full) - elif needs_control_cfg: - velocity = cond_no_control + step_control * (cond_full - cond_no_control) - elif needs_text_cfg: - velocity = uncond_full + step_guidance * (cond_full - uncond_full) - else: - velocity = cond_full - - block_state.velocity = velocity * block_state.velocity_mask - return components, block_state - - -class Cosmos3TransferLoopSchedulerStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Steps the scheduler and re-pins the conditioned frames exactly for one transfer iteration." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", UniPCMultistepScheduler)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="latents", type_hint=torch.Tensor, required=True, description="Noisy target latents to update." - ), - InputParam( - name="velocity", - type_hint=torch.Tensor, - required=True, - description="Predicted (masked) transfer velocity.", - ), - InputParam( - name="velocity_mask", - type_hint=torch.Tensor, - required=True, - description="Mask that zeroes the velocity on conditioned (clean) latent frames.", - ), - InputParam( - name="condition_latents", - type_hint=torch.Tensor, - required=True, - description="Clean target latents on the conditioned frames (the autoregressive seed).", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam("latents", type_hint=torch.Tensor, description="Updated target latents for this chunk.")] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - block_state.latents = components.scheduler.step( - block_state.velocity.unsqueeze(0), t, block_state.latents.unsqueeze(0), return_dict=False - )[0].squeeze(0) - # Re-pin conditioned frames exactly (the autoregressive seed), guarding multistep drift. - block_state.latents = ( - block_state.velocity_mask * block_state.latents - + (1.0 - block_state.velocity_mask) * block_state.condition_latents - ) - return components, block_state - - -# auto_docstring -class Cosmos3TransferDenoiseStep(Cosmos3DenoiseLoopWrapper): - """ - Runs the per-chunk transfer denoising loop over scheduler timesteps. - - Components: - transformer (`Cosmos3OmniTransformer`) scheduler (`UniPCMultistepScheduler`) - - Inputs: - timesteps (`Tensor`): - Timesteps for the denoising process. - num_inference_steps (`int`): - The number of denoising steps. - num_warmup_steps (`int`): - Number of scheduler warmup steps. - control_latents (`list`): - Clean control latents for this chunk, one per hint in canonical order. - latents (`Tensor`): - Noisy target latents to denoise. - num_noisy_vision_tokens (`int`): - Number of noisy target vision tokens denoised each step. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - velocity_mask (`Tensor`): - Mask that zeroes the velocity on conditioned (clean) latent frames. - guidance_scale (`float`, *optional*, defaults to 6.0): - Scale for text classifier-free guidance. - control_guidance (`float`, *optional*, defaults to 1.0): - Scale for the control (structural) guidance axis. - guidance_interval (`tuple`, *optional*): - Timestep interval [lo, hi] over which text guidance is active (None = always). - control_guidance_interval (`tuple`, *optional*): - Timestep interval [lo, hi] over which control guidance is active (None = always). - latents (`Tensor`): - Noisy target latents to update. - condition_latents (`Tensor`): - Clean target latents on the conditioned frames (the autoregressive seed). - - Outputs: - latents (`Tensor`): - Updated target latents for this chunk. - """ - - block_classes = [ - Cosmos3TransferLoopPrepareStep, - Cosmos3TransferLoopDenoiser, - Cosmos3TransferLoopSchedulerStep, - ] - block_names = ["prepare_transfer", "denoiser", "update_transfer"] - - @property - def description(self) -> str: - return "Runs the per-chunk transfer denoising loop over scheduler timesteps." diff --git a/diffusers/modular_pipelines/cosmos/encoders.py b/diffusers/modular_pipelines/cosmos/encoders.py deleted file mode 100644 index 81f181d4e5a28769f075dd2c46bdce7b479c220e..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/cosmos/encoders.py +++ /dev/null @@ -1,1056 +0,0 @@ -import torch -from transformers import AutoTokenizer - -from ...configuration_utils import FrozenDict -from ...models.autoencoders.autoencoder_kl_wan import AutoencoderKLWan -from ...pipelines.cosmos.pipeline_cosmos3_omni import ( - _ACTION_RESOLUTION_BINS, - CosmosActionCondition, -) -from ...utils import logging -from ...video_processor import VideoProcessor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, ConfigSpec, InputParam, OutputParam -from .modular_pipeline import Cosmos3OmniModularPipeline - - -logger = logging.get_logger(__name__) - - -# Transfer conditions on control signals (edge/blur/depth/seg/wsm), so it uses its own system prompt instead of the -# plain image/video ones. Defined here (not on the task pipeline) so the transfer text block is self-contained. -_SYSTEM_PROMPT_TRANSFER = ( - "You are a helpful assistant that generates images or videos following the user's instructions" - " and control signals (edge maps, blur, depth, or segmentation)." -) - - -class Cosmos3TextEncoderStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Prepares non-action prompt token IDs for downstream text-segment packing." - - @staticmethod - def _check_inputs(block_state) -> None: - prompt = block_state.prompt - negative_prompt = block_state.negative_prompt - - if not isinstance(prompt, str): - raise ValueError( - f"`prompt` must be a str; batched prompts are not supported, got {type(prompt).__name__}." - ) - if negative_prompt is not None and not isinstance(negative_prompt, str): - raise ValueError( - "`negative_prompt` must be a str or None; batched prompts are not supported, " - f"got {type(negative_prompt).__name__}." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_tokenizer", AutoTokenizer), - ] - - @property - def expected_configs(self) -> list[ConfigSpec]: - return [ - ConfigSpec(name="default_use_system_prompt", default=True), - ConfigSpec(name="enable_safety_checker", default=True), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt", description="The text prompt that guides Cosmos3 generation."), - InputParam.template( - "negative_prompt", description="The negative text prompt used for classifier-free guidance." - ), - InputParam(name="num_frames", type_hint=int, default=None, description="Number of frames to generate."), - InputParam( - name="height", - type_hint=int, - default=None, - description="Height of the generated video or image in pixels.", - ), - InputParam( - name="width", - type_hint=int, - default=None, - description="Width of the generated video or image in pixels.", - ), - InputParam(name="fps", type_hint=float, default=24.0, description="Frame rate of the generated video."), - InputParam( - name="use_system_prompt", - type_hint=bool | None, - default=None, - description="Whether to prepend the Cosmos3 system prompt.", - ), - InputParam( - name="add_resolution_template", - type_hint=bool, - default=True, - description="Whether to add resolution metadata to the prompt.", - ), - InputParam( - name="add_duration_template", - type_hint=bool, - default=True, - description="Whether to add duration metadata to the prompt.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("num_frames", type_hint=int, description="Number of frames to generate."), - OutputParam("height", type_hint=int, description="Height of the generated video or image in pixels."), - OutputParam("width", type_hint=int, description="Width of the generated video or image in pixels."), - OutputParam("cond_input_ids", type_hint=torch.Tensor, description="Token IDs for the conditional prompt."), - OutputParam( - "uncond_input_ids", type_hint=torch.Tensor, description="Token IDs for the unconditional prompt." - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - if block_state.num_frames is None: - block_state.num_frames = 189 - if block_state.height is None: - block_state.height = 720 - if block_state.width is None: - block_state.width = 1280 - if block_state.use_system_prompt is None: - block_state.use_system_prompt = components.config.default_use_system_prompt - - self._check_inputs(block_state) - if components.requires_safety_checker: - if getattr(components, "safety_checker", None) is None: - raise ValueError( - "Cosmos3 requires a safety checker by default. Call `pipe.enable_safety_checker()` to load it " - "(or pass your own), or opt out explicitly with `pipe.disable_safety_checker()`." - ) - device = components._execution_device - components.safety_checker.to(device) - try: - if not components.safety_checker.check_text_safety(block_state.prompt): - raise ValueError( - f"Cosmos Guardrail detected unsafe text in the prompt: {block_state.prompt}. " - "Please ensure that the prompt abides by the NVIDIA Open Model License Agreement." - ) - finally: - components.safety_checker.to("cpu") - - block_state.cond_input_ids, block_state.uncond_input_ids = components.tokenize_prompt( - block_state.prompt, - block_state.negative_prompt, - num_frames=block_state.num_frames, - height=block_state.height, - width=block_state.width, - fps=block_state.fps, - use_system_prompt=block_state.use_system_prompt, - add_resolution_template=block_state.add_resolution_template, - add_duration_template=block_state.add_duration_template, - action_mode=None, - action_view_point=None, - ) - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3TransferTextStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Tokenizes the transfer prompt with the transfer system prompt. Transfer prompts are pre-upsampled JSON " - "captions passed through verbatim (no resolution/duration templates), so this is self-contained and does " - "not reuse the standard text step." - ) - - @staticmethod - def _check_inputs(block_state) -> None: - prompt = block_state.prompt - negative_prompt = block_state.negative_prompt - - if not isinstance(prompt, (str, list)) or ( - isinstance(prompt, list) and not all(isinstance(p, str) for p in prompt) - ): - raise ValueError(f"`prompt` must be a str or list of str, got {type(prompt).__name__}.") - if negative_prompt is not None and not isinstance(negative_prompt, (str, list)): - raise ValueError( - f"`negative_prompt` must be a str, list of str, or None, got {type(negative_prompt).__name__}." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_tokenizer", AutoTokenizer), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="prompt", - type_hint=str, - required=True, - description="The text prompt that guides Cosmos3 generation.", - ), - InputParam( - name="negative_prompt", - type_hint=str, - default=None, - description="The negative text prompt used for classifier-free guidance.", - ), - InputParam( - name="use_system_prompt", - type_hint=bool, - default=True, - description="Whether to prepend the Cosmos3 transfer system prompt.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("cond_input_ids", type_hint=torch.Tensor, description="Token IDs for the conditional prompt."), - OutputParam( - "uncond_input_ids", type_hint=torch.Tensor, description="Token IDs for the unconditional prompt." - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - self._check_inputs(block_state) - - if isinstance(block_state.prompt, list): - block_state.prompt = block_state.prompt[0] - if isinstance(block_state.negative_prompt, list): - block_state.negative_prompt = block_state.negative_prompt[0] - - if components.requires_safety_checker: - if getattr(components, "safety_checker", None) is None: - raise ValueError( - "Cosmos3 requires a safety checker by default. Call `pipe.enable_safety_checker()` to load it " - "(or pass your own), or opt out explicitly with `pipe.disable_safety_checker()`." - ) - device = components._execution_device - components.safety_checker.to(device) - try: - if not components.safety_checker.check_text_safety(block_state.prompt): - raise ValueError( - f"Cosmos Guardrail detected unsafe text in the prompt: {block_state.prompt}. " - "Please ensure that the prompt abides by the NVIDIA Open Model License Agreement." - ) - finally: - components.safety_checker.to("cpu") - - # Transfer prompts are pre-upsampled JSON captions: tokenize them verbatim (no resolution/duration templates) - # under the transfer system prompt. Kept self-contained here rather than adding a flag to the standard step. - negative_prompt = block_state.negative_prompt if block_state.negative_prompt is not None else "" - special_tokens = components.llm_special_tokens - - def _tokenize(text: str) -> list[int]: - conversations = [] - if block_state.use_system_prompt: - conversations.append({"role": "system", "content": _SYSTEM_PROMPT_TRANSFER}) - conversations.append({"role": "user", "content": text}) - encoding = components.text_tokenizer.apply_chat_template( - conversations, - tokenize=True, - add_generation_prompt=True, - add_vision_id=False, - return_dict=True, - ) - return list(encoding.input_ids) + [ - special_tokens["eos_token_id"], - special_tokens["start_of_generation"], - ] - - block_state.cond_input_ids = _tokenize(block_state.prompt) - block_state.uncond_input_ids = _tokenize(negative_prompt) - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3DistilledTextEncoderStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Prepares distilled prompt token IDs. Classifier-free guidance is baked into the weights, so " - "`negative_prompt` is not exposed and the unconditional branch is derived from an empty prompt." - ) - - @staticmethod - def _check_inputs(block_state) -> None: - prompt = block_state.prompt - if not isinstance(prompt, str): - raise ValueError( - f"`prompt` must be a str; batched prompts are not supported, got {type(prompt).__name__}." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_tokenizer", AutoTokenizer), - ] - - @property - def expected_configs(self) -> list[ConfigSpec]: - return [ - ConfigSpec(name="default_use_system_prompt", default=True), - ConfigSpec(name="enable_safety_checker", default=True), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt", description="The text prompt that guides Cosmos3 generation."), - InputParam(name="num_frames", type_hint=int, default=None, description="Number of frames to generate."), - InputParam( - name="height", - type_hint=int, - default=None, - description="Height of the generated video or image in pixels.", - ), - InputParam( - name="width", - type_hint=int, - default=None, - description="Width of the generated video or image in pixels.", - ), - InputParam(name="fps", type_hint=float, default=24.0, description="Frame rate of the generated video."), - InputParam( - name="use_system_prompt", - type_hint=bool, - default=True, - description="Whether to prepend the Cosmos3 system prompt.", - ), - InputParam( - name="add_resolution_template", - type_hint=bool, - default=True, - description="Whether to add resolution metadata to the prompt.", - ), - InputParam( - name="add_duration_template", - type_hint=bool, - default=True, - description="Whether to add duration metadata to the prompt.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("num_frames", type_hint=int, description="Number of frames to generate."), - OutputParam("height", type_hint=int, description="Height of the generated video or image in pixels."), - OutputParam("width", type_hint=int, description="Width of the generated video or image in pixels."), - OutputParam("cond_input_ids", type_hint=torch.Tensor, description="Token IDs for the conditional prompt."), - OutputParam( - "uncond_input_ids", - type_hint=torch.Tensor, - description="Token IDs for the unconditional prompt (empty prompt; guidance is baked in).", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - if block_state.num_frames is None: - block_state.num_frames = 189 - if block_state.height is None: - block_state.height = 720 - if block_state.width is None: - block_state.width = 1280 - - self._check_inputs(block_state) - if components.requires_safety_checker: - if getattr(components, "safety_checker", None) is None: - raise ValueError( - "Cosmos3 requires a safety checker by default. Call `pipe.enable_safety_checker()` to load it " - "(or pass your own), or opt out explicitly with `pipe.disable_safety_checker()`." - ) - device = components._execution_device - components.safety_checker.to(device) - try: - if not components.safety_checker.check_text_safety(block_state.prompt): - raise ValueError( - f"Cosmos Guardrail detected unsafe text in the prompt: {block_state.prompt}. " - "Please ensure that the prompt abides by the NVIDIA Open Model License Agreement." - ) - finally: - components.safety_checker.to("cpu") - - # Guidance is baked into distilled weights: the unconditional branch is built from an empty prompt - # (negative_prompt is not a user-facing input) so the downstream text-segment packing contract still holds. - block_state.cond_input_ids, block_state.uncond_input_ids = components.tokenize_prompt( - block_state.prompt, - None, - num_frames=block_state.num_frames, - height=block_state.height, - width=block_state.width, - fps=block_state.fps, - use_system_prompt=block_state.use_system_prompt, - add_resolution_template=block_state.add_resolution_template, - add_duration_template=block_state.add_duration_template, - action_mode=None, - action_view_point=None, - ) - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3ActionTextStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Prepares action prompt token IDs from prompt + action metadata." - - @staticmethod - def _check_inputs(block_state) -> None: - prompt = block_state.prompt - negative_prompt = block_state.negative_prompt - action = block_state.action - num_frames = block_state.num_frames - height = block_state.height - width = block_state.width - if not isinstance(prompt, str): - raise ValueError( - f"`prompt` must be a str; batched prompts are not supported, got {type(prompt).__name__}." - ) - if negative_prompt is not None and not isinstance(negative_prompt, str): - raise ValueError( - "`negative_prompt` must be a str or None; batched prompts are not supported, " - f"got {type(negative_prompt).__name__}." - ) - if action is None: - raise ValueError("`action` is required for Cosmos3ActionTextStep.") - if action.image is None and action.video is None: - raise ValueError("`action.image` or `action.video` must be provided for action-conditioned generation.") - if num_frames is not None: - raise ValueError("`num_frames` has to be None if action is not None.") - if height is not None or width is not None: - raise ValueError("`height` and `width` have to be None if action is not None.") - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_tokenizer", AutoTokenizer), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 16, "resample": "bilinear"}), - default_creation_method="from_config", - ), - ] - - @property - def expected_configs(self) -> list[ConfigSpec]: - return [ - ConfigSpec(name="default_use_system_prompt", default=True), - ConfigSpec(name="enable_safety_checker", default=True), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt", description="The text prompt that guides Cosmos3 generation."), - InputParam.template( - "negative_prompt", description="The negative text prompt used for classifier-free guidance." - ), - InputParam( - name="action", - type_hint=CosmosActionCondition, - required=True, - description="Action-conditioning metadata and its reference visual input.", - ), - InputParam(name="num_frames", type_hint=int, default=None, description="Number of frames to generate."), - InputParam( - name="height", - type_hint=int, - default=None, - description="Height of the generated video or image in pixels.", - ), - InputParam( - name="width", - type_hint=int, - default=None, - description="Width of the generated video or image in pixels.", - ), - InputParam(name="fps", type_hint=float, default=24.0, description="Frame rate of the generated video."), - InputParam( - name="use_system_prompt", - type_hint=bool | None, - default=None, - description="Whether to prepend the Cosmos3 system prompt.", - ), - InputParam( - name="add_resolution_template", - type_hint=bool, - default=True, - description="Whether to add resolution metadata to the prompt.", - ), - InputParam( - name="add_duration_template", - type_hint=bool, - default=True, - description="Whether to add duration metadata to the prompt.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("action_mode", type_hint=str, description="Requested action-generation mode."), - OutputParam("num_frames", type_hint=int, description="Number of frames to generate."), - OutputParam("height", type_hint=int, description="Height of the generated video or image in pixels."), - OutputParam("width", type_hint=int, description="Width of the generated video or image in pixels."), - OutputParam("cond_input_ids", type_hint=torch.Tensor, description="Token IDs for the conditional prompt."), - OutputParam( - "uncond_input_ids", type_hint=torch.Tensor, description="Token IDs for the unconditional prompt." - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - self._check_inputs(block_state) - if block_state.use_system_prompt is None: - block_state.use_system_prompt = components.config.default_use_system_prompt - - action = block_state.action - block_state.action_mode = action.mode - block_state.num_frames = action.chunk_size + 1 - conditioning_clip = [action.image] if action.image is not None else action.video - probe = components.video_processor.preprocess_video(conditioning_clip) - source_h, source_w = int(probe.shape[-2]), int(probe.shape[-1]) - resolution_key = str(action.resolution_tier) - block_state.height, block_state.width = VideoProcessor.classify_height_width_bin( - source_h, source_w, ratios=_ACTION_RESOLUTION_BINS[resolution_key] - ) - - if components.requires_safety_checker: - if getattr(components, "safety_checker", None) is None: - raise ValueError( - "Cosmos3 requires a safety checker by default. Call `pipe.enable_safety_checker()` to load it " - "(or pass your own), or opt out explicitly with `pipe.disable_safety_checker()`." - ) - device = components._execution_device - components.safety_checker.to(device) - try: - if not components.safety_checker.check_text_safety(block_state.prompt): - raise ValueError( - f"Cosmos Guardrail detected unsafe text in the prompt: {block_state.prompt}. " - "Please ensure that the prompt abides by the NVIDIA Open Model License Agreement." - ) - finally: - components.safety_checker.to("cpu") - - block_state.cond_input_ids, block_state.uncond_input_ids = components.tokenize_prompt( - block_state.prompt, - block_state.negative_prompt, - num_frames=block_state.num_frames, - height=block_state.height, - width=block_state.width, - fps=block_state.fps, - use_system_prompt=block_state.use_system_prompt, - add_resolution_template=block_state.add_resolution_template, - add_duration_template=block_state.add_duration_template, - action_mode=block_state.action_mode, - action_view_point=action.view_point, - ) - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3ImageVaeEncoderStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Encodes non-action image-to-video conditioning into Cosmos3 vision latents." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLWan), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 16, "resample": "bilinear"}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="image", default=None, description="Reference image for image-to-video conditioning."), - InputParam(name="num_frames", type_hint=int, required=True, description="Number of frames to generate."), - InputParam( - name="height", type_hint=int, required=True, description="Height of the generated video in pixels." - ), - InputParam( - name="width", type_hint=int, required=True, description="Width of the generated video in pixels." - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "x0_tokens_vision", - type_hint=torch.Tensor, - description="Vision latents encoded from the conditioning image or video.", - ), - OutputParam( - "vision_condition_frames", - type_hint=list[int], - description="Latent-frame indexes fixed by visual conditioning.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - dtype = components.vae.dtype - - if block_state.image is None: - raise ValueError("`Cosmos3ImageVaeEncoderStep` requires an `image` input.") - if block_state.num_frames == 1: - raise ValueError( - "`image` conditioning requires `num_frames` > 1; image-to-image generation is not supported." - ) - if block_state.num_frames < 1: - raise ValueError(f"`num_frames` must be >= 1, got {block_state.num_frames}.") - - sf = int(components.vae.config.scale_factor_spatial) - if block_state.height % sf != 0 or block_state.width % sf != 0: - raise ValueError( - f"`height` and `width` must be multiples of {sf}, got ({block_state.height}, {block_state.width})." - ) - - conditioning_frame_2d = components.video_processor.preprocess( - block_state.image, height=block_state.height, width=block_state.width - ).to(device=device, dtype=dtype) - - vision_tensor = torch.zeros( - 1, - 3, - block_state.num_frames, - block_state.height, - block_state.width, - dtype=dtype, - device=device, - ) - vision_tensor[:, :, 0] = conditioning_frame_2d - vision_tensor[:, :, 1:] = conditioning_frame_2d.unsqueeze(2).expand(-1, -1, block_state.num_frames - 1, -1, -1) - - block_state.x0_tokens_vision = components._encode_video(vision_tensor).contiguous().float() - block_state.vision_condition_frames = [0] - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3VideoVaeEncoderStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return "Encodes non-action video conditioning into Cosmos3 vision latents." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLWan), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 16, "resample": "bilinear"}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="video", default=None, description="Reference video for video-to-video conditioning."), - InputParam( - name="condition_frame_indexes_vision", - type_hint=tuple[int, ...] | list[int], - default=(0, 1), - description="Latent-frame indexes to preserve from the conditioning video.", - ), - InputParam( - name="condition_video_keep", - type_hint=str, - default="first", - description="Which end of a longer conditioning video to use: `first` or `last`.", - ), - InputParam(name="num_frames", type_hint=int, required=True, description="Number of frames to generate."), - InputParam( - name="height", type_hint=int, required=True, description="Height of the generated video in pixels." - ), - InputParam( - name="width", type_hint=int, required=True, description="Width of the generated video in pixels." - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "x0_tokens_vision", - type_hint=torch.Tensor, - description="Vision latents encoded from the conditioning image or video.", - ), - OutputParam( - "vision_condition_frames", - type_hint=list[int], - description="Latent-frame indexes fixed by visual conditioning.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - dtype = components.vae.dtype - - if block_state.video is None: - raise ValueError("`Cosmos3VideoVaeEncoderStep` requires a `video` input.") - if block_state.num_frames == 1: - raise ValueError("`video` conditioning requires `num_frames` > 1.") - if block_state.num_frames < 1: - raise ValueError(f"`num_frames` must be >= 1, got {block_state.num_frames}.") - - sf = int(components.vae.config.scale_factor_spatial) - if block_state.height % sf != 0 or block_state.width % sf != 0: - raise ValueError( - f"`height` and `width` must be multiples of {sf}, got ({block_state.height}, {block_state.width})." - ) - - if not isinstance(block_state.condition_frame_indexes_vision, (list, tuple)) or isinstance( - block_state.condition_frame_indexes_vision, (str, bytes) - ): - raise ValueError( - "`condition_frame_indexes_vision` must be a list/tuple of non-negative ints, e.g. [0, 1]; got " - f"{block_state.condition_frame_indexes_vision!r}." - ) - if not all(isinstance(index, int) and index >= 0 for index in block_state.condition_frame_indexes_vision): - raise ValueError( - "`condition_frame_indexes_vision` must be a list/tuple of non-negative ints, e.g. [0, 1]; got " - f"{block_state.condition_frame_indexes_vision!r}." - ) - if block_state.condition_video_keep not in {"first", "last"}: - raise ValueError("`condition_video_keep` must be either 'first' or 'last'.") - - indexes = tuple(block_state.condition_frame_indexes_vision) - if not indexes: - raise ValueError("`condition_frame_indexes_vision` must contain at least one index.") - latent_t = (block_state.num_frames - 1) // int(components.vae.config.scale_factor_temporal) + 1 - if max(indexes) >= latent_t: - raise ValueError( - f"`condition_frame_indexes_vision` {indexes} contains an index outside the latent timeline " - f"(latent_frames={latent_t} for num_frames={block_state.num_frames})." - ) - - condition_indexes_vision = indexes - conditioning_frames_3d = components.video_processor.preprocess_video( - block_state.video, height=block_state.height, width=block_state.width - ).to(device=device, dtype=dtype) - temporal_compression = int(components.vae.config.scale_factor_temporal) - max_cond_frames = max(condition_indexes_vision) * temporal_compression + 1 - if block_state.condition_video_keep == "first": - conditioning_frames_3d = conditioning_frames_3d[:, :, :max_cond_frames] - else: - conditioning_frames_3d = conditioning_frames_3d[:, :, -max_cond_frames:] - - vision_tensor = torch.zeros( - 1, - 3, - block_state.num_frames, - block_state.height, - block_state.width, - dtype=dtype, - device=device, - ) - t_fill = min(conditioning_frames_3d.shape[2], block_state.num_frames) - vision_tensor[:, :, :t_fill] = conditioning_frames_3d[:, :, :t_fill] - if t_fill < block_state.num_frames: - vision_tensor[:, :, t_fill:] = vision_tensor[:, :, t_fill - 1 : t_fill].expand( - -1, -1, block_state.num_frames - t_fill, -1, -1 - ) - vision_condition_frames = list(condition_indexes_vision) - - block_state.x0_tokens_vision = components._encode_video(vision_tensor).contiguous().float() - block_state.vision_condition_frames = vision_condition_frames - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3TransferChunkVaeEncoderStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Per-chunk transfer VAE encode: slices + pads this chunk's control maps, seeds the target's conditioning " - "frames (first chunk from the input video, later chunks from the previous chunk's tail), and encodes both " - "the controls and the seeded target into clean Cosmos3 vision latents. Runs inside the autoregressive " - "chunk loop." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLWan), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 16, "resample": "bilinear"}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="chunk_id", type_hint=int, default=0, description="Index of the current chunk."), - InputParam( - name="previous_output", - default=None, - description="Decoded pixels of the previous chunk, used to seed later chunks.", - ), - InputParam( - name="control_frames", - type_hint=dict, - required=True, - description="Preprocessed, time-padded control maps in canonical hint order.", - ), - InputParam(name="chunk_frames", type_hint=int, required=True, description="Pixel frames per chunk."), - InputParam( - name="total_frames", type_hint=int, required=True, description="Total number of output frames." - ), - InputParam(name="stride", type_hint=int, required=True, description="Frame stride between chunks."), - InputParam( - name="height", type_hint=int, required=True, description="Height of the generated video in pixels." - ), - InputParam( - name="width", type_hint=int, required=True, description="Width of the generated video in pixels." - ), - InputParam( - name="video", - default=None, - description="Optional input video that seeds the first chunk's conditioning.", - ), - InputParam( - name="num_first_chunk_conditional_frames", - type_hint=int, - default=0, - description="Number of frames the first chunk reuses from the input video.", - ), - InputParam( - name="num_conditional_frames", - type_hint=int, - default=1, - description="Number of frames each later chunk reuses from the previous chunk's tail.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "control_latents", - type_hint=list[torch.Tensor], - description="Clean control latents for this chunk, one per hint in canonical order.", - ), - OutputParam( - "x0_tokens_vision", - type_hint=torch.Tensor, - description="Clean target vision latents encoded from the seeded target frames.", - ), - OutputParam( - "current_conditional_frames", - type_hint=int, - description="Number of pixel frames actually used to seed this chunk's target.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - dtype = components.vae.dtype - - chunk_id = block_state.chunk_id - chunk_frames = block_state.chunk_frames - height = block_state.height - width = block_state.width - - # Slice this chunk's window out of the (padded) control maps and reflect-pad it up to a full chunk (repeat the - # last frame once too short to keep reflecting). control_frames is already in canonical hint order. - start_frame = chunk_id * block_state.stride - end_frame = min(start_frame + chunk_frames, block_state.total_frames) - chunk_controls = [] - for frames in block_state.control_frames.values(): - frames = frames[:, :, start_frame:end_frame] - while frames.shape[2] < chunk_frames: - pad_len = min(frames.shape[2] - 1, chunk_frames - frames.shape[2]) - if pad_len <= 0: - pad_frame = frames[:, :, -1:].repeat(1, 1, chunk_frames - frames.shape[2], 1, 1) - frames = torch.cat([frames, pad_frame], dim=2) - break - frames = torch.cat([frames, frames.flip(dims=[2])[:, :, :pad_len]], dim=2) - chunk_controls.append(frames) - - # Seed the target with conditioning frames (first chunk from the input video, later chunks from the - # previous chunk's tail), repeat-padding the remaining frames so the whole clip is well-defined. - target = torch.zeros(1, 3, chunk_frames, height, width, device=device, dtype=dtype) - current_conditional_frames = 0 - if chunk_id == 0 and block_state.num_first_chunk_conditional_frames > 0 and block_state.video is not None: - input_frames = components.video_processor.preprocess_video( - block_state.video, height=height, width=width - ).to(device=device, dtype=dtype) - current_conditional_frames = min( - block_state.num_first_chunk_conditional_frames, input_frames.shape[2], chunk_frames - ) - if current_conditional_frames > 0: - target[:, :, :current_conditional_frames] = input_frames[:, :, :current_conditional_frames] - elif chunk_id > 0 and block_state.previous_output is not None: - current_conditional_frames = min( - block_state.num_conditional_frames, block_state.previous_output.shape[2], chunk_frames - ) - if current_conditional_frames > 0: - target[:, :, :current_conditional_frames] = block_state.previous_output[ - :, :, -current_conditional_frames: - ].to(device=device, dtype=dtype) - if 0 < current_conditional_frames < chunk_frames: - fill = target[:, :, current_conditional_frames - 1 : current_conditional_frames] - target[:, :, current_conditional_frames:] = fill.expand( - -1, -1, chunk_frames - current_conditional_frames, -1, -1 - ) - - block_state.control_latents = [components._encode_video(ctrl).contiguous().float() for ctrl in chunk_controls] - block_state.x0_tokens_vision = components._encode_video(target).contiguous().float() - block_state.current_conditional_frames = current_conditional_frames - - self.set_block_state(state, block_state) - return components, state - - -class Cosmos3ActionVisionVaeEncoderStep(ModularPipelineBlocks): - model_name = "cosmos3-omni" - - @property - def description(self) -> str: - return ( - "Prepares action-conditioned vision latents and action frame metadata. " - "Only the action visual reference (image/video) is VAE-encoded; action vectors are handled separately." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLWan), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 16, "resample": "bilinear"}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="action", - type_hint=CosmosActionCondition, - required=True, - description="Action-conditioning metadata and its reference visual input.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "x0_tokens_vision", - type_hint=torch.Tensor, - description="Vision latents encoded from the conditioning image or video.", - ), - OutputParam( - "vision_condition_frames", - type_hint=list[int], - description="Latent-frame indexes fixed by visual conditioning.", - ), - OutputParam( - "action_condition_frame_indexes", - type_hint=list[int], - description="Action-frame indexes fixed by action conditioning.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - dtype = components.vae.dtype - - action = block_state.action - target_frames = action.chunk_size + 1 - conditioning_clip = [action.image] if action.image is not None else action.video - vision_tensor, action_image_size, _, _ = components._prepare_action_video_conditioning( - conditioning_clip, - action.resolution_tier, - target_frames, - device=device, - dtype=dtype, - ) - - if action.mode == "forward_dynamics": - vision_condition_frames = [0] - action_condition_frame_indexes = list(range(action.chunk_size)) - elif action.mode == "policy": - vision_condition_frames = [0] - action_condition_frame_indexes = [] - elif action.mode == "inverse_dynamics": - latent_frames = (target_frames - 1) // int(components.vae.config.scale_factor_temporal) + 1 - vision_condition_frames = list(range(latent_frames)) - action_condition_frame_indexes = [] - else: - raise ValueError( - f"Unsupported action_mode={action.mode!r}; expected one of ['forward_dynamics', 'inverse_dynamics', 'policy']." - ) - - x0_tokens_vision = components._encode_video(vision_tensor).contiguous().float() - if action_image_size is not None: - x0_tokens_vision = components._remove_action_video_padding_from_latent(x0_tokens_vision, action_image_size) - - block_state.x0_tokens_vision = x0_tokens_vision - block_state.vision_condition_frames = vision_condition_frames - block_state.action_condition_frame_indexes = action_condition_frame_indexes - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py b/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py deleted file mode 100644 index 205b0256d8f6ed212080dc56d644f837c2a77b77..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py +++ /dev/null @@ -1,1301 +0,0 @@ -import torch - -from ..modular_pipeline import ( - AutoPipelineBlocks, - ConditionalPipelineBlocks, - PipelineState, - SequentialPipelineBlocks, -) -from ..modular_pipeline_utils import InputParam, OutputParam -from .after_decode import Cosmos3ActionOutputStep -from .before_denoise import ( - Cosmos3ActionDenoiseInputStep, - Cosmos3ActionPackSequenceStep, - Cosmos3ActionPrepareLatentsStep, - Cosmos3PrepareTextSegmentsStep, - Cosmos3SetTimestepsStep, - Cosmos3SoundDenoiseInputStep, - Cosmos3SoundPackSequenceStep, - Cosmos3SoundPrepareLatentsStep, - Cosmos3TransferPackSequenceStep, - Cosmos3TransferPrepareLatentsStep, - Cosmos3TransferSetTimestepsStep, - Cosmos3VisionDenoiseInputStep, - Cosmos3VisionPackSequenceStep, - Cosmos3VisionPrepareLatentsStep, -) -from .before_encoder import Cosmos3TransferSetupStep -from .decoders import ( - Cosmos3SoundDecodeStep, - Cosmos3TransferDecodeChunkStep, - Cosmos3TransferStitchStep, - Cosmos3VideoDecodeStep, -) -from .denoise import ( - Cosmos3TransferDenoiseStep, - Cosmos3VisionActionDenoiseStep, - Cosmos3VisionDenoiseStep, - Cosmos3VisionSoundActionDenoiseStep, - Cosmos3VisionSoundDenoiseStep, -) -from .encoders import ( - Cosmos3ActionTextStep, - Cosmos3ActionVisionVaeEncoderStep, - Cosmos3ImageVaeEncoderStep, - Cosmos3TextEncoderStep, - Cosmos3TransferChunkVaeEncoderStep, - Cosmos3TransferTextStep, - Cosmos3VideoVaeEncoderStep, -) -from .modular_pipeline import Cosmos3OmniModularPipeline - - -# auto_docstring -class Cosmos3TransferTextBlocks(SequentialPipelineBlocks): - """ - Transfer text branch: resolves the control-video chunk geometry, then tokenizes the (pre-upsampled) prompt in - transfer mode using the per-chunk frame count. - - Components: - video_processor (`VideoProcessor`) text_tokenizer (`AutoTokenizer`) - - Inputs: - control_videos (`dict`): - Mapping of hint name (edge/blur/depth/seg/wsm) to the control video for that modality. - height (`int`, *optional*): - Height of the generated video in pixels. - width (`int`, *optional*): - Width of the generated video in pixels. - num_frames (`int`, *optional*): - Optional cap on the number of output frames (defaults to the control video length). - num_video_frames_per_chunk (`int`, *optional*): - Number of pixel frames generated per autoregressive chunk. - num_conditional_frames (`int`, *optional*, defaults to 1): - Number of frames each chunk reuses from the previous chunk's tail. - prompt (`str`): - The text prompt that guides Cosmos3 generation. - negative_prompt (`str`, *optional*): - The negative text prompt used for classifier-free guidance. - use_system_prompt (`bool`, *optional*, defaults to True): - Whether to prepend the Cosmos3 transfer system prompt. - - Outputs: - height (`int`): - Resolved output height in pixels. - width (`int`): - Resolved output width in pixels. - control_frames (`dict`): - Preprocessed, time-padded control maps in canonical hint order. - total_frames (`int`): - Total number of output frames to generate. - chunk_frames (`int`): - Number of pixel frames per autoregressive chunk. - num_chunks (`int`): - Number of autoregressive chunks. - stride (`int`): - Frame stride between consecutive chunks. - cond_input_ids (`Tensor`): - Token IDs for the conditional prompt. - uncond_input_ids (`Tensor`): - Token IDs for the unconditional prompt. - """ - - model_name = "cosmos3-omni" - block_classes = [Cosmos3TransferSetupStep, Cosmos3TransferTextStep] - block_names = ["setup", "transfer_text"] - - @property - def description(self): - return ( - "Transfer text branch: resolves the control-video chunk geometry, then tokenizes the (pre-upsampled) " - "prompt in transfer mode using the per-chunk frame count." - ) - - -# auto_docstring -class Cosmos3AutoTextEncoderStep(AutoPipelineBlocks): - """ - Auto text encoder block for Cosmos3. - - Cosmos3TransferTextBlocks runs when control_videos are provided. - - Cosmos3ActionTextStep runs when action is provided. - - Cosmos3TextEncoderStep runs otherwise. - - Components: - video_processor (`VideoProcessor`) text_tokenizer (`AutoTokenizer`) - - Configs: - default_use_system_prompt (default: True) enable_safety_checker (default: True) - - Inputs: - control_videos (`dict`, *optional*): - Mapping of hint name (edge/blur/depth/seg/wsm) to the control video for that modality. - height (`int`, *optional*): - Height of the generated video in pixels. - width (`int`, *optional*): - Width of the generated video in pixels. - num_frames (`int`, *optional*): - Optional cap on the number of output frames (defaults to the control video length). - num_video_frames_per_chunk (`int`, *optional*): - Number of pixel frames generated per autoregressive chunk. - num_conditional_frames (`int`, *optional*, defaults to 1): - Number of frames each chunk reuses from the previous chunk's tail. - prompt (`str`): - The text prompt that guides Cosmos3 generation. - negative_prompt (`str`, *optional*): - The negative text prompt used for classifier-free guidance. - use_system_prompt (`bool`, *optional*, defaults to True): - Whether to prepend the Cosmos3 transfer system prompt. - action (`CosmosActionCondition`, *optional*): - Action-conditioning metadata and its reference visual input. - fps (`float`, *optional*, defaults to 24.0): - Frame rate of the generated video. - add_resolution_template (`bool`, *optional*, defaults to True): - Whether to add resolution metadata to the prompt. - add_duration_template (`bool`, *optional*, defaults to True): - Whether to add duration metadata to the prompt. - - Outputs: - height (`int`): - Resolved output height in pixels. - width (`int`): - Resolved output width in pixels. - control_frames (`dict`): - Preprocessed, time-padded control maps in canonical hint order. - total_frames (`int`): - Total number of output frames to generate. - chunk_frames (`int`): - Number of pixel frames per autoregressive chunk. - num_chunks (`int`): - Number of autoregressive chunks. - stride (`int`): - Frame stride between consecutive chunks. - cond_input_ids (`Tensor`): - Token IDs for the conditional prompt. - uncond_input_ids (`Tensor`): - Token IDs for the unconditional prompt. - action_mode (`str`): - Requested action-generation mode. - num_frames (`int`): - Number of frames to generate. - """ - - model_name = "cosmos3-omni" - block_classes = [Cosmos3TransferTextBlocks, Cosmos3ActionTextStep, Cosmos3TextEncoderStep] - block_names = ["transfer_text", "action_text", "text"] - block_trigger_inputs = ["control_videos", "action", None] - - @property - def description(self): - return ( - "Auto text encoder block for Cosmos3.\n" - + " - Cosmos3TransferTextBlocks runs when control_videos are provided.\n" - + " - Cosmos3ActionTextStep runs when action is provided.\n" - + " - Cosmos3TextEncoderStep runs otherwise." - ) - - -# auto_docstring -class Cosmos3AutoVaeEncoderStep(ConditionalPipelineBlocks): - """ - Auto VAE conditioning block for Cosmos3. - - Cosmos3ActionVisionVaeEncoderStep runs when action is provided. - - Cosmos3VideoVaeEncoderStep runs for the non-action video path. - - Cosmos3ImageVaeEncoderStep runs for the non-action image path. - - when no action, image, or video conditioning is provided, this block is skipped. - - Components: - vae (`AutoencoderKLWan`) video_processor (`VideoProcessor`) - - Inputs: - action (`CosmosActionCondition`, *optional*): - Action-conditioning metadata and its reference visual input. - video (`None`, *optional*): - Reference video for video-to-video conditioning. - condition_frame_indexes_vision (`tuple | list`, *optional*, defaults to (0, 1)): - Latent-frame indexes to preserve from the conditioning video. - condition_video_keep (`str`, *optional*, defaults to first): - Which end of a longer conditioning video to use: `first` or `last`. - num_frames (`int`, *optional*): - Number of frames to generate. - height (`int`, *optional*): - Height of the generated video in pixels. - width (`int`, *optional*): - Width of the generated video in pixels. - image (`None`, *optional*): - Reference image for image-to-video conditioning. - - Outputs: - x0_tokens_vision (`Tensor`): - Vision latents encoded from the conditioning image or video. - vision_condition_frames (`list`): - Latent-frame indexes fixed by visual conditioning. - action_condition_frame_indexes (`list`): - Action-frame indexes fixed by action conditioning. - """ - - model_name = "cosmos3-omni" - block_classes = [Cosmos3ActionVisionVaeEncoderStep, Cosmos3VideoVaeEncoderStep, Cosmos3ImageVaeEncoderStep] - block_names = ["action_conditioning", "video_conditioning", "image_conditioning"] - block_trigger_inputs = ["action", "video", "image", "control_videos"] - default_block_name = None - - def select_block(self, **kwargs) -> str | None: - action = kwargs.get("action") - image = kwargs.get("image") - video = kwargs.get("video") - # Transfer preprocesses/encodes its control maps inside the denoise chunk loop, so the standard VAE - # conditioning stage is skipped when control_videos drive the workflow. - if kwargs.get("control_videos") is not None: - return None - if action is not None: - if image is not None or video is not None: - raise ValueError( - "Pass action conditioning via `action.image` / `action.video`, not top-level image/video." - ) - return "action_conditioning" - if image is not None and video is not None: - raise ValueError("Pass either image or video, not both.") - if video is not None: - return "video_conditioning" - if image is not None: - return "image_conditioning" - return None - - @property - def description(self): - return ( - "Auto VAE conditioning block for Cosmos3.\n" - + " - Cosmos3ActionVisionVaeEncoderStep runs when action is provided.\n" - + " - Cosmos3VideoVaeEncoderStep runs for the non-action video path.\n" - + " - Cosmos3ImageVaeEncoderStep runs for the non-action image path.\n" - + " - when no action, image, or video conditioning is provided, this block is skipped." - ) - - -# auto_docstring -class Cosmos3AutoSoundDecodeStep(AutoPipelineBlocks): - """ - Auto sound decoder block for Cosmos3. - - Cosmos3SoundDecodeStep runs when sound_latents are present. - - if sound_latents are not provided, this block is skipped. - - Components: - sound_tokenizer (`Cosmos3AVAEAudioTokenizer`) - - Inputs: - sound_latents (`Tensor`, *optional*): - Denoised sound latents to decode. - - Outputs: - sound (`Tensor`): - Generated waveform. - sampling_rate (`int`): - Sample rate of the generated waveform in Hz. - """ - - model_name = "cosmos3-omni" - block_classes = [Cosmos3SoundDecodeStep] - block_names = ["decode"] - block_trigger_inputs = ["sound_latents"] - - @property - def description(self): - return ( - "Auto sound decoder block for Cosmos3.\n" - + " - Cosmos3SoundDecodeStep runs when sound_latents are present.\n" - + " - if sound_latents are not provided, this block is skipped." - ) - - -# auto_docstring -class Cosmos3DecodeStep(SequentialPipelineBlocks): - """ - Decodes denoised latents into modality outputs. - - Components: - vae (`AutoencoderKLWan`) video_processor (`VideoProcessor`) sound_tokenizer (`Cosmos3AVAEAudioTokenizer`) - - Inputs: - latents (`Tensor`): - Denoised vision latents to decode. - output_type (`str`, *optional*, defaults to pil): - Output format: 'pil', 'np', 'pt'. - sound_latents (`Tensor`, *optional*): - Denoised sound latents to decode. - - Outputs: - videos (`list`): - The generated videos. - sound (`Tensor`): - Generated waveform. - sampling_rate (`int`): - Sample rate of the generated waveform in Hz. - """ - - model_name = "cosmos3-omni" - block_classes = [Cosmos3VideoDecodeStep, Cosmos3AutoSoundDecodeStep] - block_names = ["video", "sound"] - - @property - def description(self) -> str: - return "Decodes denoised latents into modality outputs." - - -class Cosmos3AutoDecodeStep(ConditionalPipelineBlocks): - model_name = "cosmos3-omni" - block_classes = [Cosmos3TransferStitchStep, Cosmos3DecodeStep] - block_names = ["transfer", "standard"] - block_trigger_inputs = ["control_videos"] - default_block_name = "standard" - - def select_block(self, **kwargs) -> str | None: - if kwargs.get("control_videos") is not None: - return "transfer" - return "standard" - - @property - def description(self) -> str: - return ( - "Selects the Cosmos3 decode workflow.\n" - + " - Cosmos3TransferStitchStep stitches the decoded transfer chunks when control_videos are provided.\n" - + " - Cosmos3DecodeStep decodes the denoised latents otherwise." - ) - - -# auto_docstring -class Cosmos3VisionCoreDenoiseStep(SequentialPipelineBlocks): - """ - Runs the text-and-vision Cosmos3 denoising workflow. - - Components: - transformer (`Cosmos3OmniTransformer`) scheduler (`UniPCMultistepScheduler`) - - Configs: - use_native_flow_schedule (default: False) - - Inputs: - cond_input_ids (`None`): - Token IDs for the conditional prompt. - uncond_input_ids (`None`): - Token IDs for the unconditional prompt. - x0_tokens_vision (`Tensor`, *optional*): - Vision latents encoded from the conditioning image or video. - vision_condition_frames (`list`, *optional*): - Latent-frame indexes fixed by visual conditioning. - num_frames (`int`): - Number of frames to generate. - height (`int`): - Height of the generated video in pixels. - width (`int`): - Width of the generated video in pixels. - fps (`float`, *optional*, defaults to 24.0): - Frame rate of the generated video. - latents (`Tensor`, *optional*): - Pre-generated noisy vision latents. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`): - The number of denoising steps. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - guidance_scale (`float`, *optional*, defaults to 6.0): - Scale for classifier-free guidance. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "cosmos3-omni" - block_classes = [ - Cosmos3PrepareTextSegmentsStep, - Cosmos3VisionPrepareLatentsStep, - Cosmos3VisionPackSequenceStep, - Cosmos3VisionDenoiseInputStep, - Cosmos3SetTimestepsStep, - Cosmos3VisionDenoiseStep, - ] - block_names = [ - "prepare_text_segments", - "prepare_vision_latents", - "pack_vision_sequence", - "prepare_vision_denoiser_inputs", - "set_timesteps", - "denoise", - ] - - @property - def description(self): - return "Runs the text-and-vision Cosmos3 denoising workflow." - - @property - def outputs(self): - return [OutputParam.template("latents")] - - -# auto_docstring -class Cosmos3VisionSoundCoreDenoiseStep(SequentialPipelineBlocks): - """ - Runs the text, vision, and sound Cosmos3 denoising workflow. - - Components: - transformer (`Cosmos3OmniTransformer`) scheduler (`UniPCMultistepScheduler`) - - Configs: - use_native_flow_schedule (default: False) - - Inputs: - cond_input_ids (`None`): - Token IDs for the conditional prompt. - uncond_input_ids (`None`): - Token IDs for the unconditional prompt. - x0_tokens_vision (`Tensor`, *optional*): - Vision latents encoded from the conditioning image or video. - vision_condition_frames (`list`, *optional*): - Latent-frame indexes fixed by visual conditioning. - num_frames (`int`): - Number of frames to generate. - height (`int`): - Height of the generated video in pixels. - width (`int`): - Width of the generated video in pixels. - fps (`float`, *optional*, defaults to 24.0): - Frame rate of the generated video. - latents (`Tensor`, *optional*): - Pre-generated noisy vision latents. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`): - The number of denoising steps. - sound_latents (`Tensor`, *optional*): - Pre-generated noisy sound latents. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - guidance_scale (`float`, *optional*, defaults to 6.0): - Scale for classifier-free guidance. - - Outputs: - latents (`Tensor`): - Denoised latents. - sound_latents (`Tensor`): - Denoised sound latents. - """ - - model_name = "cosmos3-omni" - block_classes = [ - Cosmos3PrepareTextSegmentsStep, - Cosmos3VisionPrepareLatentsStep, - Cosmos3VisionPackSequenceStep, - Cosmos3VisionDenoiseInputStep, - Cosmos3SetTimestepsStep, - Cosmos3SoundPrepareLatentsStep, - Cosmos3SoundPackSequenceStep, - Cosmos3SoundDenoiseInputStep, - Cosmos3VisionSoundDenoiseStep, - ] - block_names = [ - "prepare_text_segments", - "prepare_vision_latents", - "pack_vision_sequence", - "prepare_vision_denoiser_inputs", - "set_timesteps", - "prepare_sound_latents", - "pack_sound_sequence", - "prepare_sound_denoiser_inputs", - "denoise", - ] - - @property - def description(self): - return "Runs the text, vision, and sound Cosmos3 denoising workflow." - - @property - def outputs(self): - return [ - OutputParam.template("latents"), - OutputParam("sound_latents", type_hint=torch.Tensor, description="Denoised sound latents."), - ] - - -# auto_docstring -class Cosmos3VisionActionCoreDenoiseStep(SequentialPipelineBlocks): - """ - Runs the text, vision, and action Cosmos3 denoising workflow. - - Components: - transformer (`Cosmos3OmniTransformer`) scheduler (`UniPCMultistepScheduler`) - - Configs: - use_native_flow_schedule (default: False) - - Inputs: - cond_input_ids (`None`): - Token IDs for the conditional prompt. - uncond_input_ids (`None`): - Token IDs for the unconditional prompt. - x0_tokens_vision (`Tensor`, *optional*): - Vision latents encoded from the conditioning image or video. - vision_condition_frames (`list`, *optional*): - Latent-frame indexes fixed by visual conditioning. - num_frames (`int`): - Number of frames to generate. - height (`int`): - Height of the generated video in pixels. - width (`int`): - Width of the generated video in pixels. - fps (`float`, *optional*, defaults to 24.0): - Frame rate of the generated video. - latents (`Tensor`, *optional*): - Pre-generated noisy vision latents. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`): - The number of denoising steps. - action (`CosmosActionCondition`): - Action-conditioning metadata. - action_condition_frame_indexes (`list`, *optional*): - Action-frame indexes fixed by action conditioning. - action_latents (`Tensor`, *optional*): - Pre-generated noisy action latents. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - guidance_scale (`float`, *optional*, defaults to 6.0): - Scale for classifier-free guidance. - - Outputs: - latents (`Tensor`): - Denoised latents. - action_latents (`Tensor`): - Denoised action latents. - """ - - model_name = "cosmos3-omni" - block_classes = [ - Cosmos3PrepareTextSegmentsStep, - Cosmos3VisionPrepareLatentsStep, - Cosmos3VisionPackSequenceStep, - Cosmos3VisionDenoiseInputStep, - Cosmos3SetTimestepsStep, - Cosmos3ActionPrepareLatentsStep, - Cosmos3ActionPackSequenceStep, - Cosmos3ActionDenoiseInputStep, - Cosmos3VisionActionDenoiseStep, - ] - block_names = [ - "prepare_text_segments", - "prepare_vision_latents", - "pack_vision_sequence", - "prepare_vision_denoiser_inputs", - "set_timesteps", - "prepare_action_latents", - "pack_action_sequence", - "prepare_action_denoiser_inputs", - "denoise", - ] - - @property - def description(self): - return "Runs the text, vision, and action Cosmos3 denoising workflow." - - @property - def outputs(self): - return [ - OutputParam.template("latents"), - OutputParam("action_latents", type_hint=torch.Tensor, description="Denoised action latents."), - ] - - -# auto_docstring -class Cosmos3VisionSoundActionCoreDenoiseStep(SequentialPipelineBlocks): - """ - Runs the text, vision, sound, and action Cosmos3 denoising workflow. - - Components: - transformer (`Cosmos3OmniTransformer`) scheduler (`UniPCMultistepScheduler`) - - Configs: - use_native_flow_schedule (default: False) - - Inputs: - cond_input_ids (`None`): - Token IDs for the conditional prompt. - uncond_input_ids (`None`): - Token IDs for the unconditional prompt. - x0_tokens_vision (`Tensor`, *optional*): - Vision latents encoded from the conditioning image or video. - vision_condition_frames (`list`, *optional*): - Latent-frame indexes fixed by visual conditioning. - num_frames (`int`): - Number of frames to generate. - height (`int`): - Height of the generated video in pixels. - width (`int`): - Width of the generated video in pixels. - fps (`float`, *optional*, defaults to 24.0): - Frame rate of the generated video. - latents (`Tensor`, *optional*): - Pre-generated noisy vision latents. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`): - The number of denoising steps. - sound_latents (`Tensor`, *optional*): - Pre-generated noisy sound latents. - action (`CosmosActionCondition`): - Action-conditioning metadata. - action_condition_frame_indexes (`list`, *optional*): - Action-frame indexes fixed by action conditioning. - action_latents (`Tensor`, *optional*): - Pre-generated noisy action latents. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - guidance_scale (`float`, *optional*, defaults to 6.0): - Scale for classifier-free guidance. - - Outputs: - latents (`Tensor`): - Denoised latents. - sound_latents (`Tensor`): - Denoised sound latents. - action_latents (`Tensor`): - Denoised action latents. - """ - - model_name = "cosmos3-omni" - block_classes = [ - Cosmos3PrepareTextSegmentsStep, - Cosmos3VisionPrepareLatentsStep, - Cosmos3VisionPackSequenceStep, - Cosmos3VisionDenoiseInputStep, - Cosmos3SetTimestepsStep, - Cosmos3SoundPrepareLatentsStep, - Cosmos3SoundPackSequenceStep, - Cosmos3SoundDenoiseInputStep, - Cosmos3ActionPrepareLatentsStep, - Cosmos3ActionPackSequenceStep, - Cosmos3ActionDenoiseInputStep, - Cosmos3VisionSoundActionDenoiseStep, - ] - block_names = [ - "prepare_text_segments", - "prepare_vision_latents", - "pack_vision_sequence", - "prepare_vision_denoiser_inputs", - "set_timesteps", - "prepare_sound_latents", - "pack_sound_sequence", - "prepare_sound_denoiser_inputs", - "prepare_action_latents", - "pack_action_sequence", - "prepare_action_denoiser_inputs", - "denoise", - ] - - @property - def description(self): - return "Runs the text, vision, sound, and action Cosmos3 denoising workflow." - - @property - def outputs(self): - return [ - OutputParam.template("latents"), - OutputParam("sound_latents", type_hint=torch.Tensor, description="Denoised sound latents."), - OutputParam("action_latents", type_hint=torch.Tensor, description="Denoised action latents."), - ] - - -# auto_docstring -class Cosmos3TransferChunkDenoiseStep(SequentialPipelineBlocks): - """ - Autoregressive transfer chunk loop. Overrides __call__ to iterate chunks (the inner timestep loop is a non-leaf - LoopSequentialPipelineBlocks, so this outer loop cannot itself be a LoopSequentialPipelineBlocks). Per-chunk - cross-carry (previous_output, output_chunks) lives on PipelineState. - - Components: - vae (`AutoencoderKLWan`) video_processor (`VideoProcessor`) transformer (`Cosmos3OmniTransformer`) scheduler - (`UniPCMultistepScheduler`) - - Inputs: - chunk_id (`int`, *optional*, defaults to 0): - Index of the current chunk. - previous_output (`None`, *optional*): - Decoded pixels of the previous chunk, used to seed later chunks. - control_frames (`dict`): - Preprocessed, time-padded control maps in canonical hint order. - chunk_frames (`int`): - Pixel frames per chunk. - total_frames (`int`): - Total number of output frames. - stride (`int`): - Frame stride between chunks. - height (`int`): - Height of the generated video in pixels. - width (`int`): - Width of the generated video in pixels. - video (`None`, *optional*): - Optional input video that seeds the first chunk's conditioning. - num_first_chunk_conditional_frames (`int`, *optional*, defaults to 0): - Number of frames the first chunk reuses from the input video. - num_conditional_frames (`int`, *optional*, defaults to 1): - Number of frames each later chunk reuses from the previous chunk's tail. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - cond_text_segment (`dict`): - Conditional text segment. - uncond_text_segment (`dict`): - Unconditional text segment. - fps (`float`, *optional*, defaults to 24.0): - Frame rate of the generated video. - num_inference_steps (`int`): - The number of denoising steps. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - guidance_scale (`float`, *optional*, defaults to 6.0): - Scale for text classifier-free guidance. - control_guidance (`float`, *optional*, defaults to 1.0): - Scale for the control (structural) guidance axis. - guidance_interval (`tuple`, *optional*): - Timestep interval [lo, hi] over which text guidance is active (None = always). - control_guidance_interval (`tuple`, *optional*): - Timestep interval [lo, hi] over which control guidance is active (None = always). - output_chunks (`list`): - Decoded pixel chunks accumulated so far. - num_chunks (`int`): - Number of autoregressive chunks. - - Outputs: - control_latents (`list`): - Clean control latents for this chunk, one per hint in canonical order. - x0_tokens_vision (`Tensor`): - Clean target vision latents encoded from the seeded target frames. - current_conditional_frames (`int`): - Number of pixel frames actually used to seed this chunk's target. - latents (`Tensor`): - Noisy target latents for this chunk. - velocity_mask (`Tensor`): - Mask that zeroes the velocity on conditioned (clean) latent frames. - condition_latents (`Tensor`): - Clean target latents on the conditioned frames (the autoregressive seed). - target_condition_indexes (`list`): - Latent-frame indexes fixed by the chunk's conditioning. - cond_full_static (`dict`): - Conditional [control..., target] transfer sequence carrying every control item. - cond_no_control_static (`dict`): - Conditional [target] transfer sequence with the control items dropped. - uncond_full_static (`dict`): - Unconditional [control..., target] transfer sequence for text CFG. - num_noisy_vision_tokens (`int`): - Number of noisy target vision tokens denoised each step. - timesteps (`Tensor`): - Scheduler timesteps for this chunk. - num_warmup_steps (`int`): - Number of scheduler warmup steps for this chunk. - vision_tokens_full (`list`): - Token list for the [control..., target] forward passes. - vision_tokens_target (`list`): - Token list for the target-only (no-control) forward pass. - vision_timesteps (`Tensor`): - Timesteps for the noisy target tokens. - velocity (`Tensor`): - Predicted (masked) transfer velocity. - previous_output (`Tensor`): - Decoded pixels of this chunk, used to seed the next chunk. - output_chunks (`list`): - Decoded pixel chunks accumulated so far (with this chunk appended). - """ - - model_name = "cosmos3-omni" - block_classes = [ - Cosmos3TransferChunkVaeEncoderStep, - Cosmos3TransferPrepareLatentsStep, - Cosmos3TransferPackSequenceStep, - Cosmos3TransferSetTimestepsStep, - Cosmos3TransferDenoiseStep, - Cosmos3TransferDecodeChunkStep, - ] - block_names = [ - "encode_transfer_chunk", - "prepare_transfer_latents", - "pack_transfer_sequence", - "set_timesteps", - "denoise", - "decode_chunk", - ] - - @property - def description(self) -> str: - return ( - "Autoregressive transfer chunk loop. Overrides __call__ to iterate chunks (the inner timestep loop is a " - "non-leaf LoopSequentialPipelineBlocks, so this outer loop cannot itself be a LoopSequentialPipelineBlocks). " - "Per-chunk cross-carry (previous_output, output_chunks) lives on PipelineState." - ) - - @property - def inputs(self) -> list[InputParam]: - return super().inputs + [ - InputParam(name="num_chunks", type_hint=int, required=True, description="Number of autoregressive chunks.") - ] - - @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: - num_chunks = state.get("num_chunks") - state.set("output_chunks", []) - state.set("previous_output", None) - for chunk_id in range(num_chunks): - state.set("chunk_id", chunk_id) - for _, block in self.sub_blocks.items(): - components, state = block(components, state) - return components, state - - -# auto_docstring -class Cosmos3TransferCoreDenoiseStep(SequentialPipelineBlocks): - """ - Transfer denoise stage: prepare shared text segments once, then run the autoregressive chunk loop. - - Components: - transformer (`Cosmos3OmniTransformer`) vae (`AutoencoderKLWan`) video_processor (`VideoProcessor`) scheduler - (`UniPCMultistepScheduler`) - - Inputs: - cond_input_ids (`None`): - Token IDs for the conditional prompt. - uncond_input_ids (`None`): - Token IDs for the unconditional prompt. - chunk_id (`int`, *optional*, defaults to 0): - Index of the current chunk. - previous_output (`None`, *optional*): - Decoded pixels of the previous chunk, used to seed later chunks. - control_frames (`dict`): - Preprocessed, time-padded control maps in canonical hint order. - chunk_frames (`int`): - Pixel frames per chunk. - total_frames (`int`): - Total number of output frames. - stride (`int`): - Frame stride between chunks. - height (`int`): - Height of the generated video in pixels. - width (`int`): - Width of the generated video in pixels. - video (`None`, *optional*): - Optional input video that seeds the first chunk's conditioning. - num_first_chunk_conditional_frames (`int`, *optional*, defaults to 0): - Number of frames the first chunk reuses from the input video. - num_conditional_frames (`int`, *optional*, defaults to 1): - Number of frames each later chunk reuses from the previous chunk's tail. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - fps (`float`, *optional*, defaults to 24.0): - Frame rate of the generated video. - num_inference_steps (`int`): - The number of denoising steps. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - guidance_scale (`float`, *optional*, defaults to 6.0): - Scale for text classifier-free guidance. - control_guidance (`float`, *optional*, defaults to 1.0): - Scale for the control (structural) guidance axis. - guidance_interval (`tuple`, *optional*): - Timestep interval [lo, hi] over which text guidance is active (None = always). - control_guidance_interval (`tuple`, *optional*): - Timestep interval [lo, hi] over which control guidance is active (None = always). - output_chunks (`list`): - Decoded pixel chunks accumulated so far. - num_chunks (`int`): - Number of autoregressive chunks. - - Outputs: - cond_text_segment (`dict`): - Conditional text segment for the denoiser. - uncond_text_segment (`dict`): - Unconditional text segment for the denoiser. - control_latents (`list`): - Clean control latents for this chunk, one per hint in canonical order. - x0_tokens_vision (`Tensor`): - Clean target vision latents encoded from the seeded target frames. - current_conditional_frames (`int`): - Number of pixel frames actually used to seed this chunk's target. - latents (`Tensor`): - Noisy target latents for this chunk. - velocity_mask (`Tensor`): - Mask that zeroes the velocity on conditioned (clean) latent frames. - condition_latents (`Tensor`): - Clean target latents on the conditioned frames (the autoregressive seed). - target_condition_indexes (`list`): - Latent-frame indexes fixed by the chunk's conditioning. - cond_full_static (`dict`): - Conditional [control..., target] transfer sequence carrying every control item. - cond_no_control_static (`dict`): - Conditional [target] transfer sequence with the control items dropped. - uncond_full_static (`dict`): - Unconditional [control..., target] transfer sequence for text CFG. - num_noisy_vision_tokens (`int`): - Number of noisy target vision tokens denoised each step. - timesteps (`Tensor`): - Scheduler timesteps for this chunk. - num_warmup_steps (`int`): - Number of scheduler warmup steps for this chunk. - vision_tokens_full (`list`): - Token list for the [control..., target] forward passes. - vision_tokens_target (`list`): - Token list for the target-only (no-control) forward pass. - vision_timesteps (`Tensor`): - Timesteps for the noisy target tokens. - velocity (`Tensor`): - Predicted (masked) transfer velocity. - previous_output (`Tensor`): - Decoded pixels of this chunk, used to seed the next chunk. - output_chunks (`list`): - Decoded pixel chunks accumulated so far (with this chunk appended). - """ - - model_name = "cosmos3-omni" - block_classes = [ - Cosmos3PrepareTextSegmentsStep, - Cosmos3TransferChunkDenoiseStep, - ] - block_names = ["prepare_text_segments", "chunk_denoise"] - - @property - def description(self) -> str: - return "Transfer denoise stage: prepare shared text segments once, then run the autoregressive chunk loop." - - -# auto_docstring -class Cosmos3AutoCoreDenoiseStep(ConditionalPipelineBlocks): - """ - Selects the Cosmos3 core denoising workflow. - - transfer runs the autoregressive control-video (ControlNet-style) chunk loop when control_videos are provided. - - vision_sound_action runs when action and enable_sound are provided. - - vision_action runs when action is provided. - - vision_sound runs when enable_sound is true. - - vision runs otherwise. - - Components: - transformer (`Cosmos3OmniTransformer`) vae (`AutoencoderKLWan`) video_processor (`VideoProcessor`) scheduler - (`UniPCMultistepScheduler`) - - Configs: - use_native_flow_schedule (default: False) - - Inputs: - cond_input_ids (`None`): - Token IDs for the conditional prompt. - uncond_input_ids (`None`): - Token IDs for the unconditional prompt. - chunk_id (`int`, *optional*, defaults to 0): - Index of the current chunk. - previous_output (`None`, *optional*): - Decoded pixels of the previous chunk, used to seed later chunks. - control_frames (`dict`, *optional*): - Preprocessed, time-padded control maps in canonical hint order. - chunk_frames (`int`, *optional*): - Pixel frames per chunk. - total_frames (`int`, *optional*): - Total number of output frames. - stride (`int`, *optional*): - Frame stride between chunks. - height (`int`): - Height of the generated video in pixels. - width (`int`): - Width of the generated video in pixels. - video (`None`, *optional*): - Optional input video that seeds the first chunk's conditioning. - num_first_chunk_conditional_frames (`int`, *optional*, defaults to 0): - Number of frames the first chunk reuses from the input video. - num_conditional_frames (`int`, *optional*, defaults to 1): - Number of frames each later chunk reuses from the previous chunk's tail. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - fps (`float`, *optional*, defaults to 24.0): - Frame rate of the generated video. - num_inference_steps (`int`): - The number of denoising steps. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - guidance_scale (`float`, *optional*, defaults to 6.0): - Scale for text classifier-free guidance. - control_guidance (`float`, *optional*, defaults to 1.0): - Scale for the control (structural) guidance axis. - guidance_interval (`tuple`, *optional*): - Timestep interval [lo, hi] over which text guidance is active (None = always). - control_guidance_interval (`tuple`, *optional*): - Timestep interval [lo, hi] over which control guidance is active (None = always). - output_chunks (`list`, *optional*): - Decoded pixel chunks accumulated so far. - num_chunks (`int`, *optional*): - Number of autoregressive chunks. - x0_tokens_vision (`Tensor`, *optional*): - Vision latents encoded from the conditioning image or video. - vision_condition_frames (`list`, *optional*): - Latent-frame indexes fixed by visual conditioning. - num_frames (`int`, *optional*): - Number of frames to generate. - latents (`Tensor`): - Pre-generated noisy vision latents. - sound_latents (`Tensor`, *optional*): - Pre-generated noisy sound latents. - action (`CosmosActionCondition`, *optional*): - Action-conditioning metadata. - action_condition_frame_indexes (`list`, *optional*): - Action-frame indexes fixed by action conditioning. - action_latents (`Tensor`, *optional*): - Pre-generated noisy action latents. - enable_sound (`bool`, *optional*, defaults to False): - Whether to generate a synchronized sound track. - - Outputs: - cond_text_segment (`dict`): - Conditional text segment for the denoiser. - uncond_text_segment (`dict`): - Unconditional text segment for the denoiser. - control_latents (`list`): - Clean control latents for this chunk, one per hint in canonical order. - x0_tokens_vision (`Tensor`): - Clean target vision latents encoded from the seeded target frames. - current_conditional_frames (`int`): - Number of pixel frames actually used to seed this chunk's target. - latents (`Tensor`): - Noisy target latents for this chunk. - velocity_mask (`Tensor`): - Mask that zeroes the velocity on conditioned (clean) latent frames. - condition_latents (`Tensor`): - Clean target latents on the conditioned frames (the autoregressive seed). - target_condition_indexes (`list`): - Latent-frame indexes fixed by the chunk's conditioning. - cond_full_static (`dict`): - Conditional [control..., target] transfer sequence carrying every control item. - cond_no_control_static (`dict`): - Conditional [target] transfer sequence with the control items dropped. - uncond_full_static (`dict`): - Unconditional [control..., target] transfer sequence for text CFG. - num_noisy_vision_tokens (`int`): - Number of noisy target vision tokens denoised each step. - timesteps (`Tensor`): - Scheduler timesteps for this chunk. - num_warmup_steps (`int`): - Number of scheduler warmup steps for this chunk. - vision_tokens_full (`list`): - Token list for the [control..., target] forward passes. - vision_tokens_target (`list`): - Token list for the target-only (no-control) forward pass. - vision_timesteps (`Tensor`): - Timesteps for the noisy target tokens. - velocity (`Tensor`): - Predicted (masked) transfer velocity. - previous_output (`Tensor`): - Decoded pixels of this chunk, used to seed the next chunk. - output_chunks (`list`): - Decoded pixel chunks accumulated so far (with this chunk appended). - sound_latents (`Tensor`): - Denoised sound latents. - action_latents (`Tensor`): - Denoised action latents. - """ - - model_name = "cosmos3-omni" - block_classes = [ - Cosmos3TransferCoreDenoiseStep, - Cosmos3VisionSoundActionCoreDenoiseStep, - Cosmos3VisionActionCoreDenoiseStep, - Cosmos3VisionSoundCoreDenoiseStep, - Cosmos3VisionCoreDenoiseStep, - ] - block_names = ["transfer", "vision_sound_action", "vision_action", "vision_sound", "vision"] - block_trigger_inputs = ["action", "enable_sound", "control_videos"] - default_block_name = "vision" - - @property - def inputs(self): - inputs = super().inputs - inputs.append( - InputParam( - name="enable_sound", - type_hint=bool, - default=False, - description="Whether to generate a synchronized sound track.", - ) - ) - return inputs - - def select_block(self, **kwargs) -> str | None: - action = kwargs.get("action") - enable_sound = kwargs.get("enable_sound") - if kwargs.get("control_videos") is not None: - return "transfer" - if action is not None and enable_sound: - return "vision_sound_action" - if action is not None: - return "vision_action" - if enable_sound: - return "vision_sound" - return "vision" - - @property - def description(self): - return ( - "Selects the Cosmos3 core denoising workflow.\n" - + " - transfer runs the autoregressive control-video (ControlNet-style) chunk loop when control_videos are provided.\n" - + " - vision_sound_action runs when action and enable_sound are provided.\n" - + " - vision_action runs when action is provided.\n" - + " - vision_sound runs when enable_sound is true.\n" - + " - vision runs otherwise." - ) - - -# auto_docstring -class Cosmos3OmniBlocks(SequentialPipelineBlocks): - """ - Modular pipeline blocks for Cosmos3 generation modes. - - Supported workflows: - - `text2image`: requires `prompt`, `num_frames` - - `text2video`: requires `prompt` - - `image2video`: requires `prompt`, `image` - - `video2video`: requires `prompt`, `video` - - `text2video_with_sound`: requires `prompt`, `enable_sound` - - `image2video_with_sound`: requires `prompt`, `image`, `enable_sound` - - `video2video_with_sound`: requires `prompt`, `video`, `enable_sound` - - `action_policy`: requires `prompt`, `action` - - `action_forward_dynamics`: requires `prompt`, `action` - - `action_inverse_dynamics`: requires `prompt`, `action` - - Components: - video_processor (`VideoProcessor`) text_tokenizer (`AutoTokenizer`) vae (`AutoencoderKLWan`) transformer - (`Cosmos3OmniTransformer`) scheduler (`UniPCMultistepScheduler`) sound_tokenizer - (`Cosmos3AVAEAudioTokenizer`) - - Configs: - default_use_system_prompt (default: True) enable_safety_checker (default: True) use_native_flow_schedule - (default: False) - - Inputs: - control_videos (`dict`, *optional*): - Mapping of hint name (edge/blur/depth/seg/wsm) to the control video for that modality. - height (`int`, *optional*): - Height of the generated video in pixels. - width (`int`, *optional*): - Width of the generated video in pixels. - num_frames (`int`, *optional*): - Optional cap on the number of output frames (defaults to the control video length). - num_video_frames_per_chunk (`int`, *optional*): - Number of pixel frames generated per autoregressive chunk. - num_conditional_frames (`int`, *optional*, defaults to 1): - Number of frames each chunk reuses from the previous chunk's tail. - prompt (`str`): - The text prompt that guides Cosmos3 generation. - negative_prompt (`str`, *optional*): - The negative text prompt used for classifier-free guidance. - use_system_prompt (`bool`, *optional*, defaults to True): - Whether to prepend the Cosmos3 transfer system prompt. - action (`CosmosActionCondition`, *optional*): - Action-conditioning metadata and its reference visual input. - fps (`float`, *optional*, defaults to 24.0): - Frame rate of the generated video. - add_resolution_template (`bool`, *optional*, defaults to True): - Whether to add resolution metadata to the prompt. - add_duration_template (`bool`, *optional*, defaults to True): - Whether to add duration metadata to the prompt. - video (`None`, *optional*): - Reference video for video-to-video conditioning. - condition_frame_indexes_vision (`tuple | list`, *optional*, defaults to (0, 1)): - Latent-frame indexes to preserve from the conditioning video. - condition_video_keep (`str`, *optional*, defaults to first): - Which end of a longer conditioning video to use: `first` or `last`. - image (`None`, *optional*): - Reference image for image-to-video conditioning. - chunk_id (`int`, *optional*, defaults to 0): - Index of the current chunk. - previous_output (`None`, *optional*): - Decoded pixels of the previous chunk, used to seed later chunks. - num_first_chunk_conditional_frames (`int`, *optional*, defaults to 0): - Number of frames the first chunk reuses from the input video. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`): - The number of denoising steps. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - guidance_scale (`float`, *optional*, defaults to 6.0): - Scale for text classifier-free guidance. - control_guidance (`float`, *optional*, defaults to 1.0): - Scale for the control (structural) guidance axis. - guidance_interval (`tuple`, *optional*): - Timestep interval [lo, hi] over which text guidance is active (None = always). - control_guidance_interval (`tuple`, *optional*): - Timestep interval [lo, hi] over which control guidance is active (None = always). - output_chunks (`list`, *optional*): - Decoded pixel chunks accumulated so far. - x0_tokens_vision (`Tensor`, *optional*): - Vision latents encoded from the conditioning image or video. - vision_condition_frames (`list`, *optional*): - Latent-frame indexes fixed by visual conditioning. - latents (`Tensor`): - Pre-generated noisy vision latents. - sound_latents (`Tensor`, *optional*): - Pre-generated noisy sound latents. - action_condition_frame_indexes (`list`, *optional*): - Action-frame indexes fixed by action conditioning. - action_latents (`Tensor`, *optional*): - Pre-generated noisy action latents. - enable_sound (`bool`, *optional*, defaults to False): - Whether to generate a synchronized sound track. - output_type (`str`, *optional*, defaults to pil): - Output format: 'pil', 'np', 'pt'. - - Outputs: - videos (`list`): - The generated videos. - sound (`Tensor`): - Generated waveform. - sampling_rate (`int`): - Sample rate of the generated waveform in Hz. - action (`list`): - Generated action vectors. - """ - - model_name = "cosmos3-omni" - block_classes = [ - Cosmos3AutoTextEncoderStep, - Cosmos3AutoVaeEncoderStep, - Cosmos3AutoCoreDenoiseStep, - Cosmos3AutoDecodeStep, - Cosmos3ActionOutputStep, - ] - block_names = ["text_encoder", "vae_encoder", "denoise", "decode", "after_decode"] - _workflow_map = { - "text2image": {"prompt": True, "num_frames": 1}, - "text2video": {"prompt": True}, - "image2video": {"prompt": True, "image": True}, - "video2video": {"prompt": True, "video": True}, - "text2video_with_sound": {"prompt": True, "enable_sound": True}, - "image2video_with_sound": {"prompt": True, "image": True, "enable_sound": True}, - "video2video_with_sound": {"prompt": True, "video": True, "enable_sound": True}, - "action_policy": {"prompt": True, "action": True}, - "action_forward_dynamics": {"prompt": True, "action": True}, - "action_inverse_dynamics": {"prompt": True, "action": True}, - } - - @property - def description(self): - return "Modular pipeline blocks for Cosmos3 generation modes." - - def get_workflow(self, workflow_name: str): - if workflow_name == "transfer": - raise NotImplementedError( - 'The standalone "transfer" workflow is temporarily unavailable because its nested autoregressive ' - "chunk and denoising loops cannot be preserved by the current workflow extraction logic. Transfer " - "remains available through the full Cosmos3OmniBlocks pipeline. The standalone workflow will be " - "enabled after migration to the upcoming composable nested-loop abstraction." - ) - return super().get_workflow(workflow_name) - - @property - def outputs(self): - return [ - OutputParam.template("videos"), - OutputParam("sound", type_hint=torch.Tensor, description="Generated waveform."), - OutputParam("sampling_rate", type_hint=int, description="Sample rate of the generated waveform in Hz."), - OutputParam("action", type_hint=list[torch.Tensor], description="Generated action vectors."), - ] diff --git a/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3_distilled.py b/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3_distilled.py deleted file mode 100644 index e168cc24d8cd01358589055f2f2fb478956ebd3e..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3_distilled.py +++ /dev/null @@ -1,240 +0,0 @@ -from ..modular_pipeline import ConditionalPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import OutputParam -from .before_denoise import ( - Cosmos3DistilledSetTimestepsStep, - Cosmos3PrepareTextSegmentsStep, - Cosmos3VisionDenoiseInputStep, - Cosmos3VisionPackSequenceStep, - Cosmos3VisionPrepareLatentsStep, -) -from .decoders import Cosmos3VideoDecodeStep -from .denoise import Cosmos3DistilledVisionDenoiseStep -from .encoders import ( - Cosmos3DistilledTextEncoderStep, - Cosmos3ImageVaeEncoderStep, - Cosmos3VideoVaeEncoderStep, -) - - -# auto_docstring -class Cosmos3DistilledAutoVaeEncoderStep(ConditionalPipelineBlocks): - """ - Auto VAE conditioning block for distilled Cosmos3. - - Cosmos3VideoVaeEncoderStep runs for the video path. - - Cosmos3ImageVaeEncoderStep runs for the image path. - - when no image or video conditioning is provided, this block is skipped. - - Components: - vae (`AutoencoderKLWan`) video_processor (`VideoProcessor`) - - Inputs: - video (`None`, *optional*): - Reference video for video-to-video conditioning. - condition_frame_indexes_vision (`tuple | list`, *optional*, defaults to (0, 1)): - Latent-frame indexes to preserve from the conditioning video. - condition_video_keep (`str`, *optional*, defaults to first): - Which end of a longer conditioning video to use: `first` or `last`. - num_frames (`int`, *optional*): - Number of frames to generate. - height (`int`, *optional*): - Height of the generated video in pixels. - width (`int`, *optional*): - Width of the generated video in pixels. - image (`None`, *optional*): - Reference image for image-to-video conditioning. - - Outputs: - x0_tokens_vision (`Tensor`): - Vision latents encoded from the conditioning image or video. - vision_condition_frames (`list`): - Latent-frame indexes fixed by visual conditioning. - """ - - model_name = "cosmos3-omni" - block_classes = [Cosmos3VideoVaeEncoderStep, Cosmos3ImageVaeEncoderStep] - block_names = ["video_conditioning", "image_conditioning"] - block_trigger_inputs = ["video", "image"] - default_block_name = None - - def select_block(self, **kwargs) -> str | None: - image = kwargs.get("image") - video = kwargs.get("video") - if image is not None and video is not None: - raise ValueError("Pass either image or video, not both.") - if video is not None: - return "video_conditioning" - if image is not None: - return "image_conditioning" - return None - - @property - def description(self): - return ( - "Auto VAE conditioning block for distilled Cosmos3.\n" - + " - Cosmos3VideoVaeEncoderStep runs for the video path.\n" - + " - Cosmos3ImageVaeEncoderStep runs for the image path.\n" - + " - when no image or video conditioning is provided, this block is skipped." - ) - - -# auto_docstring -class Cosmos3DistilledVisionCoreDenoiseStep(SequentialPipelineBlocks): - """ - Runs the text-and-vision distilled Cosmos3 denoising workflow. - - Components: - transformer (`Cosmos3OmniTransformer`) scheduler (`FlowMatchEulerDiscreteScheduler`) - - Configs: - is_distilled (default: True) distilled_sigmas (default: None) - - Inputs: - cond_input_ids (`None`): - Token IDs for the conditional prompt. - uncond_input_ids (`None`): - Token IDs for the unconditional prompt. - x0_tokens_vision (`Tensor`, *optional*): - Vision latents encoded from the conditioning image or video. - vision_condition_frames (`list`, *optional*): - Latent-frame indexes fixed by visual conditioning. - num_frames (`int`): - Number of frames to generate. - height (`int`): - Height of the generated video in pixels. - width (`int`): - Width of the generated video in pixels. - fps (`float`, *optional*, defaults to 24.0): - Frame rate of the generated video. - latents (`Tensor`, *optional*): - Pre-generated noisy vision latents. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`, *optional*): - The number of denoising steps. - guidance_scale (`float`, *optional*): - Unused for distilled checkpoints; classifier-free guidance is baked into the weights and the scale is - forced to 1.0. Passing a value other than 1.0 raises an error. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "cosmos3-omni" - block_classes = [ - Cosmos3PrepareTextSegmentsStep, - Cosmos3VisionPrepareLatentsStep, - Cosmos3VisionPackSequenceStep, - Cosmos3VisionDenoiseInputStep, - Cosmos3DistilledSetTimestepsStep, - Cosmos3DistilledVisionDenoiseStep, - ] - block_names = [ - "prepare_text_segments", - "prepare_vision_latents", - "pack_vision_sequence", - "prepare_vision_denoiser_inputs", - "set_timesteps", - "denoise", - ] - - @property - def description(self): - return "Runs the text-and-vision distilled Cosmos3 denoising workflow." - - @property - def outputs(self): - return [OutputParam.template("latents")] - - -# auto_docstring -class Cosmos3DistilledBlocks(SequentialPipelineBlocks): - """ - Modular pipeline blocks for distilled (few-step) Cosmos3 generation modes. - - Supported workflows: - - `text2image`: requires `prompt`, `num_frames` - - `text2video`: requires `prompt` - - `image2video`: requires `prompt`, `image` - - `video2video`: requires `prompt`, `video` - - Components: - text_tokenizer (`AutoTokenizer`) vae (`AutoencoderKLWan`) video_processor (`VideoProcessor`) transformer - (`Cosmos3OmniTransformer`) scheduler (`FlowMatchEulerDiscreteScheduler`) - - Configs: - default_use_system_prompt (default: True) enable_safety_checker (default: True) is_distilled (default: True) - distilled_sigmas (default: None) - - Inputs: - prompt (`str`): - The text prompt that guides Cosmos3 generation. - num_frames (`int`, *optional*): - Number of frames to generate. - height (`int`, *optional*): - Height of the generated video or image in pixels. - width (`int`, *optional*): - Width of the generated video or image in pixels. - fps (`float`, *optional*, defaults to 24.0): - Frame rate of the generated video. - use_system_prompt (`bool`, *optional*, defaults to True): - Whether to prepend the Cosmos3 system prompt. - add_resolution_template (`bool`, *optional*, defaults to True): - Whether to add resolution metadata to the prompt. - add_duration_template (`bool`, *optional*, defaults to True): - Whether to add duration metadata to the prompt. - video (`None`, *optional*): - Reference video for video-to-video conditioning. - condition_frame_indexes_vision (`tuple | list`, *optional*, defaults to (0, 1)): - Latent-frame indexes to preserve from the conditioning video. - condition_video_keep (`str`, *optional*, defaults to first): - Which end of a longer conditioning video to use: `first` or `last`. - image (`None`, *optional*): - Reference image for image-to-video conditioning. - x0_tokens_vision (`Tensor`, *optional*): - Vision latents encoded from the conditioning image or video. - vision_condition_frames (`list`, *optional*): - Latent-frame indexes fixed by visual conditioning. - latents (`Tensor`, *optional*): - Pre-generated noisy vision latents. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`, *optional*): - The number of denoising steps. - guidance_scale (`float`, *optional*): - Unused for distilled checkpoints; classifier-free guidance is baked into the weights and the scale is - forced to 1.0. Passing a value other than 1.0 raises an error. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - output_type (`str`, *optional*, defaults to pil): - Output format: 'pil', 'np', 'pt'. - - Outputs: - videos (`list`): - The generated videos. - """ - - model_name = "cosmos3-omni" - block_classes = [ - Cosmos3DistilledTextEncoderStep, - Cosmos3DistilledAutoVaeEncoderStep, - Cosmos3DistilledVisionCoreDenoiseStep, - Cosmos3VideoDecodeStep, - ] - block_names = ["text_encoder", "vae_encoder", "denoise", "decode"] - _workflow_map = { - "text2image": {"prompt": True, "num_frames": 1}, - "text2video": {"prompt": True}, - "image2video": {"prompt": True, "image": True}, - "video2video": {"prompt": True, "video": True}, - } - - @property - def description(self): - return "Modular pipeline blocks for distilled (few-step) Cosmos3 generation modes." - - @property - def outputs(self): - return [OutputParam.template("videos")] diff --git a/diffusers/modular_pipelines/cosmos/modular_pipeline.py b/diffusers/modular_pipelines/cosmos/modular_pipeline.py deleted file mode 100644 index d6c09703c12e023f22b307e3a14e4eed6cd58394..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/cosmos/modular_pipeline.py +++ /dev/null @@ -1,130 +0,0 @@ -import torch - -from ...pipelines.cosmos.pipeline_cosmos3_omni import Cosmos3OmniPipeline, CosmosSafetyChecker -from ..modular_pipeline import ModularPipeline - - -class Cosmos3OmniModularPipeline(ModularPipeline): - """ - A ModularPipeline for Cosmos 3 omni generation. - """ - - default_blocks_name = "Cosmos3OmniBlocks" - - duration_template = "The video is {duration:.1f} seconds long and is of {fps:.0f} FPS." - image_resolution_template = "This image is of {height}x{width} resolution." - video_resolution_template = "This video is of {height}x{width} resolution." - inverse_duration_template = "The video is not {duration:.1f} seconds long and is not of {fps:.0f} FPS." - inverse_image_resolution_template = "This image is not of {height}x{width} resolution." - inverse_video_resolution_template = "This video is not of {height}x{width} resolution." - - @property - def vae_scale_factor_spatial(self): - if getattr(self, "vae", None) is not None: - return int(self.vae.config.scale_factor_spatial) - return 16 - - @property - def vae_scale_factor_temporal(self): - if getattr(self, "vae", None) is not None: - return int(self.vae.config.scale_factor_temporal) - return 4 - - @property - def num_channels_latents(self): - if getattr(self, "transformer", None) is not None: - return int(self.transformer.config.latent_channel) - return 48 - - @property - def sound_sampling_rate(self): - if getattr(self, "sound_tokenizer", None) is not None: - return int(self.sound_tokenizer.config.sampling_rate) - return 48000 - - @property - def sound_hop_size(self): - if getattr(self, "sound_tokenizer", None) is not None: - return int(self.sound_tokenizer._hop_size) - return 1920 - - @property - def _vae_latents_mean(self): - return torch.tensor(self.vae.config.latents_mean, dtype=self.vae.dtype) - - @property - def _vae_latents_inv_std(self): - return 1.0 / torch.tensor(self.vae.config.latents_std, dtype=self.vae.dtype) - - @property - def llm_special_tokens(self): - if getattr(self, "text_tokenizer", None) is None: - return None - return { - "start_of_generation": self.text_tokenizer.convert_tokens_to_ids("<|vision_start|>"), - "eos_token_id": self.text_tokenizer.eos_token_id, - } - - def enable_safety_checker(self, safety_checker=None): - if safety_checker is not None: - self.safety_checker = safety_checker - elif getattr(self, "safety_checker", None) is None: - self.safety_checker = CosmosSafetyChecker() - self._is_safety_checker_enabled = True - - def disable_safety_checker(self): - self._is_safety_checker_enabled = False - - @property - def requires_safety_checker(self): - return getattr(self, "_is_safety_checker_enabled", self.config.enable_safety_checker) - - def _encode_video(self, x): - return Cosmos3OmniPipeline._encode_video(self, x) - - def decode_sound(self, latent): - return Cosmos3OmniPipeline.decode_sound(self, latent) - - def _prepare_text_segment(self, input_ids, device): - return Cosmos3OmniPipeline._prepare_text_segment(self, input_ids, device) - - def _prepare_vision_segment(self, *args, **kwargs): - return Cosmos3OmniPipeline._prepare_vision_segment(self, *args, **kwargs) - - def _prepare_sound_segment(self, *args, **kwargs): - return Cosmos3OmniPipeline._prepare_sound_segment(self, *args, **kwargs) - - def _prepare_action_segment(self, *args, **kwargs): - return Cosmos3OmniPipeline._prepare_action_segment(self, *args, **kwargs) - - def _prepare_action_video_conditioning(self, *args, **kwargs): - return Cosmos3OmniPipeline._prepare_action_video_conditioning(self, *args, **kwargs) - - def _remove_action_video_padding_from_latent(self, *args, **kwargs): - return Cosmos3OmniPipeline._remove_action_video_padding_from_latent(self, *args, **kwargs) - - @staticmethod - def _build_action_json_prompt(*args, **kwargs): - return Cosmos3OmniPipeline._build_action_json_prompt(*args, **kwargs) - - def tokenize_prompt(self, *args, **kwargs): - return Cosmos3OmniPipeline.tokenize_prompt(self, *args, **kwargs) - - @staticmethod - def _mask_velocity_predictions(*args, **kwargs): - return Cosmos3OmniPipeline._mask_velocity_predictions(*args, **kwargs) - - def _apply_video_safety_check(self, *args, **kwargs): - return Cosmos3OmniPipeline._apply_video_safety_check(self, *args, **kwargs) - - -class Cosmos3DistilledModularPipeline(Cosmos3OmniModularPipeline): - """ - A ModularPipeline for distilled (few-step) Cosmos 3 omni generation. - - Distilled checkpoints bake classifier-free guidance into the weights and sample on a fixed schedule read from the - pipeline's `distilled_sigmas` config (populated from `modular_model_index.json`), so `guidance_scale` and - `num_inference_steps` are fixed and `negative_prompt` is not supported. - """ - - default_blocks_name = "Cosmos3DistilledBlocks" diff --git a/diffusers/modular_pipelines/ernie_image/__init__.py b/diffusers/modular_pipelines/ernie_image/__init__.py deleted file mode 100644 index 68ed723c590c87c13c5aa7c115ece1321a4ca89e..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ernie_image/__init__.py +++ /dev/null @@ -1,47 +0,0 @@ -from typing import TYPE_CHECKING - -from ...utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, -) - - -_dummy_objects = {} -_import_structure = {} - -try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ...utils import dummy_torch_and_transformers_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) -else: - _import_structure["modular_blocks_ernie_image"] = ["ErnieImageAutoBlocks"] - _import_structure["modular_pipeline"] = ["ErnieImageModularPipeline"] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .modular_blocks_ernie_image import ErnieImageAutoBlocks - from .modular_pipeline import ErnieImageModularPipeline -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/ernie_image/before_denoise.py b/diffusers/modular_pipelines/ernie_image/before_denoise.py deleted file mode 100644 index 0342306323967db9709502f1cfe9cf743f78bb7f..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ernie_image/before_denoise.py +++ /dev/null @@ -1,270 +0,0 @@ -# Copyright 2025 Baidu ERNIE-Image Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ...models import ErnieImageTransformer2DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import ErnieImageModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _pad_text( - text_hiddens: list[torch.Tensor], device: torch.device, dtype: torch.dtype, text_in_dim: int -) -> tuple[torch.Tensor, torch.Tensor]: - """Pad a list of variable-length text hidden states to a common length and return (padded, lengths).""" - batch_size = len(text_hiddens) - if batch_size == 0: - return ( - torch.zeros((0, 0, text_in_dim), device=device, dtype=dtype), - torch.zeros((0,), device=device, dtype=torch.long), - ) - normalized = [t.squeeze(1).to(device).to(dtype) if t.dim() == 3 else t.to(device).to(dtype) for t in text_hiddens] - lengths = torch.tensor([t.shape[0] for t in normalized], device=device, dtype=torch.long) - max_length = int(lengths.max().item()) - padded = torch.zeros((batch_size, max_length, text_in_dim), device=device, dtype=dtype) - for i, t in enumerate(normalized): - padded[i, : t.shape[0], :] = t - return padded, lengths - - -class ErnieImageTextInputStep(ModularPipelineBlocks): - model_name = "ernie-image" - - @property - def description(self) -> str: - return ( - "Input processing step that pads the variable-length text hidden states to a common length and " - "produces `text_bth` / `text_lens` tensors consumed by the denoiser." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", ErnieImageTransformer2DModel)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "prompt_embeds", - required=True, - type_hint=list, - description="List of per-prompt text embeddings from the text encoder step.", - ), - InputParam( - "negative_prompt_embeds", - type_hint=list, - description="List of per-prompt negative text embeddings from the text encoder step.", - ), - InputParam( - "num_images_per_prompt", - type_hint=int, - default=1, - description="Number of images to generate per prompt.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("batch_size", type_hint=int, description="The number of prompts in the batch."), - OutputParam( - "text_bth", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Padded text hidden states of shape (B, T_max, H) fed into the transformer.", - ), - OutputParam( - "text_lens", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Actual per-prompt text lengths used to build the transformer attention mask.", - ), - OutputParam( - "negative_text_bth", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Padded negative text hidden states, when classifier-free guidance is enabled.", - ), - OutputParam( - "negative_text_lens", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Actual per-prompt negative text lengths, when classifier-free guidance is enabled.", - ), - ] - - @staticmethod - def _expand(hiddens: list[torch.Tensor], num_images_per_prompt: int) -> list[torch.Tensor]: - if num_images_per_prompt == 1: - return list(hiddens) - return [h for h in hiddens for _ in range(num_images_per_prompt)] - - @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - dtype = components.transformer.dtype - text_in_dim = components.text_in_dim - num_images_per_prompt = block_state.num_images_per_prompt - - prompt_embeds = block_state.prompt_embeds - block_state.batch_size = len(prompt_embeds) - - prompt_embeds = self._expand(prompt_embeds, num_images_per_prompt) - text_bth, text_lens = _pad_text(prompt_embeds, device, dtype, text_in_dim) - block_state.text_bth = text_bth - block_state.text_lens = text_lens - - negative_prompt_embeds = block_state.negative_prompt_embeds - if negative_prompt_embeds is not None: - negative_prompt_embeds = self._expand(negative_prompt_embeds, num_images_per_prompt) - negative_text_bth, negative_text_lens = _pad_text(negative_prompt_embeds, device, dtype, text_in_dim) - block_state.negative_text_bth = negative_text_bth - block_state.negative_text_lens = negative_text_lens - else: - block_state.negative_text_bth = None - block_state.negative_text_lens = None - - self.set_block_state(state, block_state) - return components, state - - -class ErnieImageSetTimestepsStep(ModularPipelineBlocks): - model_name = "ernie-image" - - @property - def description(self) -> str: - return "Step that sets the scheduler's timesteps for inference using a linear sigma schedule." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "num_inference_steps", - type_hint=int, - default=50, - description="Number of denoising steps.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("timesteps", type_hint=torch.Tensor, description="The timesteps to use for inference."), - OutputParam("num_inference_steps", type_hint=int, description="The number of denoising steps."), - ] - - @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - num_inference_steps = block_state.num_inference_steps - - sigmas = torch.linspace(1.0, 0.0, num_inference_steps + 1)[:-1] - components.scheduler.set_timesteps(sigmas=sigmas, device=device) - - block_state.timesteps = components.scheduler.timesteps - block_state.num_inference_steps = num_inference_steps - - self.set_block_state(state, block_state) - return components, state - - -class ErnieImagePrepareLatentsStep(ModularPipelineBlocks): - model_name = "ernie-image" - - @property - def description(self) -> str: - return "Prepare random noise latents for the ErnieImage text-to-image denoising process." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", ErnieImageTransformer2DModel)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("height", type_hint=int, description="The height in pixels of the generated image."), - InputParam("width", type_hint=int, description="The width in pixels of the generated image."), - InputParam( - "latents", - type_hint=torch.Tensor, - description="Pre-generated noisy latents. If provided, skips noise sampling.", - ), - InputParam( - "generator", - type_hint=torch.Generator, - description="Torch generator for deterministic noise sampling.", - ), - InputParam( - "text_bth", - required=True, - type_hint=torch.Tensor, - description="Padded text hidden states; used to derive the total batch size for the latents.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("latents", type_hint=torch.Tensor, description="The initial noise latents to denoise."), - OutputParam("height", type_hint=int, description="The resolved image height in pixels."), - OutputParam("width", type_hint=int, description="The resolved image width in pixels."), - ] - - @staticmethod - def _check_inputs(components: ErnieImageModularPipeline, height: int, width: int) -> None: - vae_scale_factor = components.vae_scale_factor - if height % vae_scale_factor != 0 or width % vae_scale_factor != 0: - raise ValueError( - f"`height` and `width` must be divisible by {vae_scale_factor}, got {height} and {width}." - ) - - @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - dtype = components.transformer.dtype - - height = block_state.height or components.default_height - width = block_state.width or components.default_width - self._check_inputs(components, height, width) - - total_batch_size = block_state.text_bth.shape[0] - latent_h = height // components.vae_scale_factor - latent_w = width // components.vae_scale_factor - num_channels_latents = components.num_channels_latents - - shape = (total_batch_size, num_channels_latents, latent_h, latent_w) - if block_state.latents is None: - block_state.latents = randn_tensor(shape, generator=block_state.generator, device=device, dtype=dtype) - else: - block_state.latents = block_state.latents.to(device=device, dtype=dtype) - - block_state.height = height - block_state.width = width - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/ernie_image/decoders.py b/diffusers/modular_pipelines/ernie_image/decoders.py deleted file mode 100644 index d7d056b825840fae2797ce025dbe18b14a80bca1..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ernie_image/decoders.py +++ /dev/null @@ -1,92 +0,0 @@ -# Copyright 2025 Baidu ERNIE-Image Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ...configuration_utils import FrozenDict -from ...image_processor import VaeImageProcessor -from ...models import AutoencoderKLFlux2 -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import ErnieImageModularPipeline, ErnieImagePachifier - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class ErnieImageVaeDecoderStep(ModularPipelineBlocks): - model_name = "ernie-image" - - @property - def description(self) -> str: - return "Step that decodes the denoised latents into images (unpachify, BN denormalization, VAE decode)." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLFlux2), - ComponentSpec( - "pachifier", - ErnieImagePachifier, - config=FrozenDict({"patch_size": 2}), - default_creation_method="from_config", - ), - ComponentSpec( - "image_processor", - VaeImageProcessor, - config=FrozenDict({"vae_scale_factor": 16}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "latents", - required=True, - type_hint=torch.Tensor, - description="The latents to decode into images.", - ), - InputParam( - "output_type", - type_hint=str, - default="pil", - description="Output format: 'pil', 'np', or 'pt'.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam("images", type_hint=list, description="The generated images.")] - - @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - vae = components.vae - device = block_state.latents.device - - latents = block_state.latents - bn_mean = vae.bn.running_mean.view(1, -1, 1, 1).to(device=device, dtype=latents.dtype) - bn_std = torch.sqrt(vae.bn.running_var.view(1, -1, 1, 1) + 1e-5).to(device=device, dtype=latents.dtype) - latents = latents * bn_std + bn_mean - - latents = components.pachifier.unpack_latents(latents) - - images = vae.decode(latents.to(vae.dtype), return_dict=False)[0] - block_state.images = components.image_processor.postprocess(images, output_type=block_state.output_type) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/ernie_image/denoise.py b/diffusers/modular_pipelines/ernie_image/denoise.py deleted file mode 100644 index 3a2a2e312486a061ad5c73ddc8a9ffbce0cdde3f..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ernie_image/denoise.py +++ /dev/null @@ -1,236 +0,0 @@ -# Copyright 2025 Baidu ERNIE-Image Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...models import ErnieImageTransformer2DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ..modular_pipeline import ( - BlockState, - LoopSequentialPipelineBlocks, - ModularPipelineBlocks, - PipelineState, -) -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import ErnieImageModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class ErnieImageLoopBeforeDenoiser(ModularPipelineBlocks): - model_name = "ernie-image" - - @property - def description(self) -> str: - return ( - "Step within the denoising loop that prepares the latent model input and timestep tensor. " - "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " - "object (e.g. `ErnieImageDenoiseLoopWrapper`)." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", ErnieImageTransformer2DModel)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "latents", - required=True, - type_hint=torch.Tensor, - description="The latents to denoise.", - ), - ] - - @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - latents = block_state.latents - block_state.latent_model_input = latents.to(components.transformer.dtype) - block_state.timestep = t.expand(latents.shape[0]).to(components.transformer.dtype) - return components, block_state - - -class ErnieImageLoopDenoiser(ModularPipelineBlocks): - model_name = "ernie-image" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("transformer", ErnieImageTransformer2DModel), - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 4.0}), - default_creation_method="from_config", - ), - ] - - @property - def description(self) -> str: - return ( - "Step within the denoising loop that runs the ErnieImage transformer with classifier-free guidance via " - "the configured guider." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "text_bth", - required=True, - type_hint=torch.Tensor, - description="Padded text hidden states fed into the transformer.", - ), - InputParam( - "text_lens", - required=True, - type_hint=torch.Tensor, - description="Per-prompt text lengths used by the transformer attention mask.", - ), - InputParam( - "negative_text_bth", - type_hint=torch.Tensor, - description="Padded negative text hidden states for classifier-free guidance.", - ), - InputParam( - "negative_text_lens", - type_hint=torch.Tensor, - description="Per-prompt negative text lengths for classifier-free guidance.", - ), - InputParam( - "num_inference_steps", - required=True, - type_hint=int, - description="Total number of denoising steps. Used by the guider for step-aware scheduling.", - ), - ] - - @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - guider_inputs = { - "text_bth": (block_state.text_bth, block_state.negative_text_bth), - "text_lens": (block_state.text_lens, block_state.negative_text_lens), - } - - components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) - guider_state = components.guider.prepare_inputs(guider_inputs) - - for guider_state_batch in guider_state: - components.guider.prepare_models(components.transformer) - cond_kwargs = {name: getattr(guider_state_batch, name) for name in guider_inputs.keys()} - noise_pred = components.transformer( - hidden_states=block_state.latent_model_input, - timestep=block_state.timestep, - return_dict=False, - **cond_kwargs, - )[0] - guider_state_batch.noise_pred = noise_pred - components.guider.cleanup_models(components.transformer) - - block_state.noise_pred = components.guider(guider_state)[0] - return components, block_state - - -class ErnieImageLoopAfterDenoiser(ModularPipelineBlocks): - model_name = "ernie-image" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def description(self) -> str: - return "Step within the denoising loop that updates the latents using the scheduler step." - - @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - latents_dtype = block_state.latents.dtype - block_state.latents = components.scheduler.step( - block_state.noise_pred, t, block_state.latents, return_dict=False - )[0] - if block_state.latents.dtype != latents_dtype and torch.backends.mps.is_available(): - block_state.latents = block_state.latents.to(latents_dtype) - return components, block_state - - -class ErnieImageDenoiseLoopWrapper(LoopSequentialPipelineBlocks): - model_name = "ernie-image" - - @property - def description(self) -> str: - return ( - "Pipeline block that iteratively denoises the latents over `timesteps`. " - "The specific steps within each iteration can be customized with `sub_blocks` attribute." - ) - - @property - def loop_expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), - ComponentSpec("transformer", ErnieImageTransformer2DModel), - ] - - @property - def loop_inputs(self) -> list[InputParam]: - return [ - InputParam( - "timesteps", - required=True, - type_hint=torch.Tensor, - description="The timesteps to use for inference.", - ), - InputParam( - "num_inference_steps", - required=True, - type_hint=int, - description="The number of denoising steps.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam("latents", type_hint=torch.Tensor, description="The denoised latents.")] - - @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) - progress_bar.update() - self.set_block_state(state, block_state) - return components, state - - -class ErnieImageDenoiseStep(ErnieImageDenoiseLoopWrapper): - block_classes = [ - ErnieImageLoopBeforeDenoiser, - ErnieImageLoopDenoiser, - ErnieImageLoopAfterDenoiser, - ] - block_names = ["before_denoiser", "denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoise step that iteratively denoises the latents. At each iteration it runs:\n" - " - `ErnieImageLoopBeforeDenoiser`\n" - " - `ErnieImageLoopDenoiser`\n" - " - `ErnieImageLoopAfterDenoiser`" - ) diff --git a/diffusers/modular_pipelines/ernie_image/encoders.py b/diffusers/modular_pipelines/ernie_image/encoders.py deleted file mode 100644 index 161646d181be4a512fda62b18bfc9d876949adf0..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ernie_image/encoders.py +++ /dev/null @@ -1,264 +0,0 @@ -# Copyright 2025 Baidu ERNIE-Image Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import json - -import torch -from transformers import AutoTokenizer, Mistral3Model - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...utils import logging -from ...utils.import_utils import is_transformers_version -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import ErnieImageModularPipeline - - -if is_transformers_version("<", "5.0.0"): - raise ImportError("`ErnieImageModularPipeline` requires `transformers>=5.0.0` for `Ministral3ForCausalLM`.") - -from transformers import Ministral3ForCausalLM # noqa: E402 - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class ErnieImagePromptEnhancerStep(ModularPipelineBlocks): - model_name = "ernie-image" - - @property - def description(self) -> str: - return "Prompt enhancer step that rewrites the input prompt using a causal language model (PE)." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("pe", Ministral3ForCausalLM), - ComponentSpec("pe_tokenizer", AutoTokenizer), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "prompt", - required=True, - type_hint=str, - description="The prompt or prompts to guide image generation.", - ), - InputParam("height", type_hint=int, description="The height in pixels of the generated image."), - InputParam("width", type_hint=int, description="The width in pixels of the generated image."), - InputParam( - "pe_system_prompt", - type_hint=str, - default=None, - description="Optional system prompt passed to the prompt enhancer.", - ), - InputParam( - "pe_temperature", - type_hint=float, - default=0.6, - description="Sampling temperature used when generating with the prompt enhancer.", - ), - InputParam( - "pe_top_p", - type_hint=float, - default=0.95, - description="Nucleus sampling `top_p` used when generating with the prompt enhancer.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("prompt", type_hint=list, description="The prompt list after prompt-enhancer rewriting."), - OutputParam("height", type_hint=int, description="The resolved image height in pixels."), - OutputParam("width", type_hint=int, description="The resolved image width in pixels."), - ] - - @staticmethod - def _enhance_prompt( - pe: Ministral3ForCausalLM, - pe_tokenizer: AutoTokenizer, - prompt: str, - device: torch.device, - width: int, - height: int, - system_prompt: str | None, - temperature: float, - top_p: float, - ) -> str: - user_content = json.dumps({"prompt": prompt, "width": width, "height": height}, ensure_ascii=False) - messages = [] - if system_prompt is not None: - messages.append({"role": "system", "content": system_prompt}) - messages.append({"role": "user", "content": user_content}) - - input_text = pe_tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=False) - inputs = pe_tokenizer(input_text, return_tensors="pt").to(device) - output_ids = pe.generate( - **inputs, - max_new_tokens=pe_tokenizer.model_max_length, - do_sample=temperature != 1.0 or top_p != 1.0, - temperature=temperature, - top_p=top_p, - pad_token_id=pe_tokenizer.pad_token_id, - eos_token_id=pe_tokenizer.eos_token_id, - ) - generated_ids = output_ids[0][inputs["input_ids"].shape[1] :] - return pe_tokenizer.decode(generated_ids, skip_special_tokens=True).strip() - - @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - prompt = block_state.prompt - if isinstance(prompt, str): - prompt = [prompt] - - height = block_state.height or components.default_height - width = block_state.width or components.default_width - - revised = [ - self._enhance_prompt( - pe=components.pe, - pe_tokenizer=components.pe_tokenizer, - prompt=p, - device=device, - width=width, - height=height, - system_prompt=block_state.pe_system_prompt, - temperature=block_state.pe_temperature, - top_p=block_state.pe_top_p, - ) - for p in prompt - ] - - block_state.prompt = revised - block_state.height = height - block_state.width = width - - self.set_block_state(state, block_state) - return components, state - - -class ErnieImageTextEncoderStep(ModularPipelineBlocks): - model_name = "ernie-image" - - @property - def description(self) -> str: - return ( - "Text encoder step that encodes prompts into variable-length hidden states for the ErnieImage transformer." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", Mistral3Model), - ComponentSpec("tokenizer", AutoTokenizer), - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 4.0}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("prompt", type_hint=str, description="The prompt or prompts to guide image generation."), - InputParam( - "negative_prompt", - type_hint=str, - description="The prompt or prompts to avoid during image generation.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "prompt_embeds", - type_hint=list, - kwargs_type="denoiser_input_fields", - description="List of per-prompt text embeddings of shape (T, H).", - ), - OutputParam( - "negative_prompt_embeds", - type_hint=list, - kwargs_type="denoiser_input_fields", - description="List of per-prompt negative text embeddings for classifier-free guidance.", - ), - ] - - @staticmethod - def _encode( - text_encoder: Mistral3Model, - tokenizer: AutoTokenizer, - prompt: list[str], - device: torch.device, - ) -> list[torch.Tensor]: - text_hiddens = [] - for p in prompt: - ids = tokenizer(p, add_special_tokens=True, truncation=True, padding=False)["input_ids"] - if len(ids) == 0: - ids = [tokenizer.bos_token_id if tokenizer.bos_token_id is not None else 0] - input_ids = torch.tensor([ids], device=device) - outputs = text_encoder(input_ids=input_ids, output_hidden_states=True) - text_hiddens.append(outputs.hidden_states[-2][0]) - return text_hiddens - - @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - prompt = block_state.prompt - if prompt is None: - prompt = [""] - if isinstance(prompt, str): - prompt = [prompt] - - block_state.prompt_embeds = self._encode( - text_encoder=components.text_encoder, - tokenizer=components.tokenizer, - prompt=prompt, - device=device, - ) - - if components.requires_unconditional_embeds: - negative_prompt = block_state.negative_prompt - if negative_prompt is None: - negative_prompt = "" - if isinstance(negative_prompt, str): - negative_prompt = [negative_prompt] * len(prompt) - if len(negative_prompt) != len(prompt): - raise ValueError( - f"`negative_prompt` must have the same length as `prompt` ({len(prompt)}), " - f"got {len(negative_prompt)}." - ) - block_state.negative_prompt_embeds = self._encode( - text_encoder=components.text_encoder, - tokenizer=components.tokenizer, - prompt=negative_prompt, - device=device, - ) - else: - block_state.negative_prompt_embeds = None - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/ernie_image/modular_blocks_ernie_image.py b/diffusers/modular_pipelines/ernie_image/modular_blocks_ernie_image.py deleted file mode 100644 index 17e4eebaffda980c126a3e53e57d1cbd0ab11fca..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ernie_image/modular_blocks_ernie_image.py +++ /dev/null @@ -1,200 +0,0 @@ -# Copyright 2025 Baidu ERNIE-Image Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from ...utils import logging -from ..modular_pipeline import ConditionalPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import OutputParam -from .before_denoise import ( - ErnieImagePrepareLatentsStep, - ErnieImageSetTimestepsStep, - ErnieImageTextInputStep, -) -from .decoders import ErnieImageVaeDecoderStep -from .denoise import ErnieImageDenoiseStep -from .encoders import ErnieImagePromptEnhancerStep, ErnieImageTextEncoderStep - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# auto_docstring -class ErnieImageAutoPromptEnhancerStep(ConditionalPipelineBlocks): - """ - Conditional block that runs the optional prompt enhancer when `use_pe` is truthy. - - `ErnieImagePromptEnhancerStep` is used when `use_pe=True`. - - If `use_pe` is `None` or `False`, the step is skipped. - - Components: - pe (`Ministral3ForCausalLM`) pe_tokenizer (`AutoTokenizer`) - - Inputs: - prompt (`str`, *optional*): - The prompt or prompts to guide image generation. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - pe_system_prompt (`str`, *optional*): - Optional system prompt passed to the prompt enhancer. - pe_temperature (`float`, *optional*, defaults to 0.6): - Sampling temperature used when generating with the prompt enhancer. - pe_top_p (`float`, *optional*, defaults to 0.95): - Nucleus sampling `top_p` used when generating with the prompt enhancer. - - Outputs: - prompt (`list`): - The prompt list after prompt-enhancer rewriting. - height (`int`): - The resolved image height in pixels. - width (`int`): - The resolved image width in pixels. - """ - - model_name = "ernie-image" - block_classes = [ErnieImagePromptEnhancerStep] - block_names = ["prompt_enhancer"] - block_trigger_inputs = ["use_pe"] - - def select_block(self, use_pe=None) -> str | None: - if use_pe: - return "prompt_enhancer" - return None - - @property - def description(self): - return ( - "Conditional block that runs the optional prompt enhancer when `use_pe` is truthy.\n" - " - `ErnieImagePromptEnhancerStep` is used when `use_pe=True`.\n" - " - If `use_pe` is `None` or `False`, the step is skipped." - ) - - -# auto_docstring -class ErnieImageCoreDenoiseStep(SequentialPipelineBlocks): - """ - Denoise block that takes encoded conditions and runs the denoising process for ErnieImage. - - Components: - transformer (`ErnieImageTransformer2DModel`) scheduler (`FlowMatchEulerDiscreteScheduler`) guider - (`ClassifierFreeGuidance`) - - Inputs: - prompt_embeds (`list`): - List of per-prompt text embeddings from the text encoder step. - negative_prompt_embeds (`list`, *optional*): - List of per-prompt negative text embeddings from the text encoder step. - num_images_per_prompt (`int`, *optional*, defaults to 1): - Number of images to generate per prompt. - num_inference_steps (`int`, *optional*, defaults to 50): - Number of denoising steps. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - latents (`Tensor`, *optional*): - Pre-generated noisy latents. If provided, skips noise sampling. - generator (`Generator`, *optional*): - Torch generator for deterministic noise sampling. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "ernie-image" - block_classes = [ - ErnieImageTextInputStep, - ErnieImageSetTimestepsStep, - ErnieImagePrepareLatentsStep, - ErnieImageDenoiseStep, - ] - block_names = ["input", "set_timesteps", "prepare_latents", "denoise"] - - @property - def description(self): - return "Denoise block that takes encoded conditions and runs the denoising process for ErnieImage." - - @property - def outputs(self): - return [OutputParam.template("latents")] - - -# auto_docstring -class ErnieImageAutoBlocks(SequentialPipelineBlocks): - """ - Auto modular pipeline for ErnieImage text-to-image generation. Supports an optional prompt enhancer when the `pe` - components are loaded and `use_pe=True`. - - Supported workflows: - - `text2image`: requires `prompt` - - Components: - pe (`Ministral3ForCausalLM`) pe_tokenizer (`AutoTokenizer`) text_encoder (`Mistral3Model`) tokenizer - (`AutoTokenizer`) guider (`ClassifierFreeGuidance`) transformer (`ErnieImageTransformer2DModel`) scheduler - (`FlowMatchEulerDiscreteScheduler`) vae (`AutoencoderKLFlux2`) pachifier (`ErnieImagePachifier`) - image_processor (`VaeImageProcessor`) - - Inputs: - prompt (`str`, *optional*): - The prompt or prompts to guide image generation. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - pe_system_prompt (`str`, *optional*): - Optional system prompt passed to the prompt enhancer. - pe_temperature (`float`, *optional*, defaults to 0.6): - Sampling temperature used when generating with the prompt enhancer. - pe_top_p (`float`, *optional*, defaults to 0.95): - Nucleus sampling `top_p` used when generating with the prompt enhancer. - negative_prompt (`str`, *optional*): - The prompt or prompts to avoid during image generation. - num_images_per_prompt (`int`, *optional*, defaults to 1): - Number of images to generate per prompt. - num_inference_steps (`int`, *optional*, defaults to 50): - Number of denoising steps. - latents (`Tensor`, *optional*): - Pre-generated noisy latents. If provided, skips noise sampling. - generator (`Generator`, *optional*): - Torch generator for deterministic noise sampling. - output_type (`str`, *optional*, defaults to pil): - Output format: 'pil', 'np', or 'pt'. - - Outputs: - images (`list`): - Generated images. - """ - - model_name = "ernie-image" - block_classes = [ - ErnieImageAutoPromptEnhancerStep, - ErnieImageTextEncoderStep, - ErnieImageCoreDenoiseStep, - ErnieImageVaeDecoderStep, - ] - block_names = ["prompt_enhancer", "text_encoder", "denoise", "decode"] - _workflow_map = { - "text2image": {"prompt": True}, - } - - @property - def description(self): - return ( - "Auto modular pipeline for ErnieImage text-to-image generation. Supports an optional prompt enhancer " - "when the `pe` components are loaded and `use_pe=True`." - ) - - @property - def outputs(self): - return [OutputParam.template("images")] diff --git a/diffusers/modular_pipelines/ernie_image/modular_pipeline.py b/diffusers/modular_pipelines/ernie_image/modular_pipeline.py deleted file mode 100644 index f4cb2204369c9b69e4242b2e185ee2de4c3aec6f..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ernie_image/modular_pipeline.py +++ /dev/null @@ -1,110 +0,0 @@ -# Copyright 2025 Baidu ERNIE-Image Team and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import ErnieImageLoraLoaderMixin -from ...utils import logging -from ..modular_pipeline import ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class ErnieImagePachifier(ConfigMixin): - """ - A class to pack and unpack latents for ErnieImage. - """ - - config_name = "config.json" - - @register_to_config - def __init__(self, patch_size: int = 2): - super().__init__() - - def pack_latents(self, latents: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, height, width = latents.shape - patch_size = self.config.patch_size - - if height % patch_size != 0 or width % patch_size != 0: - raise ValueError( - f"Latent height and width must be divisible by {patch_size}, but got {height} and {width}" - ) - - latents = latents.view( - batch_size, num_channels, height // patch_size, patch_size, width // patch_size, patch_size - ) - latents = latents.permute(0, 1, 3, 5, 2, 4) - return latents.reshape( - batch_size, num_channels * patch_size * patch_size, height // patch_size, width // patch_size - ) - - def unpack_latents(self, latents: torch.Tensor) -> torch.Tensor: - batch_size, num_channels, height, width = latents.shape - patch_size = self.config.patch_size - - latents = latents.reshape( - batch_size, num_channels // (patch_size * patch_size), patch_size, patch_size, height, width - ) - latents = latents.permute(0, 1, 4, 2, 5, 3) - return latents.reshape( - batch_size, num_channels // (patch_size * patch_size), height * patch_size, width * patch_size - ) - - -class ErnieImageModularPipeline(ModularPipeline, ErnieImageLoraLoaderMixin): - """ - A ModularPipeline for ErnieImage. - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "ErnieImageAutoBlocks" - - @property - def default_height(self): - return 1024 - - @property - def default_width(self): - return 1024 - - @property - def vae_scale_factor(self): - vae_scale_factor = 16 - if hasattr(self, "vae") and self.vae is not None: - vae_scale_factor = 2 ** len(self.vae.config.block_out_channels) - return vae_scale_factor - - @property - def num_channels_latents(self): - num_channels_latents = 128 - if hasattr(self, "transformer") and self.transformer is not None: - num_channels_latents = self.transformer.config.in_channels - return num_channels_latents - - @property - def text_in_dim(self): - text_in_dim = 3584 - if hasattr(self, "transformer") and self.transformer is not None: - text_in_dim = self.transformer.config.text_in_dim - return text_in_dim - - @property - def requires_unconditional_embeds(self): - requires_unconditional_embeds = False - if hasattr(self, "guider") and self.guider is not None: - requires_unconditional_embeds = self.guider._enabled and self.guider.num_conditions > 1 - return requires_unconditional_embeds diff --git a/diffusers/modular_pipelines/flux/__init__.py b/diffusers/modular_pipelines/flux/__init__.py deleted file mode 100644 index 4754ed01ce6aee9c85259fef1774ea96fbd27009..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux/__init__.py +++ /dev/null @@ -1,49 +0,0 @@ -from typing import TYPE_CHECKING - -from ...utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, -) - - -_dummy_objects = {} -_import_structure = {} - -try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ...utils import dummy_torch_and_transformers_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) -else: - _import_structure["modular_blocks_flux"] = ["FluxAutoBlocks"] - _import_structure["modular_blocks_flux_kontext"] = ["FluxKontextAutoBlocks"] - _import_structure["modular_pipeline"] = ["FluxKontextModularPipeline", "FluxModularPipeline"] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .modular_blocks_flux import FluxAutoBlocks - from .modular_blocks_flux_kontext import FluxKontextAutoBlocks - from .modular_pipeline import FluxKontextModularPipeline, FluxModularPipeline -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/flux/before_denoise.py b/diffusers/modular_pipelines/flux/before_denoise.py deleted file mode 100644 index 2d41cd76cd93563d5f4aa0d2f273ee0449127686..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux/before_denoise.py +++ /dev/null @@ -1,618 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect - -import numpy as np -import torch - -from ...pipelines import FluxPipeline -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import FluxModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps -def retrieve_timesteps( - scheduler, - num_inference_steps: int | None = None, - device: str | torch.device | None = None, - timesteps: list[int] | None = None, - sigmas: list[float] | None = None, - **kwargs, -): - r""" - Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles - custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`. - - Args: - scheduler (`SchedulerMixin`): - The scheduler to get timesteps from. - num_inference_steps (`int`): - The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps` - must be `None`. - device (`str` or `torch.device`, *optional*): - The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. - timesteps (`list[int]`, *optional*): - Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed, - `num_inference_steps` and `sigmas` must be `None`. - sigmas (`list[float]`, *optional*): - Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed, - `num_inference_steps` and `timesteps` must be `None`. - - Returns: - `tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the - second element is the number of inference steps. - """ - if timesteps is not None and sigmas is not None: - raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values") - if timesteps is not None: - accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) - if not accepts_timesteps: - raise ValueError( - f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" - f" timestep schedules. Please check whether you are using the correct scheduler." - ) - scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) - timesteps = scheduler.timesteps - num_inference_steps = len(timesteps) - elif sigmas is not None: - accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) - if not accept_sigmas: - raise ValueError( - f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" - f" sigmas schedules. Please check whether you are using the correct scheduler." - ) - scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs) - timesteps = scheduler.timesteps - num_inference_steps = len(timesteps) - else: - scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) - timesteps = scheduler.timesteps - return timesteps, num_inference_steps - - -# Copied from diffusers.pipelines.flux.pipeline_flux.calculate_shift -def calculate_shift( - image_seq_len, - base_seq_len: int = 256, - max_seq_len: int = 4096, - base_shift: float = 0.5, - max_shift: float = 1.15, -): - m = (max_shift - base_shift) / (max_seq_len - base_seq_len) - b = base_shift - m * base_seq_len - mu = image_seq_len * m + b - return mu - - -# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.retrieve_latents -def retrieve_latents( - encoder_output: torch.Tensor, generator: torch.Generator | None = None, sample_mode: str = "sample" -): - if hasattr(encoder_output, "latent_dist") and sample_mode == "sample": - return encoder_output.latent_dist.sample(generator) - elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax": - return encoder_output.latent_dist.mode() - elif hasattr(encoder_output, "latents"): - return encoder_output.latents - else: - raise AttributeError("Could not access latents of provided encoder_output") - - -def _get_initial_timesteps_and_optionals( - transformer, - scheduler, - batch_size, - height, - width, - vae_scale_factor, - num_inference_steps, - guidance_scale, - sigmas, - device, -): - image_seq_len = (int(height) // vae_scale_factor // 2) * (int(width) // vae_scale_factor // 2) - - sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) if sigmas is None else sigmas - if hasattr(scheduler.config, "use_flow_sigmas") and scheduler.config.use_flow_sigmas: - sigmas = None - mu = calculate_shift( - image_seq_len, - scheduler.config.get("base_image_seq_len", 256), - scheduler.config.get("max_image_seq_len", 4096), - scheduler.config.get("base_shift", 0.5), - scheduler.config.get("max_shift", 1.15), - ) - timesteps, num_inference_steps = retrieve_timesteps(scheduler, num_inference_steps, device, sigmas=sigmas, mu=mu) - if transformer.config.guidance_embeds: - guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32) - guidance = guidance.expand(batch_size) - else: - guidance = None - - return timesteps, num_inference_steps, sigmas, guidance - - -class FluxSetTimestepsStep(ModularPipelineBlocks): - model_name = "flux" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def description(self) -> str: - return "Step that sets the scheduler's timesteps for inference" - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("num_inference_steps", default=50), - InputParam("timesteps"), - InputParam("sigmas"), - InputParam("guidance_scale", default=3.5), - InputParam("latents", type_hint=torch.Tensor), - InputParam("num_images_per_prompt", default=1), - InputParam("height", type_hint=int), - InputParam("width", type_hint=int), - InputParam( - "batch_size", - required=True, - type_hint=int, - description="Number of prompts, the final batch size of model inputs should be `batch_size * num_images_per_prompt`. Can be generated in input step.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("timesteps", type_hint=torch.Tensor, description="The timesteps to use for inference"), - OutputParam( - "num_inference_steps", - type_hint=int, - description="The number of denoising steps to perform at inference time", - ), - OutputParam("guidance", type_hint=torch.Tensor, description="Optional guidance to be used."), - ] - - @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - block_state.device = components._execution_device - - scheduler = components.scheduler - transformer = components.transformer - - batch_size = block_state.batch_size * block_state.num_images_per_prompt - timesteps, num_inference_steps, sigmas, guidance = _get_initial_timesteps_and_optionals( - transformer, - scheduler, - batch_size, - block_state.height, - block_state.width, - components.vae_scale_factor, - block_state.num_inference_steps, - block_state.guidance_scale, - block_state.sigmas, - block_state.device, - ) - block_state.timesteps = timesteps - block_state.num_inference_steps = num_inference_steps - block_state.sigmas = sigmas - block_state.guidance = guidance - - # We set the index here to remove DtoH sync, helpful especially during compilation. - # Check out more details here: https://github.com/huggingface/diffusers/pull/11696 - components.scheduler.set_begin_index(0) - - self.set_block_state(state, block_state) - return components, state - - -class FluxImg2ImgSetTimestepsStep(ModularPipelineBlocks): - model_name = "flux" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def description(self) -> str: - return "Step that sets the scheduler's timesteps for inference" - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("num_inference_steps", default=50), - InputParam("timesteps"), - InputParam("sigmas"), - InputParam("strength", default=0.6), - InputParam("guidance_scale", default=3.5), - InputParam("num_images_per_prompt", default=1), - InputParam("height", type_hint=int), - InputParam("width", type_hint=int), - InputParam( - "batch_size", - required=True, - type_hint=int, - description="Number of prompts, the final batch size of model inputs should be `batch_size * num_images_per_prompt`. Can be generated in input step.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("timesteps", type_hint=torch.Tensor, description="The timesteps to use for inference"), - OutputParam( - "num_inference_steps", - type_hint=int, - description="The number of denoising steps to perform at inference time", - ), - OutputParam("guidance", type_hint=torch.Tensor, description="Optional guidance to be used."), - ] - - @staticmethod - # Copied from diffusers.pipelines.stable_diffusion_3.pipeline_stable_diffusion_3_img2img.StableDiffusion3Img2ImgPipeline.get_timesteps with self.scheduler->scheduler - def get_timesteps(scheduler, num_inference_steps, strength, device): - # get the original timestep using init_timestep - init_timestep = min(num_inference_steps * strength, num_inference_steps) - - t_start = int(max(num_inference_steps - init_timestep, 0)) - timesteps = scheduler.timesteps[t_start * scheduler.order :] - if hasattr(scheduler, "set_begin_index"): - scheduler.set_begin_index(t_start * scheduler.order) - - return timesteps, num_inference_steps - t_start - - @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - block_state.device = components._execution_device - - block_state.height = block_state.height or components.default_height - block_state.width = block_state.width or components.default_width - - scheduler = components.scheduler - transformer = components.transformer - batch_size = block_state.batch_size * block_state.num_images_per_prompt - timesteps, num_inference_steps, sigmas, guidance = _get_initial_timesteps_and_optionals( - transformer, - scheduler, - batch_size, - block_state.height, - block_state.width, - components.vae_scale_factor, - block_state.num_inference_steps, - block_state.guidance_scale, - block_state.sigmas, - block_state.device, - ) - timesteps, num_inference_steps = self.get_timesteps( - scheduler, num_inference_steps, block_state.strength, block_state.device - ) - block_state.timesteps = timesteps - block_state.num_inference_steps = num_inference_steps - block_state.sigmas = sigmas - block_state.guidance = guidance - - self.set_block_state(state, block_state) - return components, state - - -class FluxPrepareLatentsStep(ModularPipelineBlocks): - model_name = "flux" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [] - - @property - def description(self) -> str: - return "Prepare latents step that prepares the latents for the text-to-image generation process" - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("height", type_hint=int), - InputParam("width", type_hint=int), - InputParam("latents", type_hint=torch.Tensor | None), - InputParam("num_images_per_prompt", type_hint=int, default=1), - InputParam("generator"), - InputParam( - "batch_size", - required=True, - type_hint=int, - description="Number of prompts, the final batch size of model inputs should be `batch_size * num_images_per_prompt`. Can be generated in input step.", - ), - InputParam("dtype", type_hint=torch.dtype, description="The dtype of the model inputs"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "latents", type_hint=torch.Tensor, description="The initial latents to use for the denoising process" - ), - ] - - @staticmethod - def check_inputs(components, block_state): - if (block_state.height is not None and block_state.height % (components.vae_scale_factor * 2) != 0) or ( - block_state.width is not None and block_state.width % (components.vae_scale_factor * 2) != 0 - ): - logger.warning( - f"`height` and `width` have to be divisible by {components.vae_scale_factor} but are {block_state.height} and {block_state.width}." - ) - - @staticmethod - def prepare_latents( - comp, - batch_size, - num_channels_latents, - height, - width, - dtype, - device, - generator, - latents=None, - ): - height = 2 * (int(height) // (comp.vae_scale_factor * 2)) - width = 2 * (int(width) // (comp.vae_scale_factor * 2)) - - shape = (batch_size, num_channels_latents, height, width) - - if latents is not None: - return latents.to(device=device, dtype=dtype) - - if isinstance(generator, list) and len(generator) != batch_size: - raise ValueError( - f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" - f" size of {batch_size}. Make sure the batch size matches the length of the generators." - ) - - # TODO: move packing latents code to a patchifier similar to Qwen - latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) - latents = FluxPipeline._pack_latents(latents, batch_size, num_channels_latents, height, width) - - return latents - - @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - block_state.height = block_state.height or components.default_height - block_state.width = block_state.width or components.default_width - block_state.device = components._execution_device - block_state.num_channels_latents = components.num_channels_latents - - self.check_inputs(components, block_state) - batch_size = block_state.batch_size * block_state.num_images_per_prompt - block_state.latents = self.prepare_latents( - components, - batch_size, - block_state.num_channels_latents, - block_state.height, - block_state.width, - block_state.dtype, - block_state.device, - block_state.generator, - block_state.latents, - ) - - self.set_block_state(state, block_state) - - return components, state - - -class FluxImg2ImgPrepareLatentsStep(ModularPipelineBlocks): - model_name = "flux" - - @property - def description(self) -> str: - return "Step that adds noise to image latents for image-to-image. Should be run after `set_timesteps`," - " `prepare_latents`. Both noise and image latents should already be patchified." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="latents", - required=True, - type_hint=torch.Tensor, - description="The initial random noised, can be generated in prepare latent step.", - ), - InputParam( - name="image_latents", - required=True, - type_hint=torch.Tensor, - description="The image latents to use for the denoising process. Can be generated in vae encoder and packed in input step.", - ), - InputParam( - name="timesteps", - required=True, - type_hint=torch.Tensor, - description="The timesteps to use for the denoising process. Can be generated in set_timesteps step.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="initial_noise", - type_hint=torch.Tensor, - description="The initial random noised used for inpainting denoising.", - ), - ] - - @staticmethod - def check_inputs(image_latents, latents): - if image_latents.shape[0] != latents.shape[0]: - raise ValueError( - f"`image_latents` must have have same batch size as `latents`, but got {image_latents.shape[0]} and {latents.shape[0]}" - ) - - if image_latents.ndim != 3: - raise ValueError(f"`image_latents` must have 3 dimensions (patchified), but got {image_latents.ndim}") - - @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - self.check_inputs(image_latents=block_state.image_latents, latents=block_state.latents) - - # prepare latent timestep - latent_timestep = block_state.timesteps[:1].repeat(block_state.latents.shape[0]) - - # make copy of initial_noise - block_state.initial_noise = block_state.latents - - # scale noise - block_state.latents = components.scheduler.scale_noise( - block_state.image_latents, latent_timestep, block_state.latents - ) - - self.set_block_state(state, block_state) - - return components, state - - -class FluxRoPEInputsStep(ModularPipelineBlocks): - model_name = "flux" - - @property - def description(self) -> str: - return "Step that prepares the RoPE inputs for the denoising process. Should be placed after text encoder and latent preparation steps." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="height", required=True), - InputParam(name="width", required=True), - InputParam(name="prompt_embeds"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="txt_ids", - kwargs_type="denoiser_input_fields", - type_hint=list[int], - description="The sequence lengths of the prompt embeds, used for RoPE calculation.", - ), - OutputParam( - name="img_ids", - kwargs_type="denoiser_input_fields", - type_hint=list[int], - description="The sequence lengths of the image latents, used for RoPE calculation.", - ), - ] - - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - prompt_embeds = block_state.prompt_embeds - device, dtype = prompt_embeds.device, prompt_embeds.dtype - block_state.txt_ids = torch.zeros(prompt_embeds.shape[1], 3).to( - device=prompt_embeds.device, dtype=prompt_embeds.dtype - ) - - height = 2 * (int(block_state.height) // (components.vae_scale_factor * 2)) - width = 2 * (int(block_state.width) // (components.vae_scale_factor * 2)) - block_state.img_ids = FluxPipeline._prepare_latent_image_ids(None, height // 2, width // 2, device, dtype) - - self.set_block_state(state, block_state) - - return components, state - - -class FluxKontextRoPEInputsStep(ModularPipelineBlocks): - model_name = "flux-kontext" - - @property - def description(self) -> str: - return "Step that prepares the RoPE inputs for the denoising process of Flux Kontext. Should be placed after text encoder and latent preparation steps." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="image_height"), - InputParam(name="image_width"), - InputParam(name="height"), - InputParam(name="width"), - InputParam(name="prompt_embeds"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="txt_ids", - kwargs_type="denoiser_input_fields", - type_hint=list[int], - description="The sequence lengths of the prompt embeds, used for RoPE calculation.", - ), - OutputParam( - name="img_ids", - kwargs_type="denoiser_input_fields", - type_hint=list[int], - description="The sequence lengths of the image latents, used for RoPE calculation.", - ), - ] - - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - prompt_embeds = block_state.prompt_embeds - device, dtype = prompt_embeds.device, prompt_embeds.dtype - block_state.txt_ids = torch.zeros(prompt_embeds.shape[1], 3).to( - device=prompt_embeds.device, dtype=prompt_embeds.dtype - ) - - img_ids = None - if ( - getattr(block_state, "image_height", None) is not None - and getattr(block_state, "image_width", None) is not None - ): - image_latent_height = 2 * (int(block_state.image_height) // (components.vae_scale_factor * 2)) - image_latent_width = 2 * (int(block_state.image_width) // (components.vae_scale_factor * 2)) - img_ids = FluxPipeline._prepare_latent_image_ids( - None, image_latent_height // 2, image_latent_width // 2, device, dtype - ) - # image ids are the same as latent ids with the first dimension set to 1 instead of 0 - img_ids[..., 0] = 1 - - height = 2 * (int(block_state.height) // (components.vae_scale_factor * 2)) - width = 2 * (int(block_state.width) // (components.vae_scale_factor * 2)) - latent_ids = FluxPipeline._prepare_latent_image_ids(None, height // 2, width // 2, device, dtype) - - if img_ids is not None: - latent_ids = torch.cat([latent_ids, img_ids], dim=0) - - block_state.img_ids = latent_ids - - self.set_block_state(state, block_state) - - return components, state diff --git a/diffusers/modular_pipelines/flux/decoders.py b/diffusers/modular_pipelines/flux/decoders.py deleted file mode 100644 index 5fcde50086807401b59047b8214a8b674bf6ed14..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux/decoders.py +++ /dev/null @@ -1,109 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import numpy as np -import PIL -import torch - -from ...configuration_utils import FrozenDict -from ...models import AutoencoderKL -from ...utils import logging -from ...video_processor import VaeImageProcessor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _unpack_latents(latents, height, width, vae_scale_factor): - batch_size, num_patches, channels = latents.shape - - # VAE applies 8x compression on images but we must also account for packing which requires - # latent height and width to be divisible by 2. - height = 2 * (int(height) // (vae_scale_factor * 2)) - width = 2 * (int(width) // (vae_scale_factor * 2)) - - latents = latents.view(batch_size, height // 2, width // 2, channels // 4, 2, 2) - latents = latents.permute(0, 3, 1, 4, 2, 5) - - latents = latents.reshape(batch_size, channels // (2 * 2), height, width) - - return latents - - -class FluxDecodeStep(ModularPipelineBlocks): - model_name = "flux" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKL), - ComponentSpec( - "image_processor", - VaeImageProcessor, - config=FrozenDict({"vae_scale_factor": 16}), - default_creation_method="from_config", - ), - ] - - @property - def description(self) -> str: - return "Step that decodes the denoised latents into images" - - @property - def inputs(self) -> list[tuple[str, Any]]: - return [ - InputParam("output_type", default="pil"), - InputParam("height", default=1024), - InputParam("width", default=1024), - InputParam( - "latents", - required=True, - type_hint=torch.Tensor, - description="The denoised latents from the denoising step", - ), - ] - - @property - def intermediate_outputs(self) -> list[str]: - return [ - OutputParam( - "images", - type_hint=list[PIL.Image.Image] | torch.Tensor | np.ndarray, - description="The generated images, can be a list of PIL.Image.Image, torch.Tensor or a numpy array", - ) - ] - - @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - vae = components.vae - - if not block_state.output_type == "latent": - latents = block_state.latents - latents = _unpack_latents(latents, block_state.height, block_state.width, components.vae_scale_factor) - latents = (latents / vae.config.scaling_factor) + vae.config.shift_factor - block_state.images = vae.decode(latents, return_dict=False)[0] - block_state.images = components.image_processor.postprocess( - block_state.images, output_type=block_state.output_type - ) - else: - block_state.images = block_state.latents - - self.set_block_state(state, block_state) - - return components, state diff --git a/diffusers/modular_pipelines/flux/denoise.py b/diffusers/modular_pipelines/flux/denoise.py deleted file mode 100644 index 490ef6d88f57da290e61063750f3ac4b049fde07..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux/denoise.py +++ /dev/null @@ -1,322 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch - -from ...models import FluxTransformer2DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ..modular_pipeline import ( - BlockState, - LoopSequentialPipelineBlocks, - ModularPipelineBlocks, - PipelineState, -) -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import FluxModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class FluxLoopDenoiser(ModularPipelineBlocks): - model_name = "flux" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", FluxTransformer2DModel)] - - @property - def description(self) -> str: - return ( - "Step within the denoising loop that denoise the latents. " - "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " - "object (e.g. `FluxDenoiseLoopWrapper`)" - ) - - @property - def inputs(self) -> list[tuple[str, Any]]: - return [ - InputParam("joint_attention_kwargs"), - InputParam( - "latents", - required=True, - type_hint=torch.Tensor, - description="The initial latents to use for the denoising process. Can be generated in prepare_latent step.", - ), - InputParam( - "guidance", - required=False, - type_hint=torch.Tensor, - description="Guidance scale as a tensor", - ), - InputParam( - "prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Prompt embeddings", - ), - InputParam( - "pooled_prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Pooled prompt embeddings", - ), - InputParam( - "txt_ids", - required=True, - type_hint=torch.Tensor, - description="IDs computed from text sequence needed for RoPE", - ), - InputParam( - "img_ids", - required=True, - type_hint=torch.Tensor, - description="IDs computed from image sequence needed for RoPE", - ), - ] - - @torch.no_grad() - def __call__( - self, components: FluxModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: - noise_pred = components.transformer( - hidden_states=block_state.latents, - timestep=t.flatten() / 1000, - guidance=block_state.guidance, - encoder_hidden_states=block_state.prompt_embeds, - pooled_projections=block_state.pooled_prompt_embeds, - joint_attention_kwargs=block_state.joint_attention_kwargs, - txt_ids=block_state.txt_ids, - img_ids=block_state.img_ids, - return_dict=False, - )[0] - block_state.noise_pred = noise_pred - - return components, block_state - - -class FluxKontextLoopDenoiser(ModularPipelineBlocks): - model_name = "flux-kontext" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", FluxTransformer2DModel)] - - @property - def description(self) -> str: - return ( - "Step within the denoising loop that denoise the latents for Flux Kontext. " - "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " - "object (e.g. `FluxDenoiseLoopWrapper`)" - ) - - @property - def inputs(self) -> list[tuple[str, Any]]: - return [ - InputParam("joint_attention_kwargs"), - InputParam( - "latents", - required=True, - type_hint=torch.Tensor, - description="The initial latents to use for the denoising process. Can be generated in prepare_latent step.", - ), - InputParam( - "image_latents", - type_hint=torch.Tensor, - description="Image latents to use for the denoising process. Can be generated in prepare_latent step.", - ), - InputParam( - "guidance", - required=False, - type_hint=torch.Tensor, - description="Guidance scale as a tensor", - ), - InputParam( - "prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Prompt embeddings", - ), - InputParam( - "pooled_prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Pooled prompt embeddings", - ), - InputParam( - "txt_ids", - required=True, - type_hint=torch.Tensor, - description="IDs computed from text sequence needed for RoPE", - ), - InputParam( - "img_ids", - required=True, - type_hint=torch.Tensor, - description="IDs computed from latent sequence needed for RoPE", - ), - ] - - @torch.no_grad() - def __call__( - self, components: FluxModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: - latents = block_state.latents - latent_model_input = latents - image_latents = block_state.image_latents - if image_latents is not None: - latent_model_input = torch.cat([latent_model_input, image_latents], dim=1) - - timestep = t.expand(latents.shape[0]).to(latents.dtype) - noise_pred = components.transformer( - hidden_states=latent_model_input, - timestep=timestep / 1000, - guidance=block_state.guidance, - encoder_hidden_states=block_state.prompt_embeds, - pooled_projections=block_state.pooled_prompt_embeds, - joint_attention_kwargs=block_state.joint_attention_kwargs, - txt_ids=block_state.txt_ids, - img_ids=block_state.img_ids, - return_dict=False, - )[0] - noise_pred = noise_pred[:, : latents.size(1)] - block_state.noise_pred = noise_pred - - return components, block_state - - -class FluxLoopAfterDenoiser(ModularPipelineBlocks): - model_name = "flux" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def description(self) -> str: - return ( - "step within the denoising loop that update the latents. " - "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " - "object (e.g. `FluxDenoiseLoopWrapper`)" - ) - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam("latents", type_hint=torch.Tensor, description="The denoised latents")] - - @torch.no_grad() - def __call__(self, components: FluxModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - # Perform scheduler step using the predicted output - latents_dtype = block_state.latents.dtype - block_state.latents = components.scheduler.step( - block_state.noise_pred, - t, - block_state.latents, - return_dict=False, - )[0] - - if block_state.latents.dtype != latents_dtype: - block_state.latents = block_state.latents.to(latents_dtype) - - return components, block_state - - -class FluxDenoiseLoopWrapper(LoopSequentialPipelineBlocks): - model_name = "flux" - - @property - def description(self) -> str: - return ( - "Pipeline block that iteratively denoise the latents over `timesteps`. " - "The specific steps with each iteration can be customized with `sub_blocks` attributes" - ) - - @property - def loop_expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), - ComponentSpec("transformer", FluxTransformer2DModel), - ] - - @property - def loop_inputs(self) -> list[InputParam]: - return [ - InputParam( - "timesteps", - required=True, - type_hint=torch.Tensor, - description="The timesteps to use for the denoising process. Can be generated in set_timesteps step.", - ), - InputParam( - "num_inference_steps", - required=True, - type_hint=int, - description="The number of inference steps to use for the denoising process. Can be generated in set_timesteps step.", - ), - ] - - @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - block_state.num_warmup_steps = max( - len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order, 0 - ) - with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) - if i == len(block_state.timesteps) - 1 or ( - (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 - ): - progress_bar.update() - - self.set_block_state(state, block_state) - - return components, state - - -class FluxDenoiseStep(FluxDenoiseLoopWrapper): - block_classes = [FluxLoopDenoiser, FluxLoopAfterDenoiser] - block_names = ["denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoise step that iteratively denoise the latents. \n" - "Its loop logic is defined in `FluxDenoiseLoopWrapper.__call__` method \n" - "At each iteration, it runs blocks defined in `sub_blocks` sequentially:\n" - " - `FluxLoopDenoiser`\n" - " - `FluxLoopAfterDenoiser`\n" - "This block supports both text2image and img2img tasks." - ) - - -class FluxKontextDenoiseStep(FluxDenoiseLoopWrapper): - model_name = "flux-kontext" - block_classes = [FluxKontextLoopDenoiser, FluxLoopAfterDenoiser] - block_names = ["denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoise step that iteratively denoise the latents. \n" - "Its loop logic is defined in `FluxDenoiseLoopWrapper.__call__` method \n" - "At each iteration, it runs blocks defined in `sub_blocks` sequentially:\n" - " - `FluxKontextLoopDenoiser`\n" - " - `FluxLoopAfterDenoiser`\n" - "This block supports both text2image and img2img tasks." - ) diff --git a/diffusers/modular_pipelines/flux/encoders.py b/diffusers/modular_pipelines/flux/encoders.py deleted file mode 100644 index 5f7e61a535b76121758948315de08c80bcc56683..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux/encoders.py +++ /dev/null @@ -1,480 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import html - -import regex as re -import torch -from transformers import CLIPTextModel, CLIPTokenizer, T5EncoderModel, T5TokenizerFast - -from ...configuration_utils import FrozenDict -from ...image_processor import VaeImageProcessor, is_valid_image, is_valid_image_imagelist -from ...loaders import FluxLoraLoaderMixin, TextualInversionLoaderMixin -from ...models import AutoencoderKL -from ...utils import USE_PEFT_BACKEND, is_ftfy_available, logging, scale_lora_layers, unscale_lora_layers -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import FluxModularPipeline - - -if is_ftfy_available(): - import ftfy - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def basic_clean(text): - text = ftfy.fix_text(text) - text = html.unescape(html.unescape(text)) - return text.strip() - - -def whitespace_clean(text): - text = re.sub(r"\s+", " ", text) - text = text.strip() - return text - - -def prompt_clean(text): - text = whitespace_clean(basic_clean(text)) - return text - - -# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.retrieve_latents -def retrieve_latents( - encoder_output: torch.Tensor, generator: torch.Generator | None = None, sample_mode: str = "sample" -): - if hasattr(encoder_output, "latent_dist") and sample_mode == "sample": - return encoder_output.latent_dist.sample(generator) - elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax": - return encoder_output.latent_dist.mode() - elif hasattr(encoder_output, "latents"): - return encoder_output.latents - else: - raise AttributeError("Could not access latents of provided encoder_output") - - -def encode_vae_image(vae: AutoencoderKL, image: torch.Tensor, generator: torch.Generator, sample_mode="sample"): - if isinstance(generator, list): - image_latents = [ - retrieve_latents(vae.encode(image[i : i + 1]), generator=generator[i], sample_mode=sample_mode) - for i in range(image.shape[0]) - ] - image_latents = torch.cat(image_latents, dim=0) - else: - image_latents = retrieve_latents(vae.encode(image), generator=generator, sample_mode=sample_mode) - - image_latents = (image_latents - vae.config.shift_factor) * vae.config.scaling_factor - - return image_latents - - -class FluxProcessImagesInputStep(ModularPipelineBlocks): - model_name = "flux" - - @property - def description(self) -> str: - return "Image Preprocess step." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "image_processor", - VaeImageProcessor, - config=FrozenDict({"vae_scale_factor": 16, "vae_latent_channels": 16}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [InputParam("resized_image"), InputParam("image"), InputParam("height"), InputParam("width")] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam(name="processed_image")] - - @staticmethod - def check_inputs(height, width, vae_scale_factor): - if height is not None and height % (vae_scale_factor * 2) != 0: - raise ValueError(f"Height must be divisible by {vae_scale_factor * 2} but is {height}") - - if width is not None and width % (vae_scale_factor * 2) != 0: - raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") - - @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState): - block_state = self.get_block_state(state) - - if block_state.resized_image is None and block_state.image is None: - raise ValueError("`resized_image` and `image` cannot be None at the same time") - - if block_state.resized_image is None: - image = block_state.image - self.check_inputs( - height=block_state.height, width=block_state.width, vae_scale_factor=components.vae_scale_factor - ) - height = block_state.height or components.default_height - width = block_state.width or components.default_width - else: - width, height = block_state.resized_image[0].size - image = block_state.resized_image - - block_state.processed_image = components.image_processor.preprocess(image=image, height=height, width=width) - - self.set_block_state(state, block_state) - return components, state - - -class FluxKontextProcessImagesInputStep(ModularPipelineBlocks): - model_name = "flux-kontext" - - @property - def description(self) -> str: - return ( - "Image preprocess step for Flux Kontext. The preprocessed image goes to the VAE.\n" - "Kontext works as a T2I model, too, in case no input image is provided." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "image_processor", - VaeImageProcessor, - config=FrozenDict({"vae_scale_factor": 16}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [InputParam("image"), InputParam("_auto_resize", type_hint=bool, default=True)] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam(name="processed_image")] - - @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState): - from ...pipelines.flux.pipeline_flux_kontext import PREFERRED_KONTEXT_RESOLUTIONS - - block_state = self.get_block_state(state) - images = block_state.image - - if images is None: - block_state.processed_image = None - - else: - multiple_of = components.image_processor.config.vae_scale_factor - - if not is_valid_image_imagelist(images): - raise ValueError(f"Images must be image or list of images but are {type(images)}") - - if is_valid_image(images): - images = [images] - - img = images[0] - image_height, image_width = components.image_processor.get_default_height_width(img) - aspect_ratio = image_width / image_height - _auto_resize = block_state._auto_resize - if _auto_resize: - # Kontext is trained on specific resolutions, using one of them is recommended - _, image_width, image_height = min( - (abs(aspect_ratio - w / h), w, h) for w, h in PREFERRED_KONTEXT_RESOLUTIONS - ) - image_width = image_width // multiple_of * multiple_of - image_height = image_height // multiple_of * multiple_of - images = components.image_processor.resize(images, image_height, image_width) - block_state.processed_image = components.image_processor.preprocess(images, image_height, image_width) - - self.set_block_state(state, block_state) - return components, state - - -class FluxVaeEncoderStep(ModularPipelineBlocks): - model_name = "flux" - - def __init__( - self, input_name: str = "processed_image", output_name: str = "image_latents", sample_mode: str = "sample" - ): - """Initialize a VAE encoder step for converting images to latent representations. - - Both the input and output names are configurable so this block can be configured to process to different image - inputs (e.g., "processed_image" -> "image_latents", "processed_control_image" -> "control_image_latents"). - - Args: - input_name (str, optional): Name of the input image tensor. Defaults to "processed_image". - Examples: "processed_image" or "processed_control_image" - output_name (str, optional): Name of the output latent tensor. Defaults to "image_latents". - Examples: "image_latents" or "control_image_latents" - sample_mode (str, optional): Sampling mode to be used. - - Examples: - # Basic usage with default settings (includes image processor): # FluxImageVaeEncoderDynamicStep() - - # Custom input/output names for control image: # FluxImageVaeEncoderDynamicStep( - input_name="processed_control_image", output_name="control_image_latents" - ) - """ - self._image_input_name = input_name - self._image_latents_output_name = output_name - self.sample_mode = sample_mode - super().__init__() - - @property - def description(self) -> str: - return f"Dynamic VAE Encoder step that converts {self._image_input_name} into latent representations {self._image_latents_output_name}.\n" - - @property - def expected_components(self) -> list[ComponentSpec]: - components = [ComponentSpec("vae", AutoencoderKL)] - return components - - @property - def inputs(self) -> list[InputParam]: - inputs = [InputParam(self._image_input_name), InputParam("generator")] - return inputs - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - self._image_latents_output_name, - type_hint=torch.Tensor, - description="The latents representing the reference image", - ) - ] - - @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - image = getattr(block_state, self._image_input_name) - - if image is None: - setattr(block_state, self._image_latents_output_name, None) - else: - device = components._execution_device - dtype = components.vae.dtype - image = image.to(device=device, dtype=dtype) - - # Encode image into latents - image_latents = encode_vae_image( - image=image, vae=components.vae, generator=block_state.generator, sample_mode=self.sample_mode - ) - setattr(block_state, self._image_latents_output_name, image_latents) - - self.set_block_state(state, block_state) - - return components, state - - -class FluxTextEncoderStep(ModularPipelineBlocks): - model_name = "flux" - - @property - def description(self) -> str: - return "Text Encoder step that generate text_embeddings to guide the image generation" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", CLIPTextModel), - ComponentSpec("tokenizer", CLIPTokenizer), - ComponentSpec("text_encoder_2", T5EncoderModel), - ComponentSpec("tokenizer_2", T5TokenizerFast), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("prompt"), - InputParam("prompt_2"), - InputParam("max_sequence_length", type_hint=int, default=512, required=False), - InputParam("joint_attention_kwargs"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "prompt_embeds", - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="text embeddings used to guide the image generation", - ), - OutputParam( - "pooled_prompt_embeds", - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="pooled text embeddings used to guide the image generation", - ), - ] - - @staticmethod - def check_inputs(block_state): - for prompt in [block_state.prompt, block_state.prompt_2]: - if prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): - raise ValueError(f"`prompt` or `prompt_2` has to be of type `str` or `list` but is {type(prompt)}") - - @staticmethod - def _get_t5_prompt_embeds(components, prompt: str | list[str], max_sequence_length: int, device: torch.device): - dtype = components.text_encoder_2.dtype - prompt = [prompt] if isinstance(prompt, str) else prompt - - if isinstance(components, TextualInversionLoaderMixin): - prompt = components.maybe_convert_prompt(prompt, components.tokenizer_2) - - text_inputs = components.tokenizer_2( - prompt, - padding="max_length", - max_length=max_sequence_length, - truncation=True, - return_length=False, - return_overflowing_tokens=False, - return_tensors="pt", - ) - text_input_ids = text_inputs.input_ids - - untruncated_ids = components.tokenizer_2(prompt, padding="longest", return_tensors="pt").input_ids - if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids): - removed_text = components.tokenizer_2.batch_decode(untruncated_ids[:, max_sequence_length - 1 : -1]) - logger.warning( - "The following part of your input was truncated because `max_sequence_length` is set to " - f" {max_sequence_length} tokens: {removed_text}" - ) - - prompt_embeds = components.text_encoder_2(text_input_ids.to(device), output_hidden_states=False)[0] - prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) - return prompt_embeds - - @staticmethod - def _get_clip_prompt_embeds(components, prompt: str | list[str], device: torch.device): - prompt = [prompt] if isinstance(prompt, str) else prompt - - if isinstance(components, TextualInversionLoaderMixin): - prompt = components.maybe_convert_prompt(prompt, components.tokenizer) - - text_inputs = components.tokenizer( - prompt, - padding="max_length", - max_length=components.tokenizer.model_max_length, - truncation=True, - return_overflowing_tokens=False, - return_length=False, - return_tensors="pt", - ) - - text_input_ids = text_inputs.input_ids - tokenizer_max_length = components.tokenizer.model_max_length - untruncated_ids = components.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids - if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids): - removed_text = components.tokenizer.batch_decode(untruncated_ids[:, tokenizer_max_length - 1 : -1]) - logger.warning( - "The following part of your input was truncated because CLIP can only handle sequences up to" - f" {tokenizer_max_length} tokens: {removed_text}" - ) - prompt_embeds = components.text_encoder(text_input_ids.to(device), output_hidden_states=False) - - # Use pooled output of CLIPTextModel - prompt_embeds = prompt_embeds.pooler_output - prompt_embeds = prompt_embeds.to(dtype=components.text_encoder.dtype, device=device) - - return prompt_embeds - - @staticmethod - def encode_prompt( - components, - prompt: str | list[str], - prompt_2: str | list[str], - device: torch.device | None = None, - prompt_embeds: torch.FloatTensor | None = None, - pooled_prompt_embeds: torch.FloatTensor | None = None, - max_sequence_length: int = 512, - lora_scale: float | None = None, - ): - device = device or components._execution_device - - # set lora scale so that monkey patched LoRA - # function of text encoder can correctly access it - if lora_scale is not None and isinstance(components, FluxLoraLoaderMixin): - components._lora_scale = lora_scale - - # dynamically adjust the LoRA scale - if components.text_encoder is not None and USE_PEFT_BACKEND: - scale_lora_layers(components.text_encoder, lora_scale) - if components.text_encoder_2 is not None and USE_PEFT_BACKEND: - scale_lora_layers(components.text_encoder_2, lora_scale) - - prompt = [prompt] if isinstance(prompt, str) else prompt - - if prompt_embeds is None: - prompt_2 = prompt_2 or prompt - prompt_2 = [prompt_2] if isinstance(prompt_2, str) else prompt_2 - - # We only use the pooled prompt output from the CLIPTextModel - pooled_prompt_embeds = FluxTextEncoderStep._get_clip_prompt_embeds( - components, - prompt=prompt, - device=device, - ) - prompt_embeds = FluxTextEncoderStep._get_t5_prompt_embeds( - components, - prompt=prompt_2, - max_sequence_length=max_sequence_length, - device=device, - ) - - if components.text_encoder is not None: - if isinstance(components, FluxLoraLoaderMixin) and USE_PEFT_BACKEND: - # Retrieve the original scale by scaling back the LoRA layers - unscale_lora_layers(components.text_encoder, lora_scale) - - if components.text_encoder_2 is not None: - if isinstance(components, FluxLoraLoaderMixin) and USE_PEFT_BACKEND: - # Retrieve the original scale by scaling back the LoRA layers - unscale_lora_layers(components.text_encoder_2, lora_scale) - - return prompt_embeds, pooled_prompt_embeds - - @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - # Get inputs and intermediates - block_state = self.get_block_state(state) - self.check_inputs(block_state) - - block_state.device = components._execution_device - - # Encode input prompt - block_state.text_encoder_lora_scale = ( - block_state.joint_attention_kwargs.get("scale", None) - if block_state.joint_attention_kwargs is not None - else None - ) - block_state.prompt_embeds, block_state.pooled_prompt_embeds = self.encode_prompt( - components, - prompt=block_state.prompt, - prompt_2=None, - prompt_embeds=None, - pooled_prompt_embeds=None, - device=block_state.device, - max_sequence_length=block_state.max_sequence_length, - lora_scale=block_state.text_encoder_lora_scale, - ) - - # Add outputs - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/flux/inputs.py b/diffusers/modular_pipelines/flux/inputs.py deleted file mode 100644 index c513d237bee2acf3c558c84d43dab628ccf958d4..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux/inputs.py +++ /dev/null @@ -1,363 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch - -from ...pipelines import FluxPipeline -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import InputParam, OutputParam - -# TODO: consider making these common utilities for modular if they are not pipeline-specific. -from ..qwenimage.inputs import calculate_dimension_from_latents, repeat_tensor_to_batch_size -from .modular_pipeline import FluxModularPipeline - - -logger = logging.get_logger(__name__) - - -class FluxTextInputStep(ModularPipelineBlocks): - model_name = "flux" - - @property - def description(self) -> str: - return ( - "Text input processing step that standardizes text embeddings for the pipeline.\n" - "This step:\n" - " 1. Determines `batch_size` and `dtype` based on `prompt_embeds`\n" - " 2. Ensures all text embeddings have consistent batch sizes (batch_size * num_images_per_prompt)" - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("num_images_per_prompt", default=1), - InputParam( - "prompt_embeds", - required=True, - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="Pre-generated text embeddings. Can be generated from text_encoder step.", - ), - InputParam( - "pooled_prompt_embeds", - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="Pre-generated pooled text embeddings. Can be generated from text_encoder step.", - ), - # TODO: support negative embeddings? - ] - - @property - def intermediate_outputs(self) -> list[str]: - return [ - OutputParam( - "batch_size", - type_hint=int, - description="Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt", - ), - OutputParam( - "dtype", - type_hint=torch.dtype, - description="Data type of model tensor inputs (determined by `prompt_embeds`)", - ), - OutputParam( - "prompt_embeds", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="text embeddings used to guide the image generation", - ), - OutputParam( - "pooled_prompt_embeds", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="pooled text embeddings used to guide the image generation", - ), - # TODO: support negative embeddings? - ] - - def check_inputs(self, components, block_state): - if block_state.prompt_embeds is not None and block_state.pooled_prompt_embeds is not None: - if block_state.prompt_embeds.shape[0] != block_state.pooled_prompt_embeds.shape[0]: - raise ValueError( - "`prompt_embeds` and `pooled_prompt_embeds` must have the same batch size when passed directly, but" - f" got: `prompt_embeds` {block_state.prompt_embeds.shape} != `pooled_prompt_embeds`" - f" {block_state.pooled_prompt_embeds.shape}." - ) - - @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - # TODO: consider adding negative embeddings? - block_state = self.get_block_state(state) - self.check_inputs(components, block_state) - - block_state.batch_size = block_state.prompt_embeds.shape[0] - block_state.dtype = block_state.prompt_embeds.dtype - - _, seq_len, _ = block_state.prompt_embeds.shape - block_state.prompt_embeds = block_state.prompt_embeds.repeat(1, block_state.num_images_per_prompt, 1) - block_state.prompt_embeds = block_state.prompt_embeds.view( - block_state.batch_size * block_state.num_images_per_prompt, seq_len, -1 - ) - pooled_prompt_embeds = block_state.pooled_prompt_embeds.repeat(1, block_state.num_images_per_prompt) - block_state.pooled_prompt_embeds = pooled_prompt_embeds.view( - block_state.batch_size * block_state.num_images_per_prompt, -1 - ) - self.set_block_state(state, block_state) - - return components, state - - -# Adapted from `QwenImageAdditionalInputsStep` -class FluxAdditionalInputsStep(ModularPipelineBlocks): - model_name = "flux" - - def __init__( - self, - image_latent_inputs: list[str] = ["image_latents"], - additional_batch_inputs: list[str] = [], - ): - if not isinstance(image_latent_inputs, list): - image_latent_inputs = [image_latent_inputs] - if not isinstance(additional_batch_inputs, list): - additional_batch_inputs = [additional_batch_inputs] - - self._image_latent_inputs = image_latent_inputs - self._additional_batch_inputs = additional_batch_inputs - super().__init__() - - @property - def description(self) -> str: - # Functionality section - summary_section = ( - "Input processing step that:\n" - " 1. For image latent inputs: Updates height/width if None, patchifies latents, and expands batch size\n" - " 2. For additional batch inputs: Expands batch dimensions to match final batch size" - ) - - # Inputs info - inputs_info = "" - if self._image_latent_inputs or self._additional_batch_inputs: - inputs_info = "\n\nConfigured inputs:" - if self._image_latent_inputs: - inputs_info += f"\n - Image latent inputs: {self._image_latent_inputs}" - if self._additional_batch_inputs: - inputs_info += f"\n - Additional batch inputs: {self._additional_batch_inputs}" - - # Placement guidance - placement_section = "\n\nThis block should be placed after the encoder steps and the text input step." - - return summary_section + inputs_info + placement_section - - @property - def inputs(self) -> list[InputParam]: - inputs = [ - InputParam(name="num_images_per_prompt", default=1), - InputParam(name="batch_size", required=True), - InputParam(name="height"), - InputParam(name="width"), - ] - - # Add image latent inputs - for image_latent_input_name in self._image_latent_inputs: - inputs.append(InputParam(name=image_latent_input_name)) - - # Add additional batch inputs - for input_name in self._additional_batch_inputs: - inputs.append(InputParam(name=input_name)) - - return inputs - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam(name="image_height", type_hint=int, description="The height of the image latents"), - OutputParam(name="image_width", type_hint=int, description="The width of the image latents"), - ] - - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - # Process image latent inputs (height/width calculation, patchify, and batch expansion) - for image_latent_input_name in self._image_latent_inputs: - image_latent_tensor = getattr(block_state, image_latent_input_name) - if image_latent_tensor is None: - continue - - # 1. Calculate height/width from latents - height, width = calculate_dimension_from_latents(image_latent_tensor, components.vae_scale_factor) - block_state.height = block_state.height or height - block_state.width = block_state.width or width - - if not hasattr(block_state, "image_height"): - block_state.image_height = height - if not hasattr(block_state, "image_width"): - block_state.image_width = width - - # 2. Patchify the image latent tensor - # TODO: Implement patchifier for Flux. - latent_height, latent_width = image_latent_tensor.shape[2:] - image_latent_tensor = FluxPipeline._pack_latents( - image_latent_tensor, block_state.batch_size, image_latent_tensor.shape[1], latent_height, latent_width - ) - - # 3. Expand batch size - image_latent_tensor = repeat_tensor_to_batch_size( - input_name=image_latent_input_name, - input_tensor=image_latent_tensor, - num_images_per_prompt=block_state.num_images_per_prompt, - batch_size=block_state.batch_size, - ) - - setattr(block_state, image_latent_input_name, image_latent_tensor) - - # Process additional batch inputs (only batch expansion) - for input_name in self._additional_batch_inputs: - input_tensor = getattr(block_state, input_name) - if input_tensor is None: - continue - - # Only expand batch size - input_tensor = repeat_tensor_to_batch_size( - input_name=input_name, - input_tensor=input_tensor, - num_images_per_prompt=block_state.num_images_per_prompt, - batch_size=block_state.batch_size, - ) - - setattr(block_state, input_name, input_tensor) - - self.set_block_state(state, block_state) - return components, state - - -class FluxKontextAdditionalInputsStep(FluxAdditionalInputsStep): - model_name = "flux-kontext" - - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - # Process image latent inputs (height/width calculation, patchify, and batch expansion) - for image_latent_input_name in self._image_latent_inputs: - image_latent_tensor = getattr(block_state, image_latent_input_name) - if image_latent_tensor is None: - continue - - # 1. Calculate height/width from latents - # Unlike the `FluxAdditionalInputsStep`, we don't overwrite the `block.height` and `block.width` - height, width = calculate_dimension_from_latents(image_latent_tensor, components.vae_scale_factor) - if not hasattr(block_state, "image_height"): - block_state.image_height = height - if not hasattr(block_state, "image_width"): - block_state.image_width = width - - # 2. Patchify the image latent tensor - # TODO: Implement patchifier for Flux. - latent_height, latent_width = image_latent_tensor.shape[2:] - image_latent_tensor = FluxPipeline._pack_latents( - image_latent_tensor, block_state.batch_size, image_latent_tensor.shape[1], latent_height, latent_width - ) - - # 3. Expand batch size - image_latent_tensor = repeat_tensor_to_batch_size( - input_name=image_latent_input_name, - input_tensor=image_latent_tensor, - num_images_per_prompt=block_state.num_images_per_prompt, - batch_size=block_state.batch_size, - ) - - setattr(block_state, image_latent_input_name, image_latent_tensor) - - # Process additional batch inputs (only batch expansion) - for input_name in self._additional_batch_inputs: - input_tensor = getattr(block_state, input_name) - if input_tensor is None: - continue - - # Only expand batch size - input_tensor = repeat_tensor_to_batch_size( - input_name=input_name, - input_tensor=input_tensor, - num_images_per_prompt=block_state.num_images_per_prompt, - batch_size=block_state.batch_size, - ) - - setattr(block_state, input_name, input_tensor) - - self.set_block_state(state, block_state) - return components, state - - -class FluxKontextSetResolutionStep(ModularPipelineBlocks): - model_name = "flux-kontext" - - @property - def description(self): - return ( - "Determines the height and width to be used during the subsequent computations.\n" - "It should always be placed _before_ the latent preparation step." - ) - - @property - def inputs(self) -> list[InputParam]: - inputs = [ - InputParam(name="height"), - InputParam(name="width"), - InputParam(name="max_area", type_hint=int, default=1024**2), - ] - return inputs - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam(name="height", type_hint=int, description="The height of the initial noisy latents"), - OutputParam(name="width", type_hint=int, description="The width of the initial noisy latents"), - ] - - @staticmethod - def check_inputs(height, width, vae_scale_factor): - if height is not None and height % (vae_scale_factor * 2) != 0: - raise ValueError(f"Height must be divisible by {vae_scale_factor * 2} but is {height}") - - if width is not None and width % (vae_scale_factor * 2) != 0: - raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") - - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - height = block_state.height or components.default_height - width = block_state.width or components.default_width - self.check_inputs(height, width, components.vae_scale_factor) - - original_height, original_width = height, width - max_area = block_state.max_area - aspect_ratio = width / height - width = round((max_area * aspect_ratio) ** 0.5) - height = round((max_area / aspect_ratio) ** 0.5) - - multiple_of = components.vae_scale_factor * 2 - width = width // multiple_of * multiple_of - height = height // multiple_of * multiple_of - - if height != original_height or width != original_width: - logger.warning( - f"Generation `height` and `width` have been adjusted to {height} and {width} to fit the model requirements." - ) - - block_state.height = height - block_state.width = width - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/flux/modular_blocks_flux.py b/diffusers/modular_pipelines/flux/modular_blocks_flux.py deleted file mode 100644 index 1f028f555a1bacd09bf0e16f9218c7028f5fea52..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux/modular_blocks_flux.py +++ /dev/null @@ -1,586 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from ...utils import logging -from ..modular_pipeline import AutoPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import InsertableDict, OutputParam -from .before_denoise import ( - FluxImg2ImgPrepareLatentsStep, - FluxImg2ImgSetTimestepsStep, - FluxPrepareLatentsStep, - FluxRoPEInputsStep, - FluxSetTimestepsStep, -) -from .decoders import FluxDecodeStep -from .denoise import FluxDenoiseStep -from .encoders import ( - FluxProcessImagesInputStep, - FluxTextEncoderStep, - FluxVaeEncoderStep, -) -from .inputs import ( - FluxAdditionalInputsStep, - FluxTextInputStep, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# vae encoder (run before before_denoise) - - -# auto_docstring -class FluxImg2ImgVaeEncoderStep(SequentialPipelineBlocks): - """ - Vae encoder step that preprocess andencode the image inputs into their latent representations. - - Components: - image_processor (`VaeImageProcessor`) vae (`AutoencoderKL`) - - Inputs: - resized_image (`None`, *optional*): - TODO: Add description. - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - - Outputs: - processed_image (`None`): - TODO: Add description. - image_latents (`Tensor`): - The latents representing the reference image - """ - - model_name = "flux" - - block_classes = [FluxProcessImagesInputStep(), FluxVaeEncoderStep()] - block_names = ["preprocess", "encode"] - - @property - def description(self) -> str: - return "Vae encoder step that preprocess andencode the image inputs into their latent representations." - - -# auto_docstring -class FluxAutoVaeEncoderStep(AutoPipelineBlocks): - """ - Vae encoder step that encode the image inputs into their latent representations. - This is an auto pipeline block that works for img2img tasks. - - `FluxImg2ImgVaeEncoderStep` (img2img) is used when only `image` is provided. - if `image` is not provided, - step will be skipped. - - Components: - image_processor (`VaeImageProcessor`) vae (`AutoencoderKL`) - - Inputs: - resized_image (`None`, *optional*): - TODO: Add description. - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - - Outputs: - processed_image (`None`): - TODO: Add description. - image_latents (`Tensor`): - The latents representing the reference image - """ - - model_name = "flux" - block_classes = [FluxImg2ImgVaeEncoderStep] - block_names = ["img2img"] - block_trigger_inputs = ["image"] - - @property - def description(self): - return ( - "Vae encoder step that encode the image inputs into their latent representations.\n" - + "This is an auto pipeline block that works for img2img tasks.\n" - + " - `FluxImg2ImgVaeEncoderStep` (img2img) is used when only `image` is provided." - + " - if `image` is not provided, step will be skipped." - ) - - -# before_denoise: text2img -# auto_docstring -class FluxBeforeDenoiseStep(SequentialPipelineBlocks): - """ - Before denoise step that prepares the inputs for the denoise step in text-to-image generation. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - num_images_per_prompt (`int`, *optional*, defaults to 1): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be `batch_size * num_images_per_prompt`. - Can be generated in input step. - dtype (`dtype`, *optional*): - The dtype of the model inputs - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 3.5): - TODO: Add description. - prompt_embeds (`None`, *optional*): - TODO: Add description. - - Outputs: - latents (`Tensor`): - The initial latents to use for the denoising process - timesteps (`Tensor`): - The timesteps to use for inference - num_inference_steps (`int`): - The number of denoising steps to perform at inference time - guidance (`Tensor`): - Optional guidance to be used. - txt_ids (`list`): - The sequence lengths of the prompt embeds, used for RoPE calculation. - img_ids (`list`): - The sequence lengths of the image latents, used for RoPE calculation. - """ - - model_name = "flux" - block_classes = [FluxPrepareLatentsStep(), FluxSetTimestepsStep(), FluxRoPEInputsStep()] - block_names = ["prepare_latents", "set_timesteps", "prepare_rope_inputs"] - - @property - def description(self): - return "Before denoise step that prepares the inputs for the denoise step in text-to-image generation." - - -# before_denoise: img2img -# auto_docstring -class FluxImg2ImgBeforeDenoiseStep(SequentialPipelineBlocks): - """ - Before denoise step that prepare the inputs for the denoise step for img2img task. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - num_images_per_prompt (`int`, *optional*, defaults to 1): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be `batch_size * num_images_per_prompt`. - Can be generated in input step. - dtype (`dtype`, *optional*): - The dtype of the model inputs - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - strength (`None`, *optional*, defaults to 0.6): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 3.5): - TODO: Add description. - image_latents (`Tensor`): - The image latents to use for the denoising process. Can be generated in vae encoder and packed in input - step. - prompt_embeds (`None`, *optional*): - TODO: Add description. - - Outputs: - latents (`Tensor`): - The initial latents to use for the denoising process - timesteps (`Tensor`): - The timesteps to use for inference - num_inference_steps (`int`): - The number of denoising steps to perform at inference time - guidance (`Tensor`): - Optional guidance to be used. - initial_noise (`Tensor`): - The initial random noised used for inpainting denoising. - txt_ids (`list`): - The sequence lengths of the prompt embeds, used for RoPE calculation. - img_ids (`list`): - The sequence lengths of the image latents, used for RoPE calculation. - """ - - model_name = "flux" - block_classes = [ - FluxPrepareLatentsStep(), - FluxImg2ImgSetTimestepsStep(), - FluxImg2ImgPrepareLatentsStep(), - FluxRoPEInputsStep(), - ] - block_names = ["prepare_latents", "set_timesteps", "prepare_img2img_latents", "prepare_rope_inputs"] - - @property - def description(self): - return "Before denoise step that prepare the inputs for the denoise step for img2img task." - - -# before_denoise: all task (text2img, img2img) -# auto_docstring -class FluxAutoBeforeDenoiseStep(AutoPipelineBlocks): - """ - Before denoise step that prepare the inputs for the denoise step. - This is an auto pipeline block that works for text2image. - - `FluxBeforeDenoiseStep` (text2image) is used. - - `FluxImg2ImgBeforeDenoiseStep` (img2img) is used when only `image_latents` is provided. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - height (`int`): - TODO: Add description. - width (`int`): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - num_images_per_prompt (`int`, *optional*, defaults to 1): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be `batch_size * num_images_per_prompt`. - Can be generated in input step. - dtype (`dtype`, *optional*): - The dtype of the model inputs - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - strength (`None`, *optional*, defaults to 0.6): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 3.5): - TODO: Add description. - image_latents (`Tensor`, *optional*): - The image latents to use for the denoising process. Can be generated in vae encoder and packed in input - step. - prompt_embeds (`None`, *optional*): - TODO: Add description. - - Outputs: - latents (`Tensor`): - The initial latents to use for the denoising process - timesteps (`Tensor`): - The timesteps to use for inference - num_inference_steps (`int`): - The number of denoising steps to perform at inference time - guidance (`Tensor`): - Optional guidance to be used. - initial_noise (`Tensor`): - The initial random noised used for inpainting denoising. - txt_ids (`list`): - The sequence lengths of the prompt embeds, used for RoPE calculation. - img_ids (`list`): - The sequence lengths of the image latents, used for RoPE calculation. - """ - - model_name = "flux" - block_classes = [FluxImg2ImgBeforeDenoiseStep, FluxBeforeDenoiseStep] - block_names = ["img2img", "text2image"] - block_trigger_inputs = ["image_latents", None] - - @property - def description(self): - return ( - "Before denoise step that prepare the inputs for the denoise step.\n" - + "This is an auto pipeline block that works for text2image.\n" - + " - `FluxBeforeDenoiseStep` (text2image) is used.\n" - + " - `FluxImg2ImgBeforeDenoiseStep` (img2img) is used when only `image_latents` is provided.\n" - ) - - -# inputs: text2image/img2img - - -# auto_docstring -class FluxImg2ImgInputStep(SequentialPipelineBlocks): - """ - Input step that prepares the inputs for the img2img denoising step. It: - - Inputs: - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - pooled_prompt_embeds (`Tensor`, *optional*): - Pre-generated pooled text embeddings. Can be generated from text_encoder step. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - image_latents (`None`, *optional*): - TODO: Add description. - - Outputs: - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt - dtype (`dtype`): - Data type of model tensor inputs (determined by `prompt_embeds`) - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation - pooled_prompt_embeds (`Tensor`): - pooled text embeddings used to guide the image generation - image_height (`int`): - The height of the image latents - image_width (`int`): - The width of the image latents - """ - - model_name = "flux" - block_classes = [FluxTextInputStep(), FluxAdditionalInputsStep()] - block_names = ["text_inputs", "additional_inputs"] - - @property - def description(self): - return "Input step that prepares the inputs for the img2img denoising step. It:\n" - " - make sure the text embeddings have consistent batch size as well as the additional inputs (`image_latents`).\n" - " - update height/width based `image_latents`, patchify `image_latents`." - - -# auto_docstring -class FluxAutoInputStep(AutoPipelineBlocks): - """ - Input step that standardize the inputs for the denoising step, e.g. make sure inputs have consistent batch size, - and patchified. - This is an auto pipeline block that works for text2image/img2img tasks. - - `FluxImg2ImgInputStep` (img2img) is used when `image_latents` is provided. - - `FluxTextInputStep` (text2image) is used when `image_latents` are not provided. - - Inputs: - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - pooled_prompt_embeds (`Tensor`, *optional*): - Pre-generated pooled text embeddings. Can be generated from text_encoder step. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - image_latents (`None`, *optional*): - TODO: Add description. - - Outputs: - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt - dtype (`dtype`): - Data type of model tensor inputs (determined by `prompt_embeds`) - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation - pooled_prompt_embeds (`Tensor`): - pooled text embeddings used to guide the image generation - image_height (`int`): - The height of the image latents - image_width (`int`): - The width of the image latents - """ - - model_name = "flux" - - block_classes = [FluxImg2ImgInputStep, FluxTextInputStep] - block_names = ["img2img", "text2image"] - block_trigger_inputs = ["image_latents", None] - - @property - def description(self): - return ( - "Input step that standardize the inputs for the denoising step, e.g. make sure inputs have consistent batch size, and patchified. \n" - " This is an auto pipeline block that works for text2image/img2img tasks.\n" - + " - `FluxImg2ImgInputStep` (img2img) is used when `image_latents` is provided.\n" - + " - `FluxTextInputStep` (text2image) is used when `image_latents` are not provided.\n" - ) - - -# auto_docstring -class FluxCoreDenoiseStep(SequentialPipelineBlocks): - """ - Core step that performs the denoising process for Flux. - This step supports text-to-image and image-to-image tasks for Flux: - - for image-to-image generation, you need to provide `image_latents` - - for text-to-image generation, all you need to provide is prompt embeddings. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`FluxTransformer2DModel`) - - Inputs: - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - pooled_prompt_embeds (`Tensor`, *optional*): - Pre-generated pooled text embeddings. Can be generated from text_encoder step. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - image_latents (`None`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - strength (`None`, *optional*, defaults to 0.6): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 3.5): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "flux" - block_classes = [FluxAutoInputStep, FluxAutoBeforeDenoiseStep, FluxDenoiseStep] - block_names = ["input", "before_denoise", "denoise"] - - @property - def description(self): - return ( - "Core step that performs the denoising process for Flux.\n" - + "This step supports text-to-image and image-to-image tasks for Flux:\n" - + " - for image-to-image generation, you need to provide `image_latents`\n" - + " - for text-to-image generation, all you need to provide is prompt embeddings." - ) - - @property - def outputs(self): - return [ - OutputParam.template("latents"), - ] - - -# Auto blocks (text2image and img2img) -AUTO_BLOCKS = InsertableDict( - [ - ("text_encoder", FluxTextEncoderStep()), - ("vae_encoder", FluxAutoVaeEncoderStep()), - ("denoise", FluxCoreDenoiseStep()), - ("decode", FluxDecodeStep()), - ] -) - - -# auto_docstring -class FluxAutoBlocks(SequentialPipelineBlocks): - """ - Auto Modular pipeline for text-to-image and image-to-image using Flux. - - Supported workflows: - - `text2image`: requires `prompt` - - `image2image`: requires `image`, `prompt` - - Components: - text_encoder (`CLIPTextModel`) tokenizer (`CLIPTokenizer`) text_encoder_2 (`T5EncoderModel`) tokenizer_2 - (`T5Tokenizer`) image_processor (`VaeImageProcessor`) vae (`AutoencoderKL`) scheduler - (`FlowMatchEulerDiscreteScheduler`) transformer (`FluxTransformer2DModel`) - - Inputs: - prompt (`None`, *optional*): - TODO: Add description. - prompt_2 (`None`, *optional*): - TODO: Add description. - max_sequence_length (`int`, *optional*, defaults to 512): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - resized_image (`None`, *optional*): - TODO: Add description. - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - image_latents (`None`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - strength (`None`, *optional*, defaults to 0.6): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 3.5): - TODO: Add description. - output_type (`None`, *optional*, defaults to pil): - TODO: Add description. - - Outputs: - images (`list`): - Generated images. - """ - - model_name = "flux" - - block_classes = AUTO_BLOCKS.values() - block_names = AUTO_BLOCKS.keys() - - _workflow_map = { - "text2image": {"prompt": True}, - "image2image": {"image": True, "prompt": True}, - } - - @property - def description(self): - return "Auto Modular pipeline for text-to-image and image-to-image using Flux." - - @property - def outputs(self): - return [OutputParam.template("images")] diff --git a/diffusers/modular_pipelines/flux/modular_blocks_flux_kontext.py b/diffusers/modular_pipelines/flux/modular_blocks_flux_kontext.py deleted file mode 100644 index c4f8bffffd1e673735f519c6c8ad0c37e4bca421..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux/modular_blocks_flux_kontext.py +++ /dev/null @@ -1,585 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from ...utils import logging -from ..modular_pipeline import AutoPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import InsertableDict, OutputParam -from .before_denoise import ( - FluxKontextRoPEInputsStep, - FluxPrepareLatentsStep, - FluxRoPEInputsStep, - FluxSetTimestepsStep, -) -from .decoders import FluxDecodeStep -from .denoise import FluxKontextDenoiseStep -from .encoders import ( - FluxKontextProcessImagesInputStep, - FluxTextEncoderStep, - FluxVaeEncoderStep, -) -from .inputs import ( - FluxKontextAdditionalInputsStep, - FluxKontextSetResolutionStep, - FluxTextInputStep, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# Flux Kontext vae encoder (run before before_denoise) -# auto_docstring -class FluxKontextVaeEncoderStep(SequentialPipelineBlocks): - """ - Vae encoder step that preprocess andencode the image inputs into their latent representations. - - Components: - image_processor (`VaeImageProcessor`) vae (`AutoencoderKL`) - - Inputs: - image (`None`, *optional*): - TODO: Add description. - _auto_resize (`bool`, *optional*, defaults to True): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - - Outputs: - processed_image (`None`): - TODO: Add description. - image_latents (`Tensor`): - The latents representing the reference image - """ - - model_name = "flux-kontext" - - block_classes = [FluxKontextProcessImagesInputStep(), FluxVaeEncoderStep(sample_mode="argmax")] - block_names = ["preprocess", "encode"] - - @property - def description(self) -> str: - return "Vae encoder step that preprocess andencode the image inputs into their latent representations." - - -# auto_docstring -class FluxKontextAutoVaeEncoderStep(AutoPipelineBlocks): - """ - Vae encoder step that encode the image inputs into their latent representations. - This is an auto pipeline block that works for image-conditioned tasks. - - `FluxKontextVaeEncoderStep` (image_conditioned) is used when only `image` is provided. - if `image` is not - provided, step will be skipped. - - Components: - image_processor (`VaeImageProcessor`) vae (`AutoencoderKL`) - - Inputs: - image (`None`, *optional*): - TODO: Add description. - _auto_resize (`bool`, *optional*, defaults to True): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - - Outputs: - processed_image (`None`): - TODO: Add description. - image_latents (`Tensor`): - The latents representing the reference image - """ - - model_name = "flux-kontext" - - block_classes = [FluxKontextVaeEncoderStep] - block_names = ["image_conditioned"] - block_trigger_inputs = ["image"] - - @property - def description(self): - return ( - "Vae encoder step that encode the image inputs into their latent representations.\n" - + "This is an auto pipeline block that works for image-conditioned tasks.\n" - + " - `FluxKontextVaeEncoderStep` (image_conditioned) is used when only `image` is provided." - + " - if `image` is not provided, step will be skipped." - ) - - -# before_denoise: text2img -# auto_docstring -class FluxKontextBeforeDenoiseStep(SequentialPipelineBlocks): - """ - Before denoise step that prepares the inputs for the denoise step for Flux Kontext - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - num_images_per_prompt (`int`, *optional*, defaults to 1): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be `batch_size * num_images_per_prompt`. - Can be generated in input step. - dtype (`dtype`, *optional*): - The dtype of the model inputs - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 3.5): - TODO: Add description. - prompt_embeds (`None`, *optional*): - TODO: Add description. - - Outputs: - latents (`Tensor`): - The initial latents to use for the denoising process - timesteps (`Tensor`): - The timesteps to use for inference - num_inference_steps (`int`): - The number of denoising steps to perform at inference time - guidance (`Tensor`): - Optional guidance to be used. - txt_ids (`list`): - The sequence lengths of the prompt embeds, used for RoPE calculation. - img_ids (`list`): - The sequence lengths of the image latents, used for RoPE calculation. - """ - - model_name = "flux-kontext" - - block_classes = [FluxPrepareLatentsStep(), FluxSetTimestepsStep(), FluxRoPEInputsStep()] - block_names = ["prepare_latents", "set_timesteps", "prepare_rope_inputs"] - - @property - def description(self): - return "Before denoise step that prepares the inputs for the denoise step for Flux Kontext\n" - "for text-to-image tasks." - - -# before_denoise: image-conditioned -# auto_docstring -class FluxKontextImageConditionedBeforeDenoiseStep(SequentialPipelineBlocks): - """ - Before denoise step that prepare the inputs for the denoise step for Flux Kontext - for image-conditioned tasks. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - num_images_per_prompt (`int`, *optional*, defaults to 1): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be `batch_size * num_images_per_prompt`. - Can be generated in input step. - dtype (`dtype`, *optional*): - The dtype of the model inputs - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 3.5): - TODO: Add description. - image_height (`None`, *optional*): - TODO: Add description. - image_width (`None`, *optional*): - TODO: Add description. - prompt_embeds (`None`, *optional*): - TODO: Add description. - - Outputs: - latents (`Tensor`): - The initial latents to use for the denoising process - timesteps (`Tensor`): - The timesteps to use for inference - num_inference_steps (`int`): - The number of denoising steps to perform at inference time - guidance (`Tensor`): - Optional guidance to be used. - txt_ids (`list`): - The sequence lengths of the prompt embeds, used for RoPE calculation. - img_ids (`list`): - The sequence lengths of the image latents, used for RoPE calculation. - """ - - model_name = "flux-kontext" - - block_classes = [FluxPrepareLatentsStep(), FluxSetTimestepsStep(), FluxKontextRoPEInputsStep()] - block_names = ["prepare_latents", "set_timesteps", "prepare_rope_inputs"] - - @property - def description(self): - return ( - "Before denoise step that prepare the inputs for the denoise step for Flux Kontext\n" - "for image-conditioned tasks." - ) - - -# auto_docstring -class FluxKontextAutoBeforeDenoiseStep(AutoPipelineBlocks): - """ - Before denoise step that prepare the inputs for the denoise step. - This is an auto pipeline block that works for text2image. - - `FluxKontextBeforeDenoiseStep` (text2image) is used. - - `FluxKontextImageConditionedBeforeDenoiseStep` (image_conditioned) is used when only `image_latents` is - provided. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - num_images_per_prompt (`int`, *optional*, defaults to 1): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be `batch_size * num_images_per_prompt`. - Can be generated in input step. - dtype (`dtype`, *optional*): - The dtype of the model inputs - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 3.5): - TODO: Add description. - image_height (`None`, *optional*): - TODO: Add description. - image_width (`None`, *optional*): - TODO: Add description. - prompt_embeds (`None`, *optional*): - TODO: Add description. - - Outputs: - latents (`Tensor`): - The initial latents to use for the denoising process - timesteps (`Tensor`): - The timesteps to use for inference - num_inference_steps (`int`): - The number of denoising steps to perform at inference time - guidance (`Tensor`): - Optional guidance to be used. - txt_ids (`list`): - The sequence lengths of the prompt embeds, used for RoPE calculation. - img_ids (`list`): - The sequence lengths of the image latents, used for RoPE calculation. - """ - - model_name = "flux-kontext" - - block_classes = [FluxKontextImageConditionedBeforeDenoiseStep, FluxKontextBeforeDenoiseStep] - block_names = ["image_conditioned", "text2image"] - block_trigger_inputs = ["image_latents", None] - - @property - def description(self): - return ( - "Before denoise step that prepare the inputs for the denoise step.\n" - + "This is an auto pipeline block that works for text2image.\n" - + " - `FluxKontextBeforeDenoiseStep` (text2image) is used.\n" - + " - `FluxKontextImageConditionedBeforeDenoiseStep` (image_conditioned) is used when only `image_latents` is provided.\n" - ) - - -# inputs: Flux Kontext -# auto_docstring -class FluxKontextInputStep(SequentialPipelineBlocks): - """ - Input step that prepares the inputs for the both text2img and img2img denoising step. It: - - make sure the text embeddings have consistent batch size as well as the additional inputs (`image_latents`). - - update height/width based `image_latents`, patchify `image_latents`. - - Inputs: - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - max_area (`int`, *optional*, defaults to 1048576): - TODO: Add description. - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - pooled_prompt_embeds (`Tensor`, *optional*): - Pre-generated pooled text embeddings. Can be generated from text_encoder step. - image_latents (`None`, *optional*): - TODO: Add description. - - Outputs: - height (`int`): - The height of the initial noisy latents - width (`int`): - The width of the initial noisy latents - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt - dtype (`dtype`): - Data type of model tensor inputs (determined by `prompt_embeds`) - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation - pooled_prompt_embeds (`Tensor`): - pooled text embeddings used to guide the image generation - image_height (`int`): - The height of the image latents - image_width (`int`): - The width of the image latents - """ - - model_name = "flux-kontext" - block_classes = [FluxKontextSetResolutionStep(), FluxTextInputStep(), FluxKontextAdditionalInputsStep()] - block_names = ["set_resolution", "text_inputs", "additional_inputs"] - - @property - def description(self): - return ( - "Input step that prepares the inputs for the both text2img and img2img denoising step. It:\n" - " - make sure the text embeddings have consistent batch size as well as the additional inputs (`image_latents`).\n" - " - update height/width based `image_latents`, patchify `image_latents`." - ) - - -# auto_docstring -class FluxKontextAutoInputStep(AutoPipelineBlocks): - """ - Input step that standardize the inputs for the denoising step, e.g. make sure inputs have consistent batch size, - and patchified. - This is an auto pipeline block that works for text2image/img2img tasks. - - `FluxKontextInputStep` (image_conditioned) is used when `image_latents` is provided. - - `FluxKontextInputStep` is also capable of handling text2image task when `image_latent` isn't present. - - Inputs: - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - max_area (`int`, *optional*, defaults to 1048576): - TODO: Add description. - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - pooled_prompt_embeds (`Tensor`, *optional*): - Pre-generated pooled text embeddings. Can be generated from text_encoder step. - image_latents (`None`, *optional*): - TODO: Add description. - - Outputs: - height (`int`): - The height of the initial noisy latents - width (`int`): - The width of the initial noisy latents - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt - dtype (`dtype`): - Data type of model tensor inputs (determined by `prompt_embeds`) - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation - pooled_prompt_embeds (`Tensor`): - pooled text embeddings used to guide the image generation - image_height (`int`): - The height of the image latents - image_width (`int`): - The width of the image latents - """ - - model_name = "flux-kontext" - block_classes = [FluxKontextInputStep, FluxTextInputStep] - block_names = ["image_conditioned", "text2image"] - block_trigger_inputs = ["image_latents", None] - - @property - def description(self): - return ( - "Input step that standardize the inputs for the denoising step, e.g. make sure inputs have consistent batch size, and patchified. \n" - " This is an auto pipeline block that works for text2image/img2img tasks.\n" - + " - `FluxKontextInputStep` (image_conditioned) is used when `image_latents` is provided.\n" - + " - `FluxKontextInputStep` is also capable of handling text2image task when `image_latent` isn't present." - ) - - -# auto_docstring -class FluxKontextCoreDenoiseStep(SequentialPipelineBlocks): - """ - Core step that performs the denoising process for Flux Kontext. - This step supports text-to-image and image-conditioned tasks for Flux Kontext: - - for image-conditioned generation, you need to provide `image_latents` - - for text-to-image generation, all you need to provide is prompt embeddings. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`FluxTransformer2DModel`) - - Inputs: - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - max_area (`int`, *optional*, defaults to 1048576): - TODO: Add description. - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - pooled_prompt_embeds (`Tensor`, *optional*): - Pre-generated pooled text embeddings. Can be generated from text_encoder step. - image_latents (`None`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 3.5): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "flux-kontext" - block_classes = [FluxKontextAutoInputStep, FluxKontextAutoBeforeDenoiseStep, FluxKontextDenoiseStep] - block_names = ["input", "before_denoise", "denoise"] - - @property - def description(self): - return ( - "Core step that performs the denoising process for Flux Kontext.\n" - + "This step supports text-to-image and image-conditioned tasks for Flux Kontext:\n" - + " - for image-conditioned generation, you need to provide `image_latents`\n" - + " - for text-to-image generation, all you need to provide is prompt embeddings." - ) - - @property - def outputs(self): - return [ - OutputParam.template("latents"), - ] - - -AUTO_BLOCKS_KONTEXT = InsertableDict( - [ - ("text_encoder", FluxTextEncoderStep()), - ("vae_encoder", FluxKontextAutoVaeEncoderStep()), - ("denoise", FluxKontextCoreDenoiseStep()), - ("decode", FluxDecodeStep()), - ] -) - - -# auto_docstring -class FluxKontextAutoBlocks(SequentialPipelineBlocks): - """ - Modular pipeline for image-to-image using Flux Kontext. - - Supported workflows: - - `image_conditioned`: requires `image`, `prompt` - - `text2image`: requires `prompt` - - Components: - text_encoder (`CLIPTextModel`) tokenizer (`CLIPTokenizer`) text_encoder_2 (`T5EncoderModel`) tokenizer_2 - (`T5Tokenizer`) image_processor (`VaeImageProcessor`) vae (`AutoencoderKL`) scheduler - (`FlowMatchEulerDiscreteScheduler`) transformer (`FluxTransformer2DModel`) - - Inputs: - prompt (`None`, *optional*): - TODO: Add description. - prompt_2 (`None`, *optional*): - TODO: Add description. - max_sequence_length (`int`, *optional*, defaults to 512): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - image (`None`, *optional*): - TODO: Add description. - _auto_resize (`bool`, *optional*, defaults to True): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - max_area (`int`, *optional*, defaults to 1048576): - TODO: Add description. - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - image_latents (`None`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 3.5): - TODO: Add description. - output_type (`None`, *optional*, defaults to pil): - TODO: Add description. - - Outputs: - images (`list`): - Generated images. - """ - - model_name = "flux-kontext" - - block_classes = AUTO_BLOCKS_KONTEXT.values() - block_names = AUTO_BLOCKS_KONTEXT.keys() - _workflow_map = { - "image_conditioned": {"image": True, "prompt": True}, - "text2image": {"prompt": True}, - } - - @property - def description(self): - return "Modular pipeline for image-to-image using Flux Kontext." - - @property - def outputs(self): - return [OutputParam.template("images")] diff --git a/diffusers/modular_pipelines/flux/modular_pipeline.py b/diffusers/modular_pipelines/flux/modular_pipeline.py deleted file mode 100644 index 1de59ebad242b1fac933b555015991891f2bd56b..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux/modular_pipeline.py +++ /dev/null @@ -1,67 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from ...loaders import FluxLoraLoaderMixin, TextualInversionLoaderMixin -from ...utils import logging -from ..modular_pipeline import ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class FluxModularPipeline(ModularPipeline, FluxLoraLoaderMixin, TextualInversionLoaderMixin): - """ - A ModularPipeline for Flux. - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "FluxAutoBlocks" - - @property - def default_height(self): - return self.default_sample_size * self.vae_scale_factor - - @property - def default_width(self): - return self.default_sample_size * self.vae_scale_factor - - @property - def default_sample_size(self): - return 128 - - @property - def vae_scale_factor(self): - vae_scale_factor = 8 - if getattr(self, "vae", None) is not None: - vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) - return vae_scale_factor - - @property - def num_channels_latents(self): - num_channels_latents = 16 - if getattr(self, "transformer", None): - num_channels_latents = self.transformer.config.in_channels // 4 - return num_channels_latents - - -class FluxKontextModularPipeline(FluxModularPipeline): - """ - A ModularPipeline for Flux Kontext. - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "FluxKontextAutoBlocks" diff --git a/diffusers/modular_pipelines/flux2/__init__.py b/diffusers/modular_pipelines/flux2/__init__.py deleted file mode 100644 index d7cc8badcaf7d0e51f7e6eb2923df6c04d4172cd..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux2/__init__.py +++ /dev/null @@ -1,57 +0,0 @@ -from typing import TYPE_CHECKING - -from ...utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, -) - - -_dummy_objects = {} -_import_structure = {} - -try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ...utils import dummy_torch_and_transformers_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) -else: - _import_structure["encoders"] = ["Flux2RemoteTextEncoderStep"] - _import_structure["modular_blocks_flux2"] = ["Flux2AutoBlocks"] - _import_structure["modular_blocks_flux2_klein"] = ["Flux2KleinAutoBlocks"] - _import_structure["modular_blocks_flux2_klein_base"] = ["Flux2KleinBaseAutoBlocks"] - _import_structure["modular_pipeline"] = [ - "Flux2KleinBaseModularPipeline", - "Flux2KleinModularPipeline", - "Flux2ModularPipeline", - ] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .encoders import Flux2RemoteTextEncoderStep - from .modular_blocks_flux2 import Flux2AutoBlocks - from .modular_blocks_flux2_klein import Flux2KleinAutoBlocks - from .modular_blocks_flux2_klein_base import Flux2KleinBaseAutoBlocks - from .modular_pipeline import Flux2KleinBaseModularPipeline, Flux2KleinModularPipeline, Flux2ModularPipeline -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/flux2/before_denoise.py b/diffusers/modular_pipelines/flux2/before_denoise.py deleted file mode 100644 index 87a6b568a2582ed68c56d8a9560894f2f6a7c549..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux2/before_denoise.py +++ /dev/null @@ -1,591 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect - -import numpy as np -import torch - -from ...models import Flux2Transformer2DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Flux2ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def compute_empirical_mu(image_seq_len: int, num_steps: int) -> float: - """Compute empirical mu for Flux2 timestep scheduling.""" - a1, b1 = 8.73809524e-05, 1.89833333 - a2, b2 = 0.00016927, 0.45666666 - - if image_seq_len > 4300: - mu = a2 * image_seq_len + b2 - return float(mu) - - m_200 = a2 * image_seq_len + b2 - m_10 = a1 * image_seq_len + b1 - - a = (m_200 - m_10) / 190.0 - b = m_200 - 200.0 * a - mu = a * num_steps + b - - return float(mu) - - -# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps -def retrieve_timesteps( - scheduler, - num_inference_steps: int | None = None, - device: str | torch.device | None = None, - timesteps: list[int] | None = None, - sigmas: list[float] | None = None, - **kwargs, -): - r""" - Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles - custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`. - - Args: - scheduler (`SchedulerMixin`): - The scheduler to get timesteps from. - num_inference_steps (`int`): - The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps` - must be `None`. - device (`str` or `torch.device`, *optional*): - The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. - timesteps (`list[int]`, *optional*): - Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed, - `num_inference_steps` and `sigmas` must be `None`. - sigmas (`list[float]`, *optional*): - Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed, - `num_inference_steps` and `timesteps` must be `None`. - - Returns: - `tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the - second element is the number of inference steps. - """ - if timesteps is not None and sigmas is not None: - raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values") - if timesteps is not None: - accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) - if not accepts_timesteps: - raise ValueError( - f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" - f" timestep schedules. Please check whether you are using the correct scheduler." - ) - scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) - timesteps = scheduler.timesteps - num_inference_steps = len(timesteps) - elif sigmas is not None: - accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) - if not accept_sigmas: - raise ValueError( - f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" - f" sigmas schedules. Please check whether you are using the correct scheduler." - ) - scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs) - timesteps = scheduler.timesteps - num_inference_steps = len(timesteps) - else: - scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) - timesteps = scheduler.timesteps - return timesteps, num_inference_steps - - -class Flux2SetTimestepsStep(ModularPipelineBlocks): - model_name = "flux2" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), - ComponentSpec("transformer", Flux2Transformer2DModel), - ] - - @property - def description(self) -> str: - return "Step that sets the scheduler's timesteps for Flux2 inference using empirical mu calculation" - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("num_inference_steps", default=50), - InputParam("timesteps"), - InputParam("sigmas"), - InputParam("latents", type_hint=torch.Tensor), - InputParam("height", type_hint=int), - InputParam("width", type_hint=int), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("timesteps", type_hint=torch.Tensor, description="The timesteps to use for inference"), - OutputParam( - "num_inference_steps", - type_hint=int, - description="The number of denoising steps to perform at inference time", - ), - ] - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - scheduler = components.scheduler - - height = block_state.height or components.default_height - width = block_state.width or components.default_width - vae_scale_factor = components.vae_scale_factor - - latent_height = 2 * (int(height) // (vae_scale_factor * 2)) - latent_width = 2 * (int(width) // (vae_scale_factor * 2)) - image_seq_len = (latent_height // 2) * (latent_width // 2) - - num_inference_steps = block_state.num_inference_steps - sigmas = block_state.sigmas - timesteps = block_state.timesteps - - if timesteps is None and sigmas is None: - sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) - if hasattr(scheduler.config, "use_flow_sigmas") and scheduler.config.use_flow_sigmas: - sigmas = None - - mu = compute_empirical_mu(image_seq_len=image_seq_len, num_steps=num_inference_steps) - - timesteps, num_inference_steps = retrieve_timesteps( - scheduler, - num_inference_steps, - device, - timesteps=timesteps, - sigmas=sigmas, - mu=mu, - ) - block_state.timesteps = timesteps - block_state.num_inference_steps = num_inference_steps - - components.scheduler.set_begin_index(0) - - self.set_block_state(state, block_state) - return components, state - - -class Flux2PrepareLatentsStep(ModularPipelineBlocks): - model_name = "flux2" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [] - - @property - def description(self) -> str: - return "Prepare latents step that prepares the initial noise latents for Flux2 text-to-image generation" - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("height", type_hint=int), - InputParam("width", type_hint=int), - InputParam("latents", type_hint=torch.Tensor | None), - InputParam("num_images_per_prompt", type_hint=int, default=1), - InputParam("generator"), - InputParam( - "batch_size", - required=True, - type_hint=int, - description="Number of prompts, the final batch size of model inputs should be `batch_size * num_images_per_prompt`.", - ), - InputParam("dtype", type_hint=torch.dtype, description="The dtype of the model inputs"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "latents", type_hint=torch.Tensor, description="The initial latents to use for the denoising process" - ), - OutputParam("latent_ids", type_hint=torch.Tensor, description="Position IDs for the latents (for RoPE)"), - ] - - @staticmethod - def check_inputs(components, block_state): - vae_scale_factor = components.vae_scale_factor - if (block_state.height is not None and block_state.height % (vae_scale_factor * 2) != 0) or ( - block_state.width is not None and block_state.width % (vae_scale_factor * 2) != 0 - ): - logger.warning( - f"`height` and `width` have to be divisible by {vae_scale_factor * 2} but are {block_state.height} and {block_state.width}." - ) - - @staticmethod - def _prepare_latent_ids(latents: torch.Tensor): - """ - Generates 4D position coordinates (T, H, W, L) for latent tensors. - - Args: - latents: Latent tensor of shape (B, C, H, W) - - Returns: - Position IDs tensor of shape (B, H*W, 4) - """ - batch_size, _, height, width = latents.shape - - t = torch.arange(1) - h = torch.arange(height) - w = torch.arange(width) - l = torch.arange(1) - - latent_ids = torch.cartesian_prod(t, h, w, l) - latent_ids = latent_ids.unsqueeze(0).expand(batch_size, -1, -1) - - return latent_ids - - @staticmethod - def _pack_latents(latents): - """Pack latents: (batch_size, num_channels, height, width) -> (batch_size, height * width, num_channels)""" - batch_size, num_channels, height, width = latents.shape - latents = latents.reshape(batch_size, num_channels, height * width).permute(0, 2, 1) - return latents - - @staticmethod - def prepare_latents( - comp, - batch_size, - num_channels_latents, - height, - width, - dtype, - device, - generator, - latents=None, - ): - height = 2 * (int(height) // (comp.vae_scale_factor * 2)) - width = 2 * (int(width) // (comp.vae_scale_factor * 2)) - - shape = (batch_size, num_channels_latents * 4, height // 2, width // 2) - if isinstance(generator, list) and len(generator) != batch_size: - raise ValueError( - f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" - f" size of {batch_size}. Make sure the batch size matches the length of the generators." - ) - if latents is None: - latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) - else: - latents = latents.to(device=device, dtype=dtype) - - return latents - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - block_state.height = block_state.height or components.default_height - block_state.width = block_state.width or components.default_width - block_state.device = components._execution_device - block_state.num_channels_latents = components.num_channels_latents - - self.check_inputs(components, block_state) - batch_size = block_state.batch_size * block_state.num_images_per_prompt - - latents = self.prepare_latents( - components, - batch_size, - block_state.num_channels_latents, - block_state.height, - block_state.width, - block_state.dtype, - block_state.device, - block_state.generator, - block_state.latents, - ) - - latent_ids = self._prepare_latent_ids(latents) - latent_ids = latent_ids.to(block_state.device) - - latents = self._pack_latents(latents) - - block_state.latents = latents - block_state.latent_ids = latent_ids - - self.set_block_state(state, block_state) - return components, state - - -class Flux2RoPEInputsStep(ModularPipelineBlocks): - model_name = "flux2" - - @property - def description(self) -> str: - return "Step that prepares the 4D RoPE position IDs for Flux2 denoising. Should be placed after text encoder and latent preparation steps." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="prompt_embeds", required=True), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="txt_ids", - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="4D position IDs (T, H, W, L) for text tokens, used for RoPE calculation.", - ), - ] - - @staticmethod - def _prepare_text_ids(x: torch.Tensor, t_coord: torch.Tensor | None = None): - """Prepare 4D position IDs for text tokens.""" - B, L, _ = x.shape - out_ids = [] - - for i in range(B): - t = torch.arange(1) if t_coord is None else t_coord[i] - h = torch.arange(1) - w = torch.arange(1) - seq_l = torch.arange(L) - - coords = torch.cartesian_prod(t, h, w, seq_l) - out_ids.append(coords) - - return torch.stack(out_ids) - - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - prompt_embeds = block_state.prompt_embeds - device = prompt_embeds.device - - block_state.txt_ids = self._prepare_text_ids(prompt_embeds) - block_state.txt_ids = block_state.txt_ids.to(device) - - self.set_block_state(state, block_state) - return components, state - - -class Flux2KleinBaseRoPEInputsStep(ModularPipelineBlocks): - model_name = "flux2-klein" - - @property - def description(self) -> str: - return "Step that prepares the 4D RoPE position IDs for Flux2-Klein base model denoising. Should be placed after text encoder and latent preparation steps." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="prompt_embeds", required=True), - InputParam(name="negative_prompt_embeds", required=False), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="txt_ids", - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="4D position IDs (T, H, W, L) for text tokens, used for RoPE calculation.", - ), - OutputParam( - name="negative_txt_ids", - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="4D position IDs (T, H, W, L) for negative text tokens, used for RoPE calculation.", - ), - ] - - @staticmethod - def _prepare_text_ids(x: torch.Tensor, t_coord: torch.Tensor | None = None): - """Prepare 4D position IDs for text tokens.""" - B, L, _ = x.shape - out_ids = [] - - for i in range(B): - t = torch.arange(1) if t_coord is None else t_coord[i] - h = torch.arange(1) - w = torch.arange(1) - seq_l = torch.arange(L) - - coords = torch.cartesian_prod(t, h, w, seq_l) - out_ids.append(coords) - - return torch.stack(out_ids) - - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - prompt_embeds = block_state.prompt_embeds - device = prompt_embeds.device - - block_state.txt_ids = self._prepare_text_ids(prompt_embeds) - block_state.txt_ids = block_state.txt_ids.to(device) - - block_state.negative_txt_ids = None - if block_state.negative_prompt_embeds is not None: - block_state.negative_txt_ids = self._prepare_text_ids(block_state.negative_prompt_embeds) - block_state.negative_txt_ids = block_state.negative_txt_ids.to(device) - - self.set_block_state(state, block_state) - return components, state - - -class Flux2PrepareImageLatentsStep(ModularPipelineBlocks): - model_name = "flux2" - - @property - def description(self) -> str: - return "Step that prepares image latents and their position IDs for Flux2 image conditioning." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("image_latents", type_hint=list[torch.Tensor]), - InputParam("batch_size", required=True, type_hint=int), - InputParam("num_images_per_prompt", default=1, type_hint=int), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "image_latents", - type_hint=torch.Tensor, - description="Packed image latents for conditioning", - ), - OutputParam( - "image_latent_ids", - type_hint=torch.Tensor, - description="Position IDs for image latents", - ), - ] - - @staticmethod - def _prepare_image_ids(image_latents: list[torch.Tensor], scale: int = 10): - """ - Generates 4D time-space coordinates (T, H, W, L) for a sequence of image latents. - - Args: - image_latents: A list of image latent feature tensors of shape (1, C, H, W). - scale: Factor used to define the time separation between latents. - - Returns: - Combined coordinate tensor of shape (1, N_total, 4) - """ - if not isinstance(image_latents, list): - raise ValueError(f"Expected `image_latents` to be a list, got {type(image_latents)}.") - - t_coords = [scale + scale * t for t in torch.arange(0, len(image_latents))] - t_coords = [t.view(-1) for t in t_coords] - - image_latent_ids = [] - for x, t in zip(image_latents, t_coords): - x = x.squeeze(0) - _, height, width = x.shape - - x_ids = torch.cartesian_prod(t, torch.arange(height), torch.arange(width), torch.arange(1)) - image_latent_ids.append(x_ids) - - image_latent_ids = torch.cat(image_latent_ids, dim=0) - image_latent_ids = image_latent_ids.unsqueeze(0) - - return image_latent_ids - - @staticmethod - def _pack_latents(latents): - """Pack latents: (batch_size, num_channels, height, width) -> (batch_size, height * width, num_channels)""" - batch_size, num_channels, height, width = latents.shape - latents = latents.reshape(batch_size, num_channels, height * width).permute(0, 2, 1) - return latents - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - image_latents = block_state.image_latents - - if image_latents is None: - block_state.image_latents = None - block_state.image_latent_ids = None - self.set_block_state(state, block_state) - - return components, state - - device = components._execution_device - batch_size = block_state.batch_size * block_state.num_images_per_prompt - - image_latent_ids = self._prepare_image_ids(image_latents) - - packed_latents = [] - for latent in image_latents: - packed = self._pack_latents(latent) - packed = packed.squeeze(0) - packed_latents.append(packed) - - image_latents = torch.cat(packed_latents, dim=0) - image_latents = image_latents.unsqueeze(0) - - image_latents = image_latents.repeat(batch_size, 1, 1) - image_latent_ids = image_latent_ids.repeat(batch_size, 1, 1) - image_latent_ids = image_latent_ids.to(device) - - block_state.image_latents = image_latents - block_state.image_latent_ids = image_latent_ids - - self.set_block_state(state, block_state) - return components, state - - -class Flux2PrepareGuidanceStep(ModularPipelineBlocks): - model_name = "flux2" - - @property - def description(self) -> str: - return "Step that prepares the guidance scale tensor for Flux2 inference" - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("guidance_scale", default=4.0), - InputParam("num_images_per_prompt", default=1), - InputParam( - "batch_size", - required=True, - type_hint=int, - description="Number of prompts, the final batch size of model inputs should be `batch_size * num_images_per_prompt`.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("guidance", type_hint=torch.Tensor, description="Guidance scale tensor"), - ] - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - batch_size = block_state.batch_size * block_state.num_images_per_prompt - guidance = torch.full([1], block_state.guidance_scale, device=device, dtype=torch.float32) - guidance = guidance.expand(batch_size) - block_state.guidance = guidance - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/flux2/decoders.py b/diffusers/modular_pipelines/flux2/decoders.py deleted file mode 100644 index 81f5ca00dc33dfb91c9f783f88262719fc30e7b3..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux2/decoders.py +++ /dev/null @@ -1,185 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from __future__ import annotations - -from typing import Any, Union - -import numpy as np -import PIL -import torch - -from ...configuration_utils import FrozenDict -from ...models import AutoencoderKLFlux2 -from ...pipelines.flux2.image_processor import Flux2ImageProcessor -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class Flux2UnpackLatentsStep(ModularPipelineBlocks): - model_name = "flux2" - - @property - def description(self) -> str: - return "Step that unpacks the latents from the denoising step" - - @property - def inputs(self) -> list[tuple[str, Any]]: - return [ - InputParam( - "latents", - required=True, - type_hint=torch.Tensor, - description="The denoised latents from the denoising step", - ), - InputParam( - "latent_ids", - required=True, - type_hint=torch.Tensor, - description="Position IDs for the latents, used for unpacking", - ), - ] - - @property - def intermediate_outputs(self) -> list[str]: - return [ - OutputParam( - "latents", - type_hint=torch.Tensor, - description="The denoise latents from denoising step, unpacked with position IDs.", - ) - ] - - @staticmethod - def _unpack_latents_with_ids(x: torch.Tensor, x_ids: torch.Tensor) -> torch.Tensor: - """ - Unpack latents using position IDs to scatter tokens into place. - - Args: - x: Packed latents tensor of shape (B, seq_len, C) - x_ids: Position IDs tensor of shape (B, seq_len, 4) with (T, H, W, L) coordinates - - Returns: - Unpacked latents tensor of shape (B, C, H, W) - """ - x_list = [] - for data, pos in zip(x, x_ids): - _, ch = data.shape # noqa: F841 - h_ids = pos[:, 1].to(torch.int64) - w_ids = pos[:, 2].to(torch.int64) - - h = torch.max(h_ids) + 1 - w = torch.max(w_ids) + 1 - - flat_ids = h_ids * w + w_ids - - out = torch.zeros((h * w, ch), device=data.device, dtype=data.dtype) - out.scatter_(0, flat_ids.unsqueeze(1).expand(-1, ch), data) - - out = out.view(h, w, ch).permute(2, 0, 1) - x_list.append(out) - - return torch.stack(x_list, dim=0) - - @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - latents = block_state.latents - latent_ids = block_state.latent_ids - - latents = self._unpack_latents_with_ids(latents, latent_ids) - - block_state.latents = latents - - self.set_block_state(state, block_state) - return components, state - - -class Flux2DecodeStep(ModularPipelineBlocks): - model_name = "flux2" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLFlux2), - ComponentSpec( - "image_processor", - Flux2ImageProcessor, - config=FrozenDict({"vae_scale_factor": 16, "vae_latent_channels": 32}), - default_creation_method="from_config", - ), - ] - - @property - def description(self) -> str: - return "Step that decodes the denoised latents into images using Flux2 VAE with batch norm denormalization" - - @property - def inputs(self) -> list[tuple[str, Any]]: - return [ - InputParam("output_type", default="pil"), - InputParam( - "latents", - required=True, - type_hint=torch.Tensor, - description="The denoised latents from the denoising step", - ), - ] - - @property - def intermediate_outputs(self) -> list[str]: - return [ - OutputParam( - "images", - type_hint=Union[list[PIL.Image.Image], torch.Tensor, np.ndarray], - description="The generated images, can be a list of PIL.Image.Image, torch.Tensor or a numpy array", - ) - ] - - @staticmethod - def _unpatchify_latents(latents): - """Convert patchified latents back to regular format.""" - batch_size, num_channels_latents, height, width = latents.shape - latents = latents.reshape(batch_size, num_channels_latents // (2 * 2), 2, 2, height, width) - latents = latents.permute(0, 1, 4, 2, 5, 3) - latents = latents.reshape(batch_size, num_channels_latents // (2 * 2), height * 2, width * 2) - return latents - - @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - vae = components.vae - - latents = block_state.latents - - latents_bn_mean = vae.bn.running_mean.view(1, -1, 1, 1).to(latents.device, latents.dtype) - latents_bn_std = torch.sqrt(vae.bn.running_var.view(1, -1, 1, 1) + vae.config.batch_norm_eps).to( - latents.device, latents.dtype - ) - latents = latents * latents_bn_std + latents_bn_mean - - latents = self._unpatchify_latents(latents) - - block_state.images = vae.decode(latents, return_dict=False)[0] - block_state.images = components.image_processor.postprocess( - block_state.images, output_type=block_state.output_type - ) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/flux2/denoise.py b/diffusers/modular_pipelines/flux2/denoise.py deleted file mode 100644 index 675f14b03c63a0d777dfc7ca2638641d91460403..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux2/denoise.py +++ /dev/null @@ -1,501 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...models import Flux2Transformer2DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import is_torch_xla_available, logging -from ..modular_pipeline import ( - BlockState, - LoopSequentialPipelineBlocks, - ModularPipelineBlocks, - PipelineState, -) -from ..modular_pipeline_utils import ComponentSpec, ConfigSpec, InputParam, OutputParam -from .modular_pipeline import Flux2KleinModularPipeline, Flux2ModularPipeline - - -if is_torch_xla_available(): - import torch_xla.core.xla_model as xm - - XLA_AVAILABLE = True -else: - XLA_AVAILABLE = False - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class Flux2LoopDenoiser(ModularPipelineBlocks): - model_name = "flux2" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Flux2Transformer2DModel)] - - @property - def description(self) -> str: - return ( - "Step within the denoising loop that denoises the latents for Flux2. " - "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " - "object (e.g. `Flux2DenoiseLoopWrapper`)" - ) - - @property - def inputs(self) -> list[tuple[str, Any]]: - return [ - InputParam("joint_attention_kwargs"), - InputParam( - "latents", - required=True, - type_hint=torch.Tensor, - description="The latents to denoise. Shape: (B, seq_len, C)", - ), - InputParam( - "image_latents", - type_hint=torch.Tensor, - description="Packed image latents for conditioning. Shape: (B, img_seq_len, C)", - ), - InputParam( - "image_latent_ids", - type_hint=torch.Tensor, - description="Position IDs for image latents. Shape: (B, img_seq_len, 4)", - ), - InputParam( - "guidance", - required=True, - type_hint=torch.Tensor, - description="Guidance scale as a tensor", - ), - InputParam( - "prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Text embeddings from Mistral3", - ), - InputParam( - "txt_ids", - required=True, - type_hint=torch.Tensor, - description="4D position IDs for text tokens (T, H, W, L)", - ), - InputParam( - "latent_ids", - required=True, - type_hint=torch.Tensor, - description="4D position IDs for latent tokens (T, H, W, L)", - ), - ] - - @torch.no_grad() - def __call__( - self, components: Flux2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: - latents = block_state.latents - latent_model_input = latents.to(components.transformer.dtype) - img_ids = block_state.latent_ids - - image_latents = getattr(block_state, "image_latents", None) - if image_latents is not None: - latent_model_input = torch.cat([latents, image_latents], dim=1).to(components.transformer.dtype) - image_latent_ids = block_state.image_latent_ids - img_ids = torch.cat([img_ids, image_latent_ids], dim=1) - - timestep = t.expand(latents.shape[0]).to(latents.dtype) - - noise_pred = components.transformer( - hidden_states=latent_model_input, - timestep=timestep / 1000, - guidance=block_state.guidance, - encoder_hidden_states=block_state.prompt_embeds, - txt_ids=block_state.txt_ids, - img_ids=img_ids, - joint_attention_kwargs=block_state.joint_attention_kwargs, - return_dict=False, - )[0] - - noise_pred = noise_pred[:, : latents.size(1)] - block_state.noise_pred = noise_pred - - return components, block_state - - -# same as Flux2LoopDenoiser but guidance=None -class Flux2KleinLoopDenoiser(ModularPipelineBlocks): - model_name = "flux2-klein" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Flux2Transformer2DModel)] - - @property - def description(self) -> str: - return ( - "Step within the denoising loop that denoises the latents for Flux2. " - "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " - "object (e.g. `Flux2DenoiseLoopWrapper`)" - ) - - @property - def inputs(self) -> list[tuple[str, Any]]: - return [ - InputParam("joint_attention_kwargs"), - InputParam( - "latents", - required=True, - type_hint=torch.Tensor, - description="The latents to denoise. Shape: (B, seq_len, C)", - ), - InputParam( - "image_latents", - type_hint=torch.Tensor, - description="Packed image latents for conditioning. Shape: (B, img_seq_len, C)", - ), - InputParam( - "image_latent_ids", - type_hint=torch.Tensor, - description="Position IDs for image latents. Shape: (B, img_seq_len, 4)", - ), - InputParam( - "prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Text embeddings from Qwen3", - ), - InputParam( - "txt_ids", - required=True, - type_hint=torch.Tensor, - description="4D position IDs for text tokens (T, H, W, L)", - ), - InputParam( - "latent_ids", - required=True, - type_hint=torch.Tensor, - description="4D position IDs for latent tokens (T, H, W, L)", - ), - ] - - @torch.no_grad() - def __call__( - self, components: Flux2KleinModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: - latents = block_state.latents - latent_model_input = latents.to(components.transformer.dtype) - img_ids = block_state.latent_ids - - image_latents = getattr(block_state, "image_latents", None) - if image_latents is not None: - latent_model_input = torch.cat([latents, image_latents], dim=1).to(components.transformer.dtype) - image_latent_ids = block_state.image_latent_ids - img_ids = torch.cat([img_ids, image_latent_ids], dim=1) - - timestep = t.expand(latents.shape[0]).to(latents.dtype) - - noise_pred = components.transformer( - hidden_states=latent_model_input, - timestep=timestep / 1000, - guidance=None, - encoder_hidden_states=block_state.prompt_embeds, - txt_ids=block_state.txt_ids, - img_ids=img_ids, - joint_attention_kwargs=block_state.joint_attention_kwargs, - return_dict=False, - )[0] - - noise_pred = noise_pred[:, : latents.size(1)] - block_state.noise_pred = noise_pred - - return components, block_state - - -# support CFG for Flux2-Klein base model -class Flux2KleinBaseLoopDenoiser(ModularPipelineBlocks): - model_name = "flux2-klein" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("transformer", Flux2Transformer2DModel), - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 4.0}), - default_creation_method="from_config", - ), - ] - - @property - def expected_configs(self) -> list[ConfigSpec]: - return [ - ConfigSpec(name="is_distilled", default=False), - ] - - @property - def description(self) -> str: - return ( - "Step within the denoising loop that denoises the latents for Flux2. " - "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " - "object (e.g. `Flux2DenoiseLoopWrapper`)" - ) - - @property - def inputs(self) -> list[tuple[str, Any]]: - return [ - InputParam("joint_attention_kwargs"), - InputParam( - "latents", - required=True, - type_hint=torch.Tensor, - description="The latents to denoise. Shape: (B, seq_len, C)", - ), - InputParam( - "image_latents", - type_hint=torch.Tensor, - description="Packed image latents for conditioning. Shape: (B, img_seq_len, C)", - ), - InputParam( - "image_latent_ids", - type_hint=torch.Tensor, - description="Position IDs for image latents. Shape: (B, img_seq_len, 4)", - ), - InputParam( - "prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Text embeddings from Qwen3", - ), - InputParam( - "negative_prompt_embeds", - required=False, - type_hint=torch.Tensor, - description="Negative text embeddings from Qwen3", - ), - InputParam( - "txt_ids", - required=True, - type_hint=torch.Tensor, - description="4D position IDs for text tokens (T, H, W, L)", - ), - InputParam( - "negative_txt_ids", - required=False, - type_hint=torch.Tensor, - description="4D position IDs for negative text tokens (T, H, W, L)", - ), - InputParam( - "latent_ids", - required=True, - type_hint=torch.Tensor, - description="4D position IDs for latent tokens (T, H, W, L)", - ), - ] - - @torch.no_grad() - def __call__( - self, components: Flux2KleinModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: - latents = block_state.latents - latent_model_input = latents.to(components.transformer.dtype) - img_ids = block_state.latent_ids - - image_latents = getattr(block_state, "image_latents", None) - if image_latents is not None: - latent_model_input = torch.cat([latents, image_latents], dim=1).to(components.transformer.dtype) - image_latent_ids = block_state.image_latent_ids - img_ids = torch.cat([img_ids, image_latent_ids], dim=1) - - timestep = t.expand(latents.shape[0]).to(latents.dtype) - - guider_inputs = { - "encoder_hidden_states": ( - getattr(block_state, "prompt_embeds", None), - getattr(block_state, "negative_prompt_embeds", None), - ), - "txt_ids": ( - getattr(block_state, "txt_ids", None), - getattr(block_state, "negative_txt_ids", None), - ), - } - - components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) - guider_state = components.guider.prepare_inputs(guider_inputs) - - for guider_state_batch in guider_state: - components.guider.prepare_models(components.transformer) - cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()} - - noise_pred = components.transformer( - hidden_states=latent_model_input, - timestep=timestep / 1000, - guidance=None, - img_ids=img_ids, - joint_attention_kwargs=block_state.joint_attention_kwargs, - return_dict=False, - **cond_kwargs, - )[0] - guider_state_batch.noise_pred = noise_pred[:, : latents.size(1)] - components.guider.cleanup_models(components.transformer) - - # perform guidance - block_state.noise_pred = components.guider(guider_state)[0] - - return components, block_state - - -class Flux2LoopAfterDenoiser(ModularPipelineBlocks): - model_name = "flux2" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def description(self) -> str: - return ( - "Step within the denoising loop that updates the latents after denoising. " - "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " - "object (e.g. `Flux2DenoiseLoopWrapper`)" - ) - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam("latents", type_hint=torch.Tensor, description="The denoised latents")] - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - latents_dtype = block_state.latents.dtype - block_state.latents = components.scheduler.step( - block_state.noise_pred, - t, - block_state.latents, - return_dict=False, - )[0] - - if block_state.latents.dtype != latents_dtype: - if torch.backends.mps.is_available(): - block_state.latents = block_state.latents.to(latents_dtype) - - return components, block_state - - -class Flux2DenoiseLoopWrapper(LoopSequentialPipelineBlocks): - model_name = "flux2" - - @property - def description(self) -> str: - return ( - "Pipeline block that iteratively denoises the latents over `timesteps`. " - "The specific steps within each iteration can be customized with `sub_blocks` attribute" - ) - - @property - def loop_expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), - ComponentSpec("transformer", Flux2Transformer2DModel), - ] - - @property - def loop_inputs(self) -> list[InputParam]: - return [ - InputParam( - "timesteps", - required=True, - type_hint=torch.Tensor, - description="The timesteps to use for the denoising process.", - ), - InputParam( - "num_inference_steps", - required=True, - type_hint=int, - description="The number of inference steps to use for the denoising process.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - block_state.num_warmup_steps = max( - len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order, 0 - ) - - with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) - - if i == len(block_state.timesteps) - 1 or ( - (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 - ): - progress_bar.update() - - if XLA_AVAILABLE: - xm.mark_step() - - self.set_block_state(state, block_state) - return components, state - - -class Flux2DenoiseStep(Flux2DenoiseLoopWrapper): - block_classes = [Flux2LoopDenoiser, Flux2LoopAfterDenoiser] - block_names = ["denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoise step that iteratively denoises the latents for Flux2. \n" - "Its loop logic is defined in `Flux2DenoiseLoopWrapper.__call__` method \n" - "At each iteration, it runs blocks defined in `sub_blocks` sequentially:\n" - " - `Flux2LoopDenoiser`\n" - " - `Flux2LoopAfterDenoiser`\n" - "This block supports both text-to-image and image-conditioned generation." - ) - - -class Flux2KleinDenoiseStep(Flux2DenoiseLoopWrapper): - block_classes = [Flux2KleinLoopDenoiser, Flux2LoopAfterDenoiser] - block_names = ["denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoise step that iteratively denoises the latents for Flux2. \n" - "Its loop logic is defined in `Flux2DenoiseLoopWrapper.__call__` method \n" - "At each iteration, it runs blocks defined in `sub_blocks` sequentially:\n" - " - `Flux2KleinLoopDenoiser`\n" - " - `Flux2LoopAfterDenoiser`\n" - "This block supports both text-to-image and image-conditioned generation." - ) - - -class Flux2KleinBaseDenoiseStep(Flux2DenoiseLoopWrapper): - block_classes = [Flux2KleinBaseLoopDenoiser, Flux2LoopAfterDenoiser] - block_names = ["denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoise step that iteratively denoises the latents for Flux2. \n" - "Its loop logic is defined in `Flux2DenoiseLoopWrapper.__call__` method \n" - "At each iteration, it runs blocks defined in `sub_blocks` sequentially:\n" - " - `Flux2KleinBaseLoopDenoiser`\n" - " - `Flux2LoopAfterDenoiser`\n" - "This block supports both text-to-image and image-conditioned generation." - ) diff --git a/diffusers/modular_pipelines/flux2/encoders.py b/diffusers/modular_pipelines/flux2/encoders.py deleted file mode 100644 index 09615c4becb6865ce5e76fd7576c0ee8048bd329..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux2/encoders.py +++ /dev/null @@ -1,608 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch -from transformers import AutoProcessor, Mistral3ForConditionalGeneration, Qwen2TokenizerFast, Qwen3ForCausalLM - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...models import AutoencoderKLFlux2 -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, ConfigSpec, InputParam, OutputParam -from .modular_pipeline import Flux2KleinModularPipeline, Flux2ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def format_text_input(prompts: list[str], system_message: str = None): - """Format prompts for Mistral3 chat template.""" - cleaned_txt = [prompt.replace("[IMG]", "") for prompt in prompts] - - return [ - [ - { - "role": "system", - "content": [{"type": "text", "text": system_message}], - }, - {"role": "user", "content": [{"type": "text", "text": prompt}]}, - ] - for prompt in cleaned_txt - ] - - -# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.retrieve_latents -def retrieve_latents( - encoder_output: torch.Tensor, generator: torch.Generator | None = None, sample_mode: str = "sample" -): - if hasattr(encoder_output, "latent_dist") and sample_mode == "sample": - return encoder_output.latent_dist.sample(generator) - elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax": - return encoder_output.latent_dist.mode() - elif hasattr(encoder_output, "latents"): - return encoder_output.latents - else: - raise AttributeError("Could not access latents of provided encoder_output") - - -class Flux2TextEncoderStep(ModularPipelineBlocks): - model_name = "flux2" - - # fmt: off - DEFAULT_SYSTEM_MESSAGE = "You are an AI that reasons about image descriptions. You give structured responses focusing on object relationships, object attribution and actions without speculation." - # fmt: on - - @property - def description(self) -> str: - return "Text Encoder step that generates text embeddings using Mistral3 to guide the image generation" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", Mistral3ForConditionalGeneration), - ComponentSpec("tokenizer", AutoProcessor), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("prompt"), - InputParam("max_sequence_length", type_hint=int, default=512, required=False), - InputParam("text_encoder_out_layers", type_hint=tuple[int], default=(10, 20, 30), required=False), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "prompt_embeds", - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="Text embeddings from Mistral3 used to guide the image generation", - ), - ] - - @staticmethod - def check_inputs(block_state): - prompt = block_state.prompt - if prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): - raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}") - - @staticmethod - def _get_mistral_3_prompt_embeds( - text_encoder: Mistral3ForConditionalGeneration, - tokenizer: AutoProcessor, - prompt: str | list[str], - dtype: torch.dtype | None = None, - device: torch.device | None = None, - max_sequence_length: int = 512, - # fmt: off - system_message: str = "You are an AI that reasons about image descriptions. You give structured responses focusing on object relationships, object attribution and actions without speculation.", - # fmt: on - hidden_states_layers: tuple[int] = (10, 20, 30), - ): - dtype = text_encoder.dtype if dtype is None else dtype - device = text_encoder.device if device is None else device - - prompt = [prompt] if isinstance(prompt, str) else prompt - - messages_batch = format_text_input(prompts=prompt, system_message=system_message) - - inputs = tokenizer.apply_chat_template( - messages_batch, - add_generation_prompt=False, - tokenize=True, - return_dict=True, - return_tensors="pt", - padding="max_length", - truncation=True, - max_length=max_sequence_length, - ) - - input_ids = inputs["input_ids"].to(device) - attention_mask = inputs["attention_mask"].to(device) - - output = text_encoder( - input_ids=input_ids, - attention_mask=attention_mask, - output_hidden_states=True, - use_cache=False, - ) - - out = torch.stack([output.hidden_states[k] for k in hidden_states_layers], dim=1) - out = out.to(dtype=dtype, device=device) - - batch_size, num_channels, seq_len, hidden_dim = out.shape - prompt_embeds = out.permute(0, 2, 1, 3).reshape(batch_size, seq_len, num_channels * hidden_dim) - - return prompt_embeds - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - self.check_inputs(block_state) - - block_state.device = components._execution_device - - prompt = block_state.prompt - if prompt is None: - prompt = "" - prompt = [prompt] if isinstance(prompt, str) else prompt - - block_state.prompt_embeds = self._get_mistral_3_prompt_embeds( - text_encoder=components.text_encoder, - tokenizer=components.tokenizer, - prompt=prompt, - device=block_state.device, - max_sequence_length=block_state.max_sequence_length, - system_message=self.DEFAULT_SYSTEM_MESSAGE, - hidden_states_layers=block_state.text_encoder_out_layers, - ) - - self.set_block_state(state, block_state) - return components, state - - -class Flux2RemoteTextEncoderStep(ModularPipelineBlocks): - model_name = "flux2" - - REMOTE_URL = "https://remote-text-encoder-flux-2.huggingface.co/predict" - - @property - def description(self) -> str: - return "Text Encoder step that generates text embeddings using a remote API endpoint" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("prompt"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "prompt_embeds", - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="Text embeddings from remote API used to guide the image generation", - ), - ] - - @staticmethod - def check_inputs(block_state): - prompt = block_state.prompt - if prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): - raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(block_state.prompt)}") - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - import io - - import requests - from huggingface_hub import get_token - - block_state = self.get_block_state(state) - self.check_inputs(block_state) - - block_state.device = components._execution_device - - prompt = block_state.prompt - if prompt is None: - prompt = "" - prompt = [prompt] if isinstance(prompt, str) else prompt - - response = requests.post( - self.REMOTE_URL, - json={"prompt": prompt}, - headers={ - "Authorization": f"Bearer {get_token()}", - "Content-Type": "application/json", - }, - ) - response.raise_for_status() - - block_state.prompt_embeds = torch.load(io.BytesIO(response.content), weights_only=True) - block_state.prompt_embeds = block_state.prompt_embeds.to(block_state.device) - - self.set_block_state(state, block_state) - return components, state - - -class Flux2KleinTextEncoderStep(ModularPipelineBlocks): - model_name = "flux2-klein" - - @property - def description(self) -> str: - return "Text Encoder step that generates text embeddings using Qwen3 to guide the image generation" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", Qwen3ForCausalLM), - ComponentSpec("tokenizer", Qwen2TokenizerFast), - ] - - @property - def expected_configs(self) -> list[ConfigSpec]: - return [ - ConfigSpec(name="is_distilled", default=True), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("prompt"), - InputParam("max_sequence_length", type_hint=int, default=512, required=False), - InputParam("text_encoder_out_layers", type_hint=tuple[int], default=(9, 18, 27), required=False), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "prompt_embeds", - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="Text embeddings from qwen3 used to guide the image generation", - ), - ] - - @staticmethod - def check_inputs(block_state): - prompt = block_state.prompt - - if prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): - raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}") - - @staticmethod - # Copied from diffusers.pipelines.flux2.pipeline_flux2_klein.Flux2KleinPipeline._get_qwen3_prompt_embeds - def _get_qwen3_prompt_embeds( - text_encoder: Qwen3ForCausalLM, - tokenizer: Qwen2TokenizerFast, - prompt: str | list[str], - dtype: torch.dtype | None = None, - device: torch.device | None = None, - max_sequence_length: int = 512, - hidden_states_layers: list[int] = (9, 18, 27), - ): - dtype = text_encoder.dtype if dtype is None else dtype - device = text_encoder.device if device is None else device - - prompt = [prompt] if isinstance(prompt, str) else prompt - - all_input_ids = [] - all_attention_masks = [] - - for single_prompt in prompt: - messages = [{"role": "user", "content": single_prompt}] - text = tokenizer.apply_chat_template( - messages, - tokenize=False, - add_generation_prompt=True, - enable_thinking=False, - ) - inputs = tokenizer( - text, - return_tensors="pt", - padding="max_length", - truncation=True, - max_length=max_sequence_length, - ) - - all_input_ids.append(inputs["input_ids"]) - all_attention_masks.append(inputs["attention_mask"]) - - input_ids = torch.cat(all_input_ids, dim=0).to(device) - attention_mask = torch.cat(all_attention_masks, dim=0).to(device) - - # Forward pass through the model - output = text_encoder( - input_ids=input_ids, - attention_mask=attention_mask, - output_hidden_states=True, - use_cache=False, - ) - - # Only use outputs from intermediate layers and stack them - out = torch.stack([output.hidden_states[k] for k in hidden_states_layers], dim=1) - out = out.to(dtype=dtype, device=device) - - batch_size, num_channels, seq_len, hidden_dim = out.shape - prompt_embeds = out.permute(0, 2, 1, 3).reshape(batch_size, seq_len, num_channels * hidden_dim) - - return prompt_embeds - - @torch.no_grad() - def __call__(self, components: Flux2KleinModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - self.check_inputs(block_state) - - device = components._execution_device - - prompt = block_state.prompt - if prompt is None: - prompt = "" - prompt = [prompt] if isinstance(prompt, str) else prompt - - block_state.prompt_embeds = self._get_qwen3_prompt_embeds( - text_encoder=components.text_encoder, - tokenizer=components.tokenizer, - prompt=prompt, - device=device, - max_sequence_length=block_state.max_sequence_length, - hidden_states_layers=block_state.text_encoder_out_layers, - ) - - self.set_block_state(state, block_state) - return components, state - - -class Flux2KleinBaseTextEncoderStep(ModularPipelineBlocks): - model_name = "flux2-klein" - - @property - def description(self) -> str: - return "Text Encoder step that generates text embeddings using Qwen3 to guide the image generation" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", Qwen3ForCausalLM), - ComponentSpec("tokenizer", Qwen2TokenizerFast), - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 4.0}), - default_creation_method="from_config", - ), - ] - - @property - def expected_configs(self) -> list[ConfigSpec]: - return [ - ConfigSpec(name="is_distilled", default=False), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("prompt"), - InputParam("max_sequence_length", type_hint=int, default=512, required=False), - InputParam("text_encoder_out_layers", type_hint=tuple[int], default=(9, 18, 27), required=False), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "prompt_embeds", - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="Text embeddings from qwen3 used to guide the image generation", - ), - OutputParam( - "negative_prompt_embeds", - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="Negative text embeddings from qwen3 used to guide the image generation", - ), - ] - - @staticmethod - def check_inputs(block_state): - prompt = block_state.prompt - - if prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): - raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}") - - @staticmethod - # Copied from diffusers.pipelines.flux2.pipeline_flux2_klein.Flux2KleinPipeline._get_qwen3_prompt_embeds - def _get_qwen3_prompt_embeds( - text_encoder: Qwen3ForCausalLM, - tokenizer: Qwen2TokenizerFast, - prompt: str | list[str], - dtype: torch.dtype | None = None, - device: torch.device | None = None, - max_sequence_length: int = 512, - hidden_states_layers: list[int] = (9, 18, 27), - ): - dtype = text_encoder.dtype if dtype is None else dtype - device = text_encoder.device if device is None else device - - prompt = [prompt] if isinstance(prompt, str) else prompt - - all_input_ids = [] - all_attention_masks = [] - - for single_prompt in prompt: - messages = [{"role": "user", "content": single_prompt}] - text = tokenizer.apply_chat_template( - messages, - tokenize=False, - add_generation_prompt=True, - enable_thinking=False, - ) - inputs = tokenizer( - text, - return_tensors="pt", - padding="max_length", - truncation=True, - max_length=max_sequence_length, - ) - - all_input_ids.append(inputs["input_ids"]) - all_attention_masks.append(inputs["attention_mask"]) - - input_ids = torch.cat(all_input_ids, dim=0).to(device) - attention_mask = torch.cat(all_attention_masks, dim=0).to(device) - - # Forward pass through the model - output = text_encoder( - input_ids=input_ids, - attention_mask=attention_mask, - output_hidden_states=True, - use_cache=False, - ) - - # Only use outputs from intermediate layers and stack them - out = torch.stack([output.hidden_states[k] for k in hidden_states_layers], dim=1) - out = out.to(dtype=dtype, device=device) - - batch_size, num_channels, seq_len, hidden_dim = out.shape - prompt_embeds = out.permute(0, 2, 1, 3).reshape(batch_size, seq_len, num_channels * hidden_dim) - - return prompt_embeds - - @torch.no_grad() - def __call__(self, components: Flux2KleinModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - self.check_inputs(block_state) - - device = components._execution_device - - prompt = block_state.prompt - if prompt is None: - prompt = "" - prompt = [prompt] if isinstance(prompt, str) else prompt - - block_state.prompt_embeds = self._get_qwen3_prompt_embeds( - text_encoder=components.text_encoder, - tokenizer=components.tokenizer, - prompt=prompt, - device=device, - max_sequence_length=block_state.max_sequence_length, - hidden_states_layers=block_state.text_encoder_out_layers, - ) - - if components.requires_unconditional_embeds: - negative_prompt = [""] * len(prompt) - block_state.negative_prompt_embeds = self._get_qwen3_prompt_embeds( - text_encoder=components.text_encoder, - tokenizer=components.tokenizer, - prompt=negative_prompt, - device=device, - max_sequence_length=block_state.max_sequence_length, - hidden_states_layers=block_state.text_encoder_out_layers, - ) - else: - block_state.negative_prompt_embeds = None - - self.set_block_state(state, block_state) - return components, state - - -class Flux2VaeEncoderStep(ModularPipelineBlocks): - model_name = "flux2" - - @property - def description(self) -> str: - return "VAE Encoder step that encodes preprocessed images into latent representations for Flux2." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("vae", AutoencoderKLFlux2)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("condition_images", type_hint=list[torch.Tensor]), - InputParam("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "image_latents", - type_hint=list[torch.Tensor], - description="List of latent representations for each reference image", - ), - ] - - @staticmethod - def _patchify_latents(latents): - """Convert latents to patchified format for Flux2.""" - batch_size, num_channels_latents, height, width = latents.shape - latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2) - latents = latents.permute(0, 1, 3, 5, 2, 4) - latents = latents.reshape(batch_size, num_channels_latents * 4, height // 2, width // 2) - return latents - - def _encode_vae_image(self, vae: AutoencoderKLFlux2, image: torch.Tensor, generator: torch.Generator): - """Encode a single image using Flux2 VAE with batch norm normalization.""" - if image.ndim != 4: - raise ValueError(f"Expected image dims 4, got {image.ndim}.") - - image_latents = retrieve_latents(vae.encode(image), generator=generator, sample_mode="argmax") - image_latents = self._patchify_latents(image_latents) - - latents_bn_mean = vae.bn.running_mean.view(1, -1, 1, 1).to(image_latents.device, image_latents.dtype) - latents_bn_std = torch.sqrt(vae.bn.running_var.view(1, -1, 1, 1) + vae.config.batch_norm_eps) - latents_bn_std = latents_bn_std.to(image_latents.device, image_latents.dtype) - image_latents = (image_latents - latents_bn_mean) / latents_bn_std - - return image_latents - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - condition_images = block_state.condition_images - - if condition_images is None: - return components, state - - device = components._execution_device - dtype = components.vae.dtype - - image_latents = [] - for image in condition_images: - image = image.to(device=device, dtype=dtype) - latent = self._encode_vae_image( - vae=components.vae, - image=image, - generator=block_state.generator, - ) - image_latents.append(latent) - - block_state.image_latents = image_latents - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/flux2/inputs.py b/diffusers/modular_pipelines/flux2/inputs.py deleted file mode 100644 index 6bfe6aec97fd00621ebc695d8eaa896740c29d34..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux2/inputs.py +++ /dev/null @@ -1,242 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ...configuration_utils import FrozenDict -from ...pipelines.flux2.image_processor import Flux2ImageProcessor -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Flux2ModularPipeline - - -logger = logging.get_logger(__name__) - - -class Flux2TextInputStep(ModularPipelineBlocks): - model_name = "flux2" - - @property - def description(self) -> str: - return ( - "This step:\n" - " 1. Determines `batch_size` and `dtype` based on `prompt_embeds`\n" - " 2. Ensures all text embeddings have consistent batch sizes (batch_size * num_images_per_prompt)" - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("num_images_per_prompt", default=1), - InputParam( - "prompt_embeds", - required=True, - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="Pre-generated text embeddings. Can be generated from text_encoder step.", - ), - ] - - @property - def intermediate_outputs(self) -> list[str]: - return [ - OutputParam( - "batch_size", - type_hint=int, - description="Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt", - ), - OutputParam( - "dtype", - type_hint=torch.dtype, - description="Data type of model tensor inputs (determined by `prompt_embeds`)", - ), - OutputParam( - "prompt_embeds", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Text embeddings used to guide the image generation", - ), - ] - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - block_state.batch_size = block_state.prompt_embeds.shape[0] - block_state.dtype = block_state.prompt_embeds.dtype - - _, seq_len, _ = block_state.prompt_embeds.shape - block_state.prompt_embeds = block_state.prompt_embeds.repeat(1, block_state.num_images_per_prompt, 1) - block_state.prompt_embeds = block_state.prompt_embeds.view( - block_state.batch_size * block_state.num_images_per_prompt, seq_len, -1 - ) - - self.set_block_state(state, block_state) - return components, state - - -class Flux2KleinBaseTextInputStep(ModularPipelineBlocks): - model_name = "flux2-klein" - - @property - def description(self) -> str: - return ( - "This step:\n" - " 1. Determines `batch_size` and `dtype` based on `prompt_embeds`\n" - " 2. Ensures all text embeddings have consistent batch sizes (batch_size * num_images_per_prompt)" - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("num_images_per_prompt", default=1), - InputParam( - "prompt_embeds", - required=True, - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="Pre-generated text embeddings. Can be generated from text_encoder step.", - ), - InputParam( - "negative_prompt_embeds", - required=False, - kwargs_type="denoiser_input_fields", - type_hint=torch.Tensor, - description="Pre-generated negative text embeddings. Can be generated from text_encoder step.", - ), - ] - - @property - def intermediate_outputs(self) -> list[str]: - return [ - OutputParam( - "batch_size", - type_hint=int, - description="Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt", - ), - OutputParam( - "dtype", - type_hint=torch.dtype, - description="Data type of model tensor inputs (determined by `prompt_embeds`)", - ), - OutputParam( - "prompt_embeds", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Text embeddings used to guide the image generation", - ), - OutputParam( - "negative_prompt_embeds", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Negative text embeddings used to guide the image generation", - ), - ] - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - block_state.batch_size = block_state.prompt_embeds.shape[0] - block_state.dtype = block_state.prompt_embeds.dtype - - _, seq_len, _ = block_state.prompt_embeds.shape - block_state.prompt_embeds = block_state.prompt_embeds.repeat(1, block_state.num_images_per_prompt, 1) - block_state.prompt_embeds = block_state.prompt_embeds.view( - block_state.batch_size * block_state.num_images_per_prompt, seq_len, -1 - ) - - if block_state.negative_prompt_embeds is not None: - _, seq_len, _ = block_state.negative_prompt_embeds.shape - block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.repeat( - 1, block_state.num_images_per_prompt, 1 - ) - block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.view( - block_state.batch_size * block_state.num_images_per_prompt, seq_len, -1 - ) - - self.set_block_state(state, block_state) - return components, state - - -class Flux2ProcessImagesInputStep(ModularPipelineBlocks): - model_name = "flux2" - - @property - def description(self) -> str: - return "Image preprocess step for Flux2. Validates and preprocesses reference images." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "image_processor", - Flux2ImageProcessor, - config=FrozenDict({"vae_scale_factor": 16, "vae_latent_channels": 32}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("image"), - InputParam("height"), - InputParam("width"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam(name="condition_images", type_hint=list[torch.Tensor])] - - @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState): - block_state = self.get_block_state(state) - images = block_state.image - - if images is None: - block_state.condition_images = None - self.set_block_state(state, block_state) - return components, state - - if not isinstance(images, list): - images = [images] - - condition_images = [] - for img in images: - components.image_processor.check_image_input(img) - - image_width, image_height = img.size - if image_width * image_height > 1024 * 1024: - img = components.image_processor._resize_to_target_area(img, 1024 * 1024) - image_width, image_height = img.size - - multiple_of = components.vae_scale_factor * 2 - image_width = (image_width // multiple_of) * multiple_of - image_height = (image_height // multiple_of) * multiple_of - condition_img = components.image_processor.preprocess( - img, height=image_height, width=image_width, resize_mode="crop" - ) - condition_images.append(condition_img) - - if block_state.height is None: - block_state.height = image_height - if block_state.width is None: - block_state.width = image_width - - block_state.condition_images = condition_images - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/flux2/modular_blocks_flux2.py b/diffusers/modular_pipelines/flux2/modular_blocks_flux2.py deleted file mode 100644 index 2bbb7975a9834ff84f09063cb587adf985cda8e9..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux2/modular_blocks_flux2.py +++ /dev/null @@ -1,356 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from ...utils import logging -from ..modular_pipeline import AutoPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import InsertableDict, OutputParam -from .before_denoise import ( - Flux2PrepareGuidanceStep, - Flux2PrepareImageLatentsStep, - Flux2PrepareLatentsStep, - Flux2RoPEInputsStep, - Flux2SetTimestepsStep, -) -from .decoders import Flux2DecodeStep, Flux2UnpackLatentsStep -from .denoise import Flux2DenoiseStep -from .encoders import ( - Flux2TextEncoderStep, - Flux2VaeEncoderStep, -) -from .inputs import ( - Flux2ProcessImagesInputStep, - Flux2TextInputStep, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# auto_docstring -class Flux2VaeEncoderSequentialStep(SequentialPipelineBlocks): - """ - VAE encoder step that preprocesses, encodes, and prepares image latents for Flux2 conditioning. - - Components: - image_processor (`Flux2ImageProcessor`) vae (`AutoencoderKLFlux2`) - - Inputs: - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - - Outputs: - condition_images (`list`): - TODO: Add description. - image_latents (`list`): - List of latent representations for each reference image - """ - - model_name = "flux2" - - block_classes = [Flux2ProcessImagesInputStep(), Flux2VaeEncoderStep()] - block_names = ["preprocess", "encode"] - - @property - def description(self) -> str: - return "VAE encoder step that preprocesses, encodes, and prepares image latents for Flux2 conditioning." - - -# auto_docstring -class Flux2AutoVaeEncoderStep(AutoPipelineBlocks): - """ - VAE encoder step that encodes the image inputs into their latent representations. - This is an auto pipeline block that works for image conditioning tasks. - - `Flux2VaeEncoderSequentialStep` is used when `image` is provided. - - If `image` is not provided, step will be skipped. - - Components: - image_processor (`Flux2ImageProcessor`) vae (`AutoencoderKLFlux2`) - - Inputs: - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - - Outputs: - condition_images (`list`): - TODO: Add description. - image_latents (`list`): - List of latent representations for each reference image - """ - - block_classes = [Flux2VaeEncoderSequentialStep] - block_names = ["img_conditioning"] - block_trigger_inputs = ["image"] - - @property - def description(self): - return ( - "VAE encoder step that encodes the image inputs into their latent representations.\n" - "This is an auto pipeline block that works for image conditioning tasks.\n" - " - `Flux2VaeEncoderSequentialStep` is used when `image` is provided.\n" - " - If `image` is not provided, step will be skipped." - ) - - -Flux2CoreDenoiseBlocks = InsertableDict( - [ - ("input", Flux2TextInputStep()), - ("prepare_latents", Flux2PrepareLatentsStep()), - ("set_timesteps", Flux2SetTimestepsStep()), - ("prepare_guidance", Flux2PrepareGuidanceStep()), - ("prepare_rope_inputs", Flux2RoPEInputsStep()), - ("denoise", Flux2DenoiseStep()), - ("after_denoise", Flux2UnpackLatentsStep()), - ] -) - - -# auto_docstring -class Flux2CoreDenoiseStep(SequentialPipelineBlocks): - """ - Core denoise step that performs the denoising process for Flux2-dev. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Flux2Transformer2DModel`) - - Inputs: - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 4.0): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - image_latents (`Tensor`, *optional*): - Packed image latents for conditioning. Shape: (B, img_seq_len, C) - image_latent_ids (`Tensor`, *optional*): - Position IDs for image latents. Shape: (B, img_seq_len, 4) - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "flux2" - - block_classes = Flux2CoreDenoiseBlocks.values() - block_names = Flux2CoreDenoiseBlocks.keys() - - @property - def description(self): - return "Core denoise step that performs the denoising process for Flux2-dev." - - @property - def outputs(self): - return [ - OutputParam.template("latents"), - ] - - -Flux2ImageConditionedCoreDenoiseBlocks = InsertableDict( - [ - ("input", Flux2TextInputStep()), - ("prepare_image_latents", Flux2PrepareImageLatentsStep()), - ("prepare_latents", Flux2PrepareLatentsStep()), - ("set_timesteps", Flux2SetTimestepsStep()), - ("prepare_guidance", Flux2PrepareGuidanceStep()), - ("prepare_rope_inputs", Flux2RoPEInputsStep()), - ("denoise", Flux2DenoiseStep()), - ("after_denoise", Flux2UnpackLatentsStep()), - ] -) - - -# auto_docstring -class Flux2ImageConditionedCoreDenoiseStep(SequentialPipelineBlocks): - """ - Core denoise step that performs the denoising process for Flux2-dev with image conditioning. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Flux2Transformer2DModel`) - - Inputs: - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - image_latents (`list`, *optional*): - TODO: Add description. - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 4.0): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "flux2" - - block_classes = Flux2ImageConditionedCoreDenoiseBlocks.values() - block_names = Flux2ImageConditionedCoreDenoiseBlocks.keys() - - @property - def description(self): - return "Core denoise step that performs the denoising process for Flux2-dev with image conditioning." - - @property - def outputs(self): - return [ - OutputParam.template("latents"), - ] - - -class Flux2AutoCoreDenoiseStep(AutoPipelineBlocks): - model_name = "flux2" - - block_classes = [Flux2ImageConditionedCoreDenoiseStep, Flux2CoreDenoiseStep] - block_names = ["image_conditioned", "text2image"] - block_trigger_inputs = ["image_latents", None] - - @property - def description(self): - return ( - "Auto core denoise step that performs the denoising process for Flux2-dev." - "This is an auto pipeline block that works for text-to-image and image-conditioned generation." - " - `Flux2CoreDenoiseStep` is used for text-to-image generation.\n" - " - `Flux2ImageConditionedCoreDenoiseStep` is used for image-conditioned generation.\n" - ) - - -AUTO_BLOCKS = InsertableDict( - [ - ("text_encoder", Flux2TextEncoderStep()), - ("vae_encoder", Flux2AutoVaeEncoderStep()), - ("denoise", Flux2AutoCoreDenoiseStep()), - ("decode", Flux2DecodeStep()), - ] -) - - -# auto_docstring -class Flux2AutoBlocks(SequentialPipelineBlocks): - """ - Auto Modular pipeline for text-to-image and image-conditioned generation using Flux2. - - Supported workflows: - - `text2image`: requires `prompt` - - `image_conditioned`: requires `image`, `prompt` - - Components: - text_encoder (`Mistral3ForConditionalGeneration`) tokenizer (`AutoProcessor`) image_processor - (`Flux2ImageProcessor`) vae (`AutoencoderKLFlux2`) scheduler (`FlowMatchEulerDiscreteScheduler`) transformer - (`Flux2Transformer2DModel`) - - Inputs: - prompt (`None`, *optional*): - TODO: Add description. - max_sequence_length (`int`, *optional*, defaults to 512): - TODO: Add description. - text_encoder_out_layers (`tuple`, *optional*, defaults to (10, 20, 30)): - TODO: Add description. - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - image_latents (`list`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`): - TODO: Add description. - num_inference_steps (`None`): - TODO: Add description. - timesteps (`None`): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - guidance_scale (`None`, *optional*, defaults to 4.0): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - image_latent_ids (`Tensor`, *optional*): - Position IDs for image latents. Shape: (B, img_seq_len, 4) - output_type (`None`, *optional*, defaults to pil): - TODO: Add description. - - Outputs: - images (`list`): - Generated images. - """ - - model_name = "flux2" - - block_classes = AUTO_BLOCKS.values() - block_names = AUTO_BLOCKS.keys() - _workflow_map = { - "text2image": {"prompt": True}, - "image_conditioned": {"image": True, "prompt": True}, - } - - @property - def description(self): - return "Auto Modular pipeline for text-to-image and image-conditioned generation using Flux2." - - @property - def outputs(self): - return [ - OutputParam.template("images"), - ] diff --git a/diffusers/modular_pipelines/flux2/modular_blocks_flux2_klein.py b/diffusers/modular_pipelines/flux2/modular_blocks_flux2_klein.py deleted file mode 100644 index 689cf808c4ba93ac373226e375070bc6ffa5dd64..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux2/modular_blocks_flux2_klein.py +++ /dev/null @@ -1,399 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from ...utils import logging -from ..modular_pipeline import AutoPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import InsertableDict, OutputParam -from .before_denoise import ( - Flux2PrepareImageLatentsStep, - Flux2PrepareLatentsStep, - Flux2RoPEInputsStep, - Flux2SetTimestepsStep, -) -from .decoders import Flux2DecodeStep, Flux2UnpackLatentsStep -from .denoise import Flux2KleinDenoiseStep -from .encoders import ( - Flux2KleinTextEncoderStep, - Flux2VaeEncoderStep, -) -from .inputs import ( - Flux2ProcessImagesInputStep, - Flux2TextInputStep, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - -################ -# VAE encoder -################ - - -# auto_docstring -class Flux2KleinVaeEncoderSequentialStep(SequentialPipelineBlocks): - """ - VAE encoder step that preprocesses and encodes the image inputs into their latent representations. - - Components: - image_processor (`Flux2ImageProcessor`) vae (`AutoencoderKLFlux2`) - - Inputs: - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - - Outputs: - condition_images (`list`): - TODO: Add description. - image_latents (`list`): - List of latent representations for each reference image - """ - - model_name = "flux2-klein" - - block_classes = [Flux2ProcessImagesInputStep(), Flux2VaeEncoderStep()] - block_names = ["preprocess", "encode"] - - @property - def description(self) -> str: - return "VAE encoder step that preprocesses and encodes the image inputs into their latent representations." - - -# auto_docstring -class Flux2KleinAutoVaeEncoderStep(AutoPipelineBlocks): - """ - VAE encoder step that encodes the image inputs into their latent representations. - This is an auto pipeline block that works for image conditioning tasks. - - `Flux2KleinVaeEncoderSequentialStep` is used when `image` is provided. - - If `image` is not provided, step will be skipped. - - Components: - image_processor (`Flux2ImageProcessor`) vae (`AutoencoderKLFlux2`) - - Inputs: - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - - Outputs: - condition_images (`list`): - TODO: Add description. - image_latents (`list`): - List of latent representations for each reference image - """ - - model_name = "flux2-klein" - - block_classes = [Flux2KleinVaeEncoderSequentialStep] - block_names = ["img_conditioning"] - block_trigger_inputs = ["image"] - - @property - def description(self): - return ( - "VAE encoder step that encodes the image inputs into their latent representations.\n" - "This is an auto pipeline block that works for image conditioning tasks.\n" - " - `Flux2KleinVaeEncoderSequentialStep` is used when `image` is provided.\n" - " - If `image` is not provided, step will be skipped." - ) - - -### -### Core denoise -### - -Flux2KleinCoreDenoiseBlocks = InsertableDict( - [ - ("input", Flux2TextInputStep()), - ("prepare_latents", Flux2PrepareLatentsStep()), - ("set_timesteps", Flux2SetTimestepsStep()), - ("prepare_rope_inputs", Flux2RoPEInputsStep()), - ("denoise", Flux2KleinDenoiseStep()), - ("after_denoise", Flux2UnpackLatentsStep()), - ] -) - - -# auto_docstring -class Flux2KleinCoreDenoiseStep(SequentialPipelineBlocks): - """ - Core denoise step that performs the denoising process for Flux2-Klein (distilled model), for text-to-image - generation. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Flux2Transformer2DModel`) - - Inputs: - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - image_latents (`Tensor`, *optional*): - Packed image latents for conditioning. Shape: (B, img_seq_len, C) - image_latent_ids (`Tensor`, *optional*): - Position IDs for image latents. Shape: (B, img_seq_len, 4) - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "flux2-klein" - - block_classes = Flux2KleinCoreDenoiseBlocks.values() - block_names = Flux2KleinCoreDenoiseBlocks.keys() - - @property - def description(self): - return "Core denoise step that performs the denoising process for Flux2-Klein (distilled model), for text-to-image generation." - - @property - def outputs(self): - return [ - OutputParam.template("latents"), - ] - - -Flux2KleinImageConditionedCoreDenoiseBlocks = InsertableDict( - [ - ("input", Flux2TextInputStep()), - ("prepare_image_latents", Flux2PrepareImageLatentsStep()), - ("prepare_latents", Flux2PrepareLatentsStep()), - ("set_timesteps", Flux2SetTimestepsStep()), - ("prepare_rope_inputs", Flux2RoPEInputsStep()), - ("denoise", Flux2KleinDenoiseStep()), - ("after_denoise", Flux2UnpackLatentsStep()), - ] -) - - -# auto_docstring -class Flux2KleinImageConditionedCoreDenoiseStep(SequentialPipelineBlocks): - """ - Core denoise step that performs the denoising process for Flux2-Klein (distilled model) with image conditioning. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Flux2Transformer2DModel`) - - Inputs: - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - image_latents (`list`, *optional*): - TODO: Add description. - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "flux2-klein" - - block_classes = Flux2KleinImageConditionedCoreDenoiseBlocks.values() - block_names = Flux2KleinImageConditionedCoreDenoiseBlocks.keys() - - @property - def description(self): - return "Core denoise step that performs the denoising process for Flux2-Klein (distilled model) with image conditioning." - - @property - def outputs(self): - return [ - OutputParam.template("latents"), - ] - - -# auto_docstring -class Flux2KleinAutoCoreDenoiseStep(AutoPipelineBlocks): - """ - Auto core denoise step that performs the denoising process for Flux2-Klein. - This is an auto pipeline block that works for text-to-image and image-conditioned generation. - - `Flux2KleinCoreDenoiseStep` is used for text-to-image generation. - - `Flux2KleinImageConditionedCoreDenoiseStep` is used for image-conditioned generation. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Flux2Transformer2DModel`) - - Inputs: - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - image_latents (`list`, *optional*): - TODO: Add description. - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_inference_steps (`None`): - TODO: Add description. - timesteps (`None`): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - image_latent_ids (`Tensor`, *optional*): - Position IDs for image latents. Shape: (B, img_seq_len, 4) - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "flux2-klein" - block_classes = [Flux2KleinImageConditionedCoreDenoiseStep, Flux2KleinCoreDenoiseStep] - block_names = ["image_conditioned", "text2image"] - block_trigger_inputs = ["image_latents", None] - - @property - def description(self): - return ( - "Auto core denoise step that performs the denoising process for Flux2-Klein.\n" - "This is an auto pipeline block that works for text-to-image and image-conditioned generation.\n" - " - `Flux2KleinCoreDenoiseStep` is used for text-to-image generation.\n" - " - `Flux2KleinImageConditionedCoreDenoiseStep` is used for image-conditioned generation.\n" - ) - - -### -### Auto blocks -### - - -# auto_docstring -class Flux2KleinAutoBlocks(SequentialPipelineBlocks): - """ - Auto blocks that perform the text-to-image and image-conditioned generation using Flux2-Klein. - - Supported workflows: - - `text2image`: requires `prompt` - - `image_conditioned`: requires `image`, `prompt` - - Components: - text_encoder (`Qwen3ForCausalLM`) tokenizer (`Qwen2Tokenizer`) image_processor (`Flux2ImageProcessor`) vae - (`AutoencoderKLFlux2`) scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Flux2Transformer2DModel`) - - Configs: - is_distilled (default: True) - - Inputs: - prompt (`None`, *optional*): - TODO: Add description. - max_sequence_length (`int`, *optional*, defaults to 512): - TODO: Add description. - text_encoder_out_layers (`tuple`, *optional*, defaults to (9, 18, 27)): - TODO: Add description. - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - image_latents (`list`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`): - TODO: Add description. - num_inference_steps (`None`): - TODO: Add description. - timesteps (`None`): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - image_latent_ids (`Tensor`, *optional*): - Position IDs for image latents. Shape: (B, img_seq_len, 4) - output_type (`None`, *optional*, defaults to pil): - TODO: Add description. - - Outputs: - images (`list`): - Generated images. - """ - - model_name = "flux2-klein" - block_classes = [ - Flux2KleinTextEncoderStep(), - Flux2KleinAutoVaeEncoderStep(), - Flux2KleinAutoCoreDenoiseStep(), - Flux2DecodeStep(), - ] - block_names = ["text_encoder", "vae_encoder", "denoise", "decode"] - _workflow_map = { - "text2image": {"prompt": True}, - "image_conditioned": {"image": True, "prompt": True}, - } - - @property - def description(self): - return "Auto blocks that perform the text-to-image and image-conditioned generation using Flux2-Klein." - - @property - def outputs(self): - return [ - OutputParam.template("images"), - ] diff --git a/diffusers/modular_pipelines/flux2/modular_blocks_flux2_klein_base.py b/diffusers/modular_pipelines/flux2/modular_blocks_flux2_klein_base.py deleted file mode 100644 index f3108bdadeacdb20c3fdbf22ea3ae215018c9c8c..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux2/modular_blocks_flux2_klein_base.py +++ /dev/null @@ -1,413 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from ...utils import logging -from ..modular_pipeline import AutoPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import InsertableDict, OutputParam -from .before_denoise import ( - Flux2KleinBaseRoPEInputsStep, - Flux2PrepareImageLatentsStep, - Flux2PrepareLatentsStep, - Flux2SetTimestepsStep, -) -from .decoders import Flux2DecodeStep, Flux2UnpackLatentsStep -from .denoise import Flux2KleinBaseDenoiseStep -from .encoders import ( - Flux2KleinBaseTextEncoderStep, - Flux2VaeEncoderStep, -) -from .inputs import ( - Flux2KleinBaseTextInputStep, - Flux2ProcessImagesInputStep, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - -################ -# VAE encoder -################ - - -# auto_docstring -class Flux2KleinBaseVaeEncoderSequentialStep(SequentialPipelineBlocks): - """ - VAE encoder step that preprocesses and encodes the image inputs into their latent representations. - - Components: - image_processor (`Flux2ImageProcessor`) vae (`AutoencoderKLFlux2`) - - Inputs: - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - - Outputs: - condition_images (`list`): - TODO: Add description. - image_latents (`list`): - List of latent representations for each reference image - """ - - model_name = "flux2" - - block_classes = [Flux2ProcessImagesInputStep(), Flux2VaeEncoderStep()] - block_names = ["preprocess", "encode"] - - @property - def description(self) -> str: - return "VAE encoder step that preprocesses and encodes the image inputs into their latent representations." - - -# auto_docstring -class Flux2KleinBaseAutoVaeEncoderStep(AutoPipelineBlocks): - """ - VAE encoder step that encodes the image inputs into their latent representations. - This is an auto pipeline block that works for image conditioning tasks. - - `Flux2KleinBaseVaeEncoderSequentialStep` is used when `image` is provided. - - If `image` is not provided, step will be skipped. - - Components: - image_processor (`Flux2ImageProcessor`) vae (`AutoencoderKLFlux2`) - - Inputs: - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - - Outputs: - condition_images (`list`): - TODO: Add description. - image_latents (`list`): - List of latent representations for each reference image - """ - - block_classes = [Flux2KleinBaseVaeEncoderSequentialStep] - block_names = ["img_conditioning"] - block_trigger_inputs = ["image"] - - @property - def description(self): - return ( - "VAE encoder step that encodes the image inputs into their latent representations.\n" - "This is an auto pipeline block that works for image conditioning tasks.\n" - " - `Flux2KleinBaseVaeEncoderSequentialStep` is used when `image` is provided.\n" - " - If `image` is not provided, step will be skipped." - ) - - -### -### Core denoise -### - -Flux2KleinBaseCoreDenoiseBlocks = InsertableDict( - [ - ("input", Flux2KleinBaseTextInputStep()), - ("prepare_latents", Flux2PrepareLatentsStep()), - ("set_timesteps", Flux2SetTimestepsStep()), - ("prepare_rope_inputs", Flux2KleinBaseRoPEInputsStep()), - ("denoise", Flux2KleinBaseDenoiseStep()), - ("after_denoise", Flux2UnpackLatentsStep()), - ] -) - - -# auto_docstring -class Flux2KleinBaseCoreDenoiseStep(SequentialPipelineBlocks): - """ - Core denoise step that performs the denoising process for Flux2-Klein (base model), for text-to-image generation. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Flux2Transformer2DModel`) guider - (`ClassifierFreeGuidance`) - - Configs: - is_distilled (default: False) - - Inputs: - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - Pre-generated negative text embeddings. Can be generated from text_encoder step. - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - image_latents (`Tensor`, *optional*): - Packed image latents for conditioning. Shape: (B, img_seq_len, C) - image_latent_ids (`Tensor`, *optional*): - Position IDs for image latents. Shape: (B, img_seq_len, 4) - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "flux2-klein" - block_classes = Flux2KleinBaseCoreDenoiseBlocks.values() - block_names = Flux2KleinBaseCoreDenoiseBlocks.keys() - - @property - def description(self): - return "Core denoise step that performs the denoising process for Flux2-Klein (base model), for text-to-image generation." - - @property - def outputs(self): - return [ - OutputParam.template("latents"), - ] - - -Flux2KleinBaseImageConditionedCoreDenoiseBlocks = InsertableDict( - [ - ("input", Flux2KleinBaseTextInputStep()), - ("prepare_latents", Flux2PrepareLatentsStep()), - ("prepare_image_latents", Flux2PrepareImageLatentsStep()), - ("set_timesteps", Flux2SetTimestepsStep()), - ("prepare_rope_inputs", Flux2KleinBaseRoPEInputsStep()), - ("denoise", Flux2KleinBaseDenoiseStep()), - ("after_denoise", Flux2UnpackLatentsStep()), - ] -) - - -# auto_docstring -class Flux2KleinBaseImageConditionedCoreDenoiseStep(SequentialPipelineBlocks): - """ - Core denoise step that performs the denoising process for Flux2-Klein (base model) with image conditioning. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Flux2Transformer2DModel`) guider - (`ClassifierFreeGuidance`) - - Configs: - is_distilled (default: False) - - Inputs: - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - Pre-generated negative text embeddings. Can be generated from text_encoder step. - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - image_latents (`list`, *optional*): - TODO: Add description. - num_inference_steps (`None`, *optional*, defaults to 50): - TODO: Add description. - timesteps (`None`, *optional*): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "flux2-klein" - block_classes = Flux2KleinBaseImageConditionedCoreDenoiseBlocks.values() - block_names = Flux2KleinBaseImageConditionedCoreDenoiseBlocks.keys() - - @property - def description(self): - return "Core denoise step that performs the denoising process for Flux2-Klein (base model) with image conditioning." - - @property - def outputs(self): - return [ - OutputParam.template("latents"), - ] - - -# auto_docstring -class Flux2KleinBaseAutoCoreDenoiseStep(AutoPipelineBlocks): - """ - Auto core denoise step that performs the denoising process for Flux2-Klein (base model). - This is an auto pipeline block that works for text-to-image and image-conditioned generation. - - `Flux2KleinBaseCoreDenoiseStep` is used for text-to-image generation. - - `Flux2KleinBaseImageConditionedCoreDenoiseStep` is used for image-conditioned generation. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Flux2Transformer2DModel`) guider - (`ClassifierFreeGuidance`) - - Configs: - is_distilled (default: False) - - Inputs: - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - prompt_embeds (`Tensor`): - Pre-generated text embeddings. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - Pre-generated negative text embeddings. Can be generated from text_encoder step. - height (`int`, *optional*): - TODO: Add description. - width (`int`, *optional*): - TODO: Add description. - latents (`Tensor | NoneType`): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - image_latents (`list`, *optional*): - TODO: Add description. - num_inference_steps (`None`): - TODO: Add description. - timesteps (`None`): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - image_latent_ids (`Tensor`, *optional*): - Position IDs for image latents. Shape: (B, img_seq_len, 4) - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "flux2-klein" - block_classes = [Flux2KleinBaseImageConditionedCoreDenoiseStep, Flux2KleinBaseCoreDenoiseStep] - block_names = ["image_conditioned", "text2image"] - block_trigger_inputs = ["image_latents", None] - - @property - def description(self): - return ( - "Auto core denoise step that performs the denoising process for Flux2-Klein (base model).\n" - "This is an auto pipeline block that works for text-to-image and image-conditioned generation.\n" - " - `Flux2KleinBaseCoreDenoiseStep` is used for text-to-image generation.\n" - " - `Flux2KleinBaseImageConditionedCoreDenoiseStep` is used for image-conditioned generation.\n" - ) - - -### -### Auto blocks -### - - -# auto_docstring -class Flux2KleinBaseAutoBlocks(SequentialPipelineBlocks): - """ - Auto blocks that perform the text-to-image and image-conditioned generation using Flux2-Klein (base model). - - Supported workflows: - - `text2image`: requires `prompt` - - `image_conditioned`: requires `image`, `prompt` - - Components: - text_encoder (`Qwen3ForCausalLM`) tokenizer (`Qwen2Tokenizer`) guider (`ClassifierFreeGuidance`) - image_processor (`Flux2ImageProcessor`) vae (`AutoencoderKLFlux2`) scheduler - (`FlowMatchEulerDiscreteScheduler`) transformer (`Flux2Transformer2DModel`) - - Configs: - is_distilled (default: False) - - Inputs: - prompt (`None`, *optional*): - TODO: Add description. - max_sequence_length (`int`, *optional*, defaults to 512): - TODO: Add description. - text_encoder_out_layers (`tuple`, *optional*, defaults to (9, 18, 27)): - TODO: Add description. - image (`None`, *optional*): - TODO: Add description. - height (`None`, *optional*): - TODO: Add description. - width (`None`, *optional*): - TODO: Add description. - generator (`None`, *optional*): - TODO: Add description. - num_images_per_prompt (`None`, *optional*, defaults to 1): - TODO: Add description. - latents (`Tensor | NoneType`): - TODO: Add description. - image_latents (`list`, *optional*): - TODO: Add description. - num_inference_steps (`None`): - TODO: Add description. - timesteps (`None`): - TODO: Add description. - sigmas (`None`, *optional*): - TODO: Add description. - joint_attention_kwargs (`None`, *optional*): - TODO: Add description. - image_latent_ids (`Tensor`, *optional*): - Position IDs for image latents. Shape: (B, img_seq_len, 4) - output_type (`None`, *optional*, defaults to pil): - TODO: Add description. - - Outputs: - images (`list`): - Generated images. - """ - - model_name = "flux2-klein" - block_classes = [ - Flux2KleinBaseTextEncoderStep(), - Flux2KleinBaseAutoVaeEncoderStep(), - Flux2KleinBaseAutoCoreDenoiseStep(), - Flux2DecodeStep(), - ] - block_names = ["text_encoder", "vae_encoder", "denoise", "decode"] - _workflow_map = { - "text2image": {"prompt": True}, - "image_conditioned": {"image": True, "prompt": True}, - } - - @property - def description(self): - return "Auto blocks that perform the text-to-image and image-conditioned generation using Flux2-Klein (base model)." - - @property - def outputs(self): - return [ - OutputParam.template("images"), - ] diff --git a/diffusers/modular_pipelines/flux2/modular_pipeline.py b/diffusers/modular_pipelines/flux2/modular_pipeline.py deleted file mode 100644 index ed070206c319625dba440dea9baf46037cba91c7..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/flux2/modular_pipeline.py +++ /dev/null @@ -1,99 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from ...loaders import Flux2LoraLoaderMixin -from ...utils import logging -from ..modular_pipeline import ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class Flux2ModularPipeline(ModularPipeline, Flux2LoraLoaderMixin): - """ - A ModularPipeline for Flux2. - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "Flux2AutoBlocks" - - @property - def default_height(self): - return self.default_sample_size * self.vae_scale_factor - - @property - def default_width(self): - return self.default_sample_size * self.vae_scale_factor - - @property - def default_sample_size(self): - return 128 - - @property - def vae_scale_factor(self): - vae_scale_factor = 8 - if getattr(self, "vae", None) is not None: - vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) - return vae_scale_factor - - @property - def num_channels_latents(self): - num_channels_latents = 32 - if getattr(self, "transformer", None): - num_channels_latents = self.transformer.config.in_channels // 4 - return num_channels_latents - - -class Flux2KleinModularPipeline(Flux2ModularPipeline): - """ - A ModularPipeline for Flux2-Klein (distilled model). - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "Flux2KleinAutoBlocks" - - @property - def requires_unconditional_embeds(self): - if hasattr(self.config, "is_distilled") and self.config.is_distilled: - return False - - requires_unconditional_embeds = False - if hasattr(self, "guider") and self.guider is not None: - requires_unconditional_embeds = self.guider._enabled and self.guider.num_conditions > 1 - - return requires_unconditional_embeds - - -class Flux2KleinBaseModularPipeline(Flux2ModularPipeline): - """ - A ModularPipeline for Flux2-Klein (base model). - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "Flux2KleinBaseAutoBlocks" - - @property - def requires_unconditional_embeds(self): - if hasattr(self.config, "is_distilled") and self.config.is_distilled: - return False - - requires_unconditional_embeds = False - if hasattr(self, "guider") and self.guider is not None: - requires_unconditional_embeds = self.guider._enabled and self.guider.num_conditions > 1 - - return requires_unconditional_embeds diff --git a/diffusers/modular_pipelines/helios/__init__.py b/diffusers/modular_pipelines/helios/__init__.py deleted file mode 100644 index 26551399a3e81881ac30ae9d995fe352d3afd3da..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/helios/__init__.py +++ /dev/null @@ -1,59 +0,0 @@ -from typing import TYPE_CHECKING - -from ...utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, -) - - -_dummy_objects = {} -_import_structure = {} - -try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ...utils import dummy_torch_and_transformers_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) -else: - _import_structure["modular_blocks_helios"] = ["HeliosAutoBlocks"] - _import_structure["modular_blocks_helios_pyramid"] = ["HeliosPyramidAutoBlocks"] - _import_structure["modular_blocks_helios_pyramid_distilled"] = ["HeliosPyramidDistilledAutoBlocks"] - _import_structure["modular_pipeline"] = [ - "HeliosModularPipeline", - "HeliosPyramidDistilledModularPipeline", - "HeliosPyramidModularPipeline", - ] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .modular_blocks_helios import HeliosAutoBlocks - from .modular_blocks_helios_pyramid import HeliosPyramidAutoBlocks - from .modular_blocks_helios_pyramid_distilled import HeliosPyramidDistilledAutoBlocks - from .modular_pipeline import ( - HeliosModularPipeline, - HeliosPyramidDistilledModularPipeline, - HeliosPyramidModularPipeline, - ) -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/helios/before_denoise.py b/diffusers/modular_pipelines/helios/before_denoise.py deleted file mode 100644 index 64407db63cca0e0b0e7326e17c32c8613bc886e8..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/helios/before_denoise.py +++ /dev/null @@ -1,836 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import numpy as np -import torch - -from ...models import HeliosTransformer3DModel -from ...schedulers import HeliosScheduler -from ...utils import logging -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import HeliosModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# Copied from diffusers.pipelines.flux.pipeline_flux.calculate_shift -def calculate_shift( - image_seq_len, - base_seq_len: int = 256, - max_seq_len: int = 4096, - base_shift: float = 0.5, - max_shift: float = 1.15, -): - m = (max_shift - base_shift) / (max_seq_len - base_seq_len) - b = base_shift - m * base_seq_len - mu = image_seq_len * m + b - return mu - - -class HeliosTextInputStep(ModularPipelineBlocks): - model_name = "helios" - - @property - def description(self) -> str: - return ( - "Input processing step that:\n" - " 1. Determines `batch_size` and `dtype` based on `prompt_embeds`\n" - " 2. Adjusts input tensor shapes based on `batch_size` (number of prompts) and `num_videos_per_prompt`\n\n" - "All input tensors are expected to have either batch_size=1 or match the batch_size\n" - "of prompt_embeds. The tensors will be duplicated across the batch dimension to\n" - "have a final batch_size of batch_size * num_videos_per_prompt." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "num_videos_per_prompt", - default=1, - type_hint=int, - description="Number of videos to generate per prompt.", - ), - InputParam.template("prompt_embeds"), - InputParam.template("negative_prompt_embeds"), - ] - - @property - def intermediate_outputs(self) -> list[str]: - return [ - OutputParam( - "batch_size", - type_hint=int, - description="Number of prompts, the final batch size of model inputs should be batch_size * num_videos_per_prompt", - ), - OutputParam( - "dtype", - type_hint=torch.dtype, - description="Data type of model tensor inputs (determined by `prompt_embeds.dtype`)", - ), - ] - - def check_inputs(self, components, block_state): - if block_state.prompt_embeds is not None and block_state.negative_prompt_embeds is not None: - if block_state.prompt_embeds.shape != block_state.negative_prompt_embeds.shape: - raise ValueError( - "`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but" - f" got: `prompt_embeds` {block_state.prompt_embeds.shape} != `negative_prompt_embeds`" - f" {block_state.negative_prompt_embeds.shape}." - ) - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - self.check_inputs(components, block_state) - - block_state.batch_size = block_state.prompt_embeds.shape[0] - block_state.dtype = block_state.prompt_embeds.dtype - - _, seq_len, _ = block_state.prompt_embeds.shape - block_state.prompt_embeds = block_state.prompt_embeds.repeat(1, block_state.num_videos_per_prompt, 1) - block_state.prompt_embeds = block_state.prompt_embeds.view( - block_state.batch_size * block_state.num_videos_per_prompt, seq_len, -1 - ) - - if block_state.negative_prompt_embeds is not None: - _, seq_len, _ = block_state.negative_prompt_embeds.shape - block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.repeat( - 1, block_state.num_videos_per_prompt, 1 - ) - block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.view( - block_state.batch_size * block_state.num_videos_per_prompt, seq_len, -1 - ) - - self.set_block_state(state, block_state) - - return components, state - - -# Copied from diffusers.modular_pipelines.wan.before_denoise.repeat_tensor_to_batch_size -def repeat_tensor_to_batch_size( - input_name: str, - input_tensor: torch.Tensor, - batch_size: int, - num_videos_per_prompt: int = 1, -) -> torch.Tensor: - """Repeat tensor elements to match the final batch size. - - This function expands a tensor's batch dimension to match the final batch size (batch_size * num_videos_per_prompt) - by repeating each element along dimension 0. - - The input tensor must have batch size 1 or batch_size. The function will: - - If batch size is 1: repeat each element (batch_size * num_videos_per_prompt) times - - If batch size equals batch_size: repeat each element num_videos_per_prompt times - - Args: - input_name (str): Name of the input tensor (used for error messages) - input_tensor (torch.Tensor): The tensor to repeat. Must have batch size 1 or batch_size. - batch_size (int): The base batch size (number of prompts) - num_videos_per_prompt (int, optional): Number of videos to generate per prompt. Defaults to 1. - - Returns: - torch.Tensor: The repeated tensor with final batch size (batch_size * num_videos_per_prompt) - - Raises: - ValueError: If input_tensor is not a torch.Tensor or has invalid batch size - - Examples: - tensor = torch.tensor([[1, 2, 3]]) # shape: [1, 3] repeated = repeat_tensor_to_batch_size("image", tensor, - batch_size=2, num_videos_per_prompt=2) repeated # tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]]) - shape: - [4, 3] - - tensor = torch.tensor([[1, 2, 3], [4, 5, 6]]) # shape: [2, 3] repeated = repeat_tensor_to_batch_size("image", - tensor, batch_size=2, num_videos_per_prompt=2) repeated # tensor([[1, 2, 3], [1, 2, 3], [4, 5, 6], [4, 5, 6]]) - - shape: [4, 3] - """ - # make sure input is a tensor - if not isinstance(input_tensor, torch.Tensor): - raise ValueError(f"`{input_name}` must be a tensor") - - # make sure input tensor e.g. image_latents has batch size 1 or batch_size same as prompts - if input_tensor.shape[0] == 1: - repeat_by = batch_size * num_videos_per_prompt - elif input_tensor.shape[0] == batch_size: - repeat_by = num_videos_per_prompt - else: - raise ValueError( - f"`{input_name}` must have have batch size 1 or {batch_size}, but got {input_tensor.shape[0]}" - ) - - # expand the tensor to match the batch_size * num_videos_per_prompt - input_tensor = input_tensor.repeat_interleave(repeat_by, dim=0) - - return input_tensor - - -# Copied from diffusers.modular_pipelines.wan.before_denoise.calculate_dimension_from_latents -def calculate_dimension_from_latents( - latents: torch.Tensor, vae_scale_factor_temporal: int, vae_scale_factor_spatial: int -) -> tuple[int, int]: - """Calculate image dimensions from latent tensor dimensions. - - This function converts latent temporal and spatial dimensions to image temporal and spatial dimensions by - multiplying the latent num_frames/height/width by the VAE scale factor. - - Args: - latents (torch.Tensor): The latent tensor. Must have 4 or 5 dimensions. - Expected shapes: [batch, channels, height, width] or [batch, channels, frames, height, width] - vae_scale_factor_temporal (int): The scale factor used by the VAE to compress temporal dimension. - Typically 4 for most VAEs (video is 4x larger than latents in temporal dimension) - vae_scale_factor_spatial (int): The scale factor used by the VAE to compress spatial dimension. - Typically 8 for most VAEs (image is 8x larger than latents in each dimension) - - Returns: - tuple[int, int]: The calculated image dimensions as (height, width) - - Raises: - ValueError: If latents tensor doesn't have 4 or 5 dimensions - - """ - if latents.ndim != 5: - raise ValueError(f"latents must have 5 dimensions, but got {latents.ndim}") - - _, _, num_latent_frames, latent_height, latent_width = latents.shape - - num_frames = (num_latent_frames - 1) * vae_scale_factor_temporal + 1 - height = latent_height * vae_scale_factor_spatial - width = latent_width * vae_scale_factor_spatial - - return num_frames, height, width - - -class HeliosAdditionalInputsStep(ModularPipelineBlocks): - """Configurable step that standardizes inputs for the denoising step. - - This step handles: - 1. For encoded image latents: Computes height/width from latents and expands batch size - 2. For additional_batch_inputs: Expands batch dimensions to match final batch size - """ - - model_name = "helios" - - def __init__( - self, - image_latent_inputs: list[InputParam] | None = None, - additional_batch_inputs: list[InputParam] | None = None, - ): - if image_latent_inputs is None: - image_latent_inputs = [InputParam.template("image_latents")] - if additional_batch_inputs is None: - additional_batch_inputs = [] - - if not isinstance(image_latent_inputs, list): - raise ValueError(f"image_latent_inputs must be a list, but got {type(image_latent_inputs)}") - else: - for input_param in image_latent_inputs: - if not isinstance(input_param, InputParam): - raise ValueError(f"image_latent_inputs must be a list of InputParam, but got {type(input_param)}") - - if not isinstance(additional_batch_inputs, list): - raise ValueError(f"additional_batch_inputs must be a list, but got {type(additional_batch_inputs)}") - else: - for input_param in additional_batch_inputs: - if not isinstance(input_param, InputParam): - raise ValueError( - f"additional_batch_inputs must be a list of InputParam, but got {type(input_param)}" - ) - - self._image_latent_inputs = image_latent_inputs - self._additional_batch_inputs = additional_batch_inputs - super().__init__() - - @property - def description(self) -> str: - summary_section = ( - "Input processing step that:\n" - " 1. For image latent inputs: Computes height/width from latents and expands batch size\n" - " 2. For additional batch inputs: Expands batch dimensions to match final batch size" - ) - - inputs_info = "" - if self._image_latent_inputs or self._additional_batch_inputs: - inputs_info = "\n\nConfigured inputs:" - if self._image_latent_inputs: - inputs_info += f"\n - Image latent inputs: {[p.name for p in self._image_latent_inputs]}" - if self._additional_batch_inputs: - inputs_info += f"\n - Additional batch inputs: {[p.name for p in self._additional_batch_inputs]}" - - placement_section = "\n\nThis block should be placed after the encoder steps and the text input step." - - return summary_section + inputs_info + placement_section - - @property - def inputs(self) -> list[InputParam]: - inputs = [ - InputParam(name="num_videos_per_prompt", default=1), - InputParam(name="batch_size", required=True), - ] - inputs += self._image_latent_inputs + self._additional_batch_inputs - - return inputs - - @property - def intermediate_outputs(self) -> list[OutputParam]: - outputs = [ - OutputParam("height", type_hint=int), - OutputParam("width", type_hint=int), - ] - - for input_param in self._image_latent_inputs: - outputs.append(OutputParam(input_param.name, type_hint=torch.Tensor)) - - for input_param in self._additional_batch_inputs: - outputs.append(OutputParam(input_param.name, type_hint=torch.Tensor)) - - return outputs - - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - for input_param in self._image_latent_inputs: - image_latent_tensor = getattr(block_state, input_param.name) - if image_latent_tensor is None: - continue - - # Calculate height/width from latents - _, height, width = calculate_dimension_from_latents( - image_latent_tensor, components.vae_scale_factor_temporal, components.vae_scale_factor_spatial - ) - block_state.height = height - block_state.width = width - - # Expand batch size - image_latent_tensor = repeat_tensor_to_batch_size( - input_name=input_param.name, - input_tensor=image_latent_tensor, - num_videos_per_prompt=block_state.num_videos_per_prompt, - batch_size=block_state.batch_size, - ) - - setattr(block_state, input_param.name, image_latent_tensor) - - for input_param in self._additional_batch_inputs: - input_tensor = getattr(block_state, input_param.name) - if input_tensor is None: - continue - - input_tensor = repeat_tensor_to_batch_size( - input_name=input_param.name, - input_tensor=input_tensor, - num_videos_per_prompt=block_state.num_videos_per_prompt, - batch_size=block_state.batch_size, - ) - - setattr(block_state, input_param.name, input_tensor) - - self.set_block_state(state, block_state) - return components, state - - -class HeliosAddNoiseToImageLatentsStep(ModularPipelineBlocks): - """Adds noise to image_latents and fake_image_latents for I2V conditioning. - - Applies single-sigma noise to image_latents (using image_noise_sigma range) and single-sigma noise to - fake_image_latents (using video_noise_sigma range). - """ - - model_name = "helios" - - @property - def description(self) -> str: - return ( - "Adds noise to image_latents and fake_image_latents for I2V conditioning. " - "Uses random sigma from configured ranges for each." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("image_latents"), - InputParam( - "fake_image_latents", - required=True, - type_hint=torch.Tensor, - description="Fake image latents used as history seed for I2V generation.", - ), - InputParam( - "image_noise_sigma_min", - default=0.111, - type_hint=float, - description="Minimum sigma for image latent noise.", - ), - InputParam( - "image_noise_sigma_max", - default=0.135, - type_hint=float, - description="Maximum sigma for image latent noise.", - ), - InputParam( - "video_noise_sigma_min", - default=0.111, - type_hint=float, - description="Minimum sigma for video/fake-image latent noise.", - ), - InputParam( - "video_noise_sigma_max", - default=0.135, - type_hint=float, - description="Maximum sigma for video/fake-image latent noise.", - ), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam.template("image_latents"), - OutputParam("fake_image_latents", type_hint=torch.Tensor, description="Noisy fake image latents"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - image_latents = block_state.image_latents - fake_image_latents = block_state.fake_image_latents - - # Add noise to image_latents - image_noise_sigma = ( - torch.rand(1, device=device, generator=block_state.generator) - * (block_state.image_noise_sigma_max - block_state.image_noise_sigma_min) - + block_state.image_noise_sigma_min - ) - image_latents = ( - image_noise_sigma * randn_tensor(image_latents.shape, generator=block_state.generator, device=device) - + (1 - image_noise_sigma) * image_latents - ) - - # Add noise to fake_image_latents - fake_image_noise_sigma = ( - torch.rand(1, device=device, generator=block_state.generator) - * (block_state.video_noise_sigma_max - block_state.video_noise_sigma_min) - + block_state.video_noise_sigma_min - ) - fake_image_latents = ( - fake_image_noise_sigma - * randn_tensor(fake_image_latents.shape, generator=block_state.generator, device=device) - + (1 - fake_image_noise_sigma) * fake_image_latents - ) - - block_state.image_latents = image_latents.to(device=device, dtype=torch.float32) - block_state.fake_image_latents = fake_image_latents.to(device=device, dtype=torch.float32) - - self.set_block_state(state, block_state) - return components, state - - -class HeliosAddNoiseToVideoLatentsStep(ModularPipelineBlocks): - """Adds noise to image_latents and video_latents for V2V conditioning. - - Applies single-sigma noise to image_latents (using image_noise_sigma range) and per-frame noise to video_latents in - chunks (using video_noise_sigma range). - """ - - model_name = "helios" - - @property - def description(self) -> str: - return ( - "Adds noise to image_latents and video_latents for V2V conditioning. " - "Uses single-sigma noise for image_latents and per-frame noise for video chunks." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("image_latents"), - InputParam( - "video_latents", - required=True, - type_hint=torch.Tensor, - description="Encoded video latents for V2V generation.", - ), - InputParam( - "num_latent_frames_per_chunk", - default=9, - type_hint=int, - description="Number of latent frames per temporal chunk.", - ), - InputParam( - "image_noise_sigma_min", - default=0.111, - type_hint=float, - description="Minimum sigma for image latent noise.", - ), - InputParam( - "image_noise_sigma_max", - default=0.135, - type_hint=float, - description="Maximum sigma for image latent noise.", - ), - InputParam( - "video_noise_sigma_min", - default=0.111, - type_hint=float, - description="Minimum sigma for video latent noise.", - ), - InputParam( - "video_noise_sigma_max", - default=0.135, - type_hint=float, - description="Maximum sigma for video latent noise.", - ), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam.template("image_latents"), - OutputParam("video_latents", type_hint=torch.Tensor, description="Noisy video latents"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - image_latents = block_state.image_latents - video_latents = block_state.video_latents - num_latent_frames_per_chunk = block_state.num_latent_frames_per_chunk - - # Add noise to first frame (single sigma) - image_noise_sigma = ( - torch.rand(1, device=device, generator=block_state.generator) - * (block_state.image_noise_sigma_max - block_state.image_noise_sigma_min) - + block_state.image_noise_sigma_min - ) - image_latents = ( - image_noise_sigma * randn_tensor(image_latents.shape, generator=block_state.generator, device=device) - + (1 - image_noise_sigma) * image_latents - ) - - # Add per-frame noise to video chunks - noisy_latents_chunks = [] - num_latent_chunks = video_latents.shape[2] // num_latent_frames_per_chunk - for i in range(num_latent_chunks): - chunk_start = i * num_latent_frames_per_chunk - chunk_end = chunk_start + num_latent_frames_per_chunk - latent_chunk = video_latents[:, :, chunk_start:chunk_end, :, :] - - chunk_frames = latent_chunk.shape[2] - frame_sigmas = ( - torch.rand(chunk_frames, device=device, generator=block_state.generator) - * (block_state.video_noise_sigma_max - block_state.video_noise_sigma_min) - + block_state.video_noise_sigma_min - ) - frame_sigmas = frame_sigmas.view(1, 1, chunk_frames, 1, 1) - - noisy_chunk = ( - frame_sigmas * randn_tensor(latent_chunk.shape, generator=block_state.generator, device=device) - + (1 - frame_sigmas) * latent_chunk - ) - noisy_latents_chunks.append(noisy_chunk) - video_latents = torch.cat(noisy_latents_chunks, dim=2) - - block_state.image_latents = image_latents.to(device=device, dtype=torch.float32) - block_state.video_latents = video_latents.to(device=device, dtype=torch.float32) - - self.set_block_state(state, block_state) - return components, state - - -class HeliosPrepareHistoryStep(ModularPipelineBlocks): - """Prepares chunk/history indices and initializes history state for the chunk loop.""" - - model_name = "helios" - - @property - def description(self) -> str: - return ( - "Prepares the chunk loop by computing latent dimensions, number of chunks, " - "history indices, and initializing history state (history_latents, image_latents, latent_chunks)." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("transformer", HeliosTransformer3DModel), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("height", default=384), - InputParam.template("width", default=640), - InputParam( - "num_frames", default=132, type_hint=int, description="Total number of video frames to generate." - ), - InputParam("batch_size", required=True, type_hint=int), - InputParam( - "num_latent_frames_per_chunk", - default=9, - type_hint=int, - description="Number of latent frames per temporal chunk.", - ), - InputParam( - "history_sizes", - default=[16, 2, 1], - type_hint=list, - description="Sizes of long/mid/short history buffers for temporal context.", - ), - InputParam( - "keep_first_frame", - default=True, - type_hint=bool, - description="Whether to keep the first frame as a prefix in history.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("num_latent_chunk", type_hint=int, description="Number of temporal chunks"), - OutputParam("latent_shape", type_hint=tuple, description="Shape of latent tensor per chunk"), - OutputParam("history_sizes", type_hint=list, description="Adjusted history sizes (sorted, descending)"), - OutputParam("indices_hidden_states", type_hint=torch.Tensor, kwargs_type="denoiser_input_fields"), - OutputParam("indices_latents_history_short", type_hint=torch.Tensor, kwargs_type="denoiser_input_fields"), - OutputParam("indices_latents_history_mid", type_hint=torch.Tensor, kwargs_type="denoiser_input_fields"), - OutputParam("indices_latents_history_long", type_hint=torch.Tensor, kwargs_type="denoiser_input_fields"), - OutputParam("history_latents", type_hint=torch.Tensor, description="Initialized zero history latents"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - batch_size = block_state.batch_size - device = components._execution_device - - block_state.num_frames = max(block_state.num_frames, 1) - history_sizes = sorted(block_state.history_sizes, reverse=True) - - num_channels_latents = components.num_channels_latents - h_latent = block_state.height // components.vae_scale_factor_spatial - w_latent = block_state.width // components.vae_scale_factor_spatial - - # Compute number of chunks - block_state.window_num_frames = ( - block_state.num_latent_frames_per_chunk - 1 - ) * components.vae_scale_factor_temporal + 1 - block_state.num_latent_chunk = max( - 1, (block_state.num_frames + block_state.window_num_frames - 1) // block_state.window_num_frames - ) - - # Modify history_sizes for non-keep_first_frame (matching pipeline behavior) - if not block_state.keep_first_frame: - history_sizes = history_sizes.copy() - history_sizes[-1] = history_sizes[-1] + 1 - - # Compute indices ONCE (same structure for all chunks) - if block_state.keep_first_frame: - indices = torch.arange(0, sum([1, *history_sizes, block_state.num_latent_frames_per_chunk])) - ( - indices_prefix, - indices_latents_history_long, - indices_latents_history_mid, - indices_latents_history_1x, - indices_hidden_states, - ) = indices.split([1, *history_sizes, block_state.num_latent_frames_per_chunk], dim=0) - indices_latents_history_short = torch.cat([indices_prefix, indices_latents_history_1x], dim=0) - else: - indices = torch.arange(0, sum([*history_sizes, block_state.num_latent_frames_per_chunk])) - ( - indices_latents_history_long, - indices_latents_history_mid, - indices_latents_history_short, - indices_hidden_states, - ) = indices.split([*history_sizes, block_state.num_latent_frames_per_chunk], dim=0) - - # Latent shape per chunk - block_state.latent_shape = ( - batch_size, - num_channels_latents, - block_state.num_latent_frames_per_chunk, - h_latent, - w_latent, - ) - - # Set outputs - block_state.history_sizes = history_sizes - block_state.indices_hidden_states = indices_hidden_states.unsqueeze(0) - block_state.indices_latents_history_short = indices_latents_history_short.unsqueeze(0) - block_state.indices_latents_history_mid = indices_latents_history_mid.unsqueeze(0) - block_state.indices_latents_history_long = indices_latents_history_long.unsqueeze(0) - block_state.history_latents = torch.zeros( - batch_size, - num_channels_latents, - sum(history_sizes), - h_latent, - w_latent, - device=device, - dtype=torch.float32, - ) - - self.set_block_state(state, block_state) - - return components, state - - -class HeliosI2VSeedHistoryStep(ModularPipelineBlocks): - """Seeds history_latents with fake_image_latents for I2V pipelines. - - This small additive step runs after HeliosPrepareHistoryStep and appends fake_image_latents to the initialized - history_latents tensor. - """ - - model_name = "helios" - - @property - def description(self) -> str: - return "I2V history seeding: appends fake_image_latents to history_latents." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("history_latents", required=True, type_hint=torch.Tensor), - InputParam("fake_image_latents", required=True, type_hint=torch.Tensor), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "history_latents", type_hint=torch.Tensor, description="History latents seeded with fake_image_latents" - ), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - block_state.history_latents = torch.cat([block_state.history_latents, block_state.fake_image_latents], dim=2) - - self.set_block_state(state, block_state) - return components, state - - -class HeliosV2VSeedHistoryStep(ModularPipelineBlocks): - """Seeds history_latents with video_latents for V2V pipelines. - - This step runs after HeliosPrepareHistoryStep and replaces the tail of history_latents with video_latents. If the - video has fewer frames than the history, the beginning of history is preserved. - """ - - model_name = "helios" - - @property - def description(self) -> str: - return "V2V history seeding: replaces the tail of history_latents with video_latents." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("history_latents", required=True, type_hint=torch.Tensor), - InputParam("video_latents", required=True, type_hint=torch.Tensor), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "history_latents", type_hint=torch.Tensor, description="History latents seeded with video_latents" - ), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - history_latents = block_state.history_latents - video_latents = block_state.video_latents - - history_frames = history_latents.shape[2] - video_frames = video_latents.shape[2] - if video_frames < history_frames: - keep_frames = history_frames - video_frames - history_latents = torch.cat([history_latents[:, :, :keep_frames, :, :], video_latents], dim=2) - else: - history_latents = video_latents - - block_state.history_latents = history_latents - - self.set_block_state(state, block_state) - return components, state - - -class HeliosSetTimestepsStep(ModularPipelineBlocks): - """Computes scheduler parameters (mu, sigmas) for the chunk loop.""" - - model_name = "helios" - - @property - def description(self) -> str: - return "Computes scheduler shift parameter (mu) and default sigmas for the Helios chunk loop." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("transformer", HeliosTransformer3DModel), - ComponentSpec("scheduler", HeliosScheduler), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("latent_shape", required=True, type_hint=tuple), - InputParam.template("num_inference_steps"), - InputParam.template("sigmas"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("mu", type_hint=float, description="Scheduler shift parameter"), - OutputParam("sigmas", type_hint=list, description="Sigma schedule for diffusion"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - patch_size = components.transformer.config.patch_size - latent_shape = block_state.latent_shape - image_seq_len = (latent_shape[-1] * latent_shape[-2] * latent_shape[-3]) // ( - patch_size[0] * patch_size[1] * patch_size[2] - ) - - if block_state.sigmas is None: - block_state.sigmas = np.linspace(0.999, 0.0, block_state.num_inference_steps + 1)[:-1] - - block_state.mu = calculate_shift( - image_seq_len, - components.scheduler.config.get("base_image_seq_len", 256), - components.scheduler.config.get("max_image_seq_len", 4096), - components.scheduler.config.get("base_shift", 0.5), - components.scheduler.config.get("max_shift", 1.15), - ) - - self.set_block_state(state, block_state) - - return components, state diff --git a/diffusers/modular_pipelines/helios/decoders.py b/diffusers/modular_pipelines/helios/decoders.py deleted file mode 100644 index c448d36136e6da7ab2a5a03fc9e7c4867e43759c..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/helios/decoders.py +++ /dev/null @@ -1,112 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import numpy as np -import PIL -import torch - -from ...configuration_utils import FrozenDict -from ...models import AutoencoderKLWan -from ...utils import logging -from ...video_processor import VideoProcessor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class HeliosDecodeStep(ModularPipelineBlocks): - """Decode all chunk latents with VAE, trim frames, and postprocess into final video output.""" - - model_name = "helios" - - @property - def description(self) -> str: - return ( - "Decodes all chunk latents with the VAE, concatenates them, " - "trims to the target frame count, and postprocesses into the final video output." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLWan), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 8}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "latent_chunks", required=True, type_hint=list, description="List of per-chunk denoised latent tensors" - ), - InputParam("num_frames", required=True, type_hint=int, description="The target number of output frames"), - InputParam.template("output_type", default="np"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "videos", - type_hint=list[list[PIL.Image.Image]] | list[torch.Tensor] | list[np.ndarray], - description="The generated videos, can be a PIL.Image.Image, torch.Tensor or a numpy array", - ), - ] - - @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - vae = components.vae - device = components._execution_device - decode_dtype = vae.dtype - - latents_mean = ( - torch.tensor(vae.config.latents_mean).view(1, vae.config.z_dim, 1, 1, 1).to(device, decode_dtype) - ) - latents_std = 1.0 / torch.tensor(vae.config.latents_std).view(1, vae.config.z_dim, 1, 1, 1).to( - device, decode_dtype - ) - - history_video = None - for chunk_latents in block_state.latent_chunks: - current_latents = chunk_latents.to(device=device, dtype=decode_dtype) / latents_std + latents_mean - current_video = vae.decode(current_latents, return_dict=False)[0] - - if history_video is None: - history_video = current_video - else: - history_video = torch.cat([history_video, current_video], dim=2) - - # Trim to proper frame count - generated_frames = history_video.size(2) - generated_frames = ( - generated_frames - 1 - ) // components.vae_scale_factor_temporal * components.vae_scale_factor_temporal + 1 - history_video = history_video[:, :, :generated_frames] - - block_state.videos = components.video_processor.postprocess_video( - history_video, output_type=block_state.output_type - ) - - self.set_block_state(state, block_state) - - return components, state diff --git a/diffusers/modular_pipelines/helios/denoise.py b/diffusers/modular_pipelines/helios/denoise.py deleted file mode 100644 index 5fcf01a73ffc189f2399fd43f8ba3ea163cfd448..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/helios/denoise.py +++ /dev/null @@ -1,1069 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect -import math - -import torch -import torch.nn.functional as F -from tqdm.auto import tqdm - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance, ClassifierFreeZeroStarGuidance -from ...models import HeliosTransformer3DModel -from ...schedulers import HeliosScheduler -from ...utils import logging -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ( - BlockState, - LoopSequentialPipelineBlocks, - ModularPipelineBlocks, - PipelineState, -) -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .before_denoise import calculate_shift -from .modular_pipeline import HeliosModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def sample_block_noise( - batch_size, - channel, - num_frames, - height, - width, - gamma, - patch_size=(1, 2, 2), - device=None, - generator=None, -): - """Generate spatially-correlated block noise for pyramid upsampling correction. - - Uses a multivariate normal distribution with covariance based on `gamma` to produce noise with block structure, - matching the upsampling artifacts that need correction. - """ - # NOTE: A generator must be provided to ensure correct and reproducible results. - # Creating a default generator here is a fallback only — without a fixed seed, - # the output will be non-deterministic and may produce incorrect results in CP context. - if generator is None: - generator = torch.Generator(device=device) - elif isinstance(generator, list): - generator = generator[0] - - _, ph, pw = patch_size - block_size = ph * pw - - cov = ( - torch.eye(block_size, device=device) * (1 + gamma) - torch.ones(block_size, block_size, device=device) * gamma - ) - cov += torch.eye(block_size, device=device) * 1e-8 - cov = cov.float() # Upcast to fp32 for numerical stability — cholesky is unreliable in fp16/bf16. - - L = torch.linalg.cholesky(cov) - block_number = batch_size * channel * num_frames * (height // ph) * (width // pw) - z = torch.randn(block_number, block_size, device=generator.device, generator=generator).to(device) - noise = z @ L.T - - noise = noise.view(batch_size, channel, num_frames, height // ph, width // pw, ph, pw) - noise = noise.permute(0, 1, 2, 3, 5, 4, 6).reshape(batch_size, channel, num_frames, height, width) - return noise - - -# ======================================== -# Chunk Loop Leaf Blocks -# ======================================== - - -class HeliosChunkHistorySliceStep(ModularPipelineBlocks): - """Slices history latents into short/mid/long for a T2V chunk. - - At k==0 with no image_latents, creates a zero prefix. Otherwise uses image_latents (either provided or captured - from first chunk by HeliosChunkUpdateStep). - """ - - model_name = "helios" - - @property - def description(self) -> str: - return ( - "T2V history slice: splits history into long/mid/short. At k==0 with no image_latents, " - "creates a zero prefix; otherwise uses image_latents as prefix for short history." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "keep_first_frame", - default=True, - type_hint=bool, - description="Whether to keep the first frame as a prefix in history.", - ), - InputParam( - "history_sizes", - required=True, - type_hint=list, - description="Sizes of long/mid/short history buffers for temporal context.", - ), - InputParam( - "history_latents", - required=True, - type_hint=torch.Tensor, - description="Accumulated history latents from previous chunks.", - ), - InputParam("latent_shape", required=True, type_hint=tuple), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): - keep_first_frame = block_state.keep_first_frame - history_sizes = block_state.history_sizes - image_latents = block_state.image_latents - device = components._execution_device - - batch_size, num_channels_latents, _, h_latent, w_latent = block_state.latent_shape - - if keep_first_frame: - latents_history_long, latents_history_mid, latents_history_1x = block_state.history_latents[ - :, :, -sum(history_sizes) : - ].split(history_sizes, dim=2) - if image_latents is None and k == 0: - latents_prefix = torch.zeros( - batch_size, - num_channels_latents, - 1, - h_latent, - w_latent, - device=device, - dtype=torch.float32, - ) - else: - latents_prefix = image_latents - latents_history_short = torch.cat([latents_prefix, latents_history_1x], dim=2) - else: - latents_history_long, latents_history_mid, latents_history_short = block_state.history_latents[ - :, :, -sum(history_sizes) : - ].split(history_sizes, dim=2) - - block_state.latents_history_short = latents_history_short - block_state.latents_history_mid = latents_history_mid - block_state.latents_history_long = latents_history_long - - return components, block_state - - -class HeliosI2VChunkHistorySliceStep(ModularPipelineBlocks): - """Slices history latents into short/mid/long for an I2V chunk. - - Always uses image_latents as prefix (assumes history pre-seeded with fake_image_latents). - """ - - model_name = "helios" - - @property - def description(self) -> str: - return ( - "I2V history slice: splits pre-seeded history into long/mid/short, " - "always using image_latents as prefix for short history." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "keep_first_frame", - default=True, - type_hint=bool, - description="Whether to keep the first frame as a prefix in history.", - ), - InputParam( - "history_sizes", - required=True, - type_hint=list, - description="Sizes of long/mid/short history buffers for temporal context.", - ), - InputParam( - "history_latents", - required=True, - type_hint=torch.Tensor, - description="Accumulated history latents from previous chunks.", - ), - InputParam( - "image_latents", - required=True, - type_hint=torch.Tensor, - description="First-frame latents used as prefix for short history.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): - keep_first_frame = block_state.keep_first_frame - history_sizes = block_state.history_sizes - image_latents = block_state.image_latents - - if keep_first_frame: - latents_history_long, latents_history_mid, latents_history_1x = block_state.history_latents[ - :, :, -sum(history_sizes) : - ].split(history_sizes, dim=2) - latents_history_short = torch.cat([image_latents, latents_history_1x], dim=2) - else: - latents_history_long, latents_history_mid, latents_history_short = block_state.history_latents[ - :, :, -sum(history_sizes) : - ].split(history_sizes, dim=2) - - block_state.latents_history_short = latents_history_short - block_state.latents_history_mid = latents_history_mid - block_state.latents_history_long = latents_history_long - - return components, block_state - - -class HeliosChunkNoiseGenStep(ModularPipelineBlocks): - """Generates noise latents for a chunk using randn_tensor.""" - - model_name = "helios" - - @property - def description(self) -> str: - return "Generates random noise latents at full resolution for a single chunk." - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("latent_shape", required=True, type_hint=tuple), - InputParam.template("generator"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): - device = components._execution_device - block_state.latents = randn_tensor( - block_state.latent_shape, generator=block_state.generator, device=device, dtype=torch.float32 - ) - return components, block_state - - -class HeliosPyramidChunkNoiseGenStep(ModularPipelineBlocks): - """Generates noise latents and downsamples to smallest pyramid level.""" - - model_name = "helios-pyramid" - - @property - def description(self) -> str: - return ( - "Generates random noise at full resolution, then downsamples to the smallest " - "pyramid level via bilinear interpolation." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("latent_shape", required=True, type_hint=tuple), - InputParam( - "pyramid_num_inference_steps_list", - default=[10, 10, 10], - type_hint=list, - description="Number of denoising steps per pyramid stage.", - ), - InputParam.template("generator"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): - device = components._execution_device - batch_size, num_channels_latents, num_latent_frames, h_latent, w_latent = block_state.latent_shape - - latents = randn_tensor( - block_state.latent_shape, generator=block_state.generator, device=device, dtype=torch.float32 - ) - - # Downsample to smallest pyramid level - h, w = h_latent, w_latent - latents = latents.permute(0, 2, 1, 3, 4).reshape(batch_size * num_latent_frames, num_channels_latents, h, w) - for _ in range(len(block_state.pyramid_num_inference_steps_list) - 1): - h //= 2 - w //= 2 - latents = F.interpolate(latents, size=(h, w), mode="bilinear") * 2 - block_state.latents = latents.reshape(batch_size, num_latent_frames, num_channels_latents, h, w).permute( - 0, 2, 1, 3, 4 - ) - - return components, block_state - - -class HeliosChunkSchedulerResetStep(ModularPipelineBlocks): - """Resets the scheduler with timesteps for a single chunk.""" - - model_name = "helios" - - @property - def description(self) -> str: - return "Resets the scheduler with the correct timesteps and shift parameter (mu) for this chunk." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", HeliosScheduler), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("mu", required=True, type_hint=float), - InputParam.template("sigmas", required=True), - InputParam.template("num_inference_steps"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): - device = components._execution_device - components.scheduler.set_timesteps( - block_state.num_inference_steps, device=device, sigmas=block_state.sigmas, mu=block_state.mu - ) - block_state.timesteps = components.scheduler.timesteps - - return components, block_state - - -# ======================================== -# Inner Denoising Blocks -# ======================================== - - -class HeliosChunkDenoiseInner(ModularPipelineBlocks): - """Inner timestep loop for denoising a single chunk, using guider for guidance.""" - - model_name = "helios" - - @property - def description(self) -> str: - return ( - "Inner denoising loop that iterates over timesteps for a single chunk. " - "Uses the guider to manage conditional/unconditional forward passes with cache_context, " - "applies guidance, and runs scheduler step." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("transformer", HeliosTransformer3DModel), - ComponentSpec("scheduler", HeliosScheduler), - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 5.0}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents"), - InputParam.template("timesteps"), - InputParam("prompt_embeds", type_hint=torch.Tensor), - InputParam("negative_prompt_embeds", type_hint=torch.Tensor), - InputParam.template("denoiser_input_fields"), - InputParam.template("num_inference_steps"), - InputParam.template("attention_kwargs"), - InputParam.template("generator"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): - latents = block_state.latents - timesteps = block_state.timesteps - num_inference_steps = block_state.num_inference_steps - - transformer_dtype = components.transformer.dtype - num_warmup_steps = len(timesteps) - num_inference_steps * components.scheduler.order - - # Guider inputs: only encoder_hidden_states differs between cond/uncond - guider_inputs = { - "encoder_hidden_states": (block_state.prompt_embeds, block_state.negative_prompt_embeds), - } - - # Build shared kwargs from denoiser_input_fields (excludes guider-managed ones) - transformer_args = set(inspect.signature(components.transformer.forward).parameters.keys()) - shared_kwargs = {} - for field_name, field_value in block_state.denoiser_input_fields.items(): - if field_name in transformer_args and field_name not in guider_inputs: - shared_kwargs[field_name] = field_value - - # Add loop-internal history latents with dtype casting - shared_kwargs["latents_history_short"] = block_state.latents_history_short.to(transformer_dtype) - shared_kwargs["latents_history_mid"] = block_state.latents_history_mid.to(transformer_dtype) - shared_kwargs["latents_history_long"] = block_state.latents_history_long.to(transformer_dtype) - shared_kwargs["attention_kwargs"] = block_state.attention_kwargs - - with tqdm(total=num_inference_steps) as progress_bar: - for i, t in enumerate(timesteps): - timestep = t.expand(latents.shape[0]).to(torch.int64) - latent_model_input = latents.to(transformer_dtype) - - components.guider.set_state(step=i, num_inference_steps=num_inference_steps, timestep=t) - guider_state = components.guider.prepare_inputs(guider_inputs) - - for guider_state_batch in guider_state: - components.guider.prepare_models(components.transformer) - cond_kwargs = {k: getattr(guider_state_batch, k) for k in guider_inputs.keys()} - - context_name = getattr(guider_state_batch, components.guider._identifier_key) - with components.transformer.cache_context(context_name): - guider_state_batch.noise_pred = components.transformer( - hidden_states=latent_model_input, - timestep=timestep, - return_dict=False, - **cond_kwargs, - **shared_kwargs, - )[0] - components.guider.cleanup_models(components.transformer) - - noise_pred = components.guider(guider_state)[0] - - # Scheduler step - latents = components.scheduler.step( - noise_pred, - t, - latents, - generator=block_state.generator, - return_dict=False, - )[0] - - if i == len(timesteps) - 1 or ( - (i + 1) > num_warmup_steps and (i + 1) % components.scheduler.order == 0 - ): - progress_bar.update() - - block_state.latents = latents - return components, block_state - - -class HeliosPyramidChunkDenoiseInner(ModularPipelineBlocks): - """Nested pyramid stage loop with inner timestep denoising. - - For each pyramid stage (small -> full resolution): - 1. Upsample latents + block noise correction (stages > 0) - 2. Compute mu from current resolution, set scheduler timesteps - 3. Run timestep denoising loop (same logic as HeliosChunkDenoiseInner) - """ - - model_name = "helios-pyramid" - - @property - def description(self) -> str: - return ( - "Pyramid denoising inner block: loops over pyramid stages from smallest to full resolution. " - "Each stage upsamples latents (with block noise correction), recomputes scheduler parameters, " - "and runs the timestep denoising loop." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("transformer", HeliosTransformer3DModel), - ComponentSpec("scheduler", HeliosScheduler), - ComponentSpec( - "guider", - ClassifierFreeZeroStarGuidance, - config=FrozenDict({"guidance_scale": 5.0, "zero_init_steps": 2}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents"), - InputParam("prompt_embeds", type_hint=torch.Tensor), - InputParam("negative_prompt_embeds", type_hint=torch.Tensor), - InputParam.template("denoiser_input_fields"), - InputParam( - "pyramid_num_inference_steps_list", - default=[10, 10, 10], - type_hint=list, - description="Number of denoising steps per pyramid stage.", - ), - InputParam.template("attention_kwargs"), - InputParam.template("generator"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): - device = components._execution_device - transformer_dtype = components.transformer.dtype - latents = block_state.latents - pyramid_num_stages = len(block_state.pyramid_num_inference_steps_list) - - # Guider inputs: only encoder_hidden_states differs between cond/uncond - guider_inputs = { - "encoder_hidden_states": (block_state.prompt_embeds, block_state.negative_prompt_embeds), - } - - # Build shared kwargs from denoiser_input_fields (excludes guider-managed ones) - transformer_args = set(inspect.signature(components.transformer.forward).parameters.keys()) - shared_kwargs = {} - for field_name, field_value in block_state.denoiser_input_fields.items(): - if field_name in transformer_args and field_name not in guider_inputs: - shared_kwargs[field_name] = field_value - - # Add loop-internal history latents with dtype casting - shared_kwargs["latents_history_short"] = block_state.latents_history_short.to(transformer_dtype) - shared_kwargs["latents_history_mid"] = block_state.latents_history_mid.to(transformer_dtype) - shared_kwargs["latents_history_long"] = block_state.latents_history_long.to(transformer_dtype) - shared_kwargs["attention_kwargs"] = block_state.attention_kwargs - - # Save original zero_init_steps if the guider supports it (e.g. ClassifierFreeZeroStarGuidance). - # Helios only applies zero init in pyramid stage 0 (lowest resolution), so we disable it - # for subsequent stages by temporarily setting zero_init_steps=0. - orig_zero_init_steps = getattr(components.guider, "zero_init_steps", None) - - for i_s in range(pyramid_num_stages): - # --- Stage setup --- - - # Disable zero init for stages > 0 (only stage 0 should have zero init) - if orig_zero_init_steps is not None and i_s > 0: - components.guider.zero_init_steps = 0 - - # a. Compute mu from current resolution (before upsample, matching standard pipeline) - patch_size = components.transformer.config.patch_size - image_seq_len = (latents.shape[-1] * latents.shape[-2] * latents.shape[-3]) // ( - patch_size[0] * patch_size[1] * patch_size[2] - ) - mu = calculate_shift( - image_seq_len, - components.scheduler.config.get("base_image_seq_len", 256), - components.scheduler.config.get("max_image_seq_len", 4096), - components.scheduler.config.get("base_shift", 0.5), - components.scheduler.config.get("max_shift", 1.15), - ) - - # b. Set scheduler timesteps for this stage - num_inference_steps = block_state.pyramid_num_inference_steps_list[i_s] - components.scheduler.set_timesteps( - num_inference_steps, - i_s, - device=device, - mu=mu, - ) - timesteps = components.scheduler.timesteps - - # c. Upsample + block noise correction for stages > 0 - if i_s > 0: - batch_size, num_channels_latents, num_frames, current_h, current_w = latents.shape - new_h = current_h * 2 - new_w = current_w * 2 - - latents = latents.permute(0, 2, 1, 3, 4).reshape( - batch_size * num_frames, num_channels_latents, current_h, current_w - ) - latents = F.interpolate(latents, size=(new_h, new_w), mode="nearest") - latents = latents.reshape(batch_size, num_frames, num_channels_latents, new_h, new_w).permute( - 0, 2, 1, 3, 4 - ) - - # Block noise correction - ori_sigma = 1 - components.scheduler.ori_start_sigmas[i_s] - gamma = components.scheduler.config.gamma - alpha = 1 / (math.sqrt(1 + (1 / gamma)) * (1 - ori_sigma) + ori_sigma) - beta = alpha * (1 - ori_sigma) / math.sqrt(gamma) - - batch_size, num_channels_latents, num_frames, h, w = latents.shape - noise = sample_block_noise( - batch_size, - num_channels_latents, - num_frames, - h, - w, - gamma, - patch_size, - device=device, - generator=block_state.generator, - ) - noise = noise.to(dtype=transformer_dtype) - latents = alpha * latents + beta * noise - - # --- Timestep denoising loop --- - num_warmup_steps = len(timesteps) - num_inference_steps * components.scheduler.order - - with tqdm(total=num_inference_steps) as progress_bar: - for i, t in enumerate(timesteps): - timestep = t.expand(latents.shape[0]).to(torch.int64) - latent_model_input = latents.to(transformer_dtype) - - components.guider.set_state(step=i, num_inference_steps=num_inference_steps, timestep=t) - guider_state = components.guider.prepare_inputs(guider_inputs) - - for guider_state_batch in guider_state: - components.guider.prepare_models(components.transformer) - cond_kwargs = {kk: getattr(guider_state_batch, kk) for kk in guider_inputs.keys()} - - context_name = getattr(guider_state_batch, components.guider._identifier_key) - with components.transformer.cache_context(context_name): - guider_state_batch.noise_pred = components.transformer( - hidden_states=latent_model_input, - timestep=timestep, - return_dict=False, - **cond_kwargs, - **shared_kwargs, - )[0] - components.guider.cleanup_models(components.transformer) - - noise_pred = components.guider(guider_state)[0] - - # Scheduler step - latents = components.scheduler.step( - noise_pred, - t, - latents, - generator=block_state.generator, - return_dict=False, - )[0] - - if i == len(timesteps) - 1 or ( - (i + 1) > num_warmup_steps and (i + 1) % components.scheduler.order == 0 - ): - progress_bar.update() - - # Restore original zero_init_steps - if orig_zero_init_steps is not None: - components.guider.zero_init_steps = orig_zero_init_steps - - block_state.latents = latents - return components, block_state - - -# ======================================== -# Post-Denoise Update -# ======================================== - - -class HeliosChunkUpdateStep(ModularPipelineBlocks): - """Updates chunk collection and history after denoising a single chunk.""" - - model_name = "helios" - - @property - def description(self) -> str: - return ( - "Post-denoising update step: appends the denoised latents to the chunk list, " - "captures image_latents from the first chunk if needed, and extends history_latents." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("latents", type_hint=torch.Tensor), - InputParam("history_latents", type_hint=torch.Tensor), - InputParam("keep_first_frame", default=True, type_hint=bool), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): - # e. Collect denoised latents for this chunk - block_state.latent_chunks.append(block_state.latents) - - # f. Update history - if block_state.keep_first_frame and k == 0 and block_state.image_latents is None: - block_state.image_latents = block_state.latents[:, :, 0:1, :, :] - - block_state.history_latents = torch.cat([block_state.history_latents, block_state.latents], dim=2) - - return components, block_state - - -# ======================================== -# Chunk Loop Wrapper -# ======================================== - - -class HeliosChunkLoopWrapper(LoopSequentialPipelineBlocks): - """Outer chunk loop that iterates over temporal chunks. - - History indices, scheduler params, and history state are prepared by HeliosPrepareHistoryStep and - HeliosSetTimestepsStep before this block runs. Sub-blocks handle per-chunk preparation, denoising, and history - updates. - """ - - model_name = "helios" - - @property - def description(self) -> str: - return ( - "Pipeline block that iterates over temporal chunks for progressive video generation. " - "At each chunk iteration, it runs sub-blocks for preparation, denoising, and history updates." - ) - - @property - def loop_inputs(self) -> list[InputParam]: - return [ - InputParam("num_latent_chunk", required=True, type_hint=int), - ] - - @property - def loop_intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("latent_chunks", type_hint=list, description="List of per-chunk denoised latent tensors"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - block_state.latent_chunks = [] - - if not hasattr(block_state, "image_latents"): - block_state.image_latents = None - - for k in range(block_state.num_latent_chunk): - components, block_state = self.loop_step(components, block_state, k=k) - - self.set_block_state(state, block_state) - - return components, state - - -# ======================================== -# Composed Chunk Denoise Steps -# ======================================== - - -class HeliosChunkDenoiseStep(HeliosChunkLoopWrapper): - """T2V chunk-based denoising: history slice -> noise gen -> scheduler reset -> denoise -> update.""" - - block_classes = [ - HeliosChunkHistorySliceStep, - HeliosChunkNoiseGenStep, - HeliosChunkSchedulerResetStep, - HeliosChunkDenoiseInner, - HeliosChunkUpdateStep, - ] - block_names = ["history_slice", "noise_gen", "scheduler_reset", "denoise_inner", "update_chunk"] - - @property - def description(self) -> str: - return ( - "T2V chunk denoise step that iterates over temporal chunks.\n" - "At each chunk: history_slice -> noise_gen -> scheduler_reset -> denoise_inner -> update_chunk." - ) - - -class HeliosI2VChunkDenoiseStep(HeliosChunkLoopWrapper): - """I2V chunk-based denoising: I2V history slice -> noise gen -> scheduler reset -> denoise -> update.""" - - block_classes = [ - HeliosI2VChunkHistorySliceStep, - HeliosChunkNoiseGenStep, - HeliosChunkSchedulerResetStep, - HeliosChunkDenoiseInner, - HeliosChunkUpdateStep, - ] - block_names = ["history_slice", "noise_gen", "scheduler_reset", "denoise_inner", "update_chunk"] - - @property - def description(self) -> str: - return ( - "I2V chunk denoise step that iterates over temporal chunks.\n" - "At each chunk: history_slice (I2V) -> noise_gen -> scheduler_reset -> denoise_inner -> update_chunk." - ) - - -class HeliosPyramidDistilledChunkDenoiseInner(ModularPipelineBlocks): - """Nested pyramid stage loop with DMD denoising for distilled checkpoints. - - Same progressive multi-resolution strategy as HeliosPyramidChunkDenoiseInner, but: - - Guidance is disabled (guidance_scale=1.0, no unconditional pass) - - Supports is_amplify_first_chunk (doubles first chunk's timesteps via scheduler) - - Tracks start_point_list and passes DMD-specific args to scheduler.step() - """ - - model_name = "helios-pyramid" - - @property - def description(self) -> str: - return ( - "Distilled pyramid denoising inner block for DMD checkpoints. Loops over pyramid stages " - "from smallest to full resolution with guidance disabled and DMD scheduler support." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("transformer", HeliosTransformer3DModel), - ComponentSpec("scheduler", HeliosScheduler), - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 1.0}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents"), - InputParam("prompt_embeds", type_hint=torch.Tensor), - InputParam("negative_prompt_embeds", type_hint=torch.Tensor), - InputParam.template("denoiser_input_fields"), - InputParam( - "pyramid_num_inference_steps_list", - default=[2, 2, 2], - type_hint=list, - description="Number of denoising steps per pyramid stage.", - ), - InputParam( - "is_amplify_first_chunk", - default=True, - type_hint=bool, - description="Whether to double the first chunk's timesteps via the scheduler for amplified generation.", - ), - InputParam.template("attention_kwargs"), - InputParam.template("generator"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): - device = components._execution_device - transformer_dtype = components.transformer.dtype - latents = block_state.latents - pyramid_num_stages = len(block_state.pyramid_num_inference_steps_list) - is_first_chunk = k == 0 - - # Track start points for DMD scheduler - start_point_list = [latents] - - # Guider inputs: only encoder_hidden_states differs between cond/uncond - guider_inputs = { - "encoder_hidden_states": (block_state.prompt_embeds, block_state.negative_prompt_embeds), - } - - # Build shared kwargs from denoiser_input_fields (excludes guider-managed ones) - transformer_args = set(inspect.signature(components.transformer.forward).parameters.keys()) - shared_kwargs = {} - for field_name, field_value in block_state.denoiser_input_fields.items(): - if field_name in transformer_args and field_name not in guider_inputs: - shared_kwargs[field_name] = field_value - - # Add loop-internal history latents with dtype casting - shared_kwargs["latents_history_short"] = block_state.latents_history_short.to(transformer_dtype) - shared_kwargs["latents_history_mid"] = block_state.latents_history_mid.to(transformer_dtype) - shared_kwargs["latents_history_long"] = block_state.latents_history_long.to(transformer_dtype) - shared_kwargs["attention_kwargs"] = block_state.attention_kwargs - - for i_s in range(pyramid_num_stages): - # --- Stage setup --- - patch_size = components.transformer.config.patch_size - - # a. Compute mu from current resolution (before upsample, matching standard pipeline) - image_seq_len = (latents.shape[-1] * latents.shape[-2] * latents.shape[-3]) // ( - patch_size[0] * patch_size[1] * patch_size[2] - ) - mu = calculate_shift( - image_seq_len, - components.scheduler.config.get("base_image_seq_len", 256), - components.scheduler.config.get("max_image_seq_len", 4096), - components.scheduler.config.get("base_shift", 0.5), - components.scheduler.config.get("max_shift", 1.15), - ) - - # b. Set scheduler timesteps for this stage (with DMD amplification) - num_inference_steps = block_state.pyramid_num_inference_steps_list[i_s] - components.scheduler.set_timesteps( - num_inference_steps, - i_s, - device=device, - mu=mu, - is_amplify_first_chunk=block_state.is_amplify_first_chunk and is_first_chunk, - ) - timesteps = components.scheduler.timesteps - - # c. Upsample + block noise correction for stages > 0 - if i_s > 0: - batch_size, num_channels_latents, num_frames, current_h, current_w = latents.shape - new_h = current_h * 2 - new_w = current_w * 2 - - latents = latents.permute(0, 2, 1, 3, 4).reshape( - batch_size * num_frames, num_channels_latents, current_h, current_w - ) - latents = F.interpolate(latents, size=(new_h, new_w), mode="nearest") - latents = latents.reshape(batch_size, num_frames, num_channels_latents, new_h, new_w).permute( - 0, 2, 1, 3, 4 - ) - - # Block noise correction - ori_sigma = 1 - components.scheduler.ori_start_sigmas[i_s] - gamma = components.scheduler.config.gamma - alpha = 1 / (math.sqrt(1 + (1 / gamma)) * (1 - ori_sigma) + ori_sigma) - beta = alpha * (1 - ori_sigma) / math.sqrt(gamma) - - batch_size, num_channels_latents, num_frames, h, w = latents.shape - noise = sample_block_noise( - batch_size, - num_channels_latents, - num_frames, - h, - w, - gamma, - patch_size, - device=device, - generator=block_state.generator, - ) - noise = noise.to(dtype=transformer_dtype) - latents = alpha * latents + beta * noise - - start_point_list.append(latents) - - # --- Timestep denoising loop --- - num_warmup_steps = len(timesteps) - num_inference_steps * components.scheduler.order - - with tqdm(total=num_inference_steps) as progress_bar: - for i, t in enumerate(timesteps): - timestep = t.expand(latents.shape[0]).to(torch.int64) - latent_model_input = latents.to(transformer_dtype) - - components.guider.set_state(step=i, num_inference_steps=num_inference_steps, timestep=t) - guider_state = components.guider.prepare_inputs(guider_inputs) - - for guider_state_batch in guider_state: - components.guider.prepare_models(components.transformer) - cond_kwargs = {k: getattr(guider_state_batch, k) for k in guider_inputs.keys()} - - context_name = getattr(guider_state_batch, components.guider._identifier_key) - with components.transformer.cache_context(context_name): - guider_state_batch.noise_pred = components.transformer( - hidden_states=latent_model_input, - timestep=timestep, - return_dict=False, - **cond_kwargs, - **shared_kwargs, - )[0] - components.guider.cleanup_models(components.transformer) - - noise_pred = components.guider(guider_state)[0] - - # Scheduler step with DMD args - latents = components.scheduler.step( - noise_pred, - t, - latents, - generator=block_state.generator, - return_dict=False, - cur_sampling_step=i, - dmd_noisy_tensor=start_point_list[i_s], - dmd_sigmas=components.scheduler.sigmas, - dmd_timesteps=components.scheduler.timesteps, - all_timesteps=timesteps, - )[0] - - if i == len(timesteps) - 1 or ( - (i + 1) > num_warmup_steps and (i + 1) % components.scheduler.order == 0 - ): - progress_bar.update() - - block_state.latents = latents - return components, block_state - - -class HeliosPyramidChunkDenoiseStep(HeliosChunkLoopWrapper): - """T2V pyramid chunk denoising: history slice -> pyramid noise gen -> pyramid denoise inner -> update.""" - - block_classes = [ - HeliosChunkHistorySliceStep, - HeliosPyramidChunkNoiseGenStep, - HeliosPyramidChunkDenoiseInner, - HeliosChunkUpdateStep, - ] - block_names = ["history_slice", "noise_gen", "denoise_inner", "update_chunk"] - - @property - def description(self) -> str: - return ( - "T2V pyramid chunk denoise step that iterates over temporal chunks.\n" - "At each chunk: history_slice -> noise_gen (pyramid) -> denoise_inner (pyramid stages) -> update_chunk.\n" - "Denoising starts at the smallest resolution and progressively upsamples." - ) - - -class HeliosPyramidI2VChunkDenoiseStep(HeliosChunkLoopWrapper): - """I2V pyramid chunk denoising: I2V history slice -> pyramid noise gen -> pyramid denoise inner -> update.""" - - block_classes = [ - HeliosI2VChunkHistorySliceStep, - HeliosPyramidChunkNoiseGenStep, - HeliosPyramidChunkDenoiseInner, - HeliosChunkUpdateStep, - ] - block_names = ["history_slice", "noise_gen", "denoise_inner", "update_chunk"] - - @property - def description(self) -> str: - return ( - "I2V pyramid chunk denoise step that iterates over temporal chunks.\n" - "At each chunk: history_slice (I2V) -> noise_gen (pyramid) -> denoise_inner (pyramid stages) -> update_chunk.\n" - "Denoising starts at the smallest resolution and progressively upsamples." - ) - - -class HeliosPyramidDistilledChunkDenoiseStep(HeliosChunkLoopWrapper): - """T2V distilled pyramid chunk denoising with DMD scheduler and no CFG.""" - - block_classes = [ - HeliosChunkHistorySliceStep, - HeliosPyramidChunkNoiseGenStep, - HeliosPyramidDistilledChunkDenoiseInner, - HeliosChunkUpdateStep, - ] - block_names = ["history_slice", "noise_gen", "denoise_inner", "update_chunk"] - - @property - def description(self) -> str: - return ( - "T2V distilled pyramid chunk denoise step with DMD scheduler.\n" - "At each chunk: history_slice -> noise_gen (pyramid) -> denoise_inner (distilled/DMD) -> update_chunk." - ) - - -class HeliosPyramidDistilledI2VChunkDenoiseStep(HeliosChunkLoopWrapper): - """I2V distilled pyramid chunk denoising with DMD scheduler and no CFG.""" - - block_classes = [ - HeliosI2VChunkHistorySliceStep, - HeliosPyramidChunkNoiseGenStep, - HeliosPyramidDistilledChunkDenoiseInner, - HeliosChunkUpdateStep, - ] - block_names = ["history_slice", "noise_gen", "denoise_inner", "update_chunk"] - - @property - def description(self) -> str: - return ( - "I2V distilled pyramid chunk denoise step with DMD scheduler.\n" - "At each chunk: history_slice (I2V) -> noise_gen (pyramid) -> denoise_inner (distilled/DMD) -> update_chunk." - ) diff --git a/diffusers/modular_pipelines/helios/encoders.py b/diffusers/modular_pipelines/helios/encoders.py deleted file mode 100644 index ce11f1b5876297bff1f35b36f1791e299041971a..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/helios/encoders.py +++ /dev/null @@ -1,392 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import html - -import regex as re -import torch -from transformers import AutoTokenizer, UMT5EncoderModel - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...models import AutoencoderKLWan -from ...utils import is_ftfy_available, logging -from ...video_processor import VideoProcessor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import HeliosModularPipeline - - -if is_ftfy_available(): - import ftfy - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def basic_clean(text): - text = ftfy.fix_text(text) - text = html.unescape(html.unescape(text)) - return text.strip() - - -def whitespace_clean(text): - text = re.sub(r"\s+", " ", text) - text = text.strip() - return text - - -def prompt_clean(text): - text = whitespace_clean(basic_clean(text)) - return text - - -def get_t5_prompt_embeds( - text_encoder: UMT5EncoderModel, - tokenizer: AutoTokenizer, - prompt: str | list[str], - max_sequence_length: int, - device: torch.device, - dtype: torch.dtype | None = None, -): - """Encode text prompts into T5 embeddings for Helios. - - Args: - text_encoder: The T5 text encoder model. - tokenizer: The tokenizer for the text encoder. - prompt: The prompt or prompts to encode. - max_sequence_length: Maximum sequence length for tokenization. - device: Device to place tensors on. - dtype: Optional dtype override. Defaults to `text_encoder.dtype`. - - Returns: - A tuple of `(prompt_embeds, attention_mask)` where `prompt_embeds` is the encoded text embeddings and - `attention_mask` is a boolean mask. - """ - dtype = dtype or text_encoder.dtype - - prompt = [prompt] if isinstance(prompt, str) else prompt - prompt = [prompt_clean(u) for u in prompt] - - text_inputs = tokenizer( - prompt, - padding="max_length", - max_length=max_sequence_length, - truncation=True, - add_special_tokens=True, - return_attention_mask=True, - return_tensors="pt", - ) - text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask - seq_lens = mask.gt(0).sum(dim=1).long() - - prompt_embeds = text_encoder(text_input_ids.to(device), mask.to(device)).last_hidden_state - prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) - prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)] - prompt_embeds = torch.stack( - [torch.cat([u, u.new_zeros(max_sequence_length - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0 - ) - - return prompt_embeds, text_inputs.attention_mask.bool() - - -class HeliosTextEncoderStep(ModularPipelineBlocks): - model_name = "helios" - - @property - def description(self) -> str: - return "Text Encoder step that generates text embeddings to guide the video generation" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", UMT5EncoderModel), - ComponentSpec("tokenizer", AutoTokenizer), - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 5.0}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt"), - InputParam.template("negative_prompt"), - InputParam.template("max_sequence_length"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam.template("prompt_embeds"), - OutputParam.template("negative_prompt_embeds"), - ] - - @staticmethod - def check_inputs(prompt, negative_prompt): - if prompt is not None and not isinstance(prompt, (str, list)): - raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}") - - if negative_prompt is not None and not isinstance(negative_prompt, (str, list)): - raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}") - - if prompt is not None and negative_prompt is not None: - prompt_list = [prompt] if isinstance(prompt, str) else prompt - neg_list = [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt - if type(prompt_list) is not type(neg_list): - raise TypeError( - f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" - f" {type(prompt)}." - ) - if len(prompt_list) != len(neg_list): - raise ValueError( - f"`negative_prompt` has batch size {len(neg_list)}, but `prompt` has batch size" - f" {len(prompt_list)}. Please make sure that passed `negative_prompt` matches" - " the batch size of `prompt`." - ) - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - prompt = block_state.prompt - negative_prompt = block_state.negative_prompt - max_sequence_length = block_state.max_sequence_length - device = components._execution_device - - self.check_inputs(prompt, negative_prompt) - - # Encode prompt - block_state.prompt_embeds, _ = get_t5_prompt_embeds( - text_encoder=components.text_encoder, - tokenizer=components.tokenizer, - prompt=prompt, - max_sequence_length=max_sequence_length, - device=device, - ) - - # Encode negative prompt - block_state.negative_prompt_embeds = None - if components.requires_unconditional_embeds: - negative_prompt = negative_prompt or "" - if isinstance(prompt, list) and isinstance(negative_prompt, str): - negative_prompt = len(prompt) * [negative_prompt] - - block_state.negative_prompt_embeds, _ = get_t5_prompt_embeds( - text_encoder=components.text_encoder, - tokenizer=components.tokenizer, - prompt=negative_prompt, - max_sequence_length=max_sequence_length, - device=device, - ) - - self.set_block_state(state, block_state) - return components, state - - -class HeliosImageVaeEncoderStep(ModularPipelineBlocks): - """Encodes an input image into VAE latent space for image-to-video generation.""" - - model_name = "helios" - - @property - def description(self) -> str: - return ( - "Image Encoder step that encodes an input image into VAE latent space, " - "producing image_latents (first frame prefix) and fake_image_latents (history seed) " - "for image-to-video generation." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLWan), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 8}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("image"), - InputParam.template("height", default=384), - InputParam.template("width", default=640), - InputParam( - "num_latent_frames_per_chunk", - default=9, - type_hint=int, - description="Number of latent frames per temporal chunk.", - ), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam.template("image_latents"), - OutputParam( - "fake_image_latents", type_hint=torch.Tensor, description="Fake image latents for history seeding" - ), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - vae = components.vae - device = components._execution_device - - latents_mean = ( - torch.tensor(vae.config.latents_mean).view(1, vae.config.z_dim, 1, 1, 1).to(vae.device, vae.dtype) - ) - latents_std = 1.0 / torch.tensor(vae.config.latents_std).view(1, vae.config.z_dim, 1, 1, 1).to( - vae.device, vae.dtype - ) - - # Preprocess image to 4D tensor (B, C, H, W) - image = components.video_processor.preprocess( - block_state.image, height=block_state.height, width=block_state.width - ) - image_5d = image.unsqueeze(2).to(device=device, dtype=vae.dtype) # (B, C, 1, H, W) - - # Encode image to get image_latents - image_latents = vae.encode(image_5d).latent_dist.sample(generator=block_state.generator) - image_latents = (image_latents - latents_mean) * latents_std - - # Encode fake video to get fake_image_latents - min_frames = (block_state.num_latent_frames_per_chunk - 1) * components.vae_scale_factor_temporal + 1 - fake_video = image_5d.repeat(1, 1, min_frames, 1, 1) # (B, C, min_frames, H, W) - fake_latents_full = vae.encode(fake_video).latent_dist.sample(generator=block_state.generator) - fake_latents_full = (fake_latents_full - latents_mean) * latents_std - fake_image_latents = fake_latents_full[:, :, -1:, :, :] - - block_state.image_latents = image_latents.to(device=device, dtype=torch.float32) - block_state.fake_image_latents = fake_image_latents.to(device=device, dtype=torch.float32) - - self.set_block_state(state, block_state) - return components, state - - -class HeliosVideoVaeEncoderStep(ModularPipelineBlocks): - """Encodes an input video into VAE latent space for video-to-video generation. - - Produces `image_latents` (first frame) and `video_latents` (remaining frames encoded in chunks). - """ - - model_name = "helios" - - @property - def description(self) -> str: - return ( - "Video Encoder step that encodes an input video into VAE latent space, " - "producing image_latents (first frame) and video_latents (chunked video frames) " - "for video-to-video generation." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLWan), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 8}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("video", required=True, description="Input video for video-to-video generation"), - InputParam.template("height", default=384), - InputParam.template("width", default=640), - InputParam( - "num_latent_frames_per_chunk", - default=9, - type_hint=int, - description="Number of latent frames per temporal chunk.", - ), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam.template("image_latents"), - OutputParam("video_latents", type_hint=torch.Tensor, description="Encoded video latents (chunked)"), - ] - - @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - vae = components.vae - device = components._execution_device - num_latent_frames_per_chunk = block_state.num_latent_frames_per_chunk - - latents_mean = ( - torch.tensor(vae.config.latents_mean).view(1, vae.config.z_dim, 1, 1, 1).to(vae.device, vae.dtype) - ) - latents_std = 1.0 / torch.tensor(vae.config.latents_std).view(1, vae.config.z_dim, 1, 1, 1).to( - vae.device, vae.dtype - ) - - # Preprocess video - video = components.video_processor.preprocess_video( - block_state.video, height=block_state.height, width=block_state.width - ) - video = video.to(device=device, dtype=vae.dtype) - - # Encode video into latents - num_frames = video.shape[2] - min_frames = (num_latent_frames_per_chunk - 1) * 4 + 1 - num_chunks = num_frames // min_frames - if num_chunks == 0: - raise ValueError( - f"Video must have at least {min_frames} frames " - f"(got {num_frames} frames). " - f"Required: (num_latent_frames_per_chunk - 1) * 4 + 1 = ({num_latent_frames_per_chunk} - 1) * 4 + 1 = {min_frames}" - ) - total_valid_frames = num_chunks * min_frames - start_frame = num_frames - total_valid_frames - - # Encode first frame - first_frame = video[:, :, 0:1, :, :] - image_latents = vae.encode(first_frame).latent_dist.sample(generator=block_state.generator) - image_latents = (image_latents - latents_mean) * latents_std - - # Encode remaining frames in chunks - latents_chunks = [] - for i in range(num_chunks): - chunk_start = start_frame + i * min_frames - chunk_end = chunk_start + min_frames - video_chunk = video[:, :, chunk_start:chunk_end, :, :] - chunk_latents = vae.encode(video_chunk).latent_dist.sample(generator=block_state.generator) - chunk_latents = (chunk_latents - latents_mean) * latents_std - latents_chunks.append(chunk_latents) - video_latents = torch.cat(latents_chunks, dim=2) - - block_state.image_latents = image_latents.to(device=device, dtype=torch.float32) - block_state.video_latents = video_latents.to(device=device, dtype=torch.float32) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/helios/modular_blocks_helios.py b/diffusers/modular_pipelines/helios/modular_blocks_helios.py deleted file mode 100644 index c3d5cb4efc774a8fb911e57882bf021181146ae9..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/helios/modular_blocks_helios.py +++ /dev/null @@ -1,542 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ...utils import logging -from ..modular_pipeline import AutoPipelineBlocks, ConditionalPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import InputParam, InsertableDict, OutputParam -from .before_denoise import ( - HeliosAdditionalInputsStep, - HeliosAddNoiseToImageLatentsStep, - HeliosAddNoiseToVideoLatentsStep, - HeliosI2VSeedHistoryStep, - HeliosPrepareHistoryStep, - HeliosSetTimestepsStep, - HeliosTextInputStep, - HeliosV2VSeedHistoryStep, -) -from .decoders import HeliosDecodeStep -from .denoise import HeliosChunkDenoiseStep, HeliosI2VChunkDenoiseStep -from .encoders import HeliosImageVaeEncoderStep, HeliosTextEncoderStep, HeliosVideoVaeEncoderStep - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# ==================== -# 1. Vae Encoder -# ==================== - - -# auto_docstring -class HeliosAutoVaeEncoderStep(AutoPipelineBlocks): - """ - Encoder step that encodes video or image inputs. This is an auto pipeline block. - - `HeliosVideoVaeEncoderStep` (video_encoder) is used when `video` is provided. - - `HeliosImageVaeEncoderStep` (image_encoder) is used when `image` is provided. - - If neither is provided, step will be skipped. - - Components: - vae (`AutoencoderKLWan`) video_processor (`VideoProcessor`) - - Inputs: - video (`None`, *optional*): - Input video for video-to-video generation - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - - Outputs: - image_latents (`Tensor`): - The latent representation of the input image. - video_latents (`Tensor`): - Encoded video latents (chunked) - fake_image_latents (`Tensor`): - Fake image latents for history seeding - """ - - block_classes = [HeliosVideoVaeEncoderStep, HeliosImageVaeEncoderStep] - block_names = ["video_encoder", "image_encoder"] - block_trigger_inputs = ["video", "image"] - - @property - def description(self): - return ( - "Encoder step that encodes video or image inputs. This is an auto pipeline block.\n" - " - `HeliosVideoVaeEncoderStep` (video_encoder) is used when `video` is provided.\n" - " - `HeliosImageVaeEncoderStep` (image_encoder) is used when `image` is provided.\n" - " - If neither is provided, step will be skipped." - ) - - -# ==================== -# 2. DENOISE -# ==================== - - -# DENOISE (T2V) -# auto_docstring -class HeliosCoreDenoiseStep(SequentialPipelineBlocks): - """ - Denoise block that takes encoded conditions and runs the chunk-based denoising process. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider (`ClassifierFreeGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - history_sizes (`list`, *optional*, defaults to [16, 2, 1]): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - timesteps (`Tensor`, *optional*): - Timesteps for the denoising process. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - model_name = "helios" - block_classes = [ - HeliosTextInputStep, - HeliosPrepareHistoryStep, - HeliosSetTimestepsStep, - HeliosChunkDenoiseStep, - ] - block_names = ["input", "prepare_history", "set_timesteps", "chunk_denoise"] - - @property - def description(self): - return "Denoise block that takes encoded conditions and runs the chunk-based denoising process." - - @property - def outputs(self): - return [OutputParam("latent_chunks", type_hint=list, description="List of per-chunk denoised latent tensors")] - - -# DENOISE (I2V) -# auto_docstring -class HeliosI2VCoreDenoiseStep(SequentialPipelineBlocks): - """ - I2V denoise block that seeds history with image latents and uses I2V-aware chunk preparation. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider (`ClassifierFreeGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - image_latents (`Tensor`): - image latents used to guide the image generation. Can be generated from vae_encoder step. - fake_image_latents (`Tensor`, *optional*): - Fake image latents used as history seed for I2V generation. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video/fake-image latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video/fake-image latent noise. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - history_sizes (`list`, *optional*, defaults to [16, 2, 1]): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - timesteps (`Tensor`, *optional*): - Timesteps for the denoising process. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - model_name = "helios" - block_classes = [ - HeliosTextInputStep, - HeliosAdditionalInputsStep( - image_latent_inputs=[InputParam.template("image_latents")], - additional_batch_inputs=[ - InputParam( - "fake_image_latents", - type_hint=torch.Tensor, - description="Fake image latents used as history seed for I2V generation.", - ), - ], - ), - HeliosAddNoiseToImageLatentsStep, - HeliosPrepareHistoryStep, - HeliosI2VSeedHistoryStep, - HeliosSetTimestepsStep, - HeliosI2VChunkDenoiseStep, - ] - block_names = [ - "input", - "additional_inputs", - "add_noise_image", - "prepare_history", - "seed_history", - "set_timesteps", - "chunk_denoise", - ] - - @property - def description(self): - return "I2V denoise block that seeds history with image latents and uses I2V-aware chunk preparation." - - @property - def outputs(self): - return [OutputParam("latent_chunks", type_hint=list, description="List of per-chunk denoised latent tensors")] - - -# DENOISE (V2V) -# auto_docstring -class HeliosV2VCoreDenoiseStep(SequentialPipelineBlocks): - """ - V2V denoise block that seeds history with video latents and uses I2V-aware chunk preparation. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider (`ClassifierFreeGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - image_latents (`Tensor`, *optional*): - image latents used to guide the image generation. Can be generated from vae_encoder step. - video_latents (`Tensor`, *optional*): - Encoded video latents for V2V generation. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video latent noise. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - history_sizes (`list`, *optional*, defaults to [16, 2, 1]): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - timesteps (`Tensor`, *optional*): - Timesteps for the denoising process. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - model_name = "helios" - block_classes = [ - HeliosTextInputStep, - HeliosAdditionalInputsStep( - image_latent_inputs=[InputParam.template("image_latents")], - additional_batch_inputs=[ - InputParam( - "video_latents", type_hint=torch.Tensor, description="Encoded video latents for V2V generation." - ), - ], - ), - HeliosAddNoiseToVideoLatentsStep, - HeliosPrepareHistoryStep, - HeliosV2VSeedHistoryStep, - HeliosSetTimestepsStep, - HeliosI2VChunkDenoiseStep, - ] - block_names = [ - "input", - "additional_inputs", - "add_noise_video", - "prepare_history", - "seed_history", - "set_timesteps", - "chunk_denoise", - ] - - @property - def description(self): - return "V2V denoise block that seeds history with video latents and uses I2V-aware chunk preparation." - - @property - def outputs(self): - return [OutputParam("latent_chunks", type_hint=list, description="List of per-chunk denoised latent tensors")] - - -# AUTO DENOISE -# auto_docstring -class HeliosAutoCoreDenoiseStep(ConditionalPipelineBlocks): - """ - Core denoise step that selects the appropriate denoising block. - - `HeliosV2VCoreDenoiseStep` (video2video) for video-to-video tasks. - - `HeliosI2VCoreDenoiseStep` (image2video) for image-to-video tasks. - - `HeliosCoreDenoiseStep` (text2video) for text-to-video tasks. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider (`ClassifierFreeGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - image_latents (`Tensor`, *optional*): - image latents used to guide the image generation. Can be generated from vae_encoder step. - video_latents (`Tensor`, *optional*): - Encoded video latents for V2V generation. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video latent noise. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - history_sizes (`list`): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`): - Custom sigmas for the denoising process. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - timesteps (`Tensor`, *optional*): - Timesteps for the denoising process. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - fake_image_latents (`Tensor`, *optional*): - Fake image latents used as history seed for I2V generation. - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - block_classes = [HeliosV2VCoreDenoiseStep, HeliosI2VCoreDenoiseStep, HeliosCoreDenoiseStep] - block_names = ["video2video", "image2video", "text2video"] - block_trigger_inputs = ["video_latents", "fake_image_latents"] - default_block_name = "text2video" - - def select_block(self, video_latents=None, fake_image_latents=None): - if video_latents is not None: - return "video2video" - elif fake_image_latents is not None: - return "image2video" - return None - - @property - def description(self): - return ( - "Core denoise step that selects the appropriate denoising block.\n" - " - `HeliosV2VCoreDenoiseStep` (video2video) for video-to-video tasks.\n" - " - `HeliosI2VCoreDenoiseStep` (image2video) for image-to-video tasks.\n" - " - `HeliosCoreDenoiseStep` (text2video) for text-to-video tasks." - ) - - -AUTO_BLOCKS = InsertableDict( - [ - ("text_encoder", HeliosTextEncoderStep()), - ("vae_encoder", HeliosAutoVaeEncoderStep()), - ("denoise", HeliosAutoCoreDenoiseStep()), - ("decode", HeliosDecodeStep()), - ] -) - -# ==================== -# 3. Auto Blocks -# ==================== - - -# auto_docstring -class HeliosAutoBlocks(SequentialPipelineBlocks): - """ - Auto Modular pipeline for text-to-video, image-to-video, and video-to-video tasks using Helios. - - Supported workflows: - - `text2video`: requires `prompt` - - `image2video`: requires `prompt`, `image` - - `video2video`: requires `prompt`, `video` - - Components: - text_encoder (`UMT5EncoderModel`) tokenizer (`AutoTokenizer`) guider (`ClassifierFreeGuidance`) vae - (`AutoencoderKLWan`) video_processor (`VideoProcessor`) transformer (`HeliosTransformer3DModel`) scheduler - (`HeliosScheduler`) - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The prompt or prompts not to guide the image generation. - max_sequence_length (`int`, *optional*, defaults to 512): - Maximum sequence length for prompt encoding. - video (`None`, *optional*): - Input video for video-to-video generation - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - image_latents (`Tensor`, *optional*): - image latents used to guide the image generation. Can be generated from vae_encoder step. - video_latents (`Tensor`, *optional*): - Encoded video latents for V2V generation. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video latent noise. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - history_sizes (`list`): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`): - Custom sigmas for the denoising process. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - timesteps (`Tensor`, *optional*): - Timesteps for the denoising process. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - fake_image_latents (`Tensor`, *optional*): - Fake image latents used as history seed for I2V generation. - output_type (`str`, *optional*, defaults to np): - Output format: 'pil', 'np', 'pt'. - - Outputs: - videos (`list`): - The generated videos. - """ - - model_name = "helios" - - block_classes = AUTO_BLOCKS.values() - block_names = AUTO_BLOCKS.keys() - - _workflow_map = { - "text2video": {"prompt": True}, - "image2video": {"prompt": True, "image": True}, - "video2video": {"prompt": True, "video": True}, - } - - @property - def description(self): - return "Auto Modular pipeline for text-to-video, image-to-video, and video-to-video tasks using Helios." - - @property - def outputs(self): - return [OutputParam.template("videos")] diff --git a/diffusers/modular_pipelines/helios/modular_blocks_helios_pyramid.py b/diffusers/modular_pipelines/helios/modular_blocks_helios_pyramid.py deleted file mode 100644 index fea11786de21a5242f145c3eac47a9063cb8c333..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/helios/modular_blocks_helios_pyramid.py +++ /dev/null @@ -1,520 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ...utils import logging -from ..modular_pipeline import AutoPipelineBlocks, ConditionalPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import InputParam, InsertableDict, OutputParam -from .before_denoise import ( - HeliosAdditionalInputsStep, - HeliosAddNoiseToImageLatentsStep, - HeliosAddNoiseToVideoLatentsStep, - HeliosI2VSeedHistoryStep, - HeliosPrepareHistoryStep, - HeliosTextInputStep, - HeliosV2VSeedHistoryStep, -) -from .decoders import HeliosDecodeStep -from .denoise import HeliosPyramidChunkDenoiseStep, HeliosPyramidI2VChunkDenoiseStep -from .encoders import HeliosImageVaeEncoderStep, HeliosTextEncoderStep, HeliosVideoVaeEncoderStep - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# ==================== -# 1. Vae Encoder -# ==================== - - -# auto_docstring -class HeliosPyramidAutoVaeEncoderStep(AutoPipelineBlocks): - """ - Encoder step that encodes video or image inputs. This is an auto pipeline block. - - `HeliosVideoVaeEncoderStep` (video_encoder) is used when `video` is provided. - - `HeliosImageVaeEncoderStep` (image_encoder) is used when `image` is provided. - - If neither is provided, step will be skipped. - - Components: - vae (`AutoencoderKLWan`) video_processor (`VideoProcessor`) - - Inputs: - video (`None`, *optional*): - Input video for video-to-video generation - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - - Outputs: - image_latents (`Tensor`): - The latent representation of the input image. - video_latents (`Tensor`): - Encoded video latents (chunked) - fake_image_latents (`Tensor`): - Fake image latents for history seeding - """ - - block_classes = [HeliosVideoVaeEncoderStep, HeliosImageVaeEncoderStep] - block_names = ["video_encoder", "image_encoder"] - block_trigger_inputs = ["video", "image"] - - @property - def description(self): - return ( - "Encoder step that encodes video or image inputs. This is an auto pipeline block.\n" - " - `HeliosVideoVaeEncoderStep` (video_encoder) is used when `video` is provided.\n" - " - `HeliosImageVaeEncoderStep` (image_encoder) is used when `image` is provided.\n" - " - If neither is provided, step will be skipped." - ) - - -# ==================== -# 2. DENOISE -# ==================== - - -# DENOISE (T2V) -# auto_docstring -class HeliosPyramidCoreDenoiseStep(SequentialPipelineBlocks): - """ - T2V pyramid denoise block with progressive multi-resolution denoising. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider - (`ClassifierFreeZeroStarGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - history_sizes (`list`, *optional*, defaults to [16, 2, 1]): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - pyramid_num_inference_steps_list (`list`, *optional*, defaults to [10, 10, 10]): - Number of denoising steps per pyramid stage. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - model_name = "helios-pyramid" - block_classes = [ - HeliosTextInputStep, - HeliosPrepareHistoryStep, - HeliosPyramidChunkDenoiseStep, - ] - block_names = ["input", "prepare_history", "pyramid_chunk_denoise"] - - @property - def description(self): - return "T2V pyramid denoise block with progressive multi-resolution denoising." - - @property - def outputs(self): - return [OutputParam("latent_chunks", type_hint=list, description="List of per-chunk denoised latent tensors")] - - -# DENOISE (I2V) -# auto_docstring -class HeliosPyramidI2VCoreDenoiseStep(SequentialPipelineBlocks): - """ - I2V pyramid denoise block with progressive multi-resolution denoising. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider - (`ClassifierFreeZeroStarGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - image_latents (`Tensor`): - image latents used to guide the image generation. Can be generated from vae_encoder step. - fake_image_latents (`Tensor`, *optional*): - Fake image latents used as history seed for I2V generation. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video/fake-image latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video/fake-image latent noise. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - history_sizes (`list`, *optional*, defaults to [16, 2, 1]): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - pyramid_num_inference_steps_list (`list`, *optional*, defaults to [10, 10, 10]): - Number of denoising steps per pyramid stage. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - model_name = "helios-pyramid" - block_classes = [ - HeliosTextInputStep, - HeliosAdditionalInputsStep( - image_latent_inputs=[InputParam.template("image_latents")], - additional_batch_inputs=[ - InputParam( - "fake_image_latents", - type_hint=torch.Tensor, - description="Fake image latents used as history seed for I2V generation.", - ), - ], - ), - HeliosAddNoiseToImageLatentsStep, - HeliosPrepareHistoryStep, - HeliosI2VSeedHistoryStep, - HeliosPyramidI2VChunkDenoiseStep, - ] - block_names = [ - "input", - "additional_inputs", - "add_noise_image", - "prepare_history", - "seed_history", - "pyramid_chunk_denoise", - ] - - @property - def description(self): - return "I2V pyramid denoise block with progressive multi-resolution denoising." - - @property - def outputs(self): - return [OutputParam("latent_chunks", type_hint=list, description="List of per-chunk denoised latent tensors")] - - -# DENOISE (V2V) -# auto_docstring -class HeliosPyramidV2VCoreDenoiseStep(SequentialPipelineBlocks): - """ - V2V pyramid denoise block with progressive multi-resolution denoising. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider - (`ClassifierFreeZeroStarGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - image_latents (`Tensor`, *optional*): - image latents used to guide the image generation. Can be generated from vae_encoder step. - video_latents (`Tensor`, *optional*): - Encoded video latents for V2V generation. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video latent noise. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - history_sizes (`list`, *optional*, defaults to [16, 2, 1]): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - pyramid_num_inference_steps_list (`list`, *optional*, defaults to [10, 10, 10]): - Number of denoising steps per pyramid stage. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - model_name = "helios-pyramid" - block_classes = [ - HeliosTextInputStep, - HeliosAdditionalInputsStep( - image_latent_inputs=[InputParam.template("image_latents")], - additional_batch_inputs=[ - InputParam( - "video_latents", type_hint=torch.Tensor, description="Encoded video latents for V2V generation." - ), - ], - ), - HeliosAddNoiseToVideoLatentsStep, - HeliosPrepareHistoryStep, - HeliosV2VSeedHistoryStep, - HeliosPyramidI2VChunkDenoiseStep, - ] - block_names = [ - "input", - "additional_inputs", - "add_noise_video", - "prepare_history", - "seed_history", - "pyramid_chunk_denoise", - ] - - @property - def description(self): - return "V2V pyramid denoise block with progressive multi-resolution denoising." - - @property - def outputs(self): - return [OutputParam("latent_chunks", type_hint=list, description="List of per-chunk denoised latent tensors")] - - -# AUTO DENOISE -# auto_docstring -class HeliosPyramidAutoCoreDenoiseStep(ConditionalPipelineBlocks): - """ - Pyramid core denoise step that selects the appropriate denoising block. - - `HeliosPyramidV2VCoreDenoiseStep` (video2video) for video-to-video tasks. - - `HeliosPyramidI2VCoreDenoiseStep` (image2video) for image-to-video tasks. - - `HeliosPyramidCoreDenoiseStep` (text2video) for text-to-video tasks. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider - (`ClassifierFreeZeroStarGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - image_latents (`Tensor`, *optional*): - image latents used to guide the image generation. Can be generated from vae_encoder step. - video_latents (`Tensor`, *optional*): - Encoded video latents for V2V generation. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video latent noise. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - history_sizes (`list`): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - pyramid_num_inference_steps_list (`list`, *optional*, defaults to [10, 10, 10]): - Number of denoising steps per pyramid stage. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - fake_image_latents (`Tensor`, *optional*): - Fake image latents used as history seed for I2V generation. - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - block_classes = [HeliosPyramidV2VCoreDenoiseStep, HeliosPyramidI2VCoreDenoiseStep, HeliosPyramidCoreDenoiseStep] - block_names = ["video2video", "image2video", "text2video"] - block_trigger_inputs = ["video_latents", "fake_image_latents"] - default_block_name = "text2video" - - def select_block(self, video_latents=None, fake_image_latents=None): - if video_latents is not None: - return "video2video" - elif fake_image_latents is not None: - return "image2video" - return None - - @property - def description(self): - return ( - "Pyramid core denoise step that selects the appropriate denoising block.\n" - " - `HeliosPyramidV2VCoreDenoiseStep` (video2video) for video-to-video tasks.\n" - " - `HeliosPyramidI2VCoreDenoiseStep` (image2video) for image-to-video tasks.\n" - " - `HeliosPyramidCoreDenoiseStep` (text2video) for text-to-video tasks." - ) - - -# ==================== -# 3. Auto Blocks -# ==================== - -PYRAMID_AUTO_BLOCKS = InsertableDict( - [ - ("text_encoder", HeliosTextEncoderStep()), - ("vae_encoder", HeliosPyramidAutoVaeEncoderStep()), - ("denoise", HeliosPyramidAutoCoreDenoiseStep()), - ("decode", HeliosDecodeStep()), - ] -) - - -# auto_docstring -class HeliosPyramidAutoBlocks(SequentialPipelineBlocks): - """ - Auto Modular pipeline for pyramid progressive generation (T2V/I2V/V2V) using Helios. - - Supported workflows: - - `text2video`: requires `prompt` - - `image2video`: requires `prompt`, `image` - - `video2video`: requires `prompt`, `video` - - Components: - text_encoder (`UMT5EncoderModel`) tokenizer (`AutoTokenizer`) guider (`ClassifierFreeGuidance`) vae - (`AutoencoderKLWan`) video_processor (`VideoProcessor`) transformer (`HeliosTransformer3DModel`) scheduler - (`HeliosScheduler`) - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The prompt or prompts not to guide the image generation. - max_sequence_length (`int`, *optional*, defaults to 512): - Maximum sequence length for prompt encoding. - video (`None`, *optional*): - Input video for video-to-video generation - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - image_latents (`Tensor`, *optional*): - image latents used to guide the image generation. Can be generated from vae_encoder step. - video_latents (`Tensor`, *optional*): - Encoded video latents for V2V generation. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video latent noise. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - history_sizes (`list`): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - pyramid_num_inference_steps_list (`list`, *optional*, defaults to [10, 10, 10]): - Number of denoising steps per pyramid stage. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - fake_image_latents (`Tensor`, *optional*): - Fake image latents used as history seed for I2V generation. - output_type (`str`, *optional*, defaults to np): - Output format: 'pil', 'np', 'pt'. - - Outputs: - videos (`list`): - The generated videos. - """ - - model_name = "helios-pyramid" - - block_classes = PYRAMID_AUTO_BLOCKS.values() - block_names = PYRAMID_AUTO_BLOCKS.keys() - - _workflow_map = { - "text2video": {"prompt": True}, - "image2video": {"prompt": True, "image": True}, - "video2video": {"prompt": True, "video": True}, - } - - @property - def description(self): - return "Auto Modular pipeline for pyramid progressive generation (T2V/I2V/V2V) using Helios." - - @property - def outputs(self): - return [OutputParam.template("videos")] diff --git a/diffusers/modular_pipelines/helios/modular_blocks_helios_pyramid_distilled.py b/diffusers/modular_pipelines/helios/modular_blocks_helios_pyramid_distilled.py deleted file mode 100644 index 3e0b32f0df7e23f740208f17fe29eda95534aa6e..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/helios/modular_blocks_helios_pyramid_distilled.py +++ /dev/null @@ -1,530 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ...utils import logging -from ..modular_pipeline import AutoPipelineBlocks, ConditionalPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import InputParam, InsertableDict, OutputParam -from .before_denoise import ( - HeliosAdditionalInputsStep, - HeliosAddNoiseToImageLatentsStep, - HeliosAddNoiseToVideoLatentsStep, - HeliosI2VSeedHistoryStep, - HeliosPrepareHistoryStep, - HeliosTextInputStep, - HeliosV2VSeedHistoryStep, -) -from .decoders import HeliosDecodeStep -from .denoise import HeliosPyramidDistilledChunkDenoiseStep, HeliosPyramidDistilledI2VChunkDenoiseStep -from .encoders import HeliosImageVaeEncoderStep, HeliosTextEncoderStep, HeliosVideoVaeEncoderStep - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# ==================== -# 1. Vae Encoder -# ==================== - - -# auto_docstring -class HeliosPyramidDistilledAutoVaeEncoderStep(AutoPipelineBlocks): - """ - Encoder step for distilled pyramid pipeline. - - `HeliosVideoVaeEncoderStep` (video_encoder) is used when `video` is provided. - - `HeliosImageVaeEncoderStep` (image_encoder) is used when `image` is provided. - - If neither is provided, step will be skipped. - - Components: - vae (`AutoencoderKLWan`) video_processor (`VideoProcessor`) - - Inputs: - video (`None`, *optional*): - Input video for video-to-video generation - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - - Outputs: - image_latents (`Tensor`): - The latent representation of the input image. - video_latents (`Tensor`): - Encoded video latents (chunked) - fake_image_latents (`Tensor`): - Fake image latents for history seeding - """ - - block_classes = [HeliosVideoVaeEncoderStep, HeliosImageVaeEncoderStep] - block_names = ["video_encoder", "image_encoder"] - block_trigger_inputs = ["video", "image"] - - @property - def description(self): - return ( - "Encoder step for distilled pyramid pipeline.\n" - " - `HeliosVideoVaeEncoderStep` (video_encoder) is used when `video` is provided.\n" - " - `HeliosImageVaeEncoderStep` (image_encoder) is used when `image` is provided.\n" - " - If neither is provided, step will be skipped." - ) - - -# ==================== -# 2. DENOISE -# ==================== - - -# DENOISE (T2V) -# auto_docstring -class HeliosPyramidDistilledCoreDenoiseStep(SequentialPipelineBlocks): - """ - T2V distilled pyramid denoise block with DMD scheduler and no CFG. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider (`ClassifierFreeGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - history_sizes (`list`, *optional*, defaults to [16, 2, 1]): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - pyramid_num_inference_steps_list (`list`, *optional*, defaults to [10, 10, 10]): - Number of denoising steps per pyramid stage. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - is_amplify_first_chunk (`bool`, *optional*, defaults to True): - Whether to double the first chunk's timesteps via the scheduler for amplified generation. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - model_name = "helios-pyramid" - block_classes = [ - HeliosTextInputStep, - HeliosPrepareHistoryStep, - HeliosPyramidDistilledChunkDenoiseStep, - ] - block_names = ["input", "prepare_history", "pyramid_chunk_denoise"] - - @property - def description(self): - return "T2V distilled pyramid denoise block with DMD scheduler and no CFG." - - @property - def outputs(self): - return [OutputParam("latent_chunks", type_hint=list, description="List of per-chunk denoised latent tensors")] - - -# DENOISE (I2V) -# auto_docstring -class HeliosPyramidDistilledI2VCoreDenoiseStep(SequentialPipelineBlocks): - """ - I2V distilled pyramid denoise block with DMD scheduler and no CFG. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider (`ClassifierFreeGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - image_latents (`Tensor`): - image latents used to guide the image generation. Can be generated from vae_encoder step. - fake_image_latents (`Tensor`, *optional*): - Fake image latents used as history seed for I2V generation. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video/fake-image latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video/fake-image latent noise. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - history_sizes (`list`, *optional*, defaults to [16, 2, 1]): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - pyramid_num_inference_steps_list (`list`, *optional*, defaults to [10, 10, 10]): - Number of denoising steps per pyramid stage. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - is_amplify_first_chunk (`bool`, *optional*, defaults to True): - Whether to double the first chunk's timesteps via the scheduler for amplified generation. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - model_name = "helios-pyramid" - block_classes = [ - HeliosTextInputStep, - HeliosAdditionalInputsStep( - image_latent_inputs=[InputParam.template("image_latents")], - additional_batch_inputs=[ - InputParam( - "fake_image_latents", - type_hint=torch.Tensor, - description="Fake image latents used as history seed for I2V generation.", - ), - ], - ), - HeliosAddNoiseToImageLatentsStep, - HeliosPrepareHistoryStep, - HeliosI2VSeedHistoryStep, - HeliosPyramidDistilledI2VChunkDenoiseStep, - ] - block_names = [ - "input", - "additional_inputs", - "add_noise_image", - "prepare_history", - "seed_history", - "pyramid_chunk_denoise", - ] - - @property - def description(self): - return "I2V distilled pyramid denoise block with DMD scheduler and no CFG." - - @property - def outputs(self): - return [OutputParam("latent_chunks", type_hint=list, description="List of per-chunk denoised latent tensors")] - - -# DENOISE (V2V) -# auto_docstring -class HeliosPyramidDistilledV2VCoreDenoiseStep(SequentialPipelineBlocks): - """ - V2V distilled pyramid denoise block with DMD scheduler and no CFG. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider (`ClassifierFreeGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - image_latents (`Tensor`, *optional*): - image latents used to guide the image generation. Can be generated from vae_encoder step. - video_latents (`Tensor`, *optional*): - Encoded video latents for V2V generation. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video latent noise. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - history_sizes (`list`, *optional*, defaults to [16, 2, 1]): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - pyramid_num_inference_steps_list (`list`, *optional*, defaults to [10, 10, 10]): - Number of denoising steps per pyramid stage. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - is_amplify_first_chunk (`bool`, *optional*, defaults to True): - Whether to double the first chunk's timesteps via the scheduler for amplified generation. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - model_name = "helios-pyramid" - block_classes = [ - HeliosTextInputStep, - HeliosAdditionalInputsStep( - image_latent_inputs=[InputParam.template("image_latents")], - additional_batch_inputs=[ - InputParam( - "video_latents", type_hint=torch.Tensor, description="Encoded video latents for V2V generation." - ), - ], - ), - HeliosAddNoiseToVideoLatentsStep, - HeliosPrepareHistoryStep, - HeliosV2VSeedHistoryStep, - HeliosPyramidDistilledI2VChunkDenoiseStep, - ] - block_names = [ - "input", - "additional_inputs", - "add_noise_video", - "prepare_history", - "seed_history", - "pyramid_chunk_denoise", - ] - - @property - def description(self): - return "V2V distilled pyramid denoise block with DMD scheduler and no CFG." - - @property - def outputs(self): - return [OutputParam("latent_chunks", type_hint=list, description="List of per-chunk denoised latent tensors")] - - -# AUTO DENOISE -# auto_docstring -class HeliosPyramidDistilledAutoCoreDenoiseStep(ConditionalPipelineBlocks): - """ - Distilled pyramid core denoise step that selects the appropriate denoising block. - - `HeliosPyramidDistilledV2VCoreDenoiseStep` (video2video) for video-to-video tasks. - - `HeliosPyramidDistilledI2VCoreDenoiseStep` (image2video) for image-to-video tasks. - - `HeliosPyramidDistilledCoreDenoiseStep` (text2video) for text-to-video tasks. - - Components: - transformer (`HeliosTransformer3DModel`) scheduler (`HeliosScheduler`) guider (`ClassifierFreeGuidance`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - image_latents (`Tensor`, *optional*): - image latents used to guide the image generation. Can be generated from vae_encoder step. - video_latents (`Tensor`, *optional*): - Encoded video latents for V2V generation. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video latent noise. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - history_sizes (`list`): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - pyramid_num_inference_steps_list (`list`, *optional*, defaults to [10, 10, 10]): - Number of denoising steps per pyramid stage. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - is_amplify_first_chunk (`bool`, *optional*, defaults to True): - Whether to double the first chunk's timesteps via the scheduler for amplified generation. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - fake_image_latents (`Tensor`, *optional*): - Fake image latents used as history seed for I2V generation. - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - - Outputs: - latent_chunks (`list`): - List of per-chunk denoised latent tensors - """ - - block_classes = [ - HeliosPyramidDistilledV2VCoreDenoiseStep, - HeliosPyramidDistilledI2VCoreDenoiseStep, - HeliosPyramidDistilledCoreDenoiseStep, - ] - block_names = ["video2video", "image2video", "text2video"] - block_trigger_inputs = ["video_latents", "fake_image_latents"] - default_block_name = "text2video" - - def select_block(self, video_latents=None, fake_image_latents=None): - if video_latents is not None: - return "video2video" - elif fake_image_latents is not None: - return "image2video" - return None - - @property - def description(self): - return ( - "Distilled pyramid core denoise step that selects the appropriate denoising block.\n" - " - `HeliosPyramidDistilledV2VCoreDenoiseStep` (video2video) for video-to-video tasks.\n" - " - `HeliosPyramidDistilledI2VCoreDenoiseStep` (image2video) for image-to-video tasks.\n" - " - `HeliosPyramidDistilledCoreDenoiseStep` (text2video) for text-to-video tasks." - ) - - -# ==================== -# 3. Auto Blocks -# ==================== - -DISTILLED_PYRAMID_AUTO_BLOCKS = InsertableDict( - [ - ("text_encoder", HeliosTextEncoderStep()), - ("vae_encoder", HeliosPyramidDistilledAutoVaeEncoderStep()), - ("denoise", HeliosPyramidDistilledAutoCoreDenoiseStep()), - ("decode", HeliosDecodeStep()), - ] -) - - -# auto_docstring -class HeliosPyramidDistilledAutoBlocks(SequentialPipelineBlocks): - """ - Auto Modular pipeline for distilled pyramid progressive generation (T2V/I2V/V2V) using Helios. - - Supported workflows: - - `text2video`: requires `prompt` - - `image2video`: requires `prompt`, `image` - - `video2video`: requires `prompt`, `video` - - Components: - text_encoder (`UMT5EncoderModel`) tokenizer (`AutoTokenizer`) guider (`ClassifierFreeGuidance`) vae - (`AutoencoderKLWan`) video_processor (`VideoProcessor`) transformer (`HeliosTransformer3DModel`) scheduler - (`HeliosScheduler`) - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The prompt or prompts not to guide the image generation. - max_sequence_length (`int`, *optional*, defaults to 512): - Maximum sequence length for prompt encoding. - video (`None`, *optional*): - Input video for video-to-video generation - height (`int`, *optional*, defaults to 384): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 640): - The width in pixels of the generated image. - num_latent_frames_per_chunk (`int`, *optional*, defaults to 9): - Number of latent frames per temporal chunk. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos to generate per prompt. - image_latents (`Tensor`, *optional*): - image latents used to guide the image generation. Can be generated from vae_encoder step. - video_latents (`Tensor`, *optional*): - Encoded video latents for V2V generation. - image_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for image latent noise. - image_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for image latent noise. - video_noise_sigma_min (`float`, *optional*, defaults to 0.111): - Minimum sigma for video latent noise. - video_noise_sigma_max (`float`, *optional*, defaults to 0.135): - Maximum sigma for video latent noise. - num_frames (`int`, *optional*, defaults to 132): - Total number of video frames to generate. - history_sizes (`list`): - Sizes of long/mid/short history buffers for temporal context. - keep_first_frame (`bool`, *optional*, defaults to True): - Whether to keep the first frame as a prefix in history. - pyramid_num_inference_steps_list (`list`, *optional*, defaults to [10, 10, 10]): - Number of denoising steps per pyramid stage. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - **denoiser_input_fields (`None`, *optional*): - conditional model inputs for the denoiser: e.g. prompt_embeds, negative_prompt_embeds, etc. - is_amplify_first_chunk (`bool`, *optional*, defaults to True): - Whether to double the first chunk's timesteps via the scheduler for amplified generation. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - fake_image_latents (`Tensor`, *optional*): - Fake image latents used as history seed for I2V generation. - output_type (`str`, *optional*, defaults to np): - Output format: 'pil', 'np', 'pt'. - - Outputs: - videos (`list`): - The generated videos. - """ - - model_name = "helios-pyramid" - - block_classes = DISTILLED_PYRAMID_AUTO_BLOCKS.values() - block_names = DISTILLED_PYRAMID_AUTO_BLOCKS.keys() - - _workflow_map = { - "text2video": {"prompt": True}, - "image2video": {"prompt": True, "image": True}, - "video2video": {"prompt": True, "video": True}, - } - - @property - def description(self): - return "Auto Modular pipeline for distilled pyramid progressive generation (T2V/I2V/V2V) using Helios." - - @property - def outputs(self): - return [OutputParam.template("videos")] diff --git a/diffusers/modular_pipelines/helios/modular_pipeline.py b/diffusers/modular_pipelines/helios/modular_pipeline.py deleted file mode 100644 index 1fc338e67f05c6714bdb0634c516b732bb16256d..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/helios/modular_pipeline.py +++ /dev/null @@ -1,87 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from ...loaders import HeliosLoraLoaderMixin -from ...utils import logging -from ..modular_pipeline import ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class HeliosModularPipeline( - ModularPipeline, - HeliosLoraLoaderMixin, -): - """ - A ModularPipeline for Helios text-to-video generation. - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "HeliosAutoBlocks" - - @property - def vae_scale_factor_spatial(self): - vae_scale_factor = 8 - if hasattr(self, "vae") and self.vae is not None: - vae_scale_factor = self.vae.config.scale_factor_spatial - return vae_scale_factor - - @property - def vae_scale_factor_temporal(self): - vae_scale_factor = 4 - if hasattr(self, "vae") and self.vae is not None: - vae_scale_factor = self.vae.config.scale_factor_temporal - return vae_scale_factor - - @property - def num_channels_latents(self): - # YiYi TODO: find out default value - num_channels_latents = 16 - if hasattr(self, "transformer") and self.transformer is not None: - num_channels_latents = self.transformer.config.in_channels - return num_channels_latents - - @property - def requires_unconditional_embeds(self): - requires_unconditional_embeds = False - - if hasattr(self, "guider") and self.guider is not None: - requires_unconditional_embeds = self.guider._enabled and self.guider.num_conditions > 1 - - return requires_unconditional_embeds - - -class HeliosPyramidModularPipeline(HeliosModularPipeline): - """ - A ModularPipeline for Helios pyramid (progressive resolution) video generation. - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "HeliosPyramidAutoBlocks" - - -class HeliosPyramidDistilledModularPipeline(HeliosModularPipeline): - """ - A ModularPipeline for Helios distilled pyramid video generation using DMD scheduler. - - Uses guidance_scale=1.0 (no CFG) and supports is_amplify_first_chunk for the DMD scheduler. - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "HeliosPyramidDistilledAutoBlocks" diff --git a/diffusers/modular_pipelines/hunyuan_video1_5/__init__.py b/diffusers/modular_pipelines/hunyuan_video1_5/__init__.py deleted file mode 100644 index a9c12e4a78ce2a4d5bc743176cc4199b01a7e035..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/hunyuan_video1_5/__init__.py +++ /dev/null @@ -1,49 +0,0 @@ -from typing import TYPE_CHECKING - -from ...utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, -) - - -_dummy_objects = {} -_import_structure = {} - -try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ...utils import dummy_torch_and_transformers_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) -else: - _import_structure["modular_blocks_hunyuan_video1_5"] = [ - "HunyuanVideo15AutoBlocks", - ] - _import_structure["modular_pipeline"] = ["HunyuanVideo15ModularPipeline"] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .modular_blocks_hunyuan_video1_5 import HunyuanVideo15AutoBlocks - from .modular_pipeline import HunyuanVideo15ModularPipeline -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/hunyuan_video1_5/before_denoise.py b/diffusers/modular_pipelines/hunyuan_video1_5/before_denoise.py deleted file mode 100644 index 4c02eb9dd0846976db2d731a340a4c0197f36e92..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/hunyuan_video1_5/before_denoise.py +++ /dev/null @@ -1,324 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect - -import numpy as np -import torch - -from ...configuration_utils import FrozenDict -from ...models import HunyuanVideo15Transformer3DModel -from ...pipelines.hunyuan_video1_5.image_processor import HunyuanVideo15ImageProcessor -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import HunyuanVideo15ModularPipeline - - -logger = logging.get_logger(__name__) - - -# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps -def retrieve_timesteps( - scheduler, - num_inference_steps: int | None = None, - device: str | torch.device | None = None, - timesteps: list[int] | None = None, - sigmas: list[float] | None = None, - **kwargs, -): - r""" - Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles - custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`. - - Args: - scheduler (`SchedulerMixin`): - The scheduler to get timesteps from. - num_inference_steps (`int`): - The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps` - must be `None`. - device (`str` or `torch.device`, *optional*): - The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. - timesteps (`list[int]`, *optional*): - Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed, - `num_inference_steps` and `sigmas` must be `None`. - sigmas (`list[float]`, *optional*): - Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed, - `num_inference_steps` and `timesteps` must be `None`. - - Returns: - `tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the - second element is the number of inference steps. - """ - if timesteps is not None and sigmas is not None: - raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values") - if timesteps is not None: - accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) - if not accepts_timesteps: - raise ValueError( - f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" - f" timestep schedules. Please check whether you are using the correct scheduler." - ) - scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) - timesteps = scheduler.timesteps - num_inference_steps = len(timesteps) - elif sigmas is not None: - accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) - if not accept_sigmas: - raise ValueError( - f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" - f" sigmas schedules. Please check whether you are using the correct scheduler." - ) - scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs) - timesteps = scheduler.timesteps - num_inference_steps = len(timesteps) - else: - scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) - timesteps = scheduler.timesteps - return timesteps, num_inference_steps - - -class HunyuanVideo15TextInputStep(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - @property - def description(self) -> str: - return "Input processing step that determines batch_size" - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt_embeds"), - InputParam.template("batch_size", default=None), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("batch_size", type_hint=int), - ] - - @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - block_state.batch_size = getattr(block_state, "batch_size", None) or block_state.prompt_embeds.shape[0] - self.set_block_state(state, block_state) - return components, state - - -class HunyuanVideo15SetTimestepsStep(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def description(self) -> str: - return "Step that sets the scheduler's timesteps for inference" - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_inference_steps"), - InputParam.template("sigmas"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("timesteps", type_hint=torch.Tensor), - OutputParam("num_inference_steps", type_hint=int), - ] - - @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - sigmas = block_state.sigmas - if sigmas is None: - sigmas = np.linspace(1.0, 0.0, block_state.num_inference_steps + 1)[:-1] - - block_state.timesteps, block_state.num_inference_steps = retrieve_timesteps( - components.scheduler, block_state.num_inference_steps, device, sigmas=sigmas - ) - - self.set_block_state(state, block_state) - return components, state - - -class HunyuanVideo15PrepareLatentsStep(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - @property - def description(self) -> str: - return "Prepare latents, conditioning latents, mask, and image_embeds for T2V" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("transformer", HunyuanVideo15Transformer3DModel), - ComponentSpec( - "video_processor", - HunyuanVideo15ImageProcessor, - config=FrozenDict({"vae_scale_factor": 16}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("height"), - InputParam.template("width"), - InputParam("num_frames", type_hint=int, default=121, description="Number of video frames to generate."), - InputParam.template("latents"), - InputParam.template("num_images_per_prompt", name="num_videos_per_prompt"), - InputParam.template("generator"), - InputParam.template("batch_size", required=True, default=None), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("latents", type_hint=torch.Tensor, description="Pure noise latents"), - OutputParam("cond_latents_concat", type_hint=torch.Tensor), - OutputParam("mask_concat", type_hint=torch.Tensor), - OutputParam("image_embeds", type_hint=torch.Tensor), - ] - - @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - dtype = components.transformer.dtype - - height = block_state.height - width = block_state.width - if height is None and width is None: - height, width = components.video_processor.calculate_default_height_width( - components.default_aspect_ratio[1], components.default_aspect_ratio[0], components.target_size - ) - - batch_size = block_state.batch_size * block_state.num_videos_per_prompt - num_frames = block_state.num_frames - - latents = block_state.latents - if latents is not None: - latents = latents.to(device=device, dtype=dtype) - else: - shape = ( - batch_size, - components.num_channels_latents, - (num_frames - 1) // components.vae_scale_factor_temporal + 1, - int(height) // components.vae_scale_factor_spatial, - int(width) // components.vae_scale_factor_spatial, - ) - if isinstance(block_state.generator, list) and len(block_state.generator) != batch_size: - raise ValueError( - f"You have passed a list of generators of length {len(block_state.generator)}, but requested an effective batch" - f" size of {batch_size}. Make sure the batch size matches the length of the generators." - ) - latents = randn_tensor(shape, generator=block_state.generator, device=device, dtype=dtype) - - block_state.latents = latents - - b, c, f, h, w = latents.shape - block_state.cond_latents_concat = torch.zeros(b, c, f, h, w, dtype=dtype, device=device) - block_state.mask_concat = torch.zeros(b, 1, f, h, w, dtype=dtype, device=device) - - block_state.image_embeds = torch.zeros( - block_state.batch_size, - components.vision_num_semantic_tokens, - components.vision_states_dim, - dtype=dtype, - device=device, - ) - - self.set_block_state(state, block_state) - return components, state - - -class HunyuanVideo15Image2VideoPrepareLatentsStep(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - @property - def description(self) -> str: - return ( - "Prepare I2V conditioning from image_latents and image_embeds. " - "Expects pure noise `latents` from HunyuanVideo15PrepareLatentsStep. " - "Builds cond_latents_concat and mask_concat for the denoiser." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", HunyuanVideo15Transformer3DModel)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - "image_latents", - type_hint=torch.Tensor, - required=True, - description="Pre-encoded image latents from the VAE encoder step, used as conditioning for I2V.", - ), - InputParam( - "image_embeds", - type_hint=torch.Tensor, - required=True, - description="Siglip image embeddings from the image encoder step, used as extra conditioning for I2V.", - ), - InputParam.template("latents", required=True), - InputParam.template("num_images_per_prompt", name="num_videos_per_prompt"), - InputParam.template("batch_size", required=True, default=None), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("cond_latents_concat", type_hint=torch.Tensor), - OutputParam("mask_concat", type_hint=torch.Tensor), - OutputParam("image_embeds", type_hint=torch.Tensor), - ] - - @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - dtype = components.transformer.dtype - - batch_size = block_state.batch_size * block_state.num_videos_per_prompt - - b, c, f, h, w = block_state.latents.shape - - latent_condition = block_state.image_latents.to(device=device, dtype=dtype) - latent_condition = latent_condition.repeat(batch_size, 1, f, 1, 1) - latent_condition[:, :, 1:, :, :] = 0 - block_state.cond_latents_concat = latent_condition - - latent_mask = torch.zeros(b, 1, f, h, w, dtype=dtype, device=device) - latent_mask[:, :, 0, :, :] = 1.0 - block_state.mask_concat = latent_mask - - image_embeds = block_state.image_embeds.to(device=device, dtype=dtype) - if image_embeds.shape[0] == 1 and batch_size > 1: - image_embeds = image_embeds.repeat(batch_size, 1, 1) - block_state.image_embeds = image_embeds - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/hunyuan_video1_5/decoders.py b/diffusers/modular_pipelines/hunyuan_video1_5/decoders.py deleted file mode 100644 index 630af85c1b10f48f37ac43d1c98e5b9fc4264ceb..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/hunyuan_video1_5/decoders.py +++ /dev/null @@ -1,70 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch - -from ...configuration_utils import FrozenDict -from ...models import AutoencoderKLHunyuanVideo15 -from ...pipelines.hunyuan_video1_5.image_processor import HunyuanVideo15ImageProcessor -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam - - -logger = logging.get_logger(__name__) - - -class HunyuanVideo15VaeDecoderStep(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLHunyuanVideo15), - ComponentSpec( - "video_processor", - HunyuanVideo15ImageProcessor, - config=FrozenDict({"vae_scale_factor": 16}), - default_creation_method="from_config", - ), - ] - - @property - def description(self) -> str: - return "Step that decodes the denoised latents into videos" - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents", required=True), - InputParam.template("output_type", default="np"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam.template("videos"), - ] - - @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - latents = block_state.latents.to(components.vae.dtype) / components.vae.config.scaling_factor - video = components.vae.decode(latents, return_dict=False)[0] - block_state.videos = components.video_processor.postprocess_video(video, output_type=block_state.output_type) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py b/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py deleted file mode 100644 index 293fad57c93f41f469d09dab81b565b457a3a9d8..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py +++ /dev/null @@ -1,401 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...models import HunyuanVideo15Transformer3DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ..modular_pipeline import ( - BlockState, - LoopSequentialPipelineBlocks, - ModularPipelineBlocks, - PipelineState, -) -from ..modular_pipeline_utils import ComponentSpec, InputParam -from .modular_pipeline import HunyuanVideo15ModularPipeline - - -logger = logging.get_logger(__name__) - - -class HunyuanVideo15LoopBeforeDenoiser(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - @property - def description(self) -> str: - return "Step within the denoising loop that prepares the latent input" - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents", required=True), - InputParam("cond_latents_concat", required=True, type_hint=torch.Tensor), - InputParam("mask_concat", required=True, type_hint=torch.Tensor), - ] - - @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - block_state.latent_model_input = torch.cat( - [block_state.latents, block_state.cond_latents_concat, block_state.mask_concat], dim=1 - ) - return components, block_state - - -class HunyuanVideo15LoopDenoiser(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - def __init__(self, guider_input_fields=None): - if guider_input_fields is None: - guider_input_fields = { - "encoder_hidden_states": ("prompt_embeds", "negative_prompt_embeds"), - "encoder_attention_mask": ("prompt_embeds_mask", "negative_prompt_embeds_mask"), - "encoder_hidden_states_2": ("prompt_embeds_2", "negative_prompt_embeds_2"), - "encoder_attention_mask_2": ("prompt_embeds_mask_2", "negative_prompt_embeds_mask_2"), - } - if not isinstance(guider_input_fields, dict): - raise ValueError(f"guider_input_fields must be a dictionary but is {type(guider_input_fields)}") - self._guider_input_fields = guider_input_fields - super().__init__() - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 7.5}), - default_creation_method="from_config", - ), - ComponentSpec("transformer", HunyuanVideo15Transformer3DModel), - ] - - @property - def description(self) -> str: - return "Step within the denoising loop that denoises the latents with guidance" - - @property - def inputs(self) -> list[InputParam]: - inputs = [ - InputParam.template("attention_kwargs"), - InputParam.template("num_inference_steps", required=True, default=None), - InputParam( - "image_embeds", - type_hint=torch.Tensor, - description="Siglip image embeddings used as extra conditioning for I2V. Zero-filled for T2V.", - ), - ] - for value in self._guider_input_fields.values(): - if isinstance(value, tuple): - inputs.append( - InputParam( - name=value[0], - required=True, - type_hint=torch.Tensor, - description=f"Positive branch of the {value[0]!r} field fed into the guider.", - ) - ) - for neg_name in value[1:]: - inputs.append( - InputParam( - name=neg_name, - type_hint=torch.Tensor, - description=f"Negative branch of the {neg_name!r} field fed into the guider.", - ) - ) - else: - inputs.append( - InputParam( - name=value, - required=True, - type_hint=torch.Tensor, - description=f"{value!r} field fed into the guider.", - ) - ) - return inputs - - @torch.no_grad() - def __call__( - self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: - timestep = t.expand(block_state.latent_model_input.shape[0]).to(block_state.latent_model_input.dtype) - - # Step 1: Collect model inputs - guider_inputs = { - input_name: tuple(getattr(block_state, v) for v in value) - if isinstance(value, tuple) - else getattr(block_state, value) - for input_name, value in self._guider_input_fields.items() - } - - # Step 2: Update guider state - components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) - - # Step 3: Prepare batched inputs - guider_state = components.guider.prepare_inputs(guider_inputs) - - # Step 4: Run denoiser for each batch - for guider_state_batch in guider_state: - components.guider.prepare_models(components.transformer) - - cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()} - - context_name = getattr(guider_state_batch, components.guider._identifier_key) - with components.transformer.cache_context(context_name): - guider_state_batch.noise_pred = components.transformer( - hidden_states=block_state.latent_model_input, - image_embeds=block_state.image_embeds, - timestep=timestep, - attention_kwargs=block_state.attention_kwargs, - return_dict=False, - **cond_kwargs, - )[0] - - components.guider.cleanup_models(components.transformer) - - # Step 5: Combine predictions - block_state.noise_pred = components.guider(guider_state)[0] - - return components, block_state - - -class HunyuanVideo15LoopAfterDenoiser(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def description(self) -> str: - return "Step within the denoising loop that updates the latents" - - @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - latents_dtype = block_state.latents.dtype - block_state.latents = components.scheduler.step( - block_state.noise_pred, t, block_state.latents, return_dict=False - )[0] - - if block_state.latents.dtype != latents_dtype: - if torch.backends.mps.is_available(): - block_state.latents = block_state.latents.to(latents_dtype) - - return components, block_state - - -class HunyuanVideo15DenoiseLoopWrapper(LoopSequentialPipelineBlocks): - model_name = "hunyuan-video-1.5" - - @property - def description(self) -> str: - return "Pipeline block that iteratively denoises the latents over timesteps" - - @property - def loop_expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), - ComponentSpec("transformer", HunyuanVideo15Transformer3DModel), - ] - - @property - def loop_inputs(self) -> list[InputParam]: - return [ - InputParam.template("timesteps", required=True), - InputParam.template("num_inference_steps", required=True, default=None), - ] - - @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - block_state.num_warmup_steps = max( - len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order, 0 - ) - - with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) - if i == len(block_state.timesteps) - 1 or ( - (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 - ): - progress_bar.update() - - self.set_block_state(state, block_state) - return components, state - - -class HunyuanVideo15DenoiseStep(HunyuanVideo15DenoiseLoopWrapper): - block_classes = [ - HunyuanVideo15LoopBeforeDenoiser, - HunyuanVideo15LoopDenoiser(), - HunyuanVideo15LoopAfterDenoiser, - ] - block_names = ["before_denoiser", "denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoise step that iteratively denoises the latents.\n" - "At each iteration:\n" - " - `HunyuanVideo15LoopBeforeDenoiser`\n" - " - `HunyuanVideo15LoopDenoiser`\n" - " - `HunyuanVideo15LoopAfterDenoiser`\n" - "This block supports text-to-video tasks." - ) - - -class HunyuanVideo15Image2VideoLoopDenoiser(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - def __init__(self, guider_input_fields=None): - if guider_input_fields is None: - guider_input_fields = { - "encoder_hidden_states": ("prompt_embeds", "negative_prompt_embeds"), - "encoder_attention_mask": ("prompt_embeds_mask", "negative_prompt_embeds_mask"), - "encoder_hidden_states_2": ("prompt_embeds_2", "negative_prompt_embeds_2"), - "encoder_attention_mask_2": ("prompt_embeds_mask_2", "negative_prompt_embeds_mask_2"), - } - if not isinstance(guider_input_fields, dict): - raise ValueError(f"guider_input_fields must be a dictionary but is {type(guider_input_fields)}") - self._guider_input_fields = guider_input_fields - super().__init__() - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 7.5}), - default_creation_method="from_config", - ), - ComponentSpec("transformer", HunyuanVideo15Transformer3DModel), - ] - - @property - def description(self) -> str: - return "I2V denoiser with MeanFlow timestep_r support" - - @property - def inputs(self) -> list[InputParam]: - inputs = [ - InputParam.template("attention_kwargs"), - InputParam.template("num_inference_steps", required=True, default=None), - InputParam( - "image_embeds", - type_hint=torch.Tensor, - description="Siglip image embeddings used as extra conditioning for I2V. Zero-filled for T2V.", - ), - InputParam.template("timesteps", required=True), - ] - for value in self._guider_input_fields.values(): - if isinstance(value, tuple): - inputs.append( - InputParam( - name=value[0], - required=True, - type_hint=torch.Tensor, - description=f"Positive branch of the {value[0]!r} field fed into the guider.", - ) - ) - for neg_name in value[1:]: - inputs.append( - InputParam( - name=neg_name, - type_hint=torch.Tensor, - description=f"Negative branch of the {neg_name!r} field fed into the guider.", - ) - ) - else: - inputs.append( - InputParam( - name=value, - required=True, - type_hint=torch.Tensor, - description=f"{value!r} field fed into the guider.", - ) - ) - return inputs - - @torch.no_grad() - def __call__( - self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: - timestep = t.expand(block_state.latent_model_input.shape[0]).to(block_state.latent_model_input.dtype) - - # MeanFlow timestep_r (lines 855-862) - if components.transformer.config.use_meanflow: - if i == len(block_state.timesteps) - 1: - timestep_r = torch.tensor([0.0], device=timestep.device) - else: - timestep_r = block_state.timesteps[i + 1] - timestep_r = timestep_r.expand(block_state.latents.shape[0]).to(block_state.latents.dtype) - else: - timestep_r = None - - guider_inputs = { - input_name: tuple(getattr(block_state, v) for v in value) - if isinstance(value, tuple) - else getattr(block_state, value) - for input_name, value in self._guider_input_fields.items() - } - - components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) - guider_state = components.guider.prepare_inputs(guider_inputs) - - for guider_state_batch in guider_state: - components.guider.prepare_models(components.transformer) - - cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()} - - context_name = getattr(guider_state_batch, components.guider._identifier_key) - with components.transformer.cache_context(context_name): - guider_state_batch.noise_pred = components.transformer( - hidden_states=block_state.latent_model_input, - image_embeds=block_state.image_embeds, - timestep=timestep, - timestep_r=timestep_r, - attention_kwargs=block_state.attention_kwargs, - return_dict=False, - **cond_kwargs, - )[0] - - components.guider.cleanup_models(components.transformer) - - block_state.noise_pred = components.guider(guider_state)[0] - - return components, block_state - - -class HunyuanVideo15Image2VideoDenoiseStep(HunyuanVideo15DenoiseLoopWrapper): - block_classes = [ - HunyuanVideo15LoopBeforeDenoiser, - HunyuanVideo15Image2VideoLoopDenoiser(), - HunyuanVideo15LoopAfterDenoiser, - ] - block_names = ["before_denoiser", "denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoise step for image-to-video with MeanFlow support.\n" - "At each iteration:\n" - " - `HunyuanVideo15LoopBeforeDenoiser`\n" - " - `HunyuanVideo15Image2VideoLoopDenoiser`\n" - " - `HunyuanVideo15LoopAfterDenoiser`" - ) diff --git a/diffusers/modular_pipelines/hunyuan_video1_5/encoders.py b/diffusers/modular_pipelines/hunyuan_video1_5/encoders.py deleted file mode 100644 index 9d340cc88194c0322c8696d22ae04b1e4b47856a..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/hunyuan_video1_5/encoders.py +++ /dev/null @@ -1,441 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import re - -import torch -from transformers import ( - ByT5Tokenizer, - Qwen2_5_VLTextModel, - Qwen2TokenizerFast, - SiglipImageProcessor, - SiglipVisionModel, - T5EncoderModel, -) - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...models import AutoencoderKLHunyuanVideo15 -from ...pipelines.hunyuan_video1_5.image_processor import HunyuanVideo15ImageProcessor -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import HunyuanVideo15ModularPipeline - - -logger = logging.get_logger(__name__) - - -def format_text_input(prompt, system_message): - return [ - [{"role": "system", "content": system_message}, {"role": "user", "content": p if p else " "}] for p in prompt - ] - - -def extract_glyph_texts(prompt): - pattern = r"\"(.*?)\"|\"(.*?)\"" - matches = re.findall(pattern, prompt) - result = [match[0] or match[1] for match in matches] - result = list(dict.fromkeys(result)) if len(result) > 1 else result - if result: - formatted_result = ". ".join([f'Text "{text}"' for text in result]) + ". " - else: - formatted_result = None - return formatted_result - - -def _get_mllm_prompt_embeds( - text_encoder, - tokenizer, - prompt, - device, - tokenizer_max_length=1000, - num_hidden_layers_to_skip=2, - # fmt: off - system_message="You are a helpful assistant. Describe the video by detailing the following aspects: \ - 1. The main content and theme of the video. \ - 2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects. \ - 3. Actions, events, behaviors temporal relationships, physical movement changes of the objects. \ - 4. background environment, light, style and atmosphere. \ - 5. camera angles, movements, and transitions used in the video.", - # fmt: on - crop_start=108, -): - prompt = [prompt] if isinstance(prompt, str) else prompt - prompt = format_text_input(prompt, system_message) - - text_inputs = tokenizer.apply_chat_template( - prompt, - add_generation_prompt=True, - tokenize=True, - return_dict=True, - padding="max_length", - max_length=tokenizer_max_length + crop_start, - truncation=True, - return_tensors="pt", - ) - - text_input_ids = text_inputs.input_ids.to(device=device) - prompt_attention_mask = text_inputs.attention_mask.to(device=device) - - prompt_embeds = text_encoder( - input_ids=text_input_ids, - attention_mask=prompt_attention_mask, - output_hidden_states=True, - ).hidden_states[-(num_hidden_layers_to_skip + 1)] - - if crop_start is not None and crop_start > 0: - prompt_embeds = prompt_embeds[:, crop_start:] - prompt_attention_mask = prompt_attention_mask[:, crop_start:] - - return prompt_embeds, prompt_attention_mask - - -def _get_byt5_prompt_embeds(tokenizer, text_encoder, prompt, device, tokenizer_max_length=256): - prompt = [prompt] if isinstance(prompt, str) else prompt - glyph_texts = [extract_glyph_texts(p) for p in prompt] - - prompt_embeds_list = [] - prompt_embeds_mask_list = [] - - for glyph_text in glyph_texts: - if glyph_text is None: - glyph_text_embeds = torch.zeros( - (1, tokenizer_max_length, text_encoder.config.d_model), device=device, dtype=text_encoder.dtype - ) - glyph_text_embeds_mask = torch.zeros((1, tokenizer_max_length), device=device, dtype=torch.int64) - else: - txt_tokens = tokenizer( - glyph_text, - padding="max_length", - max_length=tokenizer_max_length, - truncation=True, - add_special_tokens=True, - return_tensors="pt", - ).to(device) - - glyph_text_embeds = text_encoder( - input_ids=txt_tokens.input_ids, - attention_mask=txt_tokens.attention_mask.float(), - )[0] - glyph_text_embeds = glyph_text_embeds.to(device=device) - glyph_text_embeds_mask = txt_tokens.attention_mask.to(device=device) - - prompt_embeds_list.append(glyph_text_embeds) - prompt_embeds_mask_list.append(glyph_text_embeds_mask) - - return torch.cat(prompt_embeds_list, dim=0), torch.cat(prompt_embeds_mask_list, dim=0) - - -class HunyuanVideo15TextEncoderStep(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - @property - def description(self) -> str: - return "Dual text encoder step using Qwen2.5-VL (MLLM) and ByT5 (glyph text)" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", Qwen2_5_VLTextModel), - ComponentSpec("tokenizer", Qwen2TokenizerFast), - ComponentSpec("text_encoder_2", T5EncoderModel), - ComponentSpec("tokenizer_2", ByT5Tokenizer), - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 7.5}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt", required=False), - InputParam.template("negative_prompt"), - InputParam.template("num_images_per_prompt", name="num_videos_per_prompt"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam.template("prompt_embeds"), - OutputParam.template("prompt_embeds_mask"), - OutputParam.template("negative_prompt_embeds"), - OutputParam.template("negative_prompt_embeds_mask"), - OutputParam( - "prompt_embeds_2", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="ByT5 glyph-text embeddings used as a second conditioning stream for the transformer.", - ), - OutputParam( - "prompt_embeds_mask_2", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Attention mask for the ByT5 glyph-text embeddings.", - ), - OutputParam( - "negative_prompt_embeds_2", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="ByT5 glyph-text negative embeddings for classifier-free guidance.", - ), - OutputParam( - "negative_prompt_embeds_mask_2", - type_hint=torch.Tensor, - kwargs_type="denoiser_input_fields", - description="Attention mask for the ByT5 glyph-text negative embeddings.", - ), - ] - - @staticmethod - def encode_prompt( - components, - prompt, - device=None, - dtype=None, - batch_size=1, - num_videos_per_prompt=1, - ): - device = device or components._execution_device - dtype = dtype or components.text_encoder.dtype - - if prompt is None: - prompt = [""] * batch_size - prompt = [prompt] if isinstance(prompt, str) else prompt - - prompt_embeds, prompt_embeds_mask = _get_mllm_prompt_embeds( - tokenizer=components.tokenizer, - text_encoder=components.text_encoder, - prompt=prompt, - device=device, - tokenizer_max_length=components.tokenizer_max_length, - system_message=components.system_message, - crop_start=components.prompt_template_encode_start_idx, - ) - - prompt_embeds_2, prompt_embeds_mask_2 = _get_byt5_prompt_embeds( - tokenizer=components.tokenizer_2, - text_encoder=components.text_encoder_2, - prompt=prompt, - device=device, - tokenizer_max_length=components.tokenizer_2_max_length, - ) - - _, seq_len, _ = prompt_embeds.shape - prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1).view( - batch_size * num_videos_per_prompt, seq_len, -1 - ) - prompt_embeds_mask = prompt_embeds_mask.repeat(1, num_videos_per_prompt, 1).view( - batch_size * num_videos_per_prompt, seq_len - ) - - _, seq_len_2, _ = prompt_embeds_2.shape - prompt_embeds_2 = prompt_embeds_2.repeat(1, num_videos_per_prompt, 1).view( - batch_size * num_videos_per_prompt, seq_len_2, -1 - ) - prompt_embeds_mask_2 = prompt_embeds_mask_2.repeat(1, num_videos_per_prompt, 1).view( - batch_size * num_videos_per_prompt, seq_len_2 - ) - - prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) - prompt_embeds_mask = prompt_embeds_mask.to(dtype=dtype, device=device) - prompt_embeds_2 = prompt_embeds_2.to(dtype=dtype, device=device) - prompt_embeds_mask_2 = prompt_embeds_mask_2.to(dtype=dtype, device=device) - - return prompt_embeds, prompt_embeds_mask, prompt_embeds_2, prompt_embeds_mask_2 - - @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - dtype = components.transformer.dtype - - prompt = block_state.prompt - negative_prompt = block_state.negative_prompt - num_videos_per_prompt = block_state.num_videos_per_prompt - - if prompt is not None and isinstance(prompt, str): - batch_size = 1 - elif prompt is not None and isinstance(prompt, list): - batch_size = len(prompt) - else: - batch_size = 1 - - ( - block_state.prompt_embeds, - block_state.prompt_embeds_mask, - block_state.prompt_embeds_2, - block_state.prompt_embeds_mask_2, - ) = self.encode_prompt( - components, - prompt=prompt, - device=device, - dtype=dtype, - batch_size=batch_size, - num_videos_per_prompt=num_videos_per_prompt, - ) - - if components.requires_unconditional_embeds: - ( - block_state.negative_prompt_embeds, - block_state.negative_prompt_embeds_mask, - block_state.negative_prompt_embeds_2, - block_state.negative_prompt_embeds_mask_2, - ) = self.encode_prompt( - components, - prompt=negative_prompt, - device=device, - dtype=dtype, - batch_size=batch_size, - num_videos_per_prompt=num_videos_per_prompt, - ) - - state.set("batch_size", batch_size) - - self.set_block_state(state, block_state) - return components, state - - -def retrieve_latents( - encoder_output: torch.Tensor, generator: torch.Generator | None = None, sample_mode: str = "sample" -): - if hasattr(encoder_output, "latent_dist") and sample_mode == "sample": - return encoder_output.latent_dist.sample(generator) - elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax": - return encoder_output.latent_dist.mode() - elif hasattr(encoder_output, "latents"): - return encoder_output.latents - else: - raise AttributeError("Could not access latents of provided encoder_output") - - -class HunyuanVideo15VaeEncoderStep(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - @property - def description(self) -> str: - return "VAE Encoder step that encodes an input image into latent space for image-to-video generation" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLHunyuanVideo15), - ComponentSpec( - "video_processor", - HunyuanVideo15ImageProcessor, - config=FrozenDict({"vae_scale_factor": 16}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("image", required=True), - InputParam.template("height"), - InputParam.template("width"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "image_latents", - type_hint=torch.Tensor, - description="Encoded image latents from the VAE encoder", - ), - OutputParam("height", type_hint=int, description="Target height resolved from image"), - OutputParam("width", type_hint=int, description="Target width resolved from image"), - ] - - @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - image = block_state.image - height = block_state.height - width = block_state.width - if height is None or width is None: - height, width = components.video_processor.calculate_default_height_width( - height=image.size[1], width=image.size[0], target_size=components.target_size - ) - image = components.video_processor.resize(image, height=height, width=width, resize_mode="crop") - - vae_dtype = components.vae.dtype - image_tensor = components.video_processor.preprocess(image, height=height, width=width).to( - device=device, dtype=vae_dtype - ) - image_tensor = image_tensor.unsqueeze(2) - image_latents = retrieve_latents(components.vae.encode(image_tensor), sample_mode="argmax") - image_latents = image_latents * components.vae.config.scaling_factor - - block_state.image_latents = image_latents - block_state.height = height - block_state.width = width - state.set("image", image) - - self.set_block_state(state, block_state) - return components, state - - -class HunyuanVideo15ImageEncoderStep(ModularPipelineBlocks): - model_name = "hunyuan-video-1.5" - - @property - def description(self) -> str: - return "Siglip image encoder step that produces image_embeds for image-to-video generation" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("image_encoder", SiglipVisionModel), - ComponentSpec("feature_extractor", SiglipImageProcessor), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("image", required=True), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "image_embeds", - type_hint=torch.Tensor, - description="Image embeddings from the Siglip vision encoder", - ), - ] - - @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - image_encoder_dtype = next(components.image_encoder.parameters()).dtype - image_inputs = components.feature_extractor.preprocess( - images=block_state.image, do_resize=True, return_tensors="pt", do_convert_rgb=True - ) - image_inputs = image_inputs.to(device=device, dtype=image_encoder_dtype) - image_embeds = components.image_encoder(**image_inputs).last_hidden_state - - block_state.image_embeds = image_embeds - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/hunyuan_video1_5/modular_blocks_hunyuan_video1_5.py b/diffusers/modular_pipelines/hunyuan_video1_5/modular_blocks_hunyuan_video1_5.py deleted file mode 100644 index bdbdba1ecdd914e7301f73ca5a8a812f89c21342..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/hunyuan_video1_5/modular_blocks_hunyuan_video1_5.py +++ /dev/null @@ -1,535 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from ...utils import logging -from ..modular_pipeline import AutoPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import OutputParam -from .before_denoise import ( - HunyuanVideo15Image2VideoPrepareLatentsStep, - HunyuanVideo15PrepareLatentsStep, - HunyuanVideo15SetTimestepsStep, - HunyuanVideo15TextInputStep, -) -from .decoders import HunyuanVideo15VaeDecoderStep -from .denoise import HunyuanVideo15DenoiseStep, HunyuanVideo15Image2VideoDenoiseStep -from .encoders import ( - HunyuanVideo15ImageEncoderStep, - HunyuanVideo15TextEncoderStep, - HunyuanVideo15VaeEncoderStep, -) - - -logger = logging.get_logger(__name__) - - -# auto_docstring -class HunyuanVideo15CoreDenoiseStep(SequentialPipelineBlocks): - """ - Denoise block that takes encoded conditions and runs the denoising process. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`HunyuanVideo15Transformer3DModel`) - video_processor (`HunyuanVideo15ImageProcessor`) guider (`ClassifierFreeGuidance`) - - Inputs: - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - batch_size (`int`, *optional*): - Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt. Can - be generated in input step. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - num_frames (`int`, *optional*, defaults to 121): - Number of video frames to generate. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - negative_prompt_embeds (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds' field fed into the guider. - prompt_embeds_mask (`Tensor`): - Positive branch of the 'prompt_embeds_mask' field fed into the guider. - negative_prompt_embeds_mask (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds_mask' field fed into the guider. - prompt_embeds_2 (`Tensor`): - Positive branch of the 'prompt_embeds_2' field fed into the guider. - negative_prompt_embeds_2 (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds_2' field fed into the guider. - prompt_embeds_mask_2 (`Tensor`): - Positive branch of the 'prompt_embeds_mask_2' field fed into the guider. - negative_prompt_embeds_mask_2 (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds_mask_2' field fed into the guider. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "hunyuan-video-1.5" - block_classes = [ - HunyuanVideo15TextInputStep, - HunyuanVideo15SetTimestepsStep, - HunyuanVideo15PrepareLatentsStep, - HunyuanVideo15DenoiseStep, - ] - block_names = ["input", "set_timesteps", "prepare_latents", "denoise"] - - @property - def description(self): - return "Denoise block that takes encoded conditions and runs the denoising process." - - @property - def outputs(self): - return [OutputParam.template("latents")] - - -# auto_docstring -class HunyuanVideo15Blocks(SequentialPipelineBlocks): - """ - Modular pipeline blocks for HunyuanVideo 1.5 text-to-video. - - Components: - text_encoder (`Qwen2_5_VLTextModel`) tokenizer (`Qwen2Tokenizer`) text_encoder_2 (`T5EncoderModel`) - tokenizer_2 (`ByT5Tokenizer`) guider (`ClassifierFreeGuidance`) scheduler (`FlowMatchEulerDiscreteScheduler`) - transformer (`HunyuanVideo15Transformer3DModel`) video_processor (`HunyuanVideo15ImageProcessor`) vae - (`AutoencoderKLHunyuanVideo15`) - - Inputs: - prompt (`str`, *optional*): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The prompt or prompts not to guide the image generation. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - batch_size (`int`, *optional*): - Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt. Can - be generated in input step. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - num_frames (`int`, *optional*, defaults to 121): - Number of video frames to generate. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - output_type (`str`, *optional*, defaults to np): - Output format: 'pil', 'np', 'pt'. - - Outputs: - videos (`list`): - The generated videos. - """ - - model_name = "hunyuan-video-1.5" - block_classes = [ - HunyuanVideo15TextEncoderStep, - HunyuanVideo15CoreDenoiseStep, - HunyuanVideo15VaeDecoderStep, - ] - block_names = ["text_encoder", "denoise", "decode"] - - @property - def description(self): - return "Modular pipeline blocks for HunyuanVideo 1.5 text-to-video." - - @property - def outputs(self): - return [OutputParam.template("videos")] - - -# auto_docstring -class HunyuanVideo15Image2VideoCoreDenoiseStep(SequentialPipelineBlocks): - """ - Denoise block for image-to-video that takes encoded conditions and runs the denoising process. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`HunyuanVideo15Transformer3DModel`) - video_processor (`HunyuanVideo15ImageProcessor`) guider (`ClassifierFreeGuidance`) - - Inputs: - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - batch_size (`int`, *optional*): - Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt. Can - be generated in input step. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - num_frames (`int`, *optional*, defaults to 121): - Number of video frames to generate. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image_latents (`Tensor`): - Pre-encoded image latents from the VAE encoder step, used as conditioning for I2V. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - negative_prompt_embeds (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds' field fed into the guider. - prompt_embeds_mask (`Tensor`): - Positive branch of the 'prompt_embeds_mask' field fed into the guider. - negative_prompt_embeds_mask (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds_mask' field fed into the guider. - prompt_embeds_2 (`Tensor`): - Positive branch of the 'prompt_embeds_2' field fed into the guider. - negative_prompt_embeds_2 (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds_2' field fed into the guider. - prompt_embeds_mask_2 (`Tensor`): - Positive branch of the 'prompt_embeds_mask_2' field fed into the guider. - negative_prompt_embeds_mask_2 (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds_mask_2' field fed into the guider. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "hunyuan-video-1.5" - block_classes = [ - HunyuanVideo15TextInputStep, - HunyuanVideo15SetTimestepsStep, - HunyuanVideo15PrepareLatentsStep, - HunyuanVideo15Image2VideoPrepareLatentsStep, - HunyuanVideo15Image2VideoDenoiseStep, - ] - block_names = ["input", "set_timesteps", "prepare_latents", "prepare_i2v_latents", "denoise"] - - @property - def description(self): - return "Denoise block for image-to-video that takes encoded conditions and runs the denoising process." - - @property - def outputs(self): - return [OutputParam.template("latents")] - - -# auto_docstring -class HunyuanVideo15AutoVaeEncoderStep(AutoPipelineBlocks): - """ - VAE encoder step that encodes the image input into its latent representation. - This is an auto pipeline block that works for image-to-video tasks. - - `HunyuanVideo15VaeEncoderStep` is used when `image` is provided. - - If `image` is not provided, step will be skipped. - - Components: - vae (`AutoencoderKLHunyuanVideo15`) video_processor (`HunyuanVideo15ImageProcessor`) - - Inputs: - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - - Outputs: - image_latents (`Tensor`): - Encoded image latents from the VAE encoder - height (`int`): - Target height resolved from image - width (`int`): - Target width resolved from image - """ - - model_name = "hunyuan-video-1.5" - block_classes = [HunyuanVideo15VaeEncoderStep] - block_names = ["vae_encoder"] - block_trigger_inputs = ["image"] - - @property - def description(self): - return ( - "VAE encoder step that encodes the image input into its latent representation.\n" - "This is an auto pipeline block that works for image-to-video tasks.\n" - " - `HunyuanVideo15VaeEncoderStep` is used when `image` is provided.\n" - " - If `image` is not provided, step will be skipped." - ) - - -# auto_docstring -class HunyuanVideo15AutoImageEncoderStep(AutoPipelineBlocks): - """ - Siglip image encoder step that produces image_embeds. - This is an auto pipeline block that works for image-to-video tasks. - - `HunyuanVideo15ImageEncoderStep` is used when `image` is provided. - - If `image` is not provided, step will be skipped. - - Components: - image_encoder (`SiglipVisionModel`) feature_extractor (`SiglipImageProcessor`) - - Inputs: - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - - Outputs: - image_embeds (`Tensor`): - Image embeddings from the Siglip vision encoder - """ - - model_name = "hunyuan-video-1.5" - block_classes = [HunyuanVideo15ImageEncoderStep] - block_names = ["image_encoder"] - block_trigger_inputs = ["image"] - - @property - def description(self): - return ( - "Siglip image encoder step that produces image_embeds.\n" - "This is an auto pipeline block that works for image-to-video tasks.\n" - " - `HunyuanVideo15ImageEncoderStep` is used when `image` is provided.\n" - " - If `image` is not provided, step will be skipped." - ) - - -# auto_docstring -class HunyuanVideo15AutoCoreDenoiseStep(AutoPipelineBlocks): - """ - Auto denoise block that selects the appropriate denoise pipeline based on inputs. - - `HunyuanVideo15Image2VideoCoreDenoiseStep` is used when `image_latents` is provided. - - `HunyuanVideo15CoreDenoiseStep` is used otherwise (text-to-video). - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`HunyuanVideo15Transformer3DModel`) - video_processor (`HunyuanVideo15ImageProcessor`) guider (`ClassifierFreeGuidance`) - - Inputs: - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt. Can - be generated in input step. - num_inference_steps (`int`): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - num_frames (`int`, *optional*, defaults to 121): - Number of video frames to generate. - latents (`Tensor`): - Pre-generated noisy latents for image generation. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image_latents (`Tensor`, *optional*): - Pre-encoded image latents from the VAE encoder step, used as conditioning for I2V. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - negative_prompt_embeds (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds' field fed into the guider. - prompt_embeds_mask (`Tensor`): - Positive branch of the 'prompt_embeds_mask' field fed into the guider. - negative_prompt_embeds_mask (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds_mask' field fed into the guider. - prompt_embeds_2 (`Tensor`): - Positive branch of the 'prompt_embeds_2' field fed into the guider. - negative_prompt_embeds_2 (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds_2' field fed into the guider. - prompt_embeds_mask_2 (`Tensor`): - Positive branch of the 'prompt_embeds_mask_2' field fed into the guider. - negative_prompt_embeds_mask_2 (`Tensor`, *optional*): - Negative branch of the 'negative_prompt_embeds_mask_2' field fed into the guider. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "hunyuan-video-1.5" - block_classes = [HunyuanVideo15Image2VideoCoreDenoiseStep, HunyuanVideo15CoreDenoiseStep] - block_names = ["image2video", "text2video"] - block_trigger_inputs = ["image_latents", None] - - @property - def description(self): - return ( - "Auto denoise block that selects the appropriate denoise pipeline based on inputs.\n" - " - `HunyuanVideo15Image2VideoCoreDenoiseStep` is used when `image_latents` is provided.\n" - " - `HunyuanVideo15CoreDenoiseStep` is used otherwise (text-to-video)." - ) - - -# auto_docstring -class HunyuanVideo15AutoBlocks(SequentialPipelineBlocks): - """ - Auto blocks for HunyuanVideo 1.5 that support both text-to-video and image-to-video workflows. - - Supported workflows: - - `text2video`: requires `prompt` - - `image2video`: requires `image`, `prompt` - - Components: - text_encoder (`Qwen2_5_VLTextModel`) tokenizer (`Qwen2Tokenizer`) text_encoder_2 (`T5EncoderModel`) - tokenizer_2 (`ByT5Tokenizer`) guider (`ClassifierFreeGuidance`) vae (`AutoencoderKLHunyuanVideo15`) - video_processor (`HunyuanVideo15ImageProcessor`) image_encoder (`SiglipVisionModel`) feature_extractor - (`SiglipImageProcessor`) scheduler (`FlowMatchEulerDiscreteScheduler`) transformer - (`HunyuanVideo15Transformer3DModel`) - - Inputs: - prompt (`str`, *optional*): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The prompt or prompts not to guide the image generation. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - batch_size (`int`): - Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt. Can - be generated in input step. - num_inference_steps (`int`): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - num_frames (`int`, *optional*, defaults to 121): - Number of video frames to generate. - latents (`Tensor`): - Pre-generated noisy latents for image generation. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image_latents (`Tensor`, *optional*): - Pre-encoded image latents from the VAE encoder step, used as conditioning for I2V. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - output_type (`str`, *optional*, defaults to np): - Output format: 'pil', 'np', 'pt'. - - Outputs: - videos (`list`): - The generated videos. - """ - - model_name = "hunyuan-video-1.5" - block_classes = [ - HunyuanVideo15TextEncoderStep, - HunyuanVideo15AutoVaeEncoderStep, - HunyuanVideo15AutoImageEncoderStep, - HunyuanVideo15AutoCoreDenoiseStep, - HunyuanVideo15VaeDecoderStep, - ] - block_names = ["text_encoder", "vae_encoder", "image_encoder", "denoise", "decode"] - _workflow_map = { - "text2video": {"prompt": True}, - "image2video": {"image": True, "prompt": True}, - } - - @property - def description(self): - return "Auto blocks for HunyuanVideo 1.5 that support both text-to-video and image-to-video workflows." - - @property - def outputs(self): - return [OutputParam.template("videos")] - - -# auto_docstring -class HunyuanVideo15Image2VideoBlocks(SequentialPipelineBlocks): - """ - Modular pipeline blocks for HunyuanVideo 1.5 image-to-video. - - Components: - text_encoder (`Qwen2_5_VLTextModel`) tokenizer (`Qwen2Tokenizer`) text_encoder_2 (`T5EncoderModel`) - tokenizer_2 (`ByT5Tokenizer`) guider (`ClassifierFreeGuidance`) vae (`AutoencoderKLHunyuanVideo15`) - video_processor (`HunyuanVideo15ImageProcessor`) image_encoder (`SiglipVisionModel`) feature_extractor - (`SiglipImageProcessor`) scheduler (`FlowMatchEulerDiscreteScheduler`) transformer - (`HunyuanVideo15Transformer3DModel`) - - Inputs: - prompt (`str`, *optional*): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The prompt or prompts not to guide the image generation. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - batch_size (`int`, *optional*): - Number of prompts, the final batch size of model inputs should be batch_size * num_images_per_prompt. Can - be generated in input step. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - num_frames (`int`, *optional*, defaults to 121): - Number of video frames to generate. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image_latents (`Tensor`): - Pre-encoded image latents from the VAE encoder step, used as conditioning for I2V. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - output_type (`str`, *optional*, defaults to np): - Output format: 'pil', 'np', 'pt'. - - Outputs: - videos (`list`): - The generated videos. - """ - - model_name = "hunyuan-video-1.5" - block_classes = [ - HunyuanVideo15TextEncoderStep, - HunyuanVideo15AutoVaeEncoderStep, - HunyuanVideo15AutoImageEncoderStep, - HunyuanVideo15Image2VideoCoreDenoiseStep, - HunyuanVideo15VaeDecoderStep, - ] - block_names = ["text_encoder", "vae_encoder", "image_encoder", "denoise", "decode"] - - @property - def description(self): - return "Modular pipeline blocks for HunyuanVideo 1.5 image-to-video." - - @property - def outputs(self): - return [OutputParam.template("videos")] diff --git a/diffusers/modular_pipelines/hunyuan_video1_5/modular_pipeline.py b/diffusers/modular_pipelines/hunyuan_video1_5/modular_pipeline.py deleted file mode 100644 index e83aa33f201dc418f390d99a0427aa5fcb6f2ce6..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/hunyuan_video1_5/modular_pipeline.py +++ /dev/null @@ -1,90 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from ...loaders import HunyuanVideoLoraLoaderMixin -from ...utils import logging -from ..modular_pipeline import ModularPipeline - - -logger = logging.get_logger(__name__) - - -class HunyuanVideo15ModularPipeline( - ModularPipeline, - HunyuanVideoLoraLoaderMixin, -): - """ - A ModularPipeline for HunyuanVideo 1.5. - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "HunyuanVideo15AutoBlocks" - - @property - def vae_scale_factor_spatial(self): - return self.vae.spatial_compression_ratio if getattr(self, "vae", None) else 16 - - @property - def vae_scale_factor_temporal(self): - return self.vae.temporal_compression_ratio if getattr(self, "vae", None) else 4 - - @property - def num_channels_latents(self): - return self.vae.config.latent_channels if getattr(self, "vae", None) else 32 - - @property - def target_size(self): - return self.transformer.config.target_size if getattr(self, "transformer", None) else 640 - - @property - def default_aspect_ratio(self): - return (16, 9) - - @property - def vision_num_semantic_tokens(self): - return 729 - - @property - def vision_states_dim(self): - return self.transformer.config.image_embed_dim if getattr(self, "transformer", None) else 1152 - - @property - def tokenizer_max_length(self): - return 1000 - - @property - def tokenizer_2_max_length(self): - return 256 - - # fmt: off - @property - def system_message(self): - return "You are a helpful assistant. Describe the video by detailing the following aspects: \ - 1. The main content and theme of the video. \ - 2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects. \ - 3. Actions, events, behaviors temporal relationships, physical movement changes of the objects. \ - 4. background environment, light, style and atmosphere. \ - 5. camera angles, movements, and transitions used in the video." - # fmt: on - - @property - def prompt_template_encode_start_idx(self): - return 108 - - @property - def requires_unconditional_embeds(self): - if hasattr(self, "guider") and self.guider is not None: - return self.guider._enabled and self.guider.num_conditions > 1 - return False diff --git a/diffusers/modular_pipelines/ideogram4/__init__.py b/diffusers/modular_pipelines/ideogram4/__init__.py deleted file mode 100644 index c7c733dda1418a238edeb0408a28b85b46a9dfa0..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ideogram4/__init__.py +++ /dev/null @@ -1,47 +0,0 @@ -from typing import TYPE_CHECKING - -from ...utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, -) - - -_dummy_objects = {} -_import_structure = {} - -try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ...utils import dummy_torch_and_transformers_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) -else: - _import_structure["modular_blocks_ideogram4"] = ["Ideogram4AutoBlocks"] - _import_structure["modular_pipeline"] = ["Ideogram4ModularPipeline"] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .modular_blocks_ideogram4 import Ideogram4AutoBlocks - from .modular_pipeline import Ideogram4ModularPipeline -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/ideogram4/before_denoise.py b/diffusers/modular_pipelines/ideogram4/before_denoise.py deleted file mode 100644 index 98be3b141aecf156722a4f2d767f8d11b0ff62b2..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ideogram4/before_denoise.py +++ /dev/null @@ -1,558 +0,0 @@ -# Copyright 2026 Ideogram AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import math - -import torch - -from ...models.transformers.transformer_ideogram4 import ( - IMAGE_POSITION_OFFSET, - LLM_TOKEN_INDICATOR, - OUTPUT_IMAGE_INDICATOR, - SEQUENCE_PADDING_INDICATOR, - Ideogram4Transformer2DModel, -) -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Ideogram4ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - -# Default per-step guidance schedule (length must equal `num_inference_steps`): 7.0 for the main steps, -# dropping to 3.0 for the final 3 "polish" steps. -DEFAULT_GUIDANCE_SCHEDULE = (7.0,) * 45 + (3.0,) * 3 - - -# Copied from diffusers.pipelines.ideogram4.pipeline_ideogram4._logit_normal_sigmas -def _logit_normal_sigmas( - num_inference_steps: int, - mu: float, - std: float = 1.0, - logsnr_min: float = -15.0, - logsnr_max: float = 18.0, - device: torch.device | None = None, -) -> torch.Tensor: - r""" - Build a length-`num_inference_steps` sigma schedule using the Ideogram4 logit-normal flow-matching schedule. - - Sigmas are returned in `[0, 1]` in decreasing order (sigma close to 1 corresponds to pure noise, sigma close to 0 - to clean data), matching diffusers conventions. - - The Ideogram4 schedule applies `sigma(s) = 1 - logit_normal_cdf_inverse(1 - s)` to `s = linspace(0, 1, N + 1)` and - keeps the first `N` entries; a terminal zero is appended downstream by the scheduler. - """ - intervals = torch.linspace(0.0, 1.0, num_inference_steps + 1, dtype=torch.float64) - # Apply the inverse CDF of a normal then push through the logistic to obtain a logit-normal CDF inverse. - z = torch.special.ndtri(intervals) - y = mu + std * z - t = 1.0 - torch.special.expit(y) - t_min = 1.0 / (1.0 + math.exp(0.5 * logsnr_max)) - t_max = 1.0 / (1.0 + math.exp(0.5 * logsnr_min)) - t = t.clamp(t_min, t_max) - # Convert from model time (0 = noise, 1 = data) to diffusers sigma (1 = noise, 0 = data) and reverse. - sigmas = (1.0 - t).flip(0) - # Drop the trailing 0; FlowMatchEulerDiscreteScheduler.set_timesteps appends one back internally. - sigmas = sigmas[:-1].to(dtype=torch.float32, device=device) - return sigmas - - -# Copied from diffusers.pipelines.ideogram4.pipeline_ideogram4._resolution_aware_mu -def _resolution_aware_mu( - height: int, - width: int, - base_mu: float, - base_resolution: tuple[int, int] = (512, 512), -) -> float: - """Shift the schedule mean as a function of image resolution.""" - num_pixels = height * width - base_pixels = base_resolution[0] * base_resolution[1] - return base_mu + 0.5 * math.log(num_pixels / base_pixels) - - -# Copied from diffusers.pipelines.ideogram4.pipeline_ideogram4._expand_tensor_to_effective_batch -def _expand_tensor_to_effective_batch( - tensor: torch.Tensor, - batch_size: int, - num_per_prompt: int, - tensor_name: str | None = None, -) -> torch.Tensor: - """Replicate `tensor` along dim 0 from `batch_size` (or 1) to `batch_size * num_per_prompt`.""" - target_batch_size = batch_size * num_per_prompt - - if tensor.shape[0] == target_batch_size: - return tensor - - if tensor.shape[0] == 1: - repeat_by = target_batch_size - elif tensor.shape[0] == batch_size: - repeat_by = num_per_prompt - else: - tensor_name = f"`{tensor_name}`" if tensor_name is not None else "Tensor" - raise ValueError( - f"{tensor_name} batch size must be 1, `batch_size` ({batch_size}), or " - f"`batch_size * num_*_per_prompt` ({target_batch_size}), but got {tensor.shape[0]}." - ) - - return torch.repeat_interleave(tensor, repeats=repeat_by, dim=0, output_size=tensor.shape[0] * repeat_by) - - -# auto_docstring -class Ideogram4TextInputsStep(ModularPipelineBlocks): - """ - Input step that determines `batch_size`/`dtype` from the per-prompt `text_features` and replicates the text outputs - to `batch_size * num_images_per_prompt`. Place after the text encoder. - - Inputs: - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - text_features (`Tensor`): - Per-prompt text features from the encoder. - text_lengths (`list`): - Per-prompt text-token counts from the encoder. - - Outputs: - batch_size (`int`): - Effective batch size (num prompts * num_images_per_prompt). - dtype (`dtype`): - The dtype of the text features. - text_features (`Tensor`): - Text features, batch-expanded. - text_lengths (`list`): - Text-token counts, batch-expanded. - """ - - model_name = "ideogram4" - - @property - def description(self) -> str: - return ( - "Input step that determines `batch_size`/`dtype` from the per-prompt `text_features` and replicates the " - "text outputs to `batch_size * num_images_per_prompt`. Place after the text encoder." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_images_per_prompt", default=1), - InputParam( - name="text_features", - required=True, - type_hint=torch.Tensor, - description="Per-prompt text features from the encoder.", - ), - InputParam( - name="text_lengths", - required=True, - type_hint=list, - description="Per-prompt text-token counts from the encoder.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="batch_size", - type_hint=int, - description="Effective batch size (num prompts * num_images_per_prompt).", - ), - OutputParam(name="dtype", type_hint=torch.dtype, description="The dtype of the text features."), - OutputParam(name="text_features", type_hint=torch.Tensor, description="Text features, batch-expanded."), - OutputParam(name="text_lengths", type_hint=list, description="Text-token counts, batch-expanded."), - ] - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - prompt_batch = block_state.text_features.shape[0] - num_per_prompt = block_state.num_images_per_prompt - - block_state.dtype = block_state.text_features.dtype - block_state.text_features = _expand_tensor_to_effective_batch( - block_state.text_features, prompt_batch, num_per_prompt, "text_features" - ) - block_state.text_lengths = [n for n in block_state.text_lengths for _ in range(num_per_prompt)] - block_state.batch_size = prompt_batch * num_per_prompt - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Ideogram4PrepareLatentsStep(ModularPipelineBlocks): - """ - Step that prepares the packed image latents (B, num_image_tokens, latent_dim) for the denoising loop. - - Components: - transformer (`Ideogram4Transformer2DModel`) - - Inputs: - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - height (`int`): - The height in pixels of the generated image. - width (`int`): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - batch_size (`int`): - Effective batch size. - - Outputs: - latents (`Tensor`): - The initial packed image latents (B, num_image_tokens, latent_dim). - num_image_tokens (`int`): - Number of image tokens (grid_h * grid_w). - """ - - model_name = "ideogram4" - - @property - def description(self) -> str: - return "Step that prepares the packed image latents (B, num_image_tokens, latent_dim) for the denoising loop." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Ideogram4Transformer2DModel)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents"), - InputParam.template("height", required=True), - InputParam.template("width", required=True), - InputParam.template("generator"), - InputParam(name="batch_size", required=True, type_hint=int, description="Effective batch size."), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="latents", - type_hint=torch.Tensor, - description="The initial packed image latents (B, num_image_tokens, latent_dim).", - ), - OutputParam( - name="num_image_tokens", type_hint=int, description="Number of image tokens (grid_h * grid_w)." - ), - ] - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - patch = components.patch_size - grid_h = block_state.height // (components.vae_scale_factor * patch) - grid_w = block_state.width // (components.vae_scale_factor * patch) - num_image_tokens = grid_h * grid_w - latent_dim = components.transformer.config.in_channels - - shape = (block_state.batch_size, num_image_tokens, latent_dim) - if block_state.latents is None: - block_state.latents = randn_tensor( - shape, generator=block_state.generator, device=device, dtype=torch.float32 - ) - else: - block_state.latents = block_state.latents.to(device=device, dtype=torch.float32) - - block_state.num_image_tokens = num_image_tokens - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Ideogram4SetTimestepsStep(ModularPipelineBlocks): - """ - Step that sets the resolution-aware logit-normal sigma schedule on the scheduler and resolves the per-step guidance - weights. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - num_inference_steps (`int`, *optional*, defaults to 48): - The number of denoising steps. - height (`int`): - The height in pixels of the generated image. - width (`int`): - The width in pixels of the generated image. - mu (`float`, *optional*, defaults to 0.0): - Base mean of the logit-normal schedule. - std (`float`, *optional*, defaults to 1.5): - Std of the logit-normal schedule. - guidance_schedule (`list`, *optional*, defaults to (7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, - 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, - 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 3.0, 3.0, 3.0)): - Per-step guidance scale schedule (length num_inference_steps). - - Outputs: - timesteps (`Tensor`): - The denoising timesteps. - gw (`Tensor`): - Per-step guidance weights (num_inference_steps,). - """ - - model_name = "ideogram4" - - @property - def description(self) -> str: - return ( - "Step that sets the resolution-aware logit-normal sigma schedule on the scheduler and resolves the " - "per-step guidance weights." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_inference_steps", default=48), - InputParam.template("height", required=True), - InputParam.template("width", required=True), - InputParam(name="mu", default=0.0, type_hint=float, description="Base mean of the logit-normal schedule."), - InputParam(name="std", default=1.5, type_hint=float, description="Std of the logit-normal schedule."), - InputParam( - name="guidance_schedule", - default=DEFAULT_GUIDANCE_SCHEDULE, - type_hint=list, - description="Per-step guidance scale schedule (length num_inference_steps).", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam(name="timesteps", type_hint=torch.Tensor, description="The denoising timesteps."), - OutputParam( - name="gw", type_hint=torch.Tensor, description="Per-step guidance weights (num_inference_steps,)." - ), - ] - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - if len(block_state.guidance_schedule) != block_state.num_inference_steps: - raise ValueError( - f"`guidance_schedule` must have length `num_inference_steps` ({block_state.num_inference_steps}), " - f"got {len(block_state.guidance_schedule)}." - ) - - schedule_mu = _resolution_aware_mu(height=block_state.height, width=block_state.width, base_mu=block_state.mu) - sigmas = _logit_normal_sigmas(block_state.num_inference_steps, schedule_mu, std=block_state.std, device=device) - components.scheduler.set_timesteps(sigmas=sigmas.tolist(), device=device) - - block_state.timesteps = components.scheduler.timesteps - block_state.gw = torch.as_tensor(block_state.guidance_schedule, dtype=torch.float32, device=device) - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Ideogram4PrepareAdditionalInputsStep(ModularPipelineBlocks): - """ - Step that prepares the additional denoiser inputs from the packed-sequence layout: the conditional - encoder_hidden_states (text features packed with image padding) and the position_ids/segment_ids/indicator, plus - the unconditional (image-only) counterparts. Place after prepare_latents. - - Inputs: - height (`int`): - The height in pixels of the generated image. - width (`int`): - The width in pixels of the generated image. - text_features (`Tensor`): - Batch-expanded text features. - text_lengths (`list`): - Batch-expanded text-token counts. - batch_size (`int`): - Effective batch size. - - Outputs: - prompt_embeds (`Tensor`): - Packed conditional encoder_hidden_states (B, total_seq, dim). - position_ids (`Tensor`): - Conditional 3-axis MRoPE position ids. - segment_ids (`Tensor`): - Conditional block-diagonal segment ids. - indicator (`Tensor`): - Conditional per-token text/image/pad role. - negative_prompt_embeds (`Tensor`): - Unconditional (zeroed) text features (B, num_image_tokens, dim). - negative_position_ids (`Tensor`): - Unconditional position ids (image region). - negative_segment_ids (`Tensor`): - Unconditional segment ids (image region). - negative_indicator (`Tensor`): - Unconditional indicator (image region). - """ - - model_name = "ideogram4" - - @property - def description(self) -> str: - return ( - "Step that prepares the additional denoiser inputs from the packed-sequence layout: the conditional " - "encoder_hidden_states (text features packed with image padding) and the position_ids/segment_ids/" - "indicator, plus the unconditional (image-only) counterparts. Place after prepare_latents." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("height", required=True), - InputParam.template("width", required=True), - InputParam( - name="text_features", - required=True, - type_hint=torch.Tensor, - description="Batch-expanded text features.", - ), - InputParam( - name="text_lengths", required=True, type_hint=list, description="Batch-expanded text-token counts." - ), - InputParam(name="batch_size", required=True, type_hint=int, description="Effective batch size."), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="prompt_embeds", - type_hint=torch.Tensor, - description="Packed conditional encoder_hidden_states (B, total_seq, dim).", - ), - OutputParam( - name="position_ids", type_hint=torch.Tensor, description="Conditional 3-axis MRoPE position ids." - ), - OutputParam( - name="segment_ids", type_hint=torch.Tensor, description="Conditional block-diagonal segment ids." - ), - OutputParam( - name="indicator", type_hint=torch.Tensor, description="Conditional per-token text/image/pad role." - ), - OutputParam( - name="negative_prompt_embeds", - type_hint=torch.Tensor, - description="Unconditional (zeroed) text features (B, num_image_tokens, dim).", - ), - OutputParam( - name="negative_position_ids", - type_hint=torch.Tensor, - description="Unconditional position ids (image region).", - ), - OutputParam( - name="negative_segment_ids", - type_hint=torch.Tensor, - description="Unconditional segment ids (image region).", - ), - OutputParam( - name="negative_indicator", - type_hint=torch.Tensor, - description="Unconditional indicator (image region).", - ), - ] - - @staticmethod - # Copied from diffusers.pipelines.ideogram4.pipeline_ideogram4.Ideogram4Pipeline._prepare_ids - def _prepare_ids( - text_lengths: list[int], - grid_h: int, - grid_w: int, - max_text_tokens: int, - device: torch.device, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Build the packed `[left-pad][text][image]` layout from the per-prompt text lengths and the image grid. - - Returns `position_ids` (3-axis MRoPE), `segment_ids` (block-diagonal attention) and `indicator` (per-token - text/image/pad role). - """ - batch_size = len(text_lengths) - num_image_tokens = grid_h * grid_w - total_seq_len = max_text_tokens + num_image_tokens - - # Image position ids (t=0, h, w); offset keeps them disjoint from text positions. - h_idx = torch.arange(grid_h).view(-1, 1).expand(grid_h, grid_w).reshape(-1) - w_idx = torch.arange(grid_w).view(1, -1).expand(grid_h, grid_w).reshape(-1) - t_idx = torch.zeros_like(h_idx) - image_pos = torch.stack([t_idx, h_idx, w_idx], dim=1) + IMAGE_POSITION_OFFSET - - position_ids = torch.zeros(batch_size, total_seq_len, 3, dtype=torch.long) - segment_ids = torch.full((batch_size, total_seq_len), SEQUENCE_PADDING_INDICATOR, dtype=torch.long) - indicator = torch.zeros(batch_size, total_seq_len, dtype=torch.long) - - for b, num_text in enumerate(text_lengths): - offset = max_text_tokens - num_text - - text_pos = torch.arange(num_text) - text_pos_3d = torch.stack([text_pos, text_pos, text_pos], dim=1) - position_ids[b, offset : offset + num_text] = text_pos_3d - position_ids[b, offset + num_text :] = image_pos - - indicator[b, offset : offset + num_text] = LLM_TOKEN_INDICATOR - indicator[b, offset + num_text :] = OUTPUT_IMAGE_INDICATOR - - segment_ids[b, offset : offset + num_text + num_image_tokens] = 1 - - return position_ids.to(device), segment_ids.to(device), indicator.to(device) - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - patch = components.patch_size - grid_h = block_state.height // (components.vae_scale_factor * patch) - grid_w = block_state.width // (components.vae_scale_factor * patch) - num_image_tokens = grid_h * grid_w - - text_features = block_state.text_features - max_text_tokens = text_features.shape[1] - feature_dim = text_features.shape[-1] - - position_ids, segment_ids, indicator = self._prepare_ids( - block_state.text_lengths, grid_h, grid_w, max_text_tokens, device - ) - - # Pack the text features into the full sequence; image positions carry no text features. - image_feature_padding = torch.zeros( - block_state.batch_size, num_image_tokens, feature_dim, dtype=text_features.dtype, device=device - ) - block_state.prompt_embeds = torch.cat([text_features, image_feature_padding], dim=1) - - # Unconditional (image-only) branch, derived from the conditioning. - block_state.negative_prompt_embeds = torch.zeros( - block_state.batch_size, num_image_tokens, feature_dim, dtype=text_features.dtype, device=device - ) - block_state.position_ids = position_ids - block_state.segment_ids = segment_ids - block_state.indicator = indicator - block_state.negative_position_ids = position_ids[:, max_text_tokens:] - block_state.negative_segment_ids = segment_ids[:, max_text_tokens:] - block_state.negative_indicator = indicator[:, max_text_tokens:] - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/ideogram4/decoders.py b/diffusers/modular_pipelines/ideogram4/decoders.py deleted file mode 100644 index bf5d69270b7c15a8cfa573970b3e3d867384eded..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ideogram4/decoders.py +++ /dev/null @@ -1,112 +0,0 @@ -# Copyright 2026 Ideogram AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch - -from ...configuration_utils import FrozenDict -from ...image_processor import VaeImageProcessor -from ...models import AutoencoderKLFlux2 -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Ideogram4ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# auto_docstring -class Ideogram4DecodeStep(ModularPipelineBlocks): - """ - Step that decodes the unpatchified (B, ae_channels, H, W) latents into images: de-normalizes with the VAE - batch-norm statistics and decodes through the VAE. - - Components: - vae (`AutoencoderKLFlux2`) image_processor (`VaeImageProcessor`) - - Inputs: - output_type (`str`, *optional*, defaults to pil): - Output format: 'pil', 'np', 'pt'. - latents (`Tensor`): - The unpatchified (B, ae_channels, H, W) latents to decode, from the after-denoise step. - - Outputs: - images (`list`): - Generated images. - """ - - model_name = "ideogram4" - - @property - def description(self) -> str: - return ( - "Step that decodes the unpatchified (B, ae_channels, H, W) latents into images: de-normalizes with the " - "VAE batch-norm statistics and decodes through the VAE." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLFlux2), - ComponentSpec( - "image_processor", - VaeImageProcessor, - config=FrozenDict({"vae_scale_factor": 16}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("output_type", default="pil"), - InputParam( - name="latents", - required=True, - type_hint=torch.Tensor, - description="The unpatchified (B, ae_channels, H, W) latents to decode, from the after-denoise step.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam.template("images")] - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - z = block_state.latents - patch = components.patch_size - ae_channels = z.shape[1] - grid_h, grid_w = z.shape[2] // patch, z.shape[3] // patch - - # VAE bn stores per-channel statistics over the packed channels, laid out as (patch_row, patch_col, - # ae_channel). Reshape them into an (ae_channels, patch, patch) tile and repeat across the grid so the - # denormalization on the unpatchified latents matches the packed-space statistics. - bn_mean = components.vae.bn.running_mean.view(patch, patch, ae_channels).permute(2, 0, 1) - bn_std = torch.sqrt(components.vae.bn.running_var + components.vae.config.batch_norm_eps) - bn_std = bn_std.view(patch, patch, ae_channels).permute(2, 0, 1) - bn_mean = bn_mean.repeat(1, grid_h, grid_w).to(device=z.device, dtype=z.dtype) - bn_std = bn_std.repeat(1, grid_h, grid_w).to(device=z.device, dtype=z.dtype) - z = z * bn_std + bn_mean - - decoded = components.vae.decode(z.to(components.vae.dtype), return_dict=False)[0] - block_state.images = components.image_processor.postprocess( - decoded.float(), output_type=block_state.output_type - ) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/ideogram4/denoise.py b/diffusers/modular_pipelines/ideogram4/denoise.py deleted file mode 100644 index 871db69d344c3383fa613709c7ef01e7643bbc3c..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ideogram4/denoise.py +++ /dev/null @@ -1,363 +0,0 @@ -# Copyright 2026 Ideogram AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch - -from ...models.transformers.transformer_ideogram4 import Ideogram4Transformer2DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ..modular_pipeline import ( - BlockState, - LoopSequentialPipelineBlocks, - ModularPipelineBlocks, - PipelineState, -) -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Ideogram4ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class Ideogram4LoopBeforeDenoiser(ModularPipelineBlocks): - model_name = "ideogram4" - - @property - def description(self) -> str: - return ( - "Within the denoising loop: build the conditional packed input `[text-padding][image latents]` and the " - "model timestep. Compose into the `sub_blocks` of `Ideogram4DenoiseLoopWrapper`." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="latents", required=True, type_hint=torch.Tensor, description="Packed image latents."), - InputParam( - name="position_ids", required=True, type_hint=torch.Tensor, description="Conditional position ids." - ), - InputParam(name="batch_size", required=True, type_hint=int, description="Effective batch size."), - ] - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - # Conditional packed sequence is [text-padding][image latents]; text region length = total - image tokens. - max_text_tokens = block_state.position_ids.shape[1] - block_state.latents.shape[1] - text_z_padding = torch.zeros( - block_state.latents.shape[0], - max_text_tokens, - block_state.latents.shape[-1], - dtype=block_state.latents.dtype, - device=block_state.latents.device, - ) - block_state.pos_z = torch.cat([text_z_padding, block_state.latents], dim=1) - block_state.max_text_tokens = max_text_tokens - - # Map sigma-domain timestep to model time t in [0, 1] (0 = noise, 1 = clean data). - num_train_timesteps = components.scheduler.config.num_train_timesteps - t_model = 1.0 - (t.float() / num_train_timesteps) - block_state.t_model = t_model.expand(block_state.batch_size) - return components, block_state - - -class Ideogram4LoopDenoiser(ModularPipelineBlocks): - model_name = "ideogram4" - - @property - def description(self) -> str: - return ( - "Within the denoising loop: run the conditional `transformer` on the full packed sequence and the " - "`unconditional_transformer` on the image-only sequence, then blend with the per-step guidance weight " - "(asymmetric CFG, no guider). Compose into `Ideogram4DenoiseLoopWrapper`." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("transformer", Ideogram4Transformer2DModel), - ComponentSpec("unconditional_transformer", Ideogram4Transformer2DModel), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Packed conditional encoder_hidden_states.", - ), - InputParam( - name="position_ids", - required=True, - type_hint=torch.Tensor, - description="Conditional 3-axis MRoPE position ids.", - ), - InputParam( - name="segment_ids", - required=True, - type_hint=torch.Tensor, - description="Conditional block-diagonal segment ids.", - ), - InputParam( - name="indicator", - required=True, - type_hint=torch.Tensor, - description="Conditional per-token text/image/pad role.", - ), - InputParam( - name="negative_prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Unconditional (zeroed) text features.", - ), - InputParam( - name="negative_position_ids", - required=True, - type_hint=torch.Tensor, - description="Unconditional position ids (image region).", - ), - InputParam( - name="negative_segment_ids", - required=True, - type_hint=torch.Tensor, - description="Unconditional segment ids (image region).", - ), - InputParam( - name="negative_indicator", - required=True, - type_hint=torch.Tensor, - description="Unconditional indicator (image region).", - ), - InputParam(name="gw", required=True, type_hint=torch.Tensor, description="Per-step guidance weights."), - InputParam(name="latents", required=True, type_hint=torch.Tensor, description="Packed image latents."), - ] - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - transformer = components.transformer - unconditional_transformer = components.unconditional_transformer - - # Conditional pass operates on the full packed sequence; the velocity is the image-token region. - pos_out = transformer( - hidden_states=block_state.pos_z.to(transformer.dtype), - timestep=block_state.t_model.to(transformer.dtype), - encoder_hidden_states=block_state.prompt_embeds.to(transformer.dtype), - position_ids=block_state.position_ids, - segment_ids=block_state.segment_ids, - indicator=block_state.indicator, - return_dict=False, - )[0] - pos_v = pos_out[:, block_state.max_text_tokens :].to(torch.float32) - - # Unconditional pass uses the image-only positions with zeroed text features. - neg_v = unconditional_transformer( - hidden_states=block_state.latents.to(unconditional_transformer.dtype), - timestep=block_state.t_model.to(unconditional_transformer.dtype), - encoder_hidden_states=block_state.negative_prompt_embeds.to(unconditional_transformer.dtype), - position_ids=block_state.negative_position_ids, - segment_ids=block_state.negative_segment_ids, - indicator=block_state.negative_indicator, - return_dict=False, - )[0].to(torch.float32) - - gw_i = block_state.gw[i] - v = gw_i * pos_v + (1.0 - gw_i) * neg_v - # The scheduler integrates `-v` (Ideogram predicts velocity v = x0 - noise). - block_state.noise_pred = -v - return components, block_state - - -class Ideogram4LoopAfterDenoiser(ModularPipelineBlocks): - model_name = "ideogram4" - - @property - def description(self) -> str: - return "Within the denoising loop: scheduler step. Compose into `Ideogram4DenoiseLoopWrapper`." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam(name="latents", type_hint=torch.Tensor, description="The denoised latents.")] - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - block_state.latents = components.scheduler.step( - block_state.noise_pred, t, block_state.latents, return_dict=False - )[0] - return components, block_state - - -# auto_docstring -class Ideogram4DenoiseStep(LoopSequentialPipelineBlocks): - """ - Denoising loop that iteratively denoises the packed image latents over `timesteps`, running both the conditional - and unconditional transformers and blending with the per-step guidance schedule. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Ideogram4Transformer2DModel`) - unconditional_transformer (`Ideogram4Transformer2DModel`) - - Inputs: - timesteps (`Tensor`): - Denoising timesteps from set_timesteps. - num_inference_steps (`int`, *optional*, defaults to 48): - The number of denoising steps. - latents (`Tensor`): - Packed image latents. - position_ids (`Tensor`): - Conditional position ids. - batch_size (`int`): - Effective batch size. - prompt_embeds (`Tensor`): - Packed conditional encoder_hidden_states. - position_ids (`Tensor`): - Conditional 3-axis MRoPE position ids. - segment_ids (`Tensor`): - Conditional block-diagonal segment ids. - indicator (`Tensor`): - Conditional per-token text/image/pad role. - negative_prompt_embeds (`Tensor`): - Unconditional (zeroed) text features. - negative_position_ids (`Tensor`): - Unconditional position ids (image region). - negative_segment_ids (`Tensor`): - Unconditional segment ids (image region). - negative_indicator (`Tensor`): - Unconditional indicator (image region). - gw (`Tensor`): - Per-step guidance weights. - - Outputs: - latents (`Tensor`): - The denoised latents. - """ - - model_name = "ideogram4" - block_classes = [Ideogram4LoopBeforeDenoiser, Ideogram4LoopDenoiser, Ideogram4LoopAfterDenoiser] - block_names = ["before_denoiser", "denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoising loop that iteratively denoises the packed image latents over `timesteps`, running both the " - "conditional and unconditional transformers and blending with the per-step guidance schedule." - ) - - @property - def loop_expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def loop_inputs(self) -> list[InputParam]: - return [ - InputParam( - name="timesteps", - required=True, - type_hint=torch.Tensor, - description="Denoising timesteps from set_timesteps.", - ), - InputParam.template("num_inference_steps", default=48), - ] - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) - progress_bar.update() - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Ideogram4AfterDenoiseStep(ModularPipelineBlocks): - """ - Step that runs after the denoising loop: unpatchifies the packed image latents (B, num_image_tokens, ae_channels * - patch ** 2) into a (B, ae_channels, H, W) latent for the decoder. - - Inputs: - height (`int`): - The height in pixels of the generated image. - width (`int`): - The width in pixels of the generated image. - latents (`Tensor`): - The denoised packed image latents (B, num_image_tokens, latent_dim). - - Outputs: - latents (`Tensor`): - Unpatchified latents (B, ae_channels, H, W) ready for the VAE decoder. - """ - - model_name = "ideogram4" - - @property - def description(self) -> str: - return ( - "Step that runs after the denoising loop: unpatchifies the packed image latents " - "(B, num_image_tokens, ae_channels * patch ** 2) into a (B, ae_channels, H, W) latent for the decoder." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("height", required=True), - InputParam.template("width", required=True), - InputParam( - name="latents", - required=True, - type_hint=torch.Tensor, - description="The denoised packed image latents (B, num_image_tokens, latent_dim).", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="latents", - type_hint=torch.Tensor, - description="Unpatchified latents (B, ae_channels, H, W) ready for the VAE decoder.", - ) - ] - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - z = block_state.latents - patch = components.patch_size - grid_h = block_state.height // (components.vae_scale_factor * patch) - grid_w = block_state.width // (components.vae_scale_factor * patch) - - ae_channels = z.shape[-1] // (patch * patch) - z = z.view(z.shape[0], grid_h, grid_w, patch, patch, ae_channels) - z = z.permute(0, 5, 1, 3, 2, 4).contiguous() - z = z.view(z.shape[0], ae_channels, grid_h * patch, grid_w * patch) - - block_state.latents = z - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/ideogram4/encoders.py b/diffusers/modular_pipelines/ideogram4/encoders.py deleted file mode 100644 index 6e149fa8392e2d21ad1948154751bfd57ce8b6a7..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ideogram4/encoders.py +++ /dev/null @@ -1,327 +0,0 @@ -# Copyright 2026 Ideogram AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch -from transformers import Qwen2Tokenizer, Qwen3VLModel -from transformers.masking_utils import create_causal_mask - -from ...pipelines.ideogram4.prompt_enhancer import ( - PROMPT_UPSAMPLE_TEMPERATURE, - Ideogram4PromptEnhancerHead, - build_caption_logits_processor, - build_prompt_enhancer, - generate_captions, -) -from ...utils import is_outlines_available, logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Ideogram4ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# Hidden states of these Qwen3-VL decoder layers are concatenated to form the per-token -# text conditioning consumed by the Ideogram4 transformer. -QWEN3_VL_ACTIVATION_LAYERS = (0, 3, 6, 9, 12, 15, 18, 21, 24, 27, 30, 33, 35) - - -# auto_docstring -class Ideogram4PromptUpsampleStep(ModularPipelineBlocks): - """ - Optional step that rewrites the prompt(s) into Ideogram4's native structured JSON caption when - `prompt_upsampling=True` (the format the model is trained on). Requires a generative `text_encoder` (a - `Qwen3VLForConditionalGeneration`); install `outlines` for schema-constrained captions. - - Components: - text_encoder (`Qwen3VLModel`): The Qwen3-VL text encoder. tokenizer (`Qwen2Tokenizer`): The tokenizer paired - with the text encoder. prompt_enhancer_head (`Ideogram4PromptEnhancerHead`): LM head grafted onto the text - encoder for prompt upsampling. - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - prompt_upsampling (`bool`, *optional*, defaults to False): - If True, rewrite the prompt into Ideogram4's native JSON caption before encoding. - prompt_upsampling_temperature (`float`, *optional*, defaults to 1.0): - Sampling temperature for prompt upsampling. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - max_sequence_length (`int`, *optional*, defaults to 2048): - Maximum sequence length for prompt encoding. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - - Outputs: - prompt (`list`): - The (possibly upsampled) prompt forwarded to the text encoder. - """ - - model_name = "ideogram4" - - def __init__(self): - # Built lazily on first upsample: the head-less encoder body + `prompt_enhancer_head`, combined. - self._prompt_enhancer = None - # Outlines logits processor for schema-constrained captions; built lazily on first upsample. - self._caption_logits_processor = None - super().__init__() - - @property - def description(self) -> str: - return ( - "Optional step that rewrites the prompt(s) into Ideogram4's native structured JSON caption when " - "`prompt_upsampling=True` (the format the model is trained on). Requires a generative `text_encoder` " - "(a `Qwen3VLForConditionalGeneration`); install `outlines` for schema-constrained captions." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", Qwen3VLModel, description="The Qwen3-VL text encoder."), - ComponentSpec("tokenizer", Qwen2Tokenizer, description="The tokenizer paired with the text encoder."), - ComponentSpec( - "prompt_enhancer_head", - Ideogram4PromptEnhancerHead, - description="LM head grafted onto the text encoder for prompt upsampling.", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt", required=True), - InputParam( - name="prompt_upsampling", - type_hint=bool, - default=False, - description="If True, rewrite the prompt into Ideogram4's native JSON caption before encoding.", - ), - InputParam( - name="prompt_upsampling_temperature", - type_hint=float, - default=PROMPT_UPSAMPLE_TEMPERATURE, - description="Sampling temperature for prompt upsampling.", - ), - InputParam.template("height"), - InputParam.template("width"), - InputParam.template("max_sequence_length", default=2048), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="prompt", - type_hint=list, - description="The (possibly upsampled) prompt forwarded to the text encoder.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - if block_state.prompt_upsampling: - if components.prompt_enhancer_head is None: - raise ValueError( - "Prompt upsampling requires the `prompt_enhancer_head` component, which is not loaded. Load an " - "`Ideogram4PromptEnhancerHead` and add it to the pipeline." - ) - if self._prompt_enhancer is None: - self._prompt_enhancer = build_prompt_enhancer(components.text_encoder, components.prompt_enhancer_head) - if self._caption_logits_processor is None and is_outlines_available(): - self._caption_logits_processor = build_caption_logits_processor( - self._prompt_enhancer, components.tokenizer - ) - if self._caption_logits_processor is None: - logger.warning_once( - "`outlines` is not installed; prompt upsampling runs unconstrained and may not return " - "schema-valid JSON. Install with `pip install outlines` for structured captions." - ) - height = block_state.height or components.default_height - width = block_state.width or components.default_width - block_state.prompt = generate_captions( - self._prompt_enhancer, - components.tokenizer, - self._caption_logits_processor, - block_state.prompt, - height, - width, - temperature=block_state.prompt_upsampling_temperature, - max_new_tokens=block_state.max_sequence_length, - generator=block_state.generator, - device=components._execution_device, - ) - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Ideogram4TextEncoderStep(ModularPipelineBlocks): - """ - Text encoder step that tokenizes the prompt(s) and runs the Qwen3-VL text encoder, returning the per-token text - features (concatenated from a fixed set of activation layers). Only the text tokens are encoded; the packed image - tokens are appended later (the encoder is causal with image after text, so they never affect the text features). - - Components: - text_encoder (`Qwen3VLModel`): The Qwen3-VL text encoder. tokenizer (`Qwen2Tokenizer`): The tokenizer paired - with the text encoder. - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - max_sequence_length (`int`, *optional*, defaults to 2048): - Maximum sequence length for prompt encoding. - - Outputs: - text_features (`Tensor`): - Per-prompt text features (B, max_sequence_length, llm_features_dim), padding zeroed. - text_lengths (`list`): - Per-prompt real text-token counts, used to lay out the packed sequence. - """ - - model_name = "ideogram4" - - @property - def description(self) -> str: - return ( - "Text encoder step that tokenizes the prompt(s) and runs the Qwen3-VL text encoder, returning the " - "per-token text features (concatenated from a fixed set of activation layers). Only the text tokens are " - "encoded; the packed image tokens are appended later (the encoder is causal with image after text, so " - "they never affect the text features)." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", Qwen3VLModel, description="The Qwen3-VL text encoder."), - ComponentSpec("tokenizer", Qwen2Tokenizer, description="The tokenizer paired with the text encoder."), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt", required=True), - InputParam.template("max_sequence_length", default=2048), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="text_features", - type_hint=torch.Tensor, - description="Per-prompt text features (B, max_sequence_length, llm_features_dim), padding zeroed.", - ), - OutputParam( - name="text_lengths", - type_hint=list, - description="Per-prompt real text-token counts, used to lay out the packed sequence.", - ), - ] - - @staticmethod - # Copied from diffusers.pipelines.ideogram4.pipeline_ideogram4.Ideogram4Pipeline._get_text_encoder_hidden_states - def _get_text_encoder_hidden_states( - text_encoder, - token_ids: torch.Tensor, - attention_mask: torch.Tensor, - pos_2d: torch.Tensor, - ) -> list[torch.Tensor]: - """Run the text encoder's decoder layers, returning the hidden states tapped at each activation layer.""" - - language_model = text_encoder.language_model - - inputs_embeds = language_model.embed_tokens(token_ids) - - position_ids_4d = pos_2d[None, ...].expand(4, pos_2d.shape[0], -1) - text_position_ids = position_ids_4d[0] - mrope_position_ids = position_ids_4d[1:] - - causal_mask = create_causal_mask( - config=language_model.config, - inputs_embeds=inputs_embeds, - attention_mask=attention_mask, - past_key_values=None, - position_ids=text_position_ids, - ) - position_embeddings = language_model.rotary_emb(inputs_embeds, mrope_position_ids) - - tap_set = set(QWEN3_VL_ACTIVATION_LAYERS) - captured: dict[int, torch.Tensor] = {} - hidden_states = inputs_embeds - for layer_idx, decoder_layer in enumerate(language_model.layers): - hidden_states = decoder_layer( - hidden_states, - attention_mask=causal_mask, - position_ids=text_position_ids, - past_key_values=None, - position_embeddings=position_embeddings, - ) - if layer_idx in tap_set: - captured[layer_idx] = hidden_states - - return [captured[i] for i in QWEN3_VL_ACTIVATION_LAYERS] - - @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - tokenizer = components.tokenizer - max_text_tokens = block_state.max_sequence_length - - prompts = [block_state.prompt] if isinstance(block_state.prompt, str) else list(block_state.prompt) - batch_size = len(prompts) - - # Tokenize each chat-formatted prompt and left-pad to `max_sequence_length`. - token_ids = torch.zeros(batch_size, max_text_tokens, dtype=torch.long) - attention_mask = torch.zeros(batch_size, max_text_tokens, dtype=torch.long) - text_position_ids = torch.zeros(batch_size, max_text_tokens, dtype=torch.long) - text_lengths = [] - for b, text_prompt in enumerate(prompts): - messages = [{"role": "user", "content": [{"type": "text", "text": text_prompt}]}] - text = tokenizer.apply_chat_template(messages, add_generation_prompt=True, tokenize=False) - toks = tokenizer(text, return_tensors="pt", add_special_tokens=False)["input_ids"][0] - n = int(toks.shape[0]) - if n > max_text_tokens: - raise ValueError(f"prompt has {n} tokens, exceeds max_sequence_length={max_text_tokens}") - text_lengths.append(n) - offset = max_text_tokens - n - token_ids[b, offset:] = toks - attention_mask[b, offset:] = 1 - text_position_ids[b, offset:] = torch.arange(n) - - token_ids = token_ids.to(device) - attention_mask = attention_mask.to(device) - text_position_ids = text_position_ids.to(device) - - # Run the text encoder, tapping the activation-layer hidden states, then concatenate them into per-token - # text features (padding zeroed). - selected = self._get_text_encoder_hidden_states( - components.text_encoder, token_ids, attention_mask, text_position_ids - ) - text_features = torch.stack(selected, dim=0).permute(1, 2, 3, 0).reshape(batch_size, max_text_tokens, -1) - text_features = (text_features * attention_mask.to(text_features.dtype).unsqueeze(-1)).to(torch.float32) - - block_state.text_features = text_features - block_state.text_lengths = text_lengths - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/ideogram4/modular_blocks_ideogram4.py b/diffusers/modular_pipelines/ideogram4/modular_blocks_ideogram4.py deleted file mode 100644 index 0b788fe236be669d79e865101dd74b392db55f71..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ideogram4/modular_blocks_ideogram4.py +++ /dev/null @@ -1,185 +0,0 @@ -# Copyright 2026 Ideogram AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from ...utils import logging -from ..modular_pipeline import SequentialPipelineBlocks -from ..modular_pipeline_utils import InsertableDict, OutputParam -from .before_denoise import ( - Ideogram4PrepareAdditionalInputsStep, - Ideogram4PrepareLatentsStep, - Ideogram4SetTimestepsStep, - Ideogram4TextInputsStep, -) -from .decoders import Ideogram4DecodeStep -from .denoise import Ideogram4AfterDenoiseStep, Ideogram4DenoiseStep -from .encoders import Ideogram4PromptUpsampleStep, Ideogram4TextEncoderStep - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# Core denoise: consumes the per-prompt text features and produces the unpatchified latents -# (batch/latents/timesteps/ids inputs -> denoising loop -> unpatchify). -CORE_DENOISE_BLOCKS = InsertableDict( - [ - ("input", Ideogram4TextInputsStep()), - ("prepare_latents", Ideogram4PrepareLatentsStep()), - ("set_timesteps", Ideogram4SetTimestepsStep()), - ("prepare_additional_inputs", Ideogram4PrepareAdditionalInputsStep()), - ("denoise", Ideogram4DenoiseStep()), - ("after_denoise", Ideogram4AfterDenoiseStep()), - ] -) - - -# auto_docstring -class Ideogram4CoreDenoiseStep(SequentialPipelineBlocks): - """ - Core denoising workflow for Ideogram4 text-to-image: prepares the batch/latents/timesteps and the packed denoiser - inputs, runs the asymmetric-CFG denoising loop over the conditional and unconditional transformers, and - unpatchifies the result for the decoder. - - Components: - transformer (`Ideogram4Transformer2DModel`) scheduler (`FlowMatchEulerDiscreteScheduler`) - unconditional_transformer (`Ideogram4Transformer2DModel`) - - Inputs: - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - text_features (`Tensor`): - Per-prompt text features from the encoder. - text_lengths (`list`): - Per-prompt text-token counts from the encoder. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - height (`int`): - The height in pixels of the generated image. - width (`int`): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`, *optional*, defaults to 48): - The number of denoising steps. - mu (`float`, *optional*, defaults to 0.0): - Base mean of the logit-normal schedule. - std (`float`, *optional*, defaults to 1.5): - Std of the logit-normal schedule. - guidance_schedule (`list`, *optional*, defaults to (7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, - 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, - 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 3.0, 3.0, 3.0)): - Per-step guidance scale schedule (length num_inference_steps). - - Outputs: - latents (`Tensor`): - Unpatchified (B, ae_channels, H, W) latents. - """ - - model_name = "ideogram4" - block_classes = list(CORE_DENOISE_BLOCKS.values()) - block_names = list(CORE_DENOISE_BLOCKS.keys()) - - @property - def description(self) -> str: - return ( - "Core denoising workflow for Ideogram4 text-to-image: prepares the batch/latents/timesteps and the packed " - "denoiser inputs, runs the asymmetric-CFG denoising loop over the conditional and unconditional " - "transformers, and unpatchifies the result for the decoder." - ) - - @property - def outputs(self) -> list[OutputParam]: - # The only meaningful product of the core step is the unpatchified latents; the batch/timesteps/packed-sequence - # inputs prepared along the way are consumed within the loop and are not updated by it. - return [OutputParam.template("latents", description="Unpatchified (B, ae_channels, H, W) latents.")] - - -# auto_docstring -class Ideogram4AutoBlocks(SequentialPipelineBlocks): - """ - Auto Modular pipeline for text-to-image generation using Ideogram4: (optional) prompt upsampling -> encode text -> - core denoise (asymmetric CFG over two transformers) -> decode. - - Supported workflows: - - `text2image`: requires `prompt` - - Components: - text_encoder (`Qwen3VLModel`): The Qwen3-VL text encoder. tokenizer (`Qwen2Tokenizer`): The tokenizer paired - with the text encoder. prompt_enhancer_head (`Ideogram4PromptEnhancerHead`): LM head grafted onto the text - encoder for prompt upsampling. transformer (`Ideogram4Transformer2DModel`) scheduler - (`FlowMatchEulerDiscreteScheduler`) unconditional_transformer (`Ideogram4Transformer2DModel`) vae - (`AutoencoderKLFlux2`) image_processor (`VaeImageProcessor`) - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - prompt_upsampling (`bool`, *optional*, defaults to False): - If True, rewrite the prompt into Ideogram4's native JSON caption before encoding. - prompt_upsampling_temperature (`float`, *optional*, defaults to 1.0): - Sampling temperature for prompt upsampling. - height (`int`, *optional*): - The height in pixels of the generated image. - width (`int`, *optional*): - The width in pixels of the generated image. - max_sequence_length (`int`, *optional*, defaults to 2048): - Maximum sequence length for prompt encoding. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - num_inference_steps (`int`, *optional*, defaults to 48): - The number of denoising steps. - mu (`float`, *optional*, defaults to 0.0): - Base mean of the logit-normal schedule. - std (`float`, *optional*, defaults to 1.5): - Std of the logit-normal schedule. - guidance_schedule (`list`, *optional*, defaults to (7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, - 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, - 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 3.0, 3.0, 3.0)): - Per-step guidance scale schedule (length num_inference_steps). - output_type (`str`, *optional*, defaults to pil): - Output format: 'pil', 'np', 'pt'. - - Outputs: - images (`list`): - Generated images. - """ - - model_name = "ideogram4" - block_classes = [ - Ideogram4PromptUpsampleStep(), - Ideogram4TextEncoderStep(), - Ideogram4CoreDenoiseStep(), - Ideogram4DecodeStep(), - ] - block_names = ["prompt_upsample", "text_encoder", "denoise", "decode"] - - # Workflow map declaring the trigger conditions for each supported workflow. - # `True` means the workflow triggers when the input is not None. - _workflow_map = { - "text2image": {"prompt": True}, - } - - @property - def description(self) -> str: - return ( - "Auto Modular pipeline for text-to-image generation using Ideogram4: (optional) prompt upsampling -> " - "encode text -> core denoise (asymmetric CFG over two transformers) -> decode." - ) - - @property - def outputs(self) -> list[OutputParam]: - return [OutputParam.template("images")] diff --git a/diffusers/modular_pipelines/ideogram4/modular_pipeline.py b/diffusers/modular_pipelines/ideogram4/modular_pipeline.py deleted file mode 100644 index 9c0ff00b880ae97089127a04ebe83a4b34b772c8..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ideogram4/modular_pipeline.py +++ /dev/null @@ -1,46 +0,0 @@ -# Copyright 2026 Ideogram AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from ...loaders import Ideogram4LoraLoaderMixin -from ..modular_pipeline import ModularPipeline - - -class Ideogram4ModularPipeline(ModularPipeline, Ideogram4LoraLoaderMixin): - """ - A ModularPipeline for Ideogram4. - - > [!WARNING] > This is an experimental feature! - """ - - default_blocks_name = "Ideogram4AutoBlocks" - - # Ideogram4 patchifies the VAE output by a factor of 2 before feeding the transformer. - @property - def patch_size(self): - return 2 - - @property - def default_height(self): - return 2048 - - @property - def default_width(self): - return 2048 - - @property - def vae_scale_factor(self): - vae_scale_factor = 8 - if getattr(self, "vae", None) is not None: - vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) - return vae_scale_factor diff --git a/diffusers/modular_pipelines/krea2/__init__.py b/diffusers/modular_pipelines/krea2/__init__.py deleted file mode 100644 index 12e51c7c3018b0e9e147183d6599935024b80b9d..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/krea2/__init__.py +++ /dev/null @@ -1,49 +0,0 @@ -from typing import TYPE_CHECKING - -from ...utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, -) - - -_dummy_objects = {} -_import_structure = {} - -try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ...utils import dummy_torch_and_transformers_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) -else: - _import_structure["modular_blocks_krea2"] = ["Krea2AutoBlocks"] - _import_structure["modular_blocks_krea2_turbo"] = ["Krea2TurboAutoBlocks"] - _import_structure["modular_pipeline"] = ["Krea2ModularPipeline", "Krea2TurboModularPipeline"] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .modular_blocks_krea2 import Krea2AutoBlocks - from .modular_blocks_krea2_turbo import Krea2TurboAutoBlocks - from .modular_pipeline import Krea2ModularPipeline, Krea2TurboModularPipeline -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/krea2/before_denoise.py b/diffusers/modular_pipelines/krea2/before_denoise.py deleted file mode 100644 index 63810d30a9035fa58b564dccdafaf550d92a2b27..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/krea2/before_denoise.py +++ /dev/null @@ -1,590 +0,0 @@ -# Copyright 2026 Krea AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import numpy as np -import torch - -from ...models.transformers.transformer_krea2 import Krea2Transformer2DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Krea2ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# Copied from diffusers.pipelines.krea2.pipeline_krea2.calculate_shift -def calculate_shift( - image_seq_len, - base_seq_len: int = 256, - max_seq_len: int = 4096, - base_shift: float = 0.5, - max_shift: float = 1.15, -): - m = (max_shift - base_shift) / (max_seq_len - base_seq_len) - b = base_shift - m * base_seq_len - mu = image_seq_len * m + b - return mu - - -# auto_docstring -class Krea2TextInputsStep(ModularPipelineBlocks): - """ - Input step that determines `batch_size`/`dtype` from the per-prompt `prompt_embeds` and replicates the text - conditioning (and the optional negative branch) to `batch_size * num_images_per_prompt`. Place after the text - encoder. - - Inputs: - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - prompt_embeds (`Tensor`): - Per-prompt stacked text features (B, text_seq_len, num_text_layers, text_hidden_dim). - prompt_embeds_mask (`Tensor`): - Per-prompt boolean text mask (B, text_seq_len). - negative_prompt_embeds (`Tensor`, *optional*): - Per-prompt negative text features. - negative_prompt_embeds_mask (`Tensor`, *optional*): - Per-prompt negative text mask. - - Outputs: - batch_size (`int`): - Effective batch size (num prompts * num_images_per_prompt). - dtype (`dtype`): - The dtype of the text features. - prompt_embeds (`Tensor`): - Text features, batch-expanded. - prompt_embeds_mask (`Tensor`): - Text mask, batch-expanded. - negative_prompt_embeds (`Tensor`): - Negative text features, batch-expanded. - negative_prompt_embeds_mask (`Tensor`): - Negative text mask, batch-expanded. - """ - - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Input step that determines `batch_size`/`dtype` from the per-prompt `prompt_embeds` and replicates the " - "text conditioning (and the optional negative branch) to `batch_size * num_images_per_prompt`. Place after " - "the text encoder." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_images_per_prompt", default=1), - InputParam( - name="prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Per-prompt stacked text features (B, text_seq_len, num_text_layers, text_hidden_dim).", - ), - InputParam( - name="prompt_embeds_mask", - required=True, - type_hint=torch.Tensor, - description="Per-prompt boolean text mask (B, text_seq_len).", - ), - InputParam( - name="negative_prompt_embeds", - type_hint=torch.Tensor, - description="Per-prompt negative text features.", - ), - InputParam( - name="negative_prompt_embeds_mask", - type_hint=torch.Tensor, - description="Per-prompt negative text mask.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="batch_size", - type_hint=int, - description="Effective batch size (num prompts * num_images_per_prompt).", - ), - OutputParam(name="dtype", type_hint=torch.dtype, description="The dtype of the text features."), - OutputParam(name="prompt_embeds", type_hint=torch.Tensor, description="Text features, batch-expanded."), - OutputParam(name="prompt_embeds_mask", type_hint=torch.Tensor, description="Text mask, batch-expanded."), - OutputParam( - name="negative_prompt_embeds", - type_hint=torch.Tensor, - description="Negative text features, batch-expanded.", - ), - OutputParam( - name="negative_prompt_embeds_mask", - type_hint=torch.Tensor, - description="Negative text mask, batch-expanded.", - ), - ] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - prompt_batch, seq_len, num_layers, dim = block_state.prompt_embeds.shape - n = block_state.num_images_per_prompt - - block_state.dtype = block_state.prompt_embeds.dtype - block_state.batch_size = prompt_batch * n - - block_state.prompt_embeds = block_state.prompt_embeds.repeat(1, n, 1, 1).view( - prompt_batch * n, seq_len, num_layers, dim - ) - block_state.prompt_embeds_mask = block_state.prompt_embeds_mask.repeat(1, n).view(prompt_batch * n, seq_len) - - if block_state.negative_prompt_embeds is not None: - block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.repeat(1, n, 1, 1).view( - prompt_batch * n, seq_len, num_layers, dim - ) - block_state.negative_prompt_embeds_mask = block_state.negative_prompt_embeds_mask.repeat(1, n).view( - prompt_batch * n, seq_len - ) - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Krea2TurboTextInputsStep(ModularPipelineBlocks): - """ - Input step for the distilled Krea 2 turbo checkpoint that determines `batch_size`/`dtype` from the per-prompt - `prompt_embeds` and replicates the text conditioning to `batch_size * num_images_per_prompt`. The distilled - checkpoint runs without classifier-free guidance, so there is no negative branch. Place after the text encoder. - - Inputs: - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - prompt_embeds (`Tensor`): - Per-prompt stacked text features (B, text_seq_len, num_text_layers, text_hidden_dim). - prompt_embeds_mask (`Tensor`): - Per-prompt boolean text mask (B, text_seq_len). - - Outputs: - batch_size (`int`): - Effective batch size (num prompts * num_images_per_prompt). - dtype (`dtype`): - The dtype of the text features. - prompt_embeds (`Tensor`): - Text features, batch-expanded. - prompt_embeds_mask (`Tensor`): - Text mask, batch-expanded. - """ - - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Input step for the distilled Krea 2 turbo checkpoint that determines `batch_size`/`dtype` from the " - "per-prompt `prompt_embeds` and replicates the text conditioning to `batch_size * num_images_per_prompt`. " - "The distilled checkpoint runs without classifier-free guidance, so there is no negative branch. Place " - "after the text encoder." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_images_per_prompt", default=1), - InputParam( - name="prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Per-prompt stacked text features (B, text_seq_len, num_text_layers, text_hidden_dim).", - ), - InputParam( - name="prompt_embeds_mask", - required=True, - type_hint=torch.Tensor, - description="Per-prompt boolean text mask (B, text_seq_len).", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="batch_size", - type_hint=int, - description="Effective batch size (num prompts * num_images_per_prompt).", - ), - OutputParam(name="dtype", type_hint=torch.dtype, description="The dtype of the text features."), - OutputParam(name="prompt_embeds", type_hint=torch.Tensor, description="Text features, batch-expanded."), - OutputParam(name="prompt_embeds_mask", type_hint=torch.Tensor, description="Text mask, batch-expanded."), - ] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - prompt_batch, seq_len, num_layers, dim = block_state.prompt_embeds.shape - n = block_state.num_images_per_prompt - - block_state.dtype = block_state.prompt_embeds.dtype - block_state.batch_size = prompt_batch * n - - block_state.prompt_embeds = block_state.prompt_embeds.repeat(1, n, 1, 1).view( - prompt_batch * n, seq_len, num_layers, dim - ) - block_state.prompt_embeds_mask = block_state.prompt_embeds_mask.repeat(1, n).view(prompt_batch * n, seq_len) - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Krea2PrepareLatentsStep(ModularPipelineBlocks): - """ - Step that samples the spatial image latents and patch-packs them into (B, image_seq_len, in_channels) for the - denoising loop. - - Components: - transformer (`Krea2Transformer2DModel`) - - Inputs: - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - height (`int`, *optional*, defaults to 1024): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 1024): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - batch_size (`int`): - Effective batch size. - dtype (`dtype`): - The working dtype. - - Outputs: - latents (`Tensor`): - The initial packed image latents (B, image_seq_len, in_channels). - image_seq_len (`int`): - Number of image tokens (grid_h * grid_w). - """ - - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Step that samples the spatial image latents and patch-packs them into (B, image_seq_len, in_channels) " - "for the denoising loop." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Krea2Transformer2DModel)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents"), - InputParam.template("height", default=1024), - InputParam.template("width", default=1024), - InputParam.template("generator"), - InputParam(name="batch_size", required=True, type_hint=int, description="Effective batch size."), - InputParam(name="dtype", required=True, type_hint=torch.dtype, description="The working dtype."), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="latents", - type_hint=torch.Tensor, - description="The initial packed image latents (B, image_seq_len, in_channels).", - ), - OutputParam(name="image_seq_len", type_hint=int, description="Number of image tokens (grid_h * grid_w)."), - ] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - p = components.patch_size - num_channels_latents = components.transformer.config.in_channels // (p**2) - - multiple = components.vae_scale_factor * components.patch_size - if block_state.height % multiple != 0 or block_state.width % multiple != 0: - rounded_height = ((block_state.height + multiple - 1) // multiple) * multiple - rounded_width = ((block_state.width + multiple - 1) // multiple) * multiple - logger.warning( - f"`height` and `width` must be multiples of {multiple}; rounding up from {block_state.height}x{block_state.width} to" - f" {rounded_height}x{rounded_width}." - ) - block_state.height, block_state.width = rounded_height, rounded_width - - latent_height = block_state.height // components.vae_scale_factor - latent_width = block_state.width // components.vae_scale_factor - - if block_state.latents is not None: - block_state.latents = block_state.latents.to(device=device, dtype=block_state.dtype) - else: - latents = randn_tensor( - (block_state.batch_size, num_channels_latents, latent_height, latent_width), - generator=block_state.generator, - device=device, - dtype=block_state.dtype, - ) - latents = latents.view( - block_state.batch_size, num_channels_latents, latent_height // p, p, latent_width // p, p - ) - latents = latents.permute(0, 2, 4, 1, 3, 5) - block_state.latents = latents.reshape( - block_state.batch_size, (latent_height // p) * (latent_width // p), num_channels_latents * p * p - ) - - block_state.image_seq_len = block_state.latents.shape[1] - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Krea2SetTimestepsStep(ModularPipelineBlocks): - """ - Step that sets the Krea 2 flow-matching schedule on the scheduler: a linear sigma schedule with a resolution-aware - dynamic time shift `mu`. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - num_inference_steps (`int`, *optional*, defaults to 28): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigma schedule (defaults to a linear ramp). - image_seq_len (`int`): - Number of image tokens, used to compute the resolution-aware shift. - - Outputs: - timesteps (`Tensor`): - The denoising timesteps. - """ - - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Step that sets the Krea 2 flow-matching schedule on the scheduler: a linear sigma schedule with a " - "resolution-aware dynamic time shift `mu`." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_inference_steps", default=28), - InputParam( - name="sigmas", type_hint=list, description="Custom sigma schedule (defaults to a linear ramp)." - ), - InputParam( - name="image_seq_len", - required=True, - type_hint=int, - description="Number of image tokens, used to compute the resolution-aware shift.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam(name="timesteps", type_hint=torch.Tensor, description="The denoising timesteps.")] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - num_inference_steps = block_state.num_inference_steps - - sigmas = block_state.sigmas - if sigmas is None: - sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) - else: - block_state.num_inference_steps = len(sigmas) - - config = components.scheduler.config - mu = calculate_shift( - block_state.image_seq_len, - config.get("base_image_seq_len", 256), - config.get("max_image_seq_len", 6400), - config.get("base_shift", 0.5), - config.get("max_shift", 1.15), - ) - - components.scheduler.set_timesteps(sigmas=sigmas, mu=mu, device=device) - components.scheduler.set_begin_index(0) - block_state.timesteps = components.scheduler.timesteps - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Krea2TurboSetTimestepsStep(ModularPipelineBlocks): - """ - Step that sets the flow-matching schedule for the distilled Krea 2 turbo checkpoint on the scheduler: a linear - sigma schedule with the fixed time shift `mu=1.15` the checkpoint was distilled with. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - num_inference_steps (`int`, *optional*, defaults to 8): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigma schedule (defaults to a linear ramp). - - Outputs: - timesteps (`Tensor`): - The denoising timesteps. - """ - - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Step that sets the flow-matching schedule for the distilled Krea 2 turbo checkpoint on the scheduler: a " - "linear sigma schedule with the fixed time shift `mu=1.15` the checkpoint was distilled with." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_inference_steps", default=8), - InputParam( - name="sigmas", type_hint=list, description="Custom sigma schedule (defaults to a linear ramp)." - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam(name="timesteps", type_hint=torch.Tensor, description="The denoising timesteps.")] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - num_inference_steps = block_state.num_inference_steps - - sigmas = block_state.sigmas - if sigmas is None: - sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) - else: - block_state.num_inference_steps = len(sigmas) - - components.scheduler.set_timesteps(sigmas=sigmas, mu=1.15, device=device) - components.scheduler.set_begin_index(0) - block_state.timesteps = components.scheduler.timesteps - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Krea2PreparePositionIdsStep(ModularPipelineBlocks): - """ - Step that builds the shared rotary position ids for the combined [text | image] sequence: text at the origin, image - tokens at their (0, h, w) latent-grid coordinates. Place after prepare_latents. - - Inputs: - height (`int`, *optional*, defaults to 1024): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 1024): - The width in pixels of the generated image. - prompt_embeds (`Tensor`): - Batch-expanded text features (only text_seq_len is used). - - Outputs: - position_ids (`Tensor`): - Shared rotary coordinates (text_seq_len + grid_h * grid_w, 3). - """ - - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Step that builds the shared rotary position ids for the combined [text | image] sequence: text at the " - "origin, image tokens at their (0, h, w) latent-grid coordinates. Place after prepare_latents." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("height", default=1024), - InputParam.template("width", default=1024), - InputParam( - name="prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Batch-expanded text features (only text_seq_len is used).", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="position_ids", - type_hint=torch.Tensor, - description="Shared rotary coordinates (text_seq_len + grid_h * grid_w, 3).", - ) - ] - - @staticmethod - # Copied from diffusers.pipelines.krea2.pipeline_krea2.Krea2Pipeline.prepare_position_ids - def prepare_position_ids(text_seq_len: int, grid_height: int, grid_width: int, device: torch.device): - """Build the `(text_seq_len + grid_height * grid_width, 3)` rotary coordinates for the combined sequence: - text tokens sit at the origin, image tokens carry their `(0, h, w)` latent-grid coordinates.""" - text_ids = torch.zeros(text_seq_len, 3, device=device) - image_ids = torch.zeros(grid_height, grid_width, 3, device=device) - image_ids[..., 1] = torch.arange(grid_height, device=device)[:, None] - image_ids[..., 2] = torch.arange(grid_width, device=device)[None, :] - image_ids = image_ids.reshape(grid_height * grid_width, 3) - return torch.cat([text_ids, image_ids], dim=0) - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - p = components.patch_size - grid_h = block_state.height // (components.vae_scale_factor * p) - grid_w = block_state.width // (components.vae_scale_factor * p) - text_seq_len = block_state.prompt_embeds.shape[1] - - block_state.position_ids = self.prepare_position_ids(text_seq_len, grid_h, grid_w, device) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/krea2/decoders.py b/diffusers/modular_pipelines/krea2/decoders.py deleted file mode 100644 index fd308b5ef64844470de7ebde464cce2e46f77fec..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/krea2/decoders.py +++ /dev/null @@ -1,121 +0,0 @@ -# Copyright 2026 Krea AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch - -from ...configuration_utils import FrozenDict -from ...image_processor import VaeImageProcessor -from ...models import AutoencoderKLQwenImage -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Krea2ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# auto_docstring -class Krea2DecodeStep(ModularPipelineBlocks): - """ - Step that unpacks the denoised packed latents back to the spatial grid, de-normalizes them with the VAE's - per-channel statistics, and decodes them through the Qwen-Image VAE into images. - - Components: - vae (`AutoencoderKLQwenImage`) image_processor (`VaeImageProcessor`) - - Inputs: - output_type (`str`, *optional*, defaults to pil): - Output format: 'pil', 'np', 'pt'. - height (`int`, *optional*, defaults to 1024): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 1024): - The width in pixels of the generated image. - latents (`Tensor`): - The denoised packed latents (B, image_seq_len, in_channels) from the denoising loop. - - Outputs: - images (`list`): - Generated images. - """ - - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Step that unpacks the denoised packed latents back to the spatial grid, de-normalizes them with the " - "VAE's per-channel statistics, and decodes them through the Qwen-Image VAE into images." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLQwenImage), - ComponentSpec( - "image_processor", - VaeImageProcessor, - # Effective pixel-to-token downsampling factor: vae_scale_factor (8) * patch_size (2). - config=FrozenDict({"vae_scale_factor": 16}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("output_type", default="pil"), - InputParam.template("height", default=1024), - InputParam.template("width", default=1024), - InputParam( - name="latents", - required=True, - type_hint=torch.Tensor, - description="The denoised packed latents (B, image_seq_len, in_channels) from the denoising loop.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam.template("images")] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - vae = components.vae - p = components.patch_size - latents = block_state.latents - - batch_size, _, channels = latents.shape - height = p * (int(block_state.height) // (components.vae_scale_factor * p)) - width = p * (int(block_state.width) // (components.vae_scale_factor * p)) - latents = latents.view(batch_size, height // p, width // p, channels // (p * p), p, p) - latents = latents.permute(0, 3, 1, 4, 2, 5) - latents = latents.reshape(batch_size, channels // (p * p), 1, height, width) - - latents = latents.to(vae.dtype) - latents_mean = ( - torch.tensor(vae.config.latents_mean).view(1, vae.config.z_dim, 1, 1, 1).to(latents.device, latents.dtype) - ) - latents_std = 1.0 / torch.tensor(vae.config.latents_std).view(1, vae.config.z_dim, 1, 1, 1).to( - latents.device, latents.dtype - ) - latents = latents / latents_std + latents_mean - image = vae.decode(latents, return_dict=False)[0][:, :, 0] - block_state.images = components.image_processor.postprocess(image, output_type=block_state.output_type) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/krea2/denoise.py b/diffusers/modular_pipelines/krea2/denoise.py deleted file mode 100644 index 88c6cdca7aba093440d66b2cf65da4d278c01eec..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/krea2/denoise.py +++ /dev/null @@ -1,369 +0,0 @@ -# Copyright 2026 Krea AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...models.transformers.transformer_krea2 import Krea2Transformer2DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ..modular_pipeline import ( - BlockState, - LoopSequentialPipelineBlocks, - ModularPipelineBlocks, - PipelineState, -) -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Krea2ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class Krea2LoopBeforeDenoiser(ModularPipelineBlocks): - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Within the denoising loop: normalize the scheduler timestep into the model's flow time and broadcast it " - "across the batch. Compose into the `sub_blocks` of a `Krea2DenoiseLoopWrapper`-based step." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="latents", required=True, type_hint=torch.Tensor, description="Packed image latents."), - InputParam(name="batch_size", required=True, type_hint=int, description="Effective batch size."), - ] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - num_train_timesteps = components.scheduler.config.num_train_timesteps - block_state.timestep = (t / num_train_timesteps).expand(block_state.batch_size) - return components, block_state - - -class Krea2LoopDenoiser(ModularPipelineBlocks): - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Within the denoising loop: run the `transformer` on the conditional (and, when the guider enables CFG, " - "the negative) text features and combine them through the `guider`. Compose into `Krea2DenoiseStep`." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "guider", - ClassifierFreeGuidance, - # Krea 2 uses cond-anchored CFG (`cond + scale * (cond - uncond)`), which is the - # `use_original_formulation` branch of ClassifierFreeGuidance; scale 0 disables it (distilled TDM). - config=FrozenDict({"guidance_scale": 4.5, "use_original_formulation": True}), - default_creation_method="from_config", - ), - ComponentSpec("transformer", Krea2Transformer2DModel), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="latents", required=True, type_hint=torch.Tensor, description="Packed image latents."), - InputParam.template("num_inference_steps", required=True), - InputParam( - name="prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Conditional stacked text features.", - ), - InputParam( - name="prompt_embeds_mask", required=True, type_hint=torch.Tensor, description="Conditional text mask." - ), - InputParam( - name="position_ids", - required=True, - type_hint=torch.Tensor, - description="Shared rotary coordinates for the [text | image] sequence.", - ), - InputParam( - name="negative_prompt_embeds", type_hint=torch.Tensor, description="Negative stacked text features." - ), - InputParam(name="negative_prompt_embeds_mask", type_hint=torch.Tensor, description="Negative text mask."), - InputParam.template("attention_kwargs"), - ] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - transformer = components.transformer - - latents = block_state.latents.to(transformer.dtype) - timestep = block_state.timestep.to(transformer.dtype) - - guider_inputs = { - "encoder_hidden_states": ( - block_state.prompt_embeds.to(transformer.dtype), - block_state.negative_prompt_embeds.to(transformer.dtype) - if block_state.negative_prompt_embeds is not None - else None, - ), - "encoder_attention_mask": ( - block_state.prompt_embeds_mask, - block_state.negative_prompt_embeds_mask, - ), - } - - components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) - guider_state = components.guider.prepare_inputs(guider_inputs) - - for guider_state_batch in guider_state: - components.guider.prepare_models(components.transformer) - cond_kwargs = {name: getattr(guider_state_batch, name) for name in guider_inputs} - guider_state_batch.noise_pred = transformer( - hidden_states=latents, - timestep=timestep, - position_ids=block_state.position_ids, - attention_kwargs=block_state.attention_kwargs, - return_dict=False, - **cond_kwargs, - )[0] - components.guider.cleanup_models(components.transformer) - - block_state.noise_pred = components.guider(guider_state).pred - return components, block_state - - -class Krea2TurboLoopDenoiser(ModularPipelineBlocks): - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Within the denoising loop: run the `transformer` on the conditional text features. The distilled Krea 2 " - "turbo checkpoint runs without classifier-free guidance, so there is no negative branch or guider. Compose " - "into the `sub_blocks` of `Krea2TurboDenoiseStep`." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", Krea2Transformer2DModel)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam(name="latents", required=True, type_hint=torch.Tensor, description="Packed image latents."), - InputParam( - name="prompt_embeds", - required=True, - type_hint=torch.Tensor, - description="Conditional stacked text features.", - ), - InputParam( - name="prompt_embeds_mask", required=True, type_hint=torch.Tensor, description="Conditional text mask." - ), - InputParam( - name="position_ids", - required=True, - type_hint=torch.Tensor, - description="Shared rotary coordinates for the [text | image] sequence.", - ), - InputParam.template("attention_kwargs"), - ] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - transformer = components.transformer - - latents = block_state.latents.to(transformer.dtype) - timestep = block_state.timestep.to(transformer.dtype) - - block_state.noise_pred = transformer( - hidden_states=latents, - timestep=timestep, - position_ids=block_state.position_ids, - attention_kwargs=block_state.attention_kwargs, - encoder_hidden_states=block_state.prompt_embeds.to(transformer.dtype), - encoder_attention_mask=block_state.prompt_embeds_mask, - return_dict=False, - )[0] - return components, block_state - - -class Krea2LoopAfterDenoiser(ModularPipelineBlocks): - model_name = "krea2" - - @property - def description(self) -> str: - return "Within the denoising loop: scheduler step. Compose into a `Krea2DenoiseLoopWrapper`-based step." - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam(name="latents", type_hint=torch.Tensor, description="The denoised latents.")] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - latents_dtype = block_state.latents.dtype - block_state.latents = components.scheduler.step( - block_state.noise_pred, t, block_state.latents, return_dict=False - )[0] - block_state.latents = block_state.latents.to(latents_dtype) - return components, block_state - - -class Krea2DenoiseLoopWrapper(LoopSequentialPipelineBlocks): - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Pipeline block that iteratively denoises the packed image latents over `timesteps`. " - "The specific steps within each iteration can be customized with the `sub_blocks` attribute." - ) - - @property - def loop_expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] - - @property - def loop_inputs(self) -> list[InputParam]: - return [ - InputParam( - name="timesteps", - required=True, - type_hint=torch.Tensor, - description="Denoising timesteps from set_timesteps.", - ), - InputParam.template("num_inference_steps", required=True), - InputParam.template("attention_kwargs"), - ] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) - progress_bar.update() - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Krea2DenoiseStep(Krea2DenoiseLoopWrapper): - """ - Denoising loop that iteratively denoises the packed image latents over `timesteps`, running the transformer on the - conditional (and, when the guider enables CFG, the negative) text features and combining them through the `guider`. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) guider (`ClassifierFreeGuidance`) transformer - (`Krea2Transformer2DModel`) - - Inputs: - timesteps (`Tensor`): - Denoising timesteps from set_timesteps. - num_inference_steps (`int`): - The number of denoising steps. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - latents (`Tensor`): - Packed image latents. - batch_size (`int`): - Effective batch size. - prompt_embeds (`Tensor`): - Conditional stacked text features. - prompt_embeds_mask (`Tensor`): - Conditional text mask. - position_ids (`Tensor`): - Shared rotary coordinates for the [text | image] sequence. - negative_prompt_embeds (`Tensor`, *optional*): - Negative stacked text features. - negative_prompt_embeds_mask (`Tensor`, *optional*): - Negative text mask. - - Outputs: - latents (`Tensor`): - The denoised latents. - """ - - model_name = "krea2" - block_classes = [Krea2LoopBeforeDenoiser, Krea2LoopDenoiser, Krea2LoopAfterDenoiser] - block_names = ["before_denoiser", "denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoising loop that iteratively denoises the packed image latents over `timesteps`, running the " - "transformer on the conditional (and, when the guider enables CFG, the negative) text features and " - "combining them through the `guider`." - ) - - -# auto_docstring -class Krea2TurboDenoiseStep(Krea2DenoiseLoopWrapper): - """ - Denoising loop for the distilled Krea 2 turbo checkpoint that iteratively denoises the packed image latents over - `timesteps`, running the transformer on the conditional text features. The distilled checkpoint runs without - classifier-free guidance. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Krea2Transformer2DModel`) - - Inputs: - timesteps (`Tensor`): - Denoising timesteps from set_timesteps. - num_inference_steps (`int`): - The number of denoising steps. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - latents (`Tensor`): - Packed image latents. - batch_size (`int`): - Effective batch size. - prompt_embeds (`Tensor`): - Conditional stacked text features. - prompt_embeds_mask (`Tensor`): - Conditional text mask. - position_ids (`Tensor`): - Shared rotary coordinates for the [text | image] sequence. - - Outputs: - latents (`Tensor`): - The denoised latents. - """ - - model_name = "krea2" - block_classes = [Krea2LoopBeforeDenoiser, Krea2TurboLoopDenoiser, Krea2LoopAfterDenoiser] - block_names = ["before_denoiser", "denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoising loop for the distilled Krea 2 turbo checkpoint that iteratively denoises the packed image " - "latents over `timesteps`, running the transformer on the conditional text features. The distilled " - "checkpoint runs without classifier-free guidance." - ) diff --git a/diffusers/modular_pipelines/krea2/encoders.py b/diffusers/modular_pipelines/krea2/encoders.py deleted file mode 100644 index 7640222e9ad2bd69588447b6c8a8bee8dcbf63b4..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/krea2/encoders.py +++ /dev/null @@ -1,276 +0,0 @@ -# Copyright 2026 Krea AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch -from transformers import AutoTokenizer, Qwen3VLModel - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import Krea2ModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -# Indices into the Qwen3-VL `hidden_states` tuple (0 is the embedding output) whose states are stacked per token as the -# transformer's text conditioning. Must have `transformer.config.num_text_layers` entries. -KREA2_TEXT_ENCODER_SELECT_LAYERS = (2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 32, 35) - -# Krea 2 wraps the prompt in this Qwen-Image chat template before encoding. The prompt is padded to a fixed length -# first and the assistant suffix is appended *after* the padding (matching how the model was sampled at training time); -# the first `_PROMPT_TEMPLATE_ENCODE_START_IDX` (system prefix) tokens are dropped from the encoder outputs. -_PROMPT_TEMPLATE_ENCODE_PREFIX = ( - "<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, " - "spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n" -) -_PROMPT_TEMPLATE_ENCODE_SUFFIX = "<|im_end|>\n<|im_start|>assistant\n" -_PROMPT_TEMPLATE_ENCODE_START_IDX = 34 -_PROMPT_TEMPLATE_ENCODE_NUM_SUFFIX_TOKENS = 5 - - -# auto_docstring -class Krea2TextEncoderStep(ModularPipelineBlocks): - """ - Text encoder step that tokenizes the prompt(s) with the Krea 2 chat template, runs the Qwen3-VL text encoder, and - stacks a fixed set of decoder-layer hidden states per token as the transformer's text conditioning. The negative - prompt is encoded the same way when the guider enables CFG. - - Components: - text_encoder (`Qwen3VLModel`): The Qwen3-VL text encoder. tokenizer (`AutoTokenizer`): The tokenizer paired - with the text encoder. guider (`ClassifierFreeGuidance`) - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The negative prompt(s) for CFG. - max_sequence_length (`int`, *optional*, defaults to 512): - Maximum sequence length for prompt encoding. - - Outputs: - prompt_embeds (`Tensor`): - Per-prompt stacked text features (B, text_seq_len, num_text_layers, text_hidden_dim). - prompt_embeds_mask (`Tensor`): - Per-prompt boolean text mask (B, text_seq_len). - negative_prompt_embeds (`Tensor`): - Per-prompt negative text features (only when guidance is enabled). - negative_prompt_embeds_mask (`Tensor`): - Per-prompt negative text mask (only when guidance is enabled). - """ - - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Text encoder step that tokenizes the prompt(s) with the Krea 2 chat template, runs the Qwen3-VL text " - "encoder, and stacks a fixed set of decoder-layer hidden states per token as the transformer's text " - "conditioning. The negative prompt is encoded the same way when the guider enables CFG." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", Qwen3VLModel, description="The Qwen3-VL text encoder."), - ComponentSpec("tokenizer", AutoTokenizer, description="The tokenizer paired with the text encoder."), - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 4.5, "use_original_formulation": True}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt", required=True), - InputParam(name="negative_prompt", type_hint=str, description="The negative prompt(s) for CFG."), - InputParam.template("max_sequence_length", default=512), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="prompt_embeds", - type_hint=torch.Tensor, - description="Per-prompt stacked text features (B, text_seq_len, num_text_layers, text_hidden_dim).", - ), - OutputParam( - name="prompt_embeds_mask", - type_hint=torch.Tensor, - description="Per-prompt boolean text mask (B, text_seq_len).", - ), - OutputParam( - name="negative_prompt_embeds", - type_hint=torch.Tensor, - description="Per-prompt negative text features (only when guidance is enabled).", - ), - OutputParam( - name="negative_prompt_embeds_mask", - type_hint=torch.Tensor, - description="Per-prompt negative text mask (only when guidance is enabled).", - ), - ] - - def _encode_prompt(self, components, prompt, max_sequence_length, device): - """Tokenize `prompt` into the fixed-length Krea 2 layout and tap the selected encoder hidden states. - - Mirrors `Krea2Pipeline.get_text_hidden_states`. Returns a `(hidden_states, attention_mask)` tuple of shapes - `(batch_size, text_seq_len, num_text_layers, text_hidden_dim)` and `(batch_size, text_seq_len)` (bool). - """ - tokenizer = components.tokenizer - prompt = [prompt] if isinstance(prompt, str) else prompt - prefix_idx = _PROMPT_TEMPLATE_ENCODE_START_IDX - text = [_PROMPT_TEMPLATE_ENCODE_PREFIX + e for e in prompt] - text_tokens = tokenizer( - text, - truncation=True, - padding="max_length", - max_length=max_sequence_length + prefix_idx - _PROMPT_TEMPLATE_ENCODE_NUM_SUFFIX_TOKENS, - return_tensors="pt", - ).to(device) - suffix_tokens = tokenizer([_PROMPT_TEMPLATE_ENCODE_SUFFIX] * len(text), return_tensors="pt").to(device) - - input_ids = torch.cat([text_tokens.input_ids, suffix_tokens.input_ids], dim=1) - attention_mask = torch.cat([text_tokens.attention_mask, suffix_tokens.attention_mask], dim=1).bool() - - # Krea 2 pads in the middle of the template (`[prefix | prompt | PAD | suffix]`), so the suffix tokens sit - # downstream of the padding. The text features must use positions that count only real tokens (padding does - # not consume a position) to match how the model was trained; otherwise the suffix gets a shifted mRoPE phase. - position_ids = (attention_mask.long().cumsum(dim=-1) - 1).clamp(min=0) - position_ids = position_ids.unsqueeze(0).expand(3, -1, -1) - - outputs = components.text_encoder( - input_ids=input_ids, - attention_mask=attention_mask, - position_ids=position_ids, - output_hidden_states=True, - ) - hidden_states = torch.stack([outputs.hidden_states[i] for i in KREA2_TEXT_ENCODER_SELECT_LAYERS], dim=2) - - hidden_states = hidden_states[:, prefix_idx:] - attention_mask = attention_mask[:, prefix_idx:] - return hidden_states, attention_mask - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - prompts = [block_state.prompt] if isinstance(block_state.prompt, str) else list(block_state.prompt) - - block_state.prompt_embeds, block_state.prompt_embeds_mask = self._encode_prompt( - components, prompts, block_state.max_sequence_length, device - ) - - block_state.negative_prompt_embeds = None - block_state.negative_prompt_embeds_mask = None - if components.requires_unconditional_embeds: - negative_prompt = block_state.negative_prompt - if negative_prompt is None: - negative_prompt = "" - if isinstance(negative_prompt, str): - negative_prompt = [negative_prompt] * len(prompts) - block_state.negative_prompt_embeds, block_state.negative_prompt_embeds_mask = self._encode_prompt( - components, negative_prompt, block_state.max_sequence_length, device - ) - - self.set_block_state(state, block_state) - return components, state - - -# auto_docstring -class Krea2TurboTextEncoderStep(Krea2TextEncoderStep): - """ - Text encoder step for the distilled Krea 2 turbo checkpoint that tokenizes the prompt(s) with the Krea 2 chat - template, runs the Qwen3-VL text encoder, and stacks a fixed set of decoder-layer hidden states per token as the - transformer's text conditioning. The distilled checkpoint runs without classifier-free guidance, so it takes no - negative prompt and has no guider. - - Components: - text_encoder (`Qwen3VLModel`): The Qwen3-VL text encoder. tokenizer (`AutoTokenizer`): The tokenizer paired - with the text encoder. - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - max_sequence_length (`int`, *optional*, defaults to 512): - Maximum sequence length for prompt encoding. - - Outputs: - prompt_embeds (`Tensor`): - Per-prompt stacked text features (B, text_seq_len, num_text_layers, text_hidden_dim). - prompt_embeds_mask (`Tensor`): - Per-prompt boolean text mask (B, text_seq_len). - """ - - model_name = "krea2" - - @property - def description(self) -> str: - return ( - "Text encoder step for the distilled Krea 2 turbo checkpoint that tokenizes the prompt(s) with the Krea 2 " - "chat template, runs the Qwen3-VL text encoder, and stacks a fixed set of decoder-layer hidden states per " - "token as the transformer's text conditioning. The distilled checkpoint runs without classifier-free " - "guidance, so it takes no negative prompt and has no guider." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", Qwen3VLModel, description="The Qwen3-VL text encoder."), - ComponentSpec("tokenizer", AutoTokenizer, description="The tokenizer paired with the text encoder."), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt", required=True), - InputParam.template("max_sequence_length", default=512), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - name="prompt_embeds", - type_hint=torch.Tensor, - description="Per-prompt stacked text features (B, text_seq_len, num_text_layers, text_hidden_dim).", - ), - OutputParam( - name="prompt_embeds_mask", - type_hint=torch.Tensor, - description="Per-prompt boolean text mask (B, text_seq_len).", - ), - ] - - @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - device = components._execution_device - prompts = [block_state.prompt] if isinstance(block_state.prompt, str) else list(block_state.prompt) - - block_state.prompt_embeds, block_state.prompt_embeds_mask = self._encode_prompt( - components, prompts, block_state.max_sequence_length, device - ) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/krea2/modular_blocks_krea2.py b/diffusers/modular_pipelines/krea2/modular_blocks_krea2.py deleted file mode 100644 index ae3b2ac4fb52b1f35c29adb907c13e6939d142a7..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/krea2/modular_blocks_krea2.py +++ /dev/null @@ -1,170 +0,0 @@ -# Copyright 2026 Krea AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from ...utils import logging -from ..modular_pipeline import SequentialPipelineBlocks -from ..modular_pipeline_utils import InsertableDict, OutputParam -from .before_denoise import ( - Krea2PrepareLatentsStep, - Krea2PreparePositionIdsStep, - Krea2SetTimestepsStep, - Krea2TextInputsStep, -) -from .decoders import Krea2DecodeStep -from .denoise import Krea2DenoiseStep -from .encoders import Krea2TextEncoderStep - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -CORE_DENOISE_BLOCKS = InsertableDict( - [ - ("input", Krea2TextInputsStep()), - ("prepare_latents", Krea2PrepareLatentsStep()), - ("set_timesteps", Krea2SetTimestepsStep()), - ("prepare_position_ids", Krea2PreparePositionIdsStep()), - ("denoise", Krea2DenoiseStep()), - ] -) - - -# auto_docstring -class Krea2CoreDenoiseStep(SequentialPipelineBlocks): - """ - Core denoising workflow for Krea 2 text-to-image: prepares the batch/latents/timesteps and the shared position ids, - then runs the symmetric-CFG denoising loop, producing the denoised packed latents for the decoder. - - Components: - transformer (`Krea2Transformer2DModel`) scheduler (`FlowMatchEulerDiscreteScheduler`) guider - (`ClassifierFreeGuidance`) - - Inputs: - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - prompt_embeds (`Tensor`): - Per-prompt stacked text features (B, text_seq_len, num_text_layers, text_hidden_dim). - prompt_embeds_mask (`Tensor`): - Per-prompt boolean text mask (B, text_seq_len). - negative_prompt_embeds (`Tensor`, *optional*): - Per-prompt negative text features. - negative_prompt_embeds_mask (`Tensor`, *optional*): - Per-prompt negative text mask. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - height (`int`, *optional*, defaults to 1024): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 1024): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`, *optional*, defaults to 28): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigma schedule (defaults to a linear ramp). - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latents (`Tensor`): - The denoised packed latents (B, image_seq_len, in_channels). - """ - - model_name = "krea2" - block_classes = list(CORE_DENOISE_BLOCKS.values()) - block_names = list(CORE_DENOISE_BLOCKS.keys()) - - @property - def description(self) -> str: - return ( - "Core denoising workflow for Krea 2 text-to-image: prepares the batch/latents/timesteps and the shared " - "position ids, then runs the symmetric-CFG denoising loop, producing the denoised packed latents for the " - "decoder." - ) - - @property - def outputs(self) -> list[OutputParam]: - return [ - OutputParam.template("latents", description="The denoised packed latents (B, image_seq_len, in_channels).") - ] - - -# auto_docstring -class Krea2AutoBlocks(SequentialPipelineBlocks): - """ - Auto Modular pipeline for text-to-image generation using Krea 2: encode text -> core denoise (symmetric CFG) -> - decode. - - Supported workflows: - - `text2image`: requires `prompt` - - Components: - text_encoder (`Qwen3VLModel`): The Qwen3-VL text encoder. tokenizer (`AutoTokenizer`): The tokenizer paired - with the text encoder. guider (`ClassifierFreeGuidance`) transformer (`Krea2Transformer2DModel`) scheduler - (`FlowMatchEulerDiscreteScheduler`) vae (`AutoencoderKLQwenImage`) image_processor (`VaeImageProcessor`) - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The negative prompt(s) for CFG. - max_sequence_length (`int`, *optional*, defaults to 512): - Maximum sequence length for prompt encoding. - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - height (`int`, *optional*, defaults to 1024): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 1024): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`, *optional*, defaults to 28): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigma schedule (defaults to a linear ramp). - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - output_type (`str`, *optional*, defaults to pil): - Output format: 'pil', 'np', 'pt'. - - Outputs: - images (`list`): - Generated images. - """ - - model_name = "krea2" - block_classes = [ - Krea2TextEncoderStep, - Krea2CoreDenoiseStep, - Krea2DecodeStep, - ] - block_names = ["text_encoder", "denoise", "decode"] - - _workflow_map = { - "text2image": {"prompt": True}, - } - - @property - def description(self) -> str: - return ( - "Auto Modular pipeline for text-to-image generation using Krea 2: encode text -> core denoise " - "(symmetric CFG) -> decode." - ) - - @property - def outputs(self) -> list[OutputParam]: - return [OutputParam.template("images")] diff --git a/diffusers/modular_pipelines/krea2/modular_blocks_krea2_turbo.py b/diffusers/modular_pipelines/krea2/modular_blocks_krea2_turbo.py deleted file mode 100644 index 79fa5406c4e5af819a7a96f7b8ea5fafff14a8b7..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/krea2/modular_blocks_krea2_turbo.py +++ /dev/null @@ -1,164 +0,0 @@ -# Copyright 2026 Krea AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from ...utils import logging -from ..modular_pipeline import SequentialPipelineBlocks -from ..modular_pipeline_utils import InsertableDict, OutputParam -from .before_denoise import ( - Krea2PrepareLatentsStep, - Krea2PreparePositionIdsStep, - Krea2TurboSetTimestepsStep, - Krea2TurboTextInputsStep, -) -from .decoders import Krea2DecodeStep -from .denoise import Krea2TurboDenoiseStep -from .encoders import Krea2TurboTextEncoderStep - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -CORE_DENOISE_BLOCKS = InsertableDict( - [ - ("input", Krea2TurboTextInputsStep()), - ("prepare_latents", Krea2PrepareLatentsStep()), - ("set_timesteps", Krea2TurboSetTimestepsStep()), - ("prepare_position_ids", Krea2PreparePositionIdsStep()), - ("denoise", Krea2TurboDenoiseStep()), - ] -) - - -# auto_docstring -class Krea2TurboCoreDenoiseStep(SequentialPipelineBlocks): - """ - Core denoising workflow for the distilled Krea 2 turbo text-to-image checkpoint: prepares the - batch/latents/timesteps and the shared position ids, then runs the guidance-free denoising loop, producing the - denoised packed latents for the decoder. - - Components: - transformer (`Krea2Transformer2DModel`) scheduler (`FlowMatchEulerDiscreteScheduler`) - - Inputs: - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - prompt_embeds (`Tensor`): - Per-prompt stacked text features (B, text_seq_len, num_text_layers, text_hidden_dim). - prompt_embeds_mask (`Tensor`): - Per-prompt boolean text mask (B, text_seq_len). - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - height (`int`, *optional*, defaults to 1024): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 1024): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`, *optional*, defaults to 8): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigma schedule (defaults to a linear ramp). - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latents (`Tensor`): - The denoised packed latents (B, image_seq_len, in_channels). - """ - - model_name = "krea2" - block_classes = list(CORE_DENOISE_BLOCKS.values()) - block_names = list(CORE_DENOISE_BLOCKS.keys()) - - @property - def description(self) -> str: - return ( - "Core denoising workflow for the distilled Krea 2 turbo text-to-image checkpoint: prepares the " - "batch/latents/timesteps and the shared position ids, then runs the guidance-free denoising loop, " - "producing the denoised packed latents for the decoder." - ) - - @property - def outputs(self) -> list[OutputParam]: - return [ - OutputParam.template("latents", description="The denoised packed latents (B, image_seq_len, in_channels).") - ] - - -# auto_docstring -class Krea2TurboAutoBlocks(SequentialPipelineBlocks): - """ - Auto Modular pipeline for text-to-image generation using the distilled Krea 2 turbo checkpoint: encode text -> core - denoise (guidance-free) -> decode. - - Supported workflows: - - `text2image`: requires `prompt` - - Components: - text_encoder (`Qwen3VLModel`): The Qwen3-VL text encoder. tokenizer (`AutoTokenizer`): The tokenizer paired - with the text encoder. transformer (`Krea2Transformer2DModel`) scheduler (`FlowMatchEulerDiscreteScheduler`) - vae (`AutoencoderKLQwenImage`) image_processor (`VaeImageProcessor`) - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - max_sequence_length (`int`, *optional*, defaults to 512): - Maximum sequence length for prompt encoding. - num_images_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - height (`int`, *optional*, defaults to 1024): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 1024): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_inference_steps (`int`, *optional*, defaults to 8): - The number of denoising steps. - sigmas (`list`, *optional*): - Custom sigma schedule (defaults to a linear ramp). - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - output_type (`str`, *optional*, defaults to pil): - Output format: 'pil', 'np', 'pt'. - - Outputs: - images (`list`): - Generated images. - """ - - model_name = "krea2" - block_classes = [ - Krea2TurboTextEncoderStep, - Krea2TurboCoreDenoiseStep, - Krea2DecodeStep, - ] - block_names = ["text_encoder", "denoise", "decode"] - - _workflow_map = { - "text2image": {"prompt": True}, - } - - @property - def description(self) -> str: - return ( - "Auto Modular pipeline for text-to-image generation using the distilled Krea 2 turbo checkpoint: encode " - "text -> core denoise (guidance-free) -> decode." - ) - - @property - def outputs(self) -> list[OutputParam]: - return [OutputParam.template("images")] diff --git a/diffusers/modular_pipelines/krea2/modular_pipeline.py b/diffusers/modular_pipelines/krea2/modular_pipeline.py deleted file mode 100644 index 70d709573eecd5d8903e4d17b94be5ede54cea02..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/krea2/modular_pipeline.py +++ /dev/null @@ -1,67 +0,0 @@ -# Copyright 2026 Krea AI and The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from ...loaders import Krea2LoraLoaderMixin -from ..modular_pipeline import ModularPipeline - - -class Krea2ModularPipeline(ModularPipeline, Krea2LoraLoaderMixin): - """ - A ModularPipeline for Krea 2. - - > [!WARNING] > This is an experimental feature! - """ - - default_blocks_name = "Krea2AutoBlocks" - - @property - def patch_size(self): - return 2 - - @property - def default_height(self): - return 1024 - - @property - def default_width(self): - return 1024 - - @property - def vae_scale_factor(self): - vae_scale_factor = 8 - if getattr(self, "vae", None) is not None: - vae_scale_factor = 2 ** len(self.vae.temperal_downsample) - return vae_scale_factor - - @property - def requires_unconditional_embeds(self): - requires_unconditional_embeds = False - if hasattr(self, "guider") and self.guider is not None: - requires_unconditional_embeds = self.guider._enabled and self.guider.num_conditions > 1 - return requires_unconditional_embeds - - -class Krea2TurboModularPipeline(Krea2ModularPipeline): - """ - A ModularPipeline for the distilled Krea 2 turbo (TDM) checkpoint. It runs without classifier-free guidance, so it - takes no negative prompt and has no guider. - - > [!WARNING] > This is an experimental feature! - """ - - default_blocks_name = "Krea2TurboAutoBlocks" - - @property - def requires_unconditional_embeds(self): - return False diff --git a/diffusers/modular_pipelines/ltx/__init__.py b/diffusers/modular_pipelines/ltx/__init__.py deleted file mode 100644 index 531d9d3e4b20c786245dce67ca43e77066fc76ff..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ltx/__init__.py +++ /dev/null @@ -1,47 +0,0 @@ -from typing import TYPE_CHECKING - -from ...utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, -) - - -_dummy_objects = {} -_import_structure = {} - -try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ...utils import dummy_torch_and_transformers_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) -else: - _import_structure["modular_blocks_ltx"] = ["LTXAutoBlocks", "LTXBlocks", "LTXImage2VideoBlocks"] - _import_structure["modular_pipeline"] = ["LTXModularPipeline"] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .modular_blocks_ltx import LTXAutoBlocks, LTXBlocks, LTXImage2VideoBlocks - from .modular_pipeline import LTXModularPipeline -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/ltx/before_denoise.py b/diffusers/modular_pipelines/ltx/before_denoise.py deleted file mode 100644 index cd8b3ea82b821dccb2edff0e9deb9062fd7e9e62..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ltx/before_denoise.py +++ /dev/null @@ -1,392 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import inspect - -import numpy as np -import torch - -from ...configuration_utils import FrozenDict -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ...utils import logging -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import LTXModularPipeline, LTXVideoPachifier - - -logger = logging.get_logger(__name__) - - -def calculate_shift( - image_seq_len, - base_seq_len: int = 256, - max_seq_len: int = 4096, - base_shift: float = 0.5, - max_shift: float = 1.15, -): - m = (max_shift - base_shift) / (max_seq_len - base_seq_len) - b = base_shift - m * base_seq_len - mu = image_seq_len * m + b - return mu - - -# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps -def retrieve_timesteps( - scheduler, - num_inference_steps: int | None = None, - device: str | torch.device | None = None, - timesteps: list[int] | None = None, - sigmas: list[float] | None = None, - **kwargs, -): - r""" - Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles - custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`. - - Args: - scheduler (`SchedulerMixin`): - The scheduler to get timesteps from. - num_inference_steps (`int`): - The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps` - must be `None`. - device (`str` or `torch.device`, *optional*): - The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. - timesteps (`list[int]`, *optional*): - Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed, - `num_inference_steps` and `sigmas` must be `None`. - sigmas (`list[float]`, *optional*): - Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed, - `num_inference_steps` and `timesteps` must be `None`. - - Returns: - `tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the - second element is the number of inference steps. - """ - if timesteps is not None and sigmas is not None: - raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values") - if timesteps is not None: - accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) - if not accepts_timesteps: - raise ValueError( - f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" - f" timestep schedules. Please check whether you are using the correct scheduler." - ) - scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) - timesteps = scheduler.timesteps - num_inference_steps = len(timesteps) - elif sigmas is not None: - accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) - if not accept_sigmas: - raise ValueError( - f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" - f" sigmas schedules. Please check whether you are using the correct scheduler." - ) - scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs) - timesteps = scheduler.timesteps - num_inference_steps = len(timesteps) - else: - scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) - timesteps = scheduler.timesteps - return timesteps, num_inference_steps - - -class LTXTextInputStep(ModularPipelineBlocks): - model_name = "ltx" - - @property - def description(self) -> str: - return ( - "Input processing step that:\n" - " 1. Determines `batch_size` and `dtype` based on `prompt_embeds`\n" - " 2. Adjusts input tensor shapes based on `batch_size` and `num_videos_per_prompt`" - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_images_per_prompt", name="num_videos_per_prompt"), - InputParam.template("prompt_embeds", required=True), - InputParam.template("prompt_embeds_mask", name="prompt_attention_mask"), - InputParam.template("negative_prompt_embeds"), - InputParam.template("negative_prompt_embeds_mask", name="negative_prompt_attention_mask"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("batch_size", type_hint=int), - OutputParam("dtype", type_hint=torch.dtype), - ] - - @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - block_state.batch_size = block_state.prompt_embeds.shape[0] - block_state.dtype = block_state.prompt_embeds.dtype - num_videos = block_state.num_videos_per_prompt - - # Repeat prompt_embeds for num_videos_per_prompt - _, seq_len, _ = block_state.prompt_embeds.shape - block_state.prompt_embeds = block_state.prompt_embeds.repeat(1, num_videos, 1) - block_state.prompt_embeds = block_state.prompt_embeds.view(block_state.batch_size * num_videos, seq_len, -1) - - if block_state.prompt_attention_mask is not None: - block_state.prompt_attention_mask = block_state.prompt_attention_mask.repeat(num_videos, 1) - - if block_state.negative_prompt_embeds is not None: - _, seq_len, _ = block_state.negative_prompt_embeds.shape - block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.repeat(1, num_videos, 1) - block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.view( - block_state.batch_size * num_videos, seq_len, -1 - ) - - if block_state.negative_prompt_attention_mask is not None: - block_state.negative_prompt_attention_mask = block_state.negative_prompt_attention_mask.repeat( - num_videos, 1 - ) - - self.set_block_state(state, block_state) - return components, state - - -class LTXSetTimestepsStep(ModularPipelineBlocks): - model_name = "ltx" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), - ] - - @property - def description(self) -> str: - return "Step that sets the scheduler's timesteps for inference" - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_inference_steps"), - InputParam.template("timesteps"), - InputParam.template("sigmas"), - InputParam.template("height", default=512), - InputParam.template("width", default=704), - InputParam("num_frames", type_hint=int, default=161), - InputParam("frame_rate", type_hint=int, default=25), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("timesteps", type_hint=torch.Tensor), - OutputParam("num_inference_steps", type_hint=int), - OutputParam("rope_interpolation_scale", type_hint=tuple), - ] - - @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - height = block_state.height - width = block_state.width - num_frames = block_state.num_frames - frame_rate = block_state.frame_rate - - latent_num_frames = (num_frames - 1) // components.vae_temporal_compression_ratio + 1 - latent_height = height // components.vae_spatial_compression_ratio - latent_width = width // components.vae_spatial_compression_ratio - video_sequence_length = latent_num_frames * latent_height * latent_width - - custom_timesteps = block_state.timesteps - sigmas = block_state.sigmas - - if custom_timesteps is not None: - # User provided custom timesteps, don't compute sigmas - block_state.timesteps, block_state.num_inference_steps = retrieve_timesteps( - components.scheduler, - block_state.num_inference_steps, - device, - custom_timesteps, - ) - else: - if sigmas is None: - sigmas = np.linspace(1.0, 1 / block_state.num_inference_steps, block_state.num_inference_steps) - - mu = calculate_shift( - video_sequence_length, - components.scheduler.config.get("base_image_seq_len", 256), - components.scheduler.config.get("max_image_seq_len", 4096), - components.scheduler.config.get("base_shift", 0.5), - components.scheduler.config.get("max_shift", 1.15), - ) - - block_state.timesteps, block_state.num_inference_steps = retrieve_timesteps( - components.scheduler, - block_state.num_inference_steps, - device, - sigmas=sigmas, - mu=mu, - ) - - block_state.rope_interpolation_scale = ( - components.vae_temporal_compression_ratio / frame_rate, - components.vae_spatial_compression_ratio, - components.vae_spatial_compression_ratio, - ) - - self.set_block_state(state, block_state) - return components, state - - -class LTXPrepareLatentsStep(ModularPipelineBlocks): - model_name = "ltx" - - @property - def description(self) -> str: - return "Prepare latents step that prepares the latents for the text-to-video generation process" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "pachifier", - LTXVideoPachifier, - config=FrozenDict({"patch_size": 1, "patch_size_t": 1}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("height", default=512), - InputParam.template("width", default=704), - InputParam("num_frames", type_hint=int, default=161), - InputParam.template("latents"), - InputParam.template("num_images_per_prompt", name="num_videos_per_prompt"), - InputParam.template("generator"), - InputParam.template("batch_size", required=True), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("latents", type_hint=torch.Tensor), - ] - - @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - batch_size = block_state.batch_size * block_state.num_videos_per_prompt - num_channels_latents = components.transformer.config.in_channels - - if block_state.latents is not None: - block_state.latents = block_state.latents.to(device=device, dtype=torch.float32) - else: - height = block_state.height // components.vae_spatial_compression_ratio - width = block_state.width // components.vae_spatial_compression_ratio - num_frames = (block_state.num_frames - 1) // components.vae_temporal_compression_ratio + 1 - - shape = (batch_size, num_channels_latents, num_frames, height, width) - block_state.latents = randn_tensor( - shape, generator=block_state.generator, device=device, dtype=torch.float32 - ) - block_state.latents = components.pachifier.pack_latents(block_state.latents) - - self.set_block_state(state, block_state) - return components, state - - -class LTXImage2VideoPrepareLatentsStep(ModularPipelineBlocks): - model_name = "ltx" - - @property - def description(self) -> str: - return ( - "Prepare image-to-video latents: adds noise to pre-encoded image latents and creates a conditioning mask. " - "Expects pure noise `latents` from LTXPrepareLatentsStep." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "pachifier", - LTXVideoPachifier, - config=FrozenDict({"patch_size": 1, "patch_size_t": 1}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam("image_latents", type_hint=torch.Tensor, required=True), - InputParam.template("latents", required=True), - InputParam.template("height", default=512), - InputParam.template("width", default=704), - InputParam("num_frames", type_hint=int, default=161), - InputParam.template("num_images_per_prompt", name="num_videos_per_prompt"), - InputParam.template("batch_size", required=True), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("latents", type_hint=torch.Tensor), - OutputParam("conditioning_mask", type_hint=torch.Tensor), - ] - - @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - batch_size = block_state.batch_size * block_state.num_videos_per_prompt - - height = block_state.height // components.vae_spatial_compression_ratio - width = block_state.width // components.vae_spatial_compression_ratio - num_frames = (block_state.num_frames - 1) // components.vae_temporal_compression_ratio + 1 - - init_latents = block_state.image_latents.to(device=device, dtype=torch.float32) - if init_latents.shape[0] < batch_size: - init_latents = init_latents.repeat_interleave(batch_size // init_latents.shape[0], dim=0) - init_latents = init_latents.repeat(1, 1, num_frames, 1, 1) - - conditioning_mask = torch.zeros( - init_latents.shape[0], - 1, - init_latents.shape[2], - init_latents.shape[3], - init_latents.shape[4], - device=device, - dtype=torch.float32, - ) - conditioning_mask[:, :, 0] = 1.0 - - noise = components.pachifier.unpack_latents(block_state.latents, num_frames, height, width) - latents = init_latents * conditioning_mask + noise * (1 - conditioning_mask) - - conditioning_mask = components.pachifier.pack_latents(conditioning_mask).squeeze(-1) - latents = components.pachifier.pack_latents(latents) - - block_state.latents = latents - block_state.conditioning_mask = conditioning_mask - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/ltx/decoders.py b/diffusers/modular_pipelines/ltx/decoders.py deleted file mode 100644 index 8664dee25bfe73d841336d78c49193dfdc8c133e..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ltx/decoders.py +++ /dev/null @@ -1,132 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch - -from ...configuration_utils import FrozenDict -from ...models import AutoencoderKLLTXVideo -from ...utils import logging -from ...utils.torch_utils import randn_tensor -from ...video_processor import VideoProcessor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import LTXVideoPachifier - - -logger = logging.get_logger(__name__) - - -def _denormalize_latents( - latents: torch.Tensor, latents_mean: torch.Tensor, latents_std: torch.Tensor, scaling_factor: float = 1.0 -) -> torch.Tensor: - # Denormalize latents across the channel dimension [B, C, F, H, W] - latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) - latents_std = latents_std.view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) - latents = latents * latents_std / scaling_factor + latents_mean - return latents - - -class LTXVaeDecoderStep(ModularPipelineBlocks): - model_name = "ltx" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLLTXVideo), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 32}), - default_creation_method="from_config", - ), - ComponentSpec( - "pachifier", - LTXVideoPachifier, - config=FrozenDict({"patch_size": 1, "patch_size_t": 1}), - default_creation_method="from_config", - ), - ] - - @property - def description(self) -> str: - return "Step that decodes the denoised latents into videos" - - @property - def inputs(self) -> list[tuple[str, Any]]: - return [ - InputParam.template("latents", required=True), - InputParam.template("output_type", default="np"), - InputParam.template("height", default=512), - InputParam.template("width", default=704), - InputParam("num_frames", type_hint=int, default=161), - InputParam("decode_timestep", default=0.0), - InputParam("decode_noise_scale", default=None), - InputParam.template("generator"), - InputParam.template("batch_size"), - InputParam.template("dtype", required=True), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam.template("videos")] - - @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - vae = components.vae - - latents = block_state.latents - - height = block_state.height - width = block_state.width - num_frames = block_state.num_frames - - latent_num_frames = (num_frames - 1) // components.vae_temporal_compression_ratio + 1 - latent_height = height // components.vae_spatial_compression_ratio - latent_width = width // components.vae_spatial_compression_ratio - - latents = components.pachifier.unpack_latents(latents, latent_num_frames, latent_height, latent_width) - latents = _denormalize_latents(latents, vae.latents_mean, vae.latents_std, vae.config.scaling_factor) - latents = latents.to(block_state.dtype) - - if not vae.config.timestep_conditioning: - timestep = None - else: - device = latents.device - batch_size = block_state.batch_size - decode_timestep = block_state.decode_timestep - decode_noise_scale = block_state.decode_noise_scale - - noise = randn_tensor(latents.shape, generator=block_state.generator, device=device, dtype=latents.dtype) - if not isinstance(decode_timestep, list): - decode_timestep = [decode_timestep] * batch_size - if decode_noise_scale is None: - decode_noise_scale = decode_timestep - elif not isinstance(decode_noise_scale, list): - decode_noise_scale = [decode_noise_scale] * batch_size - - timestep = torch.tensor(decode_timestep, device=device, dtype=latents.dtype) - decode_noise_scale = torch.tensor(decode_noise_scale, device=device, dtype=latents.dtype)[ - :, None, None, None, None - ] - latents = (1 - decode_noise_scale) * latents + decode_noise_scale * noise - - latents = latents.to(vae.dtype) - video = vae.decode(latents, timestep, return_dict=False)[0] - block_state.videos = components.video_processor.postprocess_video(video, output_type=block_state.output_type) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/ltx/denoise.py b/diffusers/modular_pipelines/ltx/denoise.py deleted file mode 100644 index b3ed86b5167934c665178987d58e5623c7a3ca90..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ltx/denoise.py +++ /dev/null @@ -1,458 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...models import LTXVideoTransformer3DModel -from ...schedulers import FlowMatchEulerDiscreteScheduler -from ..modular_pipeline import ( - BlockState, - LoopSequentialPipelineBlocks, - ModularPipelineBlocks, - PipelineState, -) -from ..modular_pipeline_utils import ComponentSpec, InputParam -from .modular_pipeline import LTXModularPipeline, LTXVideoPachifier - - -class LTXLoopBeforeDenoiser(ModularPipelineBlocks): - model_name = "ltx" - - @property - def description(self) -> str: - return ( - "Step within the denoising loop that prepares the latent input for the denoiser. " - "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " - "object (e.g. `LTXDenoiseLoopWrapper`)" - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents", required=True), - InputParam.template("dtype", required=True), - ] - - @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - block_state.latent_model_input = block_state.latents.to(block_state.dtype) - return components, block_state - - -class LTXLoopDenoiser(ModularPipelineBlocks): - model_name = "ltx" - - def __init__( - self, - guider_input_fields: dict[str, Any] | None = None, - ): - if guider_input_fields is None: - guider_input_fields = { - "encoder_hidden_states": ("prompt_embeds", "negative_prompt_embeds"), - "encoder_attention_mask": ("prompt_attention_mask", "negative_prompt_attention_mask"), - } - if not isinstance(guider_input_fields, dict): - raise ValueError(f"guider_input_fields must be a dictionary but is {type(guider_input_fields)}") - self._guider_input_fields = guider_input_fields - super().__init__() - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 3.0}), - default_creation_method="from_config", - ), - ComponentSpec("transformer", LTXVideoTransformer3DModel), - ] - - @property - def description(self) -> str: - return ( - "Step within the denoising loop that denoises the latents with guidance. " - "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " - "object (e.g. `LTXDenoiseLoopWrapper`)" - ) - - @property - def inputs(self) -> list[tuple[str, Any]]: - inputs = [ - InputParam.template("attention_kwargs"), - InputParam.template("num_inference_steps", required=True), - InputParam("rope_interpolation_scale", type_hint=tuple), - InputParam.template("height"), - InputParam.template("width"), - InputParam("num_frames", type_hint=int), - ] - guider_input_names = [] - for value in self._guider_input_fields.values(): - if isinstance(value, tuple): - guider_input_names.extend(value) - else: - guider_input_names.append(value) - - for name in guider_input_names: - inputs.append(InputParam(name=name, required=True, type_hint=torch.Tensor)) - return inputs - - @torch.no_grad() - def __call__( - self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: - components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) - - latent_num_frames = (block_state.num_frames - 1) // components.vae_temporal_compression_ratio + 1 - latent_height = block_state.height // components.vae_spatial_compression_ratio - latent_width = block_state.width // components.vae_spatial_compression_ratio - - guider_state = components.guider.prepare_inputs_from_block_state(block_state, self._guider_input_fields) - - for guider_state_batch in guider_state: - components.guider.prepare_models(components.transformer) - cond_kwargs = guider_state_batch.as_dict() - cond_kwargs = { - k: v.to(block_state.dtype) if isinstance(v, torch.Tensor) else v - for k, v in cond_kwargs.items() - if k in self._guider_input_fields.keys() - } - - context_name = getattr(guider_state_batch, components.guider._identifier_key, None) - with components.transformer.cache_context(context_name): - guider_state_batch.noise_pred = components.transformer( - hidden_states=block_state.latent_model_input, - timestep=t.expand(block_state.latent_model_input.shape[0]).to(block_state.dtype), - num_frames=latent_num_frames, - height=latent_height, - width=latent_width, - rope_interpolation_scale=block_state.rope_interpolation_scale, - attention_kwargs=block_state.attention_kwargs, - return_dict=False, - **cond_kwargs, - )[0] - components.guider.cleanup_models(components.transformer) - - block_state.noise_pred = components.guider(guider_state)[0] - - return components, block_state - - -class LTXLoopAfterDenoiser(ModularPipelineBlocks): - model_name = "ltx" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), - ] - - @property - def description(self) -> str: - return ( - "Step within the denoising loop that updates the latents. " - "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " - "object (e.g. `LTXDenoiseLoopWrapper`)" - ) - - @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - latents_dtype = block_state.latents.dtype - block_state.latents = components.scheduler.step( - block_state.noise_pred, - t, - block_state.latents, - return_dict=False, - )[0] - - if block_state.latents.dtype != latents_dtype: - block_state.latents = block_state.latents.to(latents_dtype) - - return components, block_state - - -class LTXDenoiseLoopWrapper(LoopSequentialPipelineBlocks): - model_name = "ltx" - - @property - def description(self) -> str: - return ( - "Pipeline block that iteratively denoises the latents over `timesteps`. " - "The specific steps within each iteration can be customized with `sub_blocks` attributes" - ) - - @property - def loop_expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), - ComponentSpec("transformer", LTXVideoTransformer3DModel), - ] - - @property - def loop_inputs(self) -> list[InputParam]: - return [ - InputParam.template("timesteps", required=True), - InputParam.template("num_inference_steps", required=True), - ] - - @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - block_state.num_warmup_steps = max( - len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order, 0 - ) - - with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) - if i == len(block_state.timesteps) - 1 or ( - (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 - ): - progress_bar.update() - - self.set_block_state(state, block_state) - return components, state - - -class LTXDenoiseStep(LTXDenoiseLoopWrapper): - block_classes = [ - LTXLoopBeforeDenoiser, - LTXLoopDenoiser( - guider_input_fields={ - "encoder_hidden_states": ("prompt_embeds", "negative_prompt_embeds"), - "encoder_attention_mask": ("prompt_attention_mask", "negative_prompt_attention_mask"), - } - ), - LTXLoopAfterDenoiser, - ] - block_names = ["before_denoiser", "denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoise step that iteratively denoises the latents.\n" - "Its loop logic is defined in `LTXDenoiseLoopWrapper.__call__` method.\n" - "At each iteration, it runs blocks defined in `sub_blocks` sequentially:\n" - " - `LTXLoopBeforeDenoiser`\n" - " - `LTXLoopDenoiser`\n" - " - `LTXLoopAfterDenoiser`\n" - "This block supports text-to-video tasks." - ) - - -class LTXImage2VideoLoopBeforeDenoiser(ModularPipelineBlocks): - model_name = "ltx" - - @property - def description(self) -> str: - return ( - "Step within the i2v denoising loop that prepares the latent input and modulates " - "the timestep with the conditioning mask." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("latents", required=True), - InputParam("conditioning_mask", required=True, type_hint=torch.Tensor), - InputParam.template("dtype", required=True), - ] - - @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - block_state.latent_model_input = block_state.latents.to(block_state.dtype) - block_state.timestep_adjusted = t.expand(block_state.latent_model_input.shape[0]).unsqueeze(-1) * ( - 1 - block_state.conditioning_mask - ) - return components, block_state - - -class LTXImage2VideoLoopDenoiser(ModularPipelineBlocks): - model_name = "ltx" - - def __init__( - self, - guider_input_fields: dict[str, Any] | None = None, - ): - if guider_input_fields is None: - guider_input_fields = { - "encoder_hidden_states": ("prompt_embeds", "negative_prompt_embeds"), - "encoder_attention_mask": ("prompt_attention_mask", "negative_prompt_attention_mask"), - } - if not isinstance(guider_input_fields, dict): - raise ValueError(f"guider_input_fields must be a dictionary but is {type(guider_input_fields)}") - self._guider_input_fields = guider_input_fields - super().__init__() - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 3.0}), - default_creation_method="from_config", - ), - ComponentSpec("transformer", LTXVideoTransformer3DModel), - ] - - @property - def description(self) -> str: - return ( - "Step within the i2v denoising loop that denoises the latents with guidance " - "using timestep modulated by the conditioning mask." - ) - - @property - def inputs(self) -> list[tuple[str, Any]]: - inputs = [ - InputParam.template("attention_kwargs"), - InputParam.template("num_inference_steps", required=True), - InputParam("rope_interpolation_scale", type_hint=tuple), - InputParam.template("height"), - InputParam.template("width"), - InputParam("num_frames", type_hint=int), - ] - guider_input_names = [] - for value in self._guider_input_fields.values(): - if isinstance(value, tuple): - guider_input_names.extend(value) - else: - guider_input_names.append(value) - for name in guider_input_names: - inputs.append(InputParam(name=name, required=True, type_hint=torch.Tensor)) - return inputs - - @torch.no_grad() - def __call__( - self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: - components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) - - latent_num_frames = (block_state.num_frames - 1) // components.vae_temporal_compression_ratio + 1 - latent_height = block_state.height // components.vae_spatial_compression_ratio - latent_width = block_state.width // components.vae_spatial_compression_ratio - - guider_state = components.guider.prepare_inputs_from_block_state(block_state, self._guider_input_fields) - - for guider_state_batch in guider_state: - components.guider.prepare_models(components.transformer) - cond_kwargs = guider_state_batch.as_dict() - cond_kwargs = { - k: v.to(block_state.dtype) if isinstance(v, torch.Tensor) else v - for k, v in cond_kwargs.items() - if k in self._guider_input_fields.keys() - } - - context_name = getattr(guider_state_batch, components.guider._identifier_key, None) - with components.transformer.cache_context(context_name): - guider_state_batch.noise_pred = components.transformer( - hidden_states=block_state.latent_model_input, - timestep=block_state.timestep_adjusted, - num_frames=latent_num_frames, - height=latent_height, - width=latent_width, - rope_interpolation_scale=block_state.rope_interpolation_scale, - attention_kwargs=block_state.attention_kwargs, - return_dict=False, - **cond_kwargs, - )[0] - components.guider.cleanup_models(components.transformer) - - block_state.noise_pred = components.guider(guider_state)[0] - - return components, block_state - - -class LTXImage2VideoLoopAfterDenoiser(ModularPipelineBlocks): - model_name = "ltx" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), - ComponentSpec( - "pachifier", - LTXVideoPachifier, - config=FrozenDict({"patch_size": 1, "patch_size_t": 1}), - default_creation_method="from_config", - ), - ] - - @property - def description(self) -> str: - return ( - "Step within the i2v denoising loop that updates the latents, " - "applying the scheduler step only to frames after the first (conditioned) frame." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("height"), - InputParam.template("width"), - InputParam("num_frames", type_hint=int), - ] - - @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - latent_num_frames = (block_state.num_frames - 1) // components.vae_temporal_compression_ratio + 1 - latent_height = block_state.height // components.vae_spatial_compression_ratio - latent_width = block_state.width // components.vae_spatial_compression_ratio - - noise_pred = components.pachifier.unpack_latents( - block_state.noise_pred, latent_num_frames, latent_height, latent_width - ) - latents = components.pachifier.unpack_latents( - block_state.latents, latent_num_frames, latent_height, latent_width - ) - - noise_pred = noise_pred[:, :, 1:] - noise_latents = latents[:, :, 1:] - pred_latents = components.scheduler.step(noise_pred, t, noise_latents, return_dict=False)[0] - - latents = torch.cat([latents[:, :, :1], pred_latents], dim=2) - block_state.latents = components.pachifier.pack_latents(latents) - - return components, block_state - - -class LTXImage2VideoDenoiseStep(LTXDenoiseLoopWrapper): - block_classes = [ - LTXImage2VideoLoopBeforeDenoiser, - LTXImage2VideoLoopDenoiser( - guider_input_fields={ - "encoder_hidden_states": ("prompt_embeds", "negative_prompt_embeds"), - "encoder_attention_mask": ("prompt_attention_mask", "negative_prompt_attention_mask"), - } - ), - LTXImage2VideoLoopAfterDenoiser, - ] - block_names = ["before_denoiser", "denoiser", "after_denoiser"] - - @property - def description(self) -> str: - return ( - "Denoise step for image-to-video that iteratively denoises the latents.\n" - "The first frame is kept fixed via a conditioning mask.\n" - "At each iteration, it runs blocks defined in `sub_blocks` sequentially:\n" - " - `LTXImage2VideoLoopBeforeDenoiser`\n" - " - `LTXImage2VideoLoopDenoiser`\n" - " - `LTXImage2VideoLoopAfterDenoiser`" - ) diff --git a/diffusers/modular_pipelines/ltx/encoders.py b/diffusers/modular_pipelines/ltx/encoders.py deleted file mode 100644 index 55405ad0aefefd43ff58239c992fe3ec9c9d329b..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ltx/encoders.py +++ /dev/null @@ -1,273 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -from transformers import T5EncoderModel, T5TokenizerFast - -from ...configuration_utils import FrozenDict -from ...guiders import ClassifierFreeGuidance -from ...models import AutoencoderKLLTXVideo -from ...utils import logging -from ...video_processor import VideoProcessor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import LTXModularPipeline - - -logger = logging.get_logger(__name__) - - -def _get_t5_prompt_embeds( - components, - prompt: str | list[str], - max_sequence_length: int, - device: torch.device, - dtype: torch.dtype, -): - prompt = [prompt] if isinstance(prompt, str) else prompt - - text_inputs = components.tokenizer( - prompt, - padding="max_length", - max_length=max_sequence_length, - truncation=True, - add_special_tokens=True, - return_tensors="pt", - ) - text_input_ids = text_inputs.input_ids - prompt_attention_mask = text_inputs.attention_mask - prompt_attention_mask = prompt_attention_mask.bool().to(device) - - prompt_embeds = components.text_encoder(text_input_ids.to(device))[0] - prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) - - return prompt_embeds, prompt_attention_mask - - -class LTXTextEncoderStep(ModularPipelineBlocks): - model_name = "ltx" - - @property - def description(self) -> str: - return "Text Encoder step that generates text embeddings to guide the video generation" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("text_encoder", T5EncoderModel), - ComponentSpec("tokenizer", T5TokenizerFast), - ComponentSpec( - "guider", - ClassifierFreeGuidance, - config=FrozenDict({"guidance_scale": 3.0}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt"), - InputParam.template("negative_prompt"), - InputParam.template("max_sequence_length", default=128), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam.template("prompt_embeds"), - OutputParam.template("prompt_embeds_mask", name="prompt_attention_mask"), - OutputParam.template("negative_prompt_embeds"), - OutputParam.template("negative_prompt_embeds_mask", name="negative_prompt_attention_mask"), - ] - - @staticmethod - def check_inputs(block_state): - if block_state.prompt is not None and ( - not isinstance(block_state.prompt, str) and not isinstance(block_state.prompt, list) - ): - raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(block_state.prompt)}") - - @staticmethod - def encode_prompt( - components, - prompt: str, - device: torch.device | None = None, - prepare_unconditional_embeds: bool = True, - negative_prompt: str | None = None, - max_sequence_length: int = 128, - ): - device = device or components._execution_device - dtype = components.text_encoder.dtype - - if not isinstance(prompt, list): - prompt = [prompt] - batch_size = len(prompt) - - prompt_embeds, prompt_attention_mask = _get_t5_prompt_embeds( - components=components, - prompt=prompt, - max_sequence_length=max_sequence_length, - device=device, - dtype=dtype, - ) - - negative_prompt_embeds = None - negative_prompt_attention_mask = None - - if prepare_unconditional_embeds: - negative_prompt = negative_prompt or "" - negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt - - if batch_size != len(negative_prompt): - raise ValueError( - f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" - f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches" - " the batch size of `prompt`." - ) - - negative_prompt_embeds, negative_prompt_attention_mask = _get_t5_prompt_embeds( - components=components, - prompt=negative_prompt, - max_sequence_length=max_sequence_length, - device=device, - dtype=dtype, - ) - - return prompt_embeds, prompt_attention_mask, negative_prompt_embeds, negative_prompt_attention_mask - - @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - self.check_inputs(block_state) - - block_state.device = components._execution_device - - ( - block_state.prompt_embeds, - block_state.prompt_attention_mask, - block_state.negative_prompt_embeds, - block_state.negative_prompt_attention_mask, - ) = self.encode_prompt( - components=components, - prompt=block_state.prompt, - device=block_state.device, - prepare_unconditional_embeds=components.requires_unconditional_embeds, - negative_prompt=block_state.negative_prompt, - max_sequence_length=block_state.max_sequence_length, - ) - - self.set_block_state(state, block_state) - return components, state - - -# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.retrieve_latents -def retrieve_latents( - encoder_output: torch.Tensor, generator: torch.Generator | None = None, sample_mode: str = "sample" -): - if hasattr(encoder_output, "latent_dist") and sample_mode == "sample": - return encoder_output.latent_dist.sample(generator) - elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax": - return encoder_output.latent_dist.mode() - elif hasattr(encoder_output, "latents"): - return encoder_output.latents - else: - raise AttributeError("Could not access latents of provided encoder_output") - - -def _normalize_latents( - latents: torch.Tensor, latents_mean: torch.Tensor, latents_std: torch.Tensor, scaling_factor: float = 1.0 -) -> torch.Tensor: - # Normalize latents across the channel dimension [B, C, F, H, W] - latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) - latents_std = latents_std.view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) - latents = (latents - latents_mean) * scaling_factor / latents_std - return latents - - -class LTXVaeEncoderStep(ModularPipelineBlocks): - model_name = "ltx" - - @property - def description(self) -> str: - return "VAE Encoder step that encodes an input image into latent space for image-to-video generation" - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLLTXVideo), - ComponentSpec( - "video_processor", - VideoProcessor, - config=FrozenDict({"vae_scale_factor": 32}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("image", required=True), - InputParam.template("height", default=512), - InputParam.template("width", default=704), - InputParam.template("generator"), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "image_latents", - type_hint=torch.Tensor, - description="Encoded image latents from the VAE encoder", - ), - ] - - @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - image = block_state.image - if not isinstance(image, torch.Tensor): - image = components.video_processor.preprocess(image, height=block_state.height, width=block_state.width) - image = image.to(device=device, dtype=torch.float32) - - vae_dtype = components.vae.dtype - - num_images = image.shape[0] - if isinstance(block_state.generator, list): - init_latents = [ - retrieve_latents( - components.vae.encode(image[i].unsqueeze(0).unsqueeze(2).to(vae_dtype)), - block_state.generator[i], - ) - for i in range(num_images) - ] - else: - init_latents = [ - retrieve_latents( - components.vae.encode(img.unsqueeze(0).unsqueeze(2).to(vae_dtype)), - block_state.generator, - ) - for img in image - ] - - init_latents = torch.cat(init_latents, dim=0).to(torch.float32) - block_state.image_latents = _normalize_latents( - init_latents, components.vae.latents_mean, components.vae.latents_std - ) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/ltx/modular_blocks_ltx.py b/diffusers/modular_pipelines/ltx/modular_blocks_ltx.py deleted file mode 100644 index 828c79e1c72df0d0b97763bd7fd99eca1b22a64a..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ltx/modular_blocks_ltx.py +++ /dev/null @@ -1,487 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from ...utils import logging -from ..modular_pipeline import AutoPipelineBlocks, SequentialPipelineBlocks -from ..modular_pipeline_utils import OutputParam -from .before_denoise import ( - LTXImage2VideoPrepareLatentsStep, - LTXPrepareLatentsStep, - LTXSetTimestepsStep, - LTXTextInputStep, -) -from .decoders import LTXVaeDecoderStep -from .denoise import LTXDenoiseStep, LTXImage2VideoDenoiseStep -from .encoders import LTXTextEncoderStep, LTXVaeEncoderStep - - -logger = logging.get_logger(__name__) - - -# auto_docstring -class LTXCoreDenoiseStep(SequentialPipelineBlocks): - """ - Denoise block that takes encoded conditions and runs the denoising process. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) pachifier (`LTXVideoPachifier`) guider - (`ClassifierFreeGuidance`) transformer (`LTXVideoTransformer3DModel`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - prompt_attention_mask (`Tensor`): - mask for the text embeddings. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_attention_mask (`Tensor`, *optional*): - mask for the negative text embeddings. Can be generated from text_encoder step. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - timesteps (`Tensor`, *optional*): - Timesteps for the denoising process. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - height (`int`, *optional*, defaults to 512): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 704): - The width in pixels of the generated image. - num_frames (`int`, *optional*, defaults to 161): - TODO: Add description. - frame_rate (`int`, *optional*, defaults to 25): - TODO: Add description. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "ltx" - block_classes = [ - LTXTextInputStep, - LTXSetTimestepsStep, - LTXPrepareLatentsStep, - LTXDenoiseStep, - ] - block_names = ["input", "set_timesteps", "prepare_latents", "denoise"] - - @property - def description(self): - return "Denoise block that takes encoded conditions and runs the denoising process." - - @property - def outputs(self): - return [OutputParam.template("latents")] - - -# auto_docstring -class LTXImage2VideoCoreDenoiseStep(SequentialPipelineBlocks): - """ - Denoise block for image-to-video that takes encoded conditions and image latents, and runs the denoising process. - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) pachifier (`LTXVideoPachifier`) guider - (`ClassifierFreeGuidance`) transformer (`LTXVideoTransformer3DModel`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - prompt_attention_mask (`Tensor`): - mask for the text embeddings. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`, *optional*): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_attention_mask (`Tensor`, *optional*): - mask for the negative text embeddings. Can be generated from text_encoder step. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - timesteps (`Tensor`, *optional*): - Timesteps for the denoising process. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - height (`int`, *optional*, defaults to 512): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 704): - The width in pixels of the generated image. - num_frames (`int`, *optional*, defaults to 161): - TODO: Add description. - frame_rate (`int`, *optional*, defaults to 25): - TODO: Add description. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image_latents (`Tensor`): - TODO: Add description. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "ltx" - block_classes = [ - LTXTextInputStep, - LTXSetTimestepsStep, - LTXPrepareLatentsStep, - LTXImage2VideoPrepareLatentsStep, - LTXImage2VideoDenoiseStep, - ] - block_names = ["input", "set_timesteps", "prepare_latents", "prepare_i2v_latents", "denoise"] - - @property - def description(self): - return "Denoise block for image-to-video that takes encoded conditions and image latents, and runs the denoising process." - - @property - def outputs(self): - return [OutputParam.template("latents")] - - -# auto_docstring -class LTXBlocks(SequentialPipelineBlocks): - """ - Modular pipeline blocks for LTX Video text-to-video. - - Components: - text_encoder (`T5EncoderModel`) tokenizer (`T5Tokenizer`) guider (`ClassifierFreeGuidance`) scheduler - (`FlowMatchEulerDiscreteScheduler`) pachifier (`LTXVideoPachifier`) transformer - (`LTXVideoTransformer3DModel`) vae (`AutoencoderKLLTXVideo`) video_processor (`VideoProcessor`) - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The prompt or prompts not to guide the image generation. - max_sequence_length (`int`, *optional*, defaults to 128): - Maximum sequence length for prompt encoding. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - timesteps (`Tensor`, *optional*): - Timesteps for the denoising process. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - height (`int`, *optional*, defaults to 512): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 704): - The width in pixels of the generated image. - num_frames (`int`, *optional*, defaults to 161): - TODO: Add description. - frame_rate (`int`, *optional*, defaults to 25): - TODO: Add description. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - output_type (`str`, *optional*, defaults to np): - Output format: 'pil', 'np', 'pt'. - decode_timestep (`None`, *optional*, defaults to 0.0): - TODO: Add description. - decode_noise_scale (`None`, *optional*): - TODO: Add description. - - Outputs: - videos (`list`): - The generated videos. - """ - - model_name = "ltx" - block_classes = [ - LTXTextEncoderStep, - LTXCoreDenoiseStep, - LTXVaeDecoderStep, - ] - block_names = ["text_encoder", "denoise", "decode"] - - @property - def description(self): - return "Modular pipeline blocks for LTX Video text-to-video." - - @property - def outputs(self): - return [OutputParam.template("videos")] - - -# auto_docstring -class LTXAutoVaeEncoderStep(AutoPipelineBlocks): - """ - VAE encoder step that encodes the image input into its latent representation. - This is an auto pipeline block that works for image-to-video tasks. - - `LTXVaeEncoderStep` is used when `image` is provided. - - If `image` is not provided, step will be skipped. - - Components: - vae (`AutoencoderKLLTXVideo`) video_processor (`VideoProcessor`) - - Inputs: - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - height (`int`, *optional*, defaults to 512): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 704): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - - Outputs: - image_latents (`Tensor`): - Encoded image latents from the VAE encoder - """ - - model_name = "ltx" - block_classes = [LTXVaeEncoderStep] - block_names = ["vae_encoder"] - block_trigger_inputs = ["image"] - - @property - def description(self): - return ( - "VAE encoder step that encodes the image input into its latent representation.\n" - "This is an auto pipeline block that works for image-to-video tasks.\n" - " - `LTXVaeEncoderStep` is used when `image` is provided.\n" - " - If `image` is not provided, step will be skipped." - ) - - -# auto_docstring -class LTXAutoCoreDenoiseStep(AutoPipelineBlocks): - """ - Auto denoise block that selects the appropriate denoise pipeline based on inputs. - - `LTXImage2VideoCoreDenoiseStep` is used when `image_latents` is provided. - - `LTXCoreDenoiseStep` is used otherwise (text-to-video). - - Components: - scheduler (`FlowMatchEulerDiscreteScheduler`) pachifier (`LTXVideoPachifier`) guider - (`ClassifierFreeGuidance`) transformer (`LTXVideoTransformer3DModel`) - - Inputs: - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - prompt_embeds (`Tensor`): - text embeddings used to guide the image generation. Can be generated from text_encoder step. - prompt_attention_mask (`Tensor`): - mask for the text embeddings. Can be generated from text_encoder step. - negative_prompt_embeds (`Tensor`): - negative text embeddings used to guide the image generation. Can be generated from text_encoder step. - negative_prompt_attention_mask (`Tensor`): - mask for the negative text embeddings. Can be generated from text_encoder step. - num_inference_steps (`int`): - The number of denoising steps. - timesteps (`Tensor`): - Timesteps for the denoising process. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - height (`int`, *optional*, defaults to 512): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 704): - The width in pixels of the generated image. - num_frames (`int`, *optional*, defaults to 161): - TODO: Add description. - frame_rate (`int`, *optional*, defaults to 25): - TODO: Add description. - latents (`Tensor`): - Pre-generated noisy latents for image generation. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - image_latents (`Tensor`, *optional*): - TODO: Add description. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - - Outputs: - latents (`Tensor`): - Denoised latents. - """ - - model_name = "ltx" - block_classes = [LTXImage2VideoCoreDenoiseStep, LTXCoreDenoiseStep] - block_names = ["image2video", "text2video"] - block_trigger_inputs = ["image_latents", None] - - @property - def description(self): - return ( - "Auto denoise block that selects the appropriate denoise pipeline based on inputs.\n" - " - `LTXImage2VideoCoreDenoiseStep` is used when `image_latents` is provided.\n" - " - `LTXCoreDenoiseStep` is used otherwise (text-to-video)." - ) - - -# auto_docstring -class LTXAutoBlocks(SequentialPipelineBlocks): - """ - Auto blocks for LTX Video that support both text-to-video and image-to-video workflows. - - Supported workflows: - - `text2video`: requires `prompt` - - `image2video`: requires `image`, `prompt` - - Components: - text_encoder (`T5EncoderModel`) tokenizer (`T5Tokenizer`) guider (`ClassifierFreeGuidance`) vae - (`AutoencoderKLLTXVideo`) video_processor (`VideoProcessor`) scheduler (`FlowMatchEulerDiscreteScheduler`) - pachifier (`LTXVideoPachifier`) transformer (`LTXVideoTransformer3DModel`) - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The prompt or prompts not to guide the image generation. - max_sequence_length (`int`, *optional*, defaults to 128): - Maximum sequence length for prompt encoding. - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - height (`int`, *optional*, defaults to 512): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 704): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - num_inference_steps (`int`): - The number of denoising steps. - timesteps (`Tensor`): - Timesteps for the denoising process. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - num_frames (`int`, *optional*, defaults to 161): - TODO: Add description. - frame_rate (`int`, *optional*, defaults to 25): - TODO: Add description. - latents (`Tensor`): - Pre-generated noisy latents for image generation. - image_latents (`Tensor`, *optional*): - TODO: Add description. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - output_type (`str`, *optional*, defaults to np): - Output format: 'pil', 'np', 'pt'. - decode_timestep (`None`, *optional*, defaults to 0.0): - TODO: Add description. - decode_noise_scale (`None`, *optional*): - TODO: Add description. - - Outputs: - videos (`list`): - The generated videos. - """ - - model_name = "ltx" - block_classes = [ - LTXTextEncoderStep, - LTXAutoVaeEncoderStep, - LTXAutoCoreDenoiseStep, - LTXVaeDecoderStep, - ] - block_names = ["text_encoder", "vae_encoder", "denoise", "decode"] - _workflow_map = { - "text2video": {"prompt": True}, - "image2video": {"image": True, "prompt": True}, - } - - @property - def description(self): - return "Auto blocks for LTX Video that support both text-to-video and image-to-video workflows." - - @property - def outputs(self): - return [OutputParam.template("videos")] - - -# auto_docstring -class LTXImage2VideoBlocks(SequentialPipelineBlocks): - """ - Modular pipeline blocks for LTX Video image-to-video. - - Components: - text_encoder (`T5EncoderModel`) tokenizer (`T5Tokenizer`) guider (`ClassifierFreeGuidance`) vae - (`AutoencoderKLLTXVideo`) video_processor (`VideoProcessor`) scheduler (`FlowMatchEulerDiscreteScheduler`) - pachifier (`LTXVideoPachifier`) transformer (`LTXVideoTransformer3DModel`) - - Inputs: - prompt (`str`): - The prompt or prompts to guide image generation. - negative_prompt (`str`, *optional*): - The prompt or prompts not to guide the image generation. - max_sequence_length (`int`, *optional*, defaults to 128): - Maximum sequence length for prompt encoding. - image (`Image | list`, *optional*): - Reference image(s) for denoising. Can be a single image or list of images. - height (`int`, *optional*, defaults to 512): - The height in pixels of the generated image. - width (`int`, *optional*, defaults to 704): - The width in pixels of the generated image. - generator (`Generator`, *optional*): - Torch generator for deterministic generation. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - num_inference_steps (`int`, *optional*, defaults to 50): - The number of denoising steps. - timesteps (`Tensor`, *optional*): - Timesteps for the denoising process. - sigmas (`list`, *optional*): - Custom sigmas for the denoising process. - num_frames (`int`, *optional*, defaults to 161): - TODO: Add description. - frame_rate (`int`, *optional*, defaults to 25): - TODO: Add description. - latents (`Tensor`, *optional*): - Pre-generated noisy latents for image generation. - image_latents (`Tensor`): - TODO: Add description. - attention_kwargs (`dict`, *optional*): - Additional kwargs for attention processors. - output_type (`str`, *optional*, defaults to np): - Output format: 'pil', 'np', 'pt'. - decode_timestep (`None`, *optional*, defaults to 0.0): - TODO: Add description. - decode_noise_scale (`None`, *optional*): - TODO: Add description. - - Outputs: - videos (`list`): - The generated videos. - """ - - model_name = "ltx" - block_classes = [ - LTXTextEncoderStep, - LTXAutoVaeEncoderStep, - LTXImage2VideoCoreDenoiseStep, - LTXVaeDecoderStep, - ] - block_names = ["text_encoder", "vae_encoder", "denoise", "decode"] - - @property - def description(self): - return "Modular pipeline blocks for LTX Video image-to-video." - - @property - def outputs(self): - return [OutputParam.template("videos")] diff --git a/diffusers/modular_pipelines/ltx/modular_pipeline.py b/diffusers/modular_pipelines/ltx/modular_pipeline.py deleted file mode 100644 index a5771e376cd764368ae2b51a36e736f2cbc4be77..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/ltx/modular_pipeline.py +++ /dev/null @@ -1,95 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import torch - -from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import LTXVideoLoraLoaderMixin -from ...utils import logging -from ..modular_pipeline import ModularPipeline - - -logger = logging.get_logger(__name__) - - -class LTXVideoPachifier(ConfigMixin): - """ - A class to pack and unpack latents for LTX Video. - """ - - config_name = "config.json" - - @register_to_config - def __init__(self, patch_size: int = 1, patch_size_t: int = 1): - super().__init__() - - def pack_latents(self, latents: torch.Tensor) -> torch.Tensor: - batch_size, _, num_frames, height, width = latents.shape - patch_size = self.config.patch_size - patch_size_t = self.config.patch_size_t - post_patch_num_frames = num_frames // patch_size_t - post_patch_height = height // patch_size - post_patch_width = width // patch_size - latents = latents.reshape( - batch_size, - -1, - post_patch_num_frames, - patch_size_t, - post_patch_height, - patch_size, - post_patch_width, - patch_size, - ) - latents = latents.permute(0, 2, 4, 6, 1, 3, 5, 7).flatten(4, 7).flatten(1, 3) - return latents - - def unpack_latents(self, latents: torch.Tensor, num_frames: int, height: int, width: int) -> torch.Tensor: - batch_size = latents.size(0) - patch_size = self.config.patch_size - patch_size_t = self.config.patch_size_t - latents = latents.reshape(batch_size, num_frames, height, width, -1, patch_size_t, patch_size, patch_size) - latents = latents.permute(0, 4, 1, 5, 2, 6, 3, 7).flatten(6, 7).flatten(4, 5).flatten(2, 3) - return latents - - -class LTXModularPipeline( - ModularPipeline, - LTXVideoLoraLoaderMixin, -): - """ - A ModularPipeline for LTX Video. - - > [!WARNING] > This is an experimental feature and is likely to change in the future. - """ - - default_blocks_name = "LTXAutoBlocks" - - @property - def vae_spatial_compression_ratio(self): - if getattr(self, "vae", None) is not None: - return self.vae.spatial_compression_ratio - return 32 - - @property - def vae_temporal_compression_ratio(self): - if getattr(self, "vae", None) is not None: - return self.vae.temporal_compression_ratio - return 8 - - @property - def requires_unconditional_embeds(self): - if hasattr(self, "guider") and self.guider is not None: - return self.guider._enabled and self.guider.num_conditions > 1 - return False diff --git a/diffusers/modular_pipelines/mellon_node_utils.py b/diffusers/modular_pipelines/mellon_node_utils.py deleted file mode 100644 index f65459dfc99023c72d250df35f3eee7554b81ecc..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/mellon_node_utils.py +++ /dev/null @@ -1,1101 +0,0 @@ -import copy -import json -import logging -import os - -# Simple typed wrapper for parameter overrides -from dataclasses import asdict, dataclass -from typing import Any - -from huggingface_hub import create_repo, hf_hub_download, upload_file -from huggingface_hub.utils import ( - EntryNotFoundError, - HfHubHTTPError, - RepositoryNotFoundError, - RevisionNotFoundError, -) - -from ..utils import HUGGINGFACE_CO_RESOLVE_ENDPOINT -from .modular_pipeline_utils import InputParam, OutputParam - - -logger = logging.getLogger(__name__) - - -def _name_to_label(name: str) -> str: - """Convert snake_case name to Title Case label.""" - return name.replace("_", " ").title() - - -# Template definitions for standard diffuser pipeline parameters -MELLON_PARAM_TEMPLATES = { - # Image I/O - "image": {"label": "Image", "type": "image", "display": "input", "required_block_params": ["image"]}, - "images": {"label": "Images", "type": "image", "display": "output", "required_block_params": ["images"]}, - "control_image": { - "label": "Control Image", - "type": "image", - "display": "input", - "required_block_params": ["control_image"], - }, - # Latents - "latents": {"label": "Latents", "type": "latents", "display": "input", "required_block_params": ["latents"]}, - "image_latents": { - "label": "Image Latents", - "type": "latents", - "display": "input", - "required_block_params": ["image_latents"], - }, - "first_frame_latents": { - "label": "First Frame Latents", - "type": "latents", - "display": "input", - "required_block_params": ["first_frame_latents"], - }, - "latents_preview": {"label": "Latents Preview", "type": "latent", "display": "output"}, - # Image Latents with Strength - "image_latents_with_strength": { - "name": "image_latents", # name is not same as template key - "label": "Image Latents", - "type": "latents", - "display": "input", - "onChange": {"false": ["height", "width"], "true": ["strength"]}, - "required_block_params": ["image_latents", "strength"], - }, - # Embeddings - "embeddings": {"label": "Text Embeddings", "type": "embeddings", "display": "output"}, - "image_embeds": { - "label": "Image Embeddings", - "type": "image_embeds", - "display": "output", - "required_block_params": ["image_embeds"], - }, - # Text inputs - "prompt": { - "label": "Prompt", - "type": "string", - "display": "textarea", - "default": "", - "required_block_params": ["prompt"], - }, - "negative_prompt": { - "label": "Negative Prompt", - "type": "string", - "display": "textarea", - "default": "", - "required_block_params": ["negative_prompt"], - }, - # Numeric params - "guidance_scale": { - "label": "Guidance Scale", - "type": "float", - "display": "slider", - "default": 5.0, - "min": 1.0, - "max": 30.0, - "step": 0.1, - }, - "strength": { - "label": "Strength", - "type": "float", - "default": 0.5, - "min": 0.0, - "max": 1.0, - "step": 0.01, - "required_block_params": ["strength"], - }, - "height": { - "label": "Height", - "type": "int", - "default": 1024, - "min": 64, - "step": 8, - "required_block_params": ["height"], - }, - "width": { - "label": "Width", - "type": "int", - "default": 1024, - "min": 64, - "step": 8, - "required_block_params": ["width"], - }, - "seed": { - "label": "Seed", - "type": "int", - "default": 0, - "min": 0, - "max": 4294967295, - "display": "random", - "required_block_params": ["generator"], - }, - "num_inference_steps": { - "label": "Steps", - "type": "int", - "default": 25, - "min": 1, - "max": 100, - "display": "slider", - "required_block_params": ["num_inference_steps"], - }, - "num_frames": { - "label": "Frames", - "type": "int", - "default": 81, - "min": 1, - "max": 480, - "display": "slider", - "required_block_params": ["num_frames"], - }, - "layers": { - "label": "Layers", - "type": "int", - "default": 4, - "min": 1, - "max": 10, - "display": "slider", - "required_block_params": ["layers"], - }, - "output_type": { - "label": "Output Type", - "type": "dropdown", - "default": "np", - "options": ["np", "pil", "pt"], - }, - # ControlNet - "controlnet_conditioning_scale": { - "label": "Controlnet Conditioning Scale", - "type": "float", - "default": 0.5, - "min": 0.0, - "max": 1.0, - "step": 0.01, - "required_block_params": ["controlnet_conditioning_scale"], - }, - "control_guidance_start": { - "label": "Control Guidance Start", - "type": "float", - "default": 0.0, - "min": 0.0, - "max": 1.0, - "step": 0.01, - "required_block_params": ["control_guidance_start"], - }, - "control_guidance_end": { - "label": "Control Guidance End", - "type": "float", - "default": 1.0, - "min": 0.0, - "max": 1.0, - "step": 0.01, - "required_block_params": ["control_guidance_end"], - }, - # Video - "videos": {"label": "Videos", "type": "video", "display": "output", "required_block_params": ["videos"]}, - # Models - "vae": {"label": "VAE", "type": "diffusers_auto_model", "display": "input", "required_block_params": ["vae"]}, - "image_encoder": { - "label": "Image Encoder", - "type": "diffusers_auto_model", - "display": "input", - "required_block_params": ["image_encoder"], - }, - "unet": {"label": "Denoise Model", "type": "diffusers_auto_model", "display": "input"}, - "scheduler": {"label": "Scheduler", "type": "diffusers_auto_model", "display": "input"}, - "controlnet": { - "label": "ControlNet Model", - "type": "diffusers_auto_model", - "display": "input", - "required_block_params": ["controlnet"], - }, - "text_encoders": { - "label": "Text Encoders", - "type": "diffusers_auto_models", - "display": "input", - "required_block_params": ["text_encoder"], - }, - # Bundles/Custom - "controlnet_bundle": { - "label": "ControlNet", - "type": "custom_controlnet", - "display": "input", - "required_block_params": "controlnet_image", - }, - "ip_adapter": {"label": "IP Adapter", "type": "custom_ip_adapter", "display": "input"}, - "guider": { - "label": "Guider", - "type": "custom_guider", - "display": "input", - "onChange": {False: ["guidance_scale"], True: []}, - }, - "doc": {"label": "Doc", "type": "string", "display": "output"}, -} - - -class MellonParamMeta(type): - """Metaclass that enables MellonParam.template_name(**overrides) syntax.""" - - def __getattr__(cls, name: str): - if name in MELLON_PARAM_TEMPLATES: - - def factory(default=None, **overrides): - template = MELLON_PARAM_TEMPLATES[name] - # Use template's name if specified, otherwise use the key - params = {"name": template.get("name", name), **template, **overrides} - if default is not None: - params["default"] = default - return cls(**params) - - return factory - - raise AttributeError(f"type object 'MellonParam' has no attribute '{name}'") - - -@dataclass(frozen=True) -class MellonParam(metaclass=MellonParamMeta): - """ - Parameter definition for Mellon nodes. - - Usage: - ```python - # From template (standard diffuser params) - MellonParam.seed() - MellonParam.prompt(default="a cat") - MellonParam.latents(display="output") - - # Generic inputs (for custom blocks) - MellonParam.Input.slider("my_scale", default=1.0, min=0.0, max=2.0) - MellonParam.Input.dropdown("mode", options=["fast", "slow"]) - - # Generic outputs - MellonParam.Output.image("result_images") - - # Fully custom - MellonParam(name="custom", label="Custom", type="float", default=0.5) - ``` - """ - - name: str - label: str - type: str - display: str | None = None - default: Any = None - min: float | None = None - max: float | None = None - step: float | None = None - options: Any = None - value: Any = None - fieldOptions: dict[str, Any] | None = None - onChange: Any = None - onSignal: Any = None - required_block_params: str | list[str] | None = None - - def to_dict(self) -> dict[str, Any]: - """Convert to dict for Mellon schema, excluding None values and internal fields.""" - data = asdict(self) - return {k: v for k, v in data.items() if v is not None and k not in ("name", "required_block_params")} - - # ========================================================================= - # Input: Generic input parameter factories (for custom blocks) - # ========================================================================= - class Input: - """input UI elements for custom blocks.""" - - @classmethod - def image(cls, name: str) -> "MellonParam": - """image input.""" - return MellonParam(name=name, label=_name_to_label(name), type="image", display="input") - - @classmethod - def textbox(cls, name: str, default: str = "") -> "MellonParam": - """text input as textarea.""" - return MellonParam( - name=name, label=_name_to_label(name), type="string", display="textarea", default=default - ) - - @classmethod - def dropdown(cls, name: str, options: list[str] = None, default: str = None) -> "MellonParam": - """dropdown selection.""" - if options and not default: - default = options[0] - if not default: - default = "" - if not options: - options = [default] - return MellonParam(name=name, label=_name_to_label(name), type="string", options=options, value=default) - - @classmethod - def slider( - cls, name: str, default: float = 0, min: float = None, max: float = None, step: float = None - ) -> "MellonParam": - """slider input.""" - is_float = isinstance(default, float) or (step is not None and isinstance(step, float)) - param_type = "float" if is_float else "int" - if min is None: - min = default - if max is None: - max = default - if step is None: - step = 0.01 if is_float else 1 - return MellonParam( - name=name, - label=_name_to_label(name), - type=param_type, - display="slider", - default=default, - min=min, - max=max, - step=step, - ) - - @classmethod - def number( - cls, name: str, default: float = 0, min: float = None, max: float = None, step: float = None - ) -> "MellonParam": - """number input (no slider).""" - is_float = isinstance(default, float) or (step is not None and isinstance(step, float)) - param_type = "float" if is_float else "int" - return MellonParam( - name=name, label=_name_to_label(name), type=param_type, default=default, min=min, max=max, step=step - ) - - @classmethod - def seed(cls, name: str = "seed", default: int = 0) -> "MellonParam": - """seed input with randomize button.""" - return MellonParam( - name=name, - label=_name_to_label(name), - type="int", - display="random", - default=default, - min=0, - max=4294967295, - ) - - @classmethod - def checkbox(cls, name: str, default: bool = False) -> "MellonParam": - """boolean checkbox.""" - return MellonParam(name=name, label=_name_to_label(name), type="boolean", value=default) - - @classmethod - def custom_type(cls, name: str, type: str) -> "MellonParam": - """custom type input for node connections.""" - return MellonParam(name=name, label=_name_to_label(name), type=type, display="input") - - @classmethod - def model(cls, name: str) -> "MellonParam": - """model input for diffusers components.""" - return MellonParam(name=name, label=_name_to_label(name), type="diffusers_auto_model", display="input") - - # ========================================================================= - # Output: Generic output parameter factories (for custom blocks) - # ========================================================================= - class Output: - """output UI elements for custom blocks.""" - - @classmethod - def image(cls, name: str) -> "MellonParam": - """image output.""" - return MellonParam(name=name, label=_name_to_label(name), type="image", display="output") - - @classmethod - def video(cls, name: str) -> "MellonParam": - """video output.""" - return MellonParam(name=name, label=_name_to_label(name), type="video", display="output") - - @classmethod - def text(cls, name: str) -> "MellonParam": - """text output.""" - return MellonParam(name=name, label=_name_to_label(name), type="string", display="output") - - @classmethod - def custom_type(cls, name: str, type: str) -> "MellonParam": - """custom type output for node connections.""" - return MellonParam(name=name, label=_name_to_label(name), type=type, display="output") - - @classmethod - def model(cls, name: str) -> "MellonParam": - """model output for diffusers components.""" - return MellonParam(name=name, label=_name_to_label(name), type="diffusers_auto_model", display="output") - - -def input_param_to_mellon_param(input_param: "InputParam") -> MellonParam: - """ - Convert an InputParam to a MellonParam using metadata. - - Args: - input_param: An InputParam with optional metadata containing either: - - {"mellon": ""} for simple types (image, textbox, slider, etc.) - - {"mellon": MellonParam(...)} for full control over UI configuration - - Returns: - MellonParam instance - """ - name = input_param.name - metadata = input_param.metadata - mellon_value = metadata.get("mellon") if metadata else None - default = input_param.default - - # If it's already a MellonParam, return it directly - if isinstance(mellon_value, MellonParam): - return mellon_value - - mellon_type = mellon_value - - if mellon_type == "image": - return MellonParam.Input.image(name) - elif mellon_type == "textbox": - return MellonParam.Input.textbox(name, default=default or "") - elif mellon_type == "dropdown": - return MellonParam.Input.dropdown(name, default=default or "") - elif mellon_type == "slider": - return MellonParam.Input.slider(name, default=default or 0) - elif mellon_type == "number": - return MellonParam.Input.number(name, default=default or 0) - elif mellon_type == "seed": - return MellonParam.Input.seed(name, default=default or 0) - elif mellon_type == "checkbox": - return MellonParam.Input.checkbox(name, default=default or False) - elif mellon_type == "model": - return MellonParam.Input.model(name) - else: - # None or unknown -> custom - return MellonParam.Input.custom_type(name, type="custom") - - -def output_param_to_mellon_param(output_param: "OutputParam") -> MellonParam: - """ - Convert an OutputParam to a MellonParam using metadata. - - Args: - output_param: An OutputParam with optional metadata={"mellon": ""} where type is one of: - image, video, text, model. If metadata is None or unknown, maps to "custom". - - Returns: - MellonParam instance - """ - name = output_param.name - metadata = output_param.metadata - mellon_type = metadata.get("mellon") if metadata else None - - if mellon_type == "image": - return MellonParam.Output.image(name) - elif mellon_type == "video": - return MellonParam.Output.video(name) - elif mellon_type == "text": - return MellonParam.Output.text(name) - elif mellon_type == "model": - return MellonParam.Output.model(name) - else: - # None or unknown -> custom - return MellonParam.Output.custom_type(name, type="custom") - - -DEFAULT_NODE_SPECS = { - "controlnet": None, - "denoise": { - "inputs": [ - MellonParam.embeddings(display="input"), - MellonParam.width(), - MellonParam.height(), - MellonParam.seed(), - MellonParam.num_inference_steps(), - MellonParam.num_frames(), - MellonParam.guidance_scale(), - MellonParam.strength(), - MellonParam.image_latents_with_strength(), - MellonParam.image_latents(), - MellonParam.first_frame_latents(), - MellonParam.controlnet_bundle(display="input"), - ], - "model_inputs": [ - MellonParam.unet(), - MellonParam.guider(), - MellonParam.scheduler(), - ], - "outputs": [ - MellonParam.latents(display="output"), - MellonParam.latents_preview(), - MellonParam.doc(), - ], - "required_inputs": ["embeddings"], - "required_model_inputs": ["unet", "scheduler"], - "block_name": "denoise", - }, - "vae_encoder": { - "inputs": [ - MellonParam.image(), - ], - "model_inputs": [ - MellonParam.vae(), - ], - "outputs": [ - MellonParam.image_latents(display="output"), - MellonParam.doc(), - ], - "required_inputs": ["image"], - "required_model_inputs": ["vae"], - "block_name": "vae_encoder", - }, - "text_encoder": { - "inputs": [ - MellonParam.prompt(), - MellonParam.negative_prompt(), - ], - "model_inputs": [ - MellonParam.text_encoders(), - ], - "outputs": [ - MellonParam.embeddings(display="output"), - MellonParam.doc(), - ], - "required_inputs": ["prompt"], - "required_model_inputs": ["text_encoders"], - "block_name": "text_encoder", - }, - "decoder": { - "inputs": [ - MellonParam.latents(display="input"), - ], - "model_inputs": [ - MellonParam.vae(), - ], - "outputs": [ - MellonParam.images(), - MellonParam.videos(), - MellonParam.doc(), - ], - "required_inputs": ["latents"], - "required_model_inputs": ["vae"], - "block_name": "decode", - }, -} - - -def mark_required(label: str, marker: str = " *") -> str: - """Add required marker to label if not already present.""" - if label.endswith(marker): - return label - return f"{label}{marker}" - - -def node_spec_to_mellon_dict(node_spec: dict[str, Any], node_type: str) -> dict[str, Any]: - """ - Convert a node spec dict into Mellon format. - - A node spec is how we define a Mellon diffusers node in code. This function converts it into the `params` map - format that Mellon UI expects. - - The `params` map is a dict where keys are parameter names and values are UI configuration: - ```python - {"seed": {"label": "Seed", "type": "int", "default": 0}} - ``` - - For Modular Mellon nodes, we need to distinguish: - - `inputs`: Pipeline inputs (e.g., seed, prompt, image) - - `model_inputs`: Model components (e.g., unet, vae, scheduler) - - `outputs`: Node outputs (e.g., latents, images) - - The node spec also includes: - - `required_inputs` / `required_model_inputs`: Which params are required (marked with *) - - `block_name`: The modular pipeline block this node corresponds to on backend - - We provide factory methods for common parameters (e.g., `MellonParam.seed()`, `MellonParam.unet()`) so you don't - have to manually specify all the UI configuration. - - Args: - node_spec: Dict with `inputs`, `model_inputs`, `outputs` (lists of MellonParam), - plus `required_inputs`, `required_model_inputs`, `block_name`. - node_type: The node type string (e.g., "denoise", "controlnet") - - Returns: - Dict with: - - `params`: Flat dict of all params in Mellon UI format - - `input_names`: List of input parameter names - - `model_input_names`: List of model input parameter names - - `output_names`: List of output parameter names - - `block_name`: The backend block name - - `node_type`: The node type - - Example: - ```python - node_spec = { - "inputs": [MellonParam.seed(), MellonParam.prompt()], - "model_inputs": [MellonParam.unet()], - "outputs": [MellonParam.latents(display="output")], - "required_inputs": ["prompt"], - "required_model_inputs": ["unet"], - "block_name": "denoise", - } - - result = node_spec_to_mellon_dict(node_spec, "denoise") - # Returns: - # { - # "params": { - # "seed": {"label": "Seed", "type": "int", "default": 0}, - # "prompt": {"label": "Prompt *", "type": "string", "default": ""}, # * marks required - # "unet": {"label": "Denoise Model *", "type": "diffusers_auto_model", "display": "input"}, - # "latents": {"label": "Latents", "type": "latents", "display": "output"}, - # }, - # "input_names": ["seed", "prompt"], - # "model_input_names": ["unet"], - # "output_names": ["latents"], - # "block_name": "denoise", - # "node_type": "denoise", - # } - ``` - """ - params = {} - input_names = [] - model_input_names = [] - output_names = [] - - required_inputs = node_spec.get("required_inputs", []) - required_model_inputs = node_spec.get("required_model_inputs", []) - - # Process inputs - for p in node_spec.get("inputs", []): - param_dict = p.to_dict() - if p.name in required_inputs: - param_dict["label"] = mark_required(param_dict["label"]) - params[p.name] = param_dict - input_names.append(p.name) - - # Process model_inputs - for p in node_spec.get("model_inputs", []): - param_dict = p.to_dict() - if p.name in required_model_inputs: - param_dict["label"] = mark_required(param_dict["label"]) - params[p.name] = param_dict - model_input_names.append(p.name) - - # Process outputs: add a prefix to the output name if it already exists as an input - for p in node_spec.get("outputs", []): - if p.name in input_names: - # rename to out_ - output_name = f"out_{p.name}" - else: - output_name = p.name - params[output_name] = p.to_dict() - output_names.append(output_name) - - return { - "params": params, - "input_names": input_names, - "model_input_names": model_input_names, - "output_names": output_names, - "block_name": node_spec.get("block_name"), - "node_type": node_type, - } - - -class MellonPipelineConfig: - """ - Configuration for an entire Mellon pipeline containing multiple nodes. - - Accepts node specs as dicts with inputs/model_inputs/outputs lists of MellonParam, converts them to Mellon-ready - format, and handles save/load to Hub. - - Example: - ```python - config = MellonPipelineConfig( - node_specs={ - "denoise": { - "inputs": [MellonParam.seed(), MellonParam.prompt()], - "model_inputs": [MellonParam.unet()], - "outputs": [MellonParam.latents(display="output")], - "required_inputs": ["prompt"], - "required_model_inputs": ["unet"], - "block_name": "denoise", - }, - "decoder": { - "inputs": [MellonParam.latents(display="input")], - "outputs": [MellonParam.images()], - "block_name": "decoder", - }, - }, - label="My Pipeline", - default_repo="user/my-pipeline", - default_dtype="float16", - ) - - # Access Mellon format dict - denoise = config.node_params["denoise"] - input_names = denoise["input_names"] - params = denoise["params"] - - # Save to Hub - config.save("./my_config", push_to_hub=True, repo_id="user/my-pipeline") - - # Load from Hub - loaded = MellonPipelineConfig.load("user/my-pipeline") - ``` - """ - - config_name = "mellon_pipeline_config.json" - - def __init__( - self, - node_specs: dict[str, dict[str, Any] | None], - label: str = "", - default_repo: str = "", - default_dtype: str = "", - ): - """ - Args: - node_specs: Dict mapping node_type to node spec or None. - Node spec has: inputs, model_inputs, outputs, required_inputs, required_model_inputs, - block_name (all optional) - label: Human-readable label for the pipeline - default_repo: Default HuggingFace repo for this pipeline - default_dtype: Default dtype (e.g., "float16", "bfloat16") - """ - # Convert all node specs to Mellon format immediately - self.node_specs = node_specs - - self.label = label - self.default_repo = default_repo - self.default_dtype = default_dtype - - @property - def node_params(self) -> dict[str, Any]: - """Lazily compute node_params from node_specs.""" - if self.node_specs is None: - return self._node_params - - params = {} - for node_type, spec in self.node_specs.items(): - if spec is None: - params[node_type] = None - else: - params[node_type] = node_spec_to_mellon_dict(spec, node_type) - return params - - def __repr__(self) -> str: - lines = [ - f"MellonPipelineConfig(label={self.label!r}, default_repo={self.default_repo!r}, default_dtype={self.default_dtype!r})" - ] - for node_type, spec in self.node_specs.items(): - if spec is None: - lines.append(f" {node_type}: None") - else: - inputs = [p.name for p in spec.get("inputs", [])] - model_inputs = [p.name for p in spec.get("model_inputs", [])] - outputs = [p.name for p in spec.get("outputs", [])] - lines.append(f" {node_type}:") - lines.append(f" inputs: {inputs}") - lines.append(f" model_inputs: {model_inputs}") - lines.append(f" outputs: {outputs}") - return "\n".join(lines) - - def to_dict(self) -> dict[str, Any]: - """Convert to a JSON-serializable dictionary.""" - return { - "label": self.label, - "default_repo": self.default_repo, - "default_dtype": self.default_dtype, - "node_params": self.node_params, - } - - @classmethod - def from_dict(cls, data: dict[str, Any]) -> "MellonPipelineConfig": - """ - Create from a dictionary (loaded from JSON). - - Note: The mellon_params are already in Mellon format when loading from JSON. - """ - instance = cls.__new__(cls) - instance.node_specs = None - instance._node_params = data.get("node_params", {}) - instance.label = data.get("label", "") - instance.default_repo = data.get("default_repo", "") - instance.default_dtype = data.get("default_dtype", "") - return instance - - def to_json_string(self) -> str: - """Serialize to JSON string.""" - return json.dumps(self.to_dict(), indent=2, sort_keys=False) + "\n" - - def to_json_file(self, json_file_path: str | os.PathLike): - """Save to a JSON file.""" - with open(json_file_path, "w", encoding="utf-8") as writer: - writer.write(self.to_json_string()) - - @classmethod - def from_json_file(cls, json_file_path: str | os.PathLike) -> "MellonPipelineConfig": - """Load from a JSON file.""" - with open(json_file_path, "r", encoding="utf-8") as reader: - data = json.load(reader) - return cls.from_dict(data) - - def save(self, save_directory: str | os.PathLike, push_to_hub: bool = False, **kwargs): - """Save the mellon pipeline config to a directory.""" - if os.path.isfile(save_directory): - raise AssertionError(f"Provided path ({save_directory}) should be a directory, not a file") - - os.makedirs(save_directory, exist_ok=True) - output_path = os.path.join(save_directory, self.config_name) - self.to_json_file(output_path) - logger.info(f"Pipeline config saved to {output_path}") - - if push_to_hub: - commit_message = kwargs.pop("commit_message", None) - private = kwargs.pop("private", None) - create_pr = kwargs.pop("create_pr", False) - token = kwargs.pop("token", None) - repo_id = kwargs.pop("repo_id", save_directory.split(os.path.sep)[-1]) - repo_id = create_repo(repo_id, exist_ok=True, private=private, token=token).repo_id - - upload_file( - path_or_fileobj=output_path, - path_in_repo=self.config_name, - repo_id=repo_id, - token=token, - commit_message=commit_message or "Upload MellonPipelineConfig", - create_pr=create_pr, - ) - logger.info(f"Pipeline config pushed to hub: {repo_id}") - - @classmethod - def load( - cls, - pretrained_model_name_or_path: str | os.PathLike, - **kwargs, - ) -> "MellonPipelineConfig": - """Load a pipeline config from a local path or Hugging Face Hub.""" - cache_dir = kwargs.pop("cache_dir", None) - local_dir = kwargs.pop("local_dir", None) - local_dir_use_symlinks = kwargs.pop("local_dir_use_symlinks", "auto") - force_download = kwargs.pop("force_download", False) - proxies = kwargs.pop("proxies", None) - token = kwargs.pop("token", None) - local_files_only = kwargs.pop("local_files_only", False) - revision = kwargs.pop("revision", None) - subfolder = kwargs.pop("subfolder", None) - - pretrained_model_name_or_path = str(pretrained_model_name_or_path) - - if os.path.isfile(pretrained_model_name_or_path): - config_file = pretrained_model_name_or_path - elif os.path.isdir(pretrained_model_name_or_path): - config_file = os.path.join(pretrained_model_name_or_path, cls.config_name) - if not os.path.isfile(config_file): - raise EnvironmentError(f"No file named {cls.config_name} found in {pretrained_model_name_or_path}") - else: - try: - config_file = hf_hub_download( - pretrained_model_name_or_path, - filename=cls.config_name, - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - local_dir=local_dir, - local_dir_use_symlinks=local_dir_use_symlinks, - ) - except RepositoryNotFoundError: - raise EnvironmentError( - f"{pretrained_model_name_or_path} is not a local folder and is not a valid model identifier" - " listed on 'https://huggingface.co/models'\nIf this is a private repository, make sure to pass a" - " token having permission to this repo with `token` or log in with `hf auth login`." - ) - except RevisionNotFoundError: - raise EnvironmentError( - f"{revision} is not a valid git identifier (branch name, tag name or commit id) that exists for" - " this model name. Check the model page at" - f" 'https://huggingface.co/{pretrained_model_name_or_path}' for available revisions." - ) - except EntryNotFoundError: - raise EnvironmentError( - f"{pretrained_model_name_or_path} does not appear to have a file named {cls.config_name}." - ) - except HfHubHTTPError as err: - raise EnvironmentError( - "There was a specific connection error when trying to load" - f" {pretrained_model_name_or_path}:\n{err}" - ) - except ValueError: - raise EnvironmentError( - f"We couldn't connect to '{HUGGINGFACE_CO_RESOLVE_ENDPOINT}' to load this model, couldn't find it" - f" in the cached files and it looks like {pretrained_model_name_or_path} is not the path to a" - f" directory containing a {cls.config_name} file.\nCheckout your internet connection or see how to" - " run the library in offline mode at" - " 'https://huggingface.co/docs/diffusers/installation#offline-mode'." - ) - except EnvironmentError: - raise EnvironmentError( - f"Can't load config for '{pretrained_model_name_or_path}'. If you were trying to load it from " - "'https://huggingface.co/models', make sure you don't have a local directory with the same name. " - f"Otherwise, make sure '{pretrained_model_name_or_path}' is the correct path to a directory " - f"containing a {cls.config_name} file" - ) - - try: - return cls.from_json_file(config_file) - except (json.JSONDecodeError, UnicodeDecodeError): - raise EnvironmentError(f"The config file at '{config_file}' is not a valid JSON file.") - - @classmethod - def from_blocks( - cls, - blocks, - template: dict[str, dict[str, Any]] | None = None, - label: str = "", - default_repo: str = "", - default_dtype: str = "bfloat16", - ) -> "MellonPipelineConfig": - """ - Create MellonPipelineConfig by matching template against actual pipeline blocks. - """ - if template is None: - template = DEFAULT_NODE_SPECS - - sub_block_map = dict(blocks.sub_blocks) - - def filter_spec_for_block(template_spec: dict[str, Any], block) -> dict[str, Any] | None: - """Filter template spec params based on what the block actually supports.""" - block_input_names = set(block.input_names) - block_output_names = set(block.intermediate_output_names) - block_component_names = set(block.component_names) - - filtered_inputs = [ - p - for p in template_spec.get("inputs", []) - if p.required_block_params is None - or all(name in block_input_names for name in p.required_block_params) - ] - filtered_model_inputs = [ - p - for p in template_spec.get("model_inputs", []) - if p.required_block_params is None - or all(name in block_component_names for name in p.required_block_params) - ] - filtered_outputs = [ - p - for p in template_spec.get("outputs", []) - if p.required_block_params is None - or all(name in block_output_names for name in p.required_block_params) - ] - - filtered_input_names = {p.name for p in filtered_inputs} - filtered_model_input_names = {p.name for p in filtered_model_inputs} - - filtered_required_inputs = [ - r for r in template_spec.get("required_inputs", []) if r in filtered_input_names - ] - filtered_required_model_inputs = [ - r for r in template_spec.get("required_model_inputs", []) if r in filtered_model_input_names - ] - - return { - "inputs": filtered_inputs, - "model_inputs": filtered_model_inputs, - "outputs": filtered_outputs, - "required_inputs": filtered_required_inputs, - "required_model_inputs": filtered_required_model_inputs, - "block_name": template_spec.get("block_name"), - } - - # Build node specs - node_specs = {} - for node_type, template_spec in template.items(): - if template_spec is None: - node_specs[node_type] = None - continue - - block_name = template_spec.get("block_name") - if block_name is None or block_name not in sub_block_map: - node_specs[node_type] = None - continue - - node_specs[node_type] = filter_spec_for_block(template_spec, sub_block_map[block_name]) - - return cls( - node_specs=node_specs, - label=label or getattr(blocks, "model_name", ""), - default_repo=default_repo, - default_dtype=default_dtype, - ) - - @classmethod - def from_custom_block( - cls, - block, - node_label: str = None, - input_types: dict[str, Any] | None = None, - output_types: dict[str, Any] | None = None, - ) -> "MellonPipelineConfig": - """ - Create a MellonPipelineConfig from a custom block. - - Args: - block: A block instance with `inputs`, `outputs`, and `expected_components`/`component_names` properties. - Each InputParam/OutputParam should have metadata={"mellon": ""} where type is one of: image, - video, text, checkbox, number, slider, dropdown, model. If metadata is None, maps to "custom". - node_label: The display label for the node. Defaults to block class name with spaces. - input_types: - Optional dict mapping input param names to mellon types. Overrides the block's metadata if provided. - Example: {"prompt": "textbox", "image": "image"} - output_types: - Optional dict mapping output param names to mellon types. Overrides the block's metadata if provided. - Example: {"prompt": "text", "images": "image"} - - Returns: - MellonPipelineConfig instance - """ - if node_label is None: - class_name = block.__class__.__name__ - node_label = "".join([" " + c if c.isupper() else c for c in class_name]).strip() - - if input_types is None: - input_types = {} - if output_types is None: - output_types = {} - - inputs = [] - model_inputs = [] - outputs = [] - - # Process block inputs - for input_param in block.inputs: - if input_param.name is None: - continue - if input_param.name in input_types: - input_param = copy.copy(input_param) - input_param.metadata = {"mellon": input_types[input_param.name]} - print(f" processing input: {input_param.name}, metadata: {input_param.metadata}") - inputs.append(input_param_to_mellon_param(input_param)) - - # Process block outputs - for output_param in block.outputs: - if output_param.name is None: - continue - if output_param.name in output_types: - output_param = copy.copy(output_param) - output_param.metadata = {"mellon": output_types[output_param.name]} - outputs.append(output_param_to_mellon_param(output_param)) - - # Process expected components (all map to model inputs) - component_names = block.component_names - for component_name in component_names: - model_inputs.append(MellonParam.Input.model(component_name)) - - # Always add doc output - outputs.append(MellonParam.doc()) - - node_spec = { - "inputs": inputs, - "model_inputs": model_inputs, - "outputs": outputs, - "required_inputs": [], - "required_model_inputs": [], - "block_name": "custom", - } - - return cls( - node_specs={"custom": node_spec}, - label=node_label, - ) diff --git a/diffusers/modular_pipelines/minimax_h3/__init__.py b/diffusers/modular_pipelines/minimax_h3/__init__.py deleted file mode 100644 index 6f492f17eff1edd7b6590282802e27b51ee45cf0..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/minimax_h3/__init__.py +++ /dev/null @@ -1,49 +0,0 @@ -from typing import TYPE_CHECKING - -from ...utils import ( - DIFFUSERS_SLOW_IMPORT, - OptionalDependencyNotAvailable, - _LazyModule, - get_objects_from_module, - is_torch_available, - is_transformers_available, -) - - -_dummy_objects = {} -_import_structure = {} - -try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() -except OptionalDependencyNotAvailable: - from ...utils import dummy_torch_and_transformers_objects # noqa F403 - - _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) -else: - _import_structure["modular_blocks_minimax_h3"] = ["MiniMaxH3Blocks", "MiniMaxH3Ref2VABlocks"] - _import_structure["modular_pipeline"] = ["MiniMaxH3ModularPipeline", "MiniMaxH3Ref2VAModularPipeline"] - _import_structure["packing_ref2va"] = ["MiniMaxH3Reference"] - -if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: - try: - if not (is_transformers_available() and is_torch_available()): - raise OptionalDependencyNotAvailable() - except OptionalDependencyNotAvailable: - from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 - else: - from .modular_blocks_minimax_h3 import MiniMaxH3Blocks, MiniMaxH3Ref2VABlocks - from .modular_pipeline import MiniMaxH3ModularPipeline, MiniMaxH3Ref2VAModularPipeline - from .packing_ref2va import MiniMaxH3Reference -else: - import sys - - sys.modules[__name__] = _LazyModule( - __name__, - globals()["__file__"], - _import_structure, - module_spec=__spec__, - ) - - for name, value in _dummy_objects.items(): - setattr(sys.modules[__name__], name, value) diff --git a/diffusers/modular_pipelines/minimax_h3/before_denoise.py b/diffusers/modular_pipelines/minimax_h3/before_denoise.py deleted file mode 100644 index ef874b1e0100dd770ee42fefcd4e51b5414fbd24..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/minimax_h3/before_denoise.py +++ /dev/null @@ -1,425 +0,0 @@ -# Copyright 2026 The MiniMax and HuggingFace Teams. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ...schedulers import MiniMaxH3Scheduler -from ...utils import logging -from ...utils.torch_utils import randn_tensor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import MiniMaxH3ModularPipeline, MiniMaxH3Ref2VAModularPipeline -from .packing import ( - MINIMAX_H3_AUDIO_CHANNELS, - MINIMAX_H3_KEYFRAME_NOISE_AUG, - MiniMaxH3PackedSequence, - build_packed_sequence, - build_row_timesteps, - patchify_video_latents, -) -from .packing_ref2va import MiniMaxH3PreparedReference, build_ref2va_packed_sequence - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _layout_inputs() -> list[InputParam]: - r"""What both packed layouts are built from, beyond the conditioning of the task itself.""" - return [ - InputParam( - name="text_token_tags", - type_hint=torch.Tensor, - required=True, - description="The per-row modality tag of every row of `prompt_embeds`.", - ), - InputParam( - name="num_latent_frames", type_hint=int, required=True, description="Number of video latent frames." - ), - InputParam(name="latent_height", type_hint=int, required=True, description="Height of the video latents."), - InputParam(name="latent_width", type_hint=int, required=True, description="Width of the video latents."), - InputParam( - name="num_audio_latents", - type_hint=int, - required=True, - description="Number of audio latents per channel.", - ), - ] - - -def _layout_outputs() -> list[OutputParam]: - r"""The row layout of the packed sequence, shared by the two tasks.""" - return [ - OutputParam( - "layout", - type_hint=MiniMaxH3PackedSequence, - description="The structural description of the packed sequence.", - ), - OutputParam( - "position_ids", - type_hint=torch.Tensor, - description="The `(t, h, w)` rotary coordinate of every row, in float64.", - ), - OutputParam("token_tags", type_hint=torch.Tensor, description="The modality tag of every row."), - OutputParam( - "video_indices", - type_hint=torch.Tensor, - description="Sequence positions of the video rows, conditioning rows first.", - ), - OutputParam( - "audio_indices", - type_hint=torch.Tensor, - description="Sequence positions of the audio rows, reference rows first.", - ), - OutputParam("text_indices", type_hint=torch.Tensor, description="Sequence positions of the text rows."), - OutputParam( - "num_condition_video_rows", - type_hint=int, - description="How many leading video rows are conditioning rows rather than generated rows.", - ), - OutputParam( - "num_condition_audio_rows", - type_hint=int, - description="How many leading audio rows are reference rows rather than generated rows.", - ), - ] - - -def _set_layout_state(block_state, layout: MiniMaxH3PackedSequence, device: torch.device) -> None: - block_state.layout = layout - block_state.position_ids = layout.position_ids.to(device) - block_state.token_tags = layout.token_tags.to(device) - block_state.video_indices = layout.video_indices.to(device) - block_state.audio_indices = layout.audio_indices.to(device) - block_state.text_indices = layout.text_indices.to(device) - block_state.num_condition_video_rows = layout.num_condition_video_rows - block_state.num_condition_audio_rows = layout.num_condition_audio_rows - - -class MiniMaxH3PrepareLayoutStep(ModularPipelineBlocks): - model_name = "minimax-h3" - - @property - def description(self) -> str: - return ( - "Builds the packed layout of a `t2va` / `fl2va` request — `[text | keyframe conditions | target audio | " - "target video]` — and its fp64 rotary grid. MiniMax-H3 runs full self-attention over this one sequence, " - "so the layout is what every later block addresses rows through." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - *_layout_inputs(), - InputParam( - name="keyframe_anchors", - type_hint=tuple, - default=(), - description="Which end of the video every keyframe is anchored to, in packed order.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return _layout_outputs() - - @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - layout = build_packed_sequence( - block_state.text_token_tags, - block_state.num_latent_frames, - block_state.latent_height, - block_state.latent_width, - block_state.num_audio_latents, - components.patch_size, - block_state.keyframe_anchors, - ) - _set_layout_state(block_state, layout, components._execution_device) - - self.set_block_state(state, block_state) - return components, state - - -class MiniMaxH3Ref2VAPrepareLayoutStep(ModularPipelineBlocks): - model_name = "minimax-h3-ref2va" - - @property - def description(self) -> str: - return ( - "Builds the packed layout of a `ref2va` request — `[text | reference blocks | target audio | target " - "video]` — and its fp64 rotary grid. The reference order advances the shared audio/video rotary clock, so " - "it is part of the layout rather than a detail of the presentation." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - *_layout_inputs(), - InputParam( - name="prepared_references", - type_hint=list[MiniMaxH3PreparedReference], - required=True, - description="The prepared references, in packed order, with their latent geometry filled in.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return _layout_outputs() - - @torch.no_grad() - def __call__(self, components: MiniMaxH3Ref2VAModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - layout = build_ref2va_packed_sequence( - block_state.text_token_tags, - block_state.prepared_references, - block_state.num_latent_frames, - block_state.latent_height, - block_state.latent_width, - block_state.num_audio_latents, - components.patch_size, - ) - _set_layout_state(block_state, layout, components._execution_device) - - self.set_block_state(state, block_state) - return components, state - - -class MiniMaxH3PrepareLatentsStep(ModularPipelineBlocks): - model_name = "minimax-h3" - - @property - def description(self) -> str: - return ( - "Draws the initial noise of the generated rows and prepends the conditioning rows. MiniMax-H3 draws the " - "video noise as a latent tensor and patchifies it afterwards, then the audio noise directly in row " - "layout — both off the request's generator, after the conditioning noise of the encoder step." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="num_latent_frames", type_hint=int, required=True, description="Number of video latent frames." - ), - InputParam(name="latent_height", type_hint=int, required=True, description="Height of the video latents."), - InputParam(name="latent_width", type_hint=int, required=True, description="Width of the video latents."), - InputParam( - name="num_audio_latents", - type_hint=int, - required=True, - description="Number of audio latents per channel.", - ), - InputParam.template( - "generator", - description=( - "The generator of the request. The video noise is drawn from it first, then the audio noise." - ), - ), - InputParam( - name="latents", - type_hint=torch.Tensor, - description=( - "Pre-generated video noise of shape `(1, 24, num_latent_frames, latent_height, latent_width)`, " - "used instead of the draw." - ), - ), - InputParam( - name="audio_latents", - type_hint=torch.Tensor, - description="Pre-generated audio noise of shape `(2, 32, num_audio_latents)`.", - ), - InputParam( - name="condition_latents", - type_hint=torch.Tensor, - description="The video conditioning rows to prepend, or None for a request that has none.", - ), - InputParam( - name="audio_condition_latents", - type_hint=torch.Tensor, - description="The audio conditioning rows to prepend, or None for a request that has none.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "latents", - type_hint=torch.Tensor, - description="The video rows of the packed sequence, conditioning rows first.", - ), - OutputParam( - "audio_latents", - type_hint=torch.Tensor, - description="The channel-major audio rows of the packed sequence, reference rows first.", - ), - ] - - @staticmethod - def prepare_latents( - components, - num_latent_frames: int, - latent_height: int, - latent_width: int, - num_audio_latents: int, - device: torch.device, - generator: torch.Generator | list[torch.Generator] | None = None, - latents: torch.Tensor | None = None, - audio_latents: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - r""" - Draw the initial noise of both modalities and pack it into transformer rows. - - A request draws every stream from the one generator it is given, and the order is part of what that generator - reproduces: the conditioning noise of the keyframes or references first (one draw per condition, in - [`~modular_pipelines.minimax_h3.packing.keyframe_condition_noise`]), then the video noise here, as a latent tensor - that is patchified afterwards, then the audio noise, directly in row layout. Passing `latents` or - `audio_latents` skips its draw and shifts the ones after it. - - Args: - num_latent_frames (`int`): Number of video latent frames. - latent_height (`int`): Latent height. - latent_width (`int`): Latent width. - num_audio_latents (`int`): Number of audio latents per channel. - device (`torch.device`): The device the rows are drawn on. - generator (`torch.Generator`, *optional*): The generator of the request. - latents (`torch.Tensor`, *optional*): - Pre-generated video noise of shape `(1, latent_channels, num_latent_frames, latent_height, - latent_width)`, used instead of the draw. - audio_latents (`torch.Tensor`, *optional*): - Pre-generated audio noise of shape `(2, audio_latent_channels, num_audio_latents)`. - - Returns: - `tuple[torch.Tensor, torch.Tensor]`: the video rows and the channel-major audio rows. - """ - if latents is None: - latents = randn_tensor( - (1, components.vae_latent_channels, num_latent_frames, latent_height, latent_width), - generator=generator, - device=device, - dtype=torch.float32, - ) - video_rows = patchify_video_latents(latents.to(torch.float32), components.patch_size) - - if audio_latents is None: - audio_rows = randn_tensor( - (num_audio_latents * MINIMAX_H3_AUDIO_CHANNELS, components.audio_latent_channels), - generator=generator, - device=device, - dtype=torch.float32, - ) - else: - audio_rows = audio_latents.to(torch.float32).permute(0, 2, 1).reshape(-1, components.audio_latent_channels) - return video_rows.to(device), audio_rows.to(device) - - @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - - latents, audio_latents = self.prepare_latents( - components, - block_state.num_latent_frames, - block_state.latent_height, - block_state.latent_width, - block_state.num_audio_latents, - components._execution_device, - block_state.generator, - block_state.latents, - block_state.audio_latents, - ) - if block_state.condition_latents is not None: - latents = torch.cat([block_state.condition_latents, latents]) - if block_state.audio_condition_latents is not None: - audio_latents = torch.cat([block_state.audio_condition_latents, audio_latents]) - block_state.latents, block_state.audio_latents = latents, audio_latents - - self.set_block_state(state, block_state) - return components, state - - -class MiniMaxH3SetTimestepsStep(ModularPipelineBlocks): - model_name = "minimax-h3" - - @property - def description(self) -> str: - return ( - "Initializes the two schedules — `shift = 12.0` for video, `shift = 3.0` for audio — and stages the " - "row-to-timestep plan of every step. One forward serves every modality and every noise level at once: " - "the generated rows step down their own schedule while the conditioning rows stay pinned at their " - "noise-augmentation level, and that assignment is static per step." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", MiniMaxH3Scheduler), - ComponentSpec("audio_scheduler", MiniMaxH3Scheduler), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("num_inference_steps", required=True), - InputParam( - name="layout", - type_hint=MiniMaxH3PackedSequence, - required=True, - description="The structural description of the packed sequence.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("timesteps", type_hint=torch.Tensor, description="Timesteps of the video schedule."), - OutputParam("audio_timesteps", type_hint=torch.Tensor, description="Timesteps of the audio schedule."), - OutputParam( - "row_timestep_plan", - type_hint=list, - description=( - "One `(timestep, timestep_indices)` pair per step: the distinct timesteps of the sequence and the " - "index of every row into them." - ), - ), - ] - - @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - components.scheduler.set_timesteps(block_state.num_inference_steps, device=device) - components.audio_scheduler.set_timesteps(block_state.num_inference_steps, device=device) - block_state.timesteps = components.scheduler.timesteps - block_state.audio_timesteps = components.audio_scheduler.timesteps - - block_state.row_timestep_plan = [ - tuple( - tensor.to(device) - for tensor in build_row_timesteps( - block_state.layout, - float(timestep), - float(audio_timestep), - max(float(timestep), MINIMAX_H3_KEYFRAME_NOISE_AUG), - 1.0, - ) - ) - for timestep, audio_timestep in zip(block_state.timesteps, block_state.audio_timesteps) - ] - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/minimax_h3/before_encoder.py b/diffusers/modular_pipelines/minimax_h3/before_encoder.py deleted file mode 100644 index 88978bb6ff773faf0daa48faa5883eb5dd30baaf..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/minimax_h3/before_encoder.py +++ /dev/null @@ -1,408 +0,0 @@ -# Copyright 2026 The MiniMax and HuggingFace Teams. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import PIL -import torch -from PIL import Image, ImageOps - -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import InputParam, OutputParam -from .modular_pipeline import MiniMaxH3ModularPipeline, MiniMaxH3Ref2VAModularPipeline -from .packing import ( - MINIMAX_H3_CANVAS_MULTIPLE, - MINIMAX_H3_FPS, - MINIMAX_H3_MAX_DURATION, - MINIMAX_H3_MIN_DURATION, - align_num_frames, - audio_latent_num_frames, - prepare_keyframe_image, - resolve_canvas_size, - video_latent_num_frames, -) -from .packing_ref2va import ( - MINIMAX_H3_MAX_REFERENCE_AUDIOS, - MINIMAX_H3_MAX_REFERENCE_IMAGES, - MINIMAX_H3_MAX_REFERENCE_VIDEOS, - MINIMAX_H3_MAX_REFERENCES, - MiniMaxH3PreparedReference, - MiniMaxH3Reference, - prepare_reference_frames, - prepare_reference_image, - prepare_reference_waveform, - reference_kind, - reference_media_to_uint8, - resample_reference_frames, - resolve_reference_image_size, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _latent_geometry(components, height: int, width: int, num_frames: int) -> tuple[int, int, int, int]: - r"""The latent geometry the packed layout, the noise draws and the decoders all key off.""" - ratio = components.vae_spatial_compression_ratio - return video_latent_num_frames(num_frames), height // ratio, width // ratio, audio_latent_num_frames(num_frames) - - -def _latent_geometry_outputs() -> list[OutputParam]: - r"""The declaration of what [`_latent_geometry`] resolves, shared by the two setup blocks.""" - return [ - OutputParam("num_latent_frames", type_hint=int, description="Number of generated video latent frames."), - OutputParam("latent_height", type_hint=int, description="Height of the generated video latents."), - OutputParam("latent_width", type_hint=int, description="Width of the generated video latents."), - OutputParam("num_audio_latents", type_hint=int, description="Number of generated audio latents per channel."), - ] - - -class MiniMaxH3SetupStep(ModularPipelineBlocks): - model_name = "minimax-h3" - - @property - def description(self) -> str: - return ( - "Resolves the plan shared by the `t2va` and `fl2va` tasks: the canvas (MiniMax-H3's own 768-short-edge " - "geometry for the aspect ratio of the first keyframe, or 16:9 without keyframes), the `17 * n + 5` frame " - "count the video VAE can decode, the latent geometry every later block keys off, and the keyframes put " - "onto that canvas." - ) - - @staticmethod - def _check_inputs(block_state) -> None: - if (block_state.height is None) != (block_state.width is None): - raise ValueError("`height` and `width` have to be passed together, or neither of them.") - if block_state.height is not None and ( - block_state.height % MINIMAX_H3_CANVAS_MULTIPLE or block_state.width % MINIMAX_H3_CANVAS_MULTIPLE - ): - raise ValueError( - f"`height` and `width` must be multiples of {MINIMAX_H3_CANVAS_MULTIPLE}, got " - f"{block_state.height}x{block_state.width}." - ) - # The duration the request generates is the one of the *aligned* frame count, so that is what the ceiling has - # to hold for: 346 frames would otherwise pass the check and then be rounded up to 362, i.e. 15.083 seconds. - aligned_num_frames = align_num_frames(block_state.num_frames) - duration = aligned_num_frames / MINIMAX_H3_FPS - if not MINIMAX_H3_MIN_DURATION <= duration <= MINIMAX_H3_MAX_DURATION: - raise ValueError( - f"MiniMax-H3 generates between {MINIMAX_H3_MIN_DURATION} and {MINIMAX_H3_MAX_DURATION} seconds at " - f"{MINIMAX_H3_FPS} fps, so `num_frames`, rounded up to the next `17 * n + 5` the video VAE can " - f"encode, must be between {int(MINIMAX_H3_MIN_DURATION * MINIMAX_H3_FPS)} and " - f"{int(MINIMAX_H3_MAX_DURATION * MINIMAX_H3_FPS)}, got {block_state.num_frames} (rounded up to " - f"{aligned_num_frames})." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="image", - type_hint=PIL.Image.Image, - description=( - "Keyframe the video starts from. It is *stretched* onto the target canvas, which by default is " - "derived from its own aspect ratio." - ), - ), - InputParam( - name="last_image", - type_hint=PIL.Image.Image, - description=( - "Keyframe the video ends on. Can be passed on its own to generate *up to* a frame. Combined with " - "`image` it is the follower of the two and is cover-cropped onto the canvas." - ), - ), - InputParam.template("height", description="Height of the generated video in pixels, a multiple of 32."), - InputParam.template("width", description="Width of the generated video in pixels, a multiple of 32."), - InputParam( - name="num_frames", - type_hint=int, - default=124, - description=( - "Number of frames to generate, at the fixed 24 fps. Snapped up to the next `17 * n + 5` the video " - "VAE can decode; the resulting duration must stay between 5 and 15 seconds." - ), - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("height", type_hint=int, description="Resolved height of the generated video in pixels."), - OutputParam("width", type_hint=int, description="Resolved width of the generated video in pixels."), - OutputParam("num_frames", type_hint=int, description="Resolved number of frames, of the form 17 * n + 5."), - *_latent_geometry_outputs(), - OutputParam( - "keyframes", - type_hint=list, - description="The keyframes put onto the target canvas, in packed order (empty for `t2va`).", - ), - OutputParam( - "keyframe_anchors", - type_hint=tuple, - description="Which end of the video every keyframe is anchored to, in packed order.", - ), - ] - - @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - self._check_inputs(block_state) - - keyframes = [ - ImageOps.exif_transpose(keyframe).convert("RGB") - for keyframe in (block_state.image, block_state.last_image) - if keyframe is not None - ] - block_state.keyframe_anchors = tuple( - anchor - for anchor, keyframe in (("first", block_state.image), ("last", block_state.last_image)) - if keyframe is not None - ) - if block_state.height is None: - block_state.height, block_state.width = resolve_canvas_size(*(keyframes[0].size if keyframes else (16, 9))) - - aligned_num_frames = align_num_frames(block_state.num_frames) - if aligned_num_frames != block_state.num_frames: - logger.warning( - f"`num_frames` has to be of the form 17 * n + 5 for the video VAE; rounding {block_state.num_frames} " - f"up to {aligned_num_frames}." - ) - block_state.num_frames = aligned_num_frames - - ( - block_state.num_latent_frames, - block_state.latent_height, - block_state.latent_width, - block_state.num_audio_latents, - ) = _latent_geometry(components, block_state.height, block_state.width, block_state.num_frames) - - block_state.keyframes = [ - prepare_keyframe_image(keyframe, block_state.height, block_state.width, stretch=index == 0) - for index, keyframe in enumerate(keyframes) - ] - self.set_block_state(state, block_state) - return components, state - - -class MiniMaxH3Ref2VASetupStep(ModularPipelineBlocks): - model_name = "minimax-h3-ref2va" - - @property - def description(self) -> str: - return ( - "Resolves the `ref2va` plan: the canvas (MiniMax-H3's own 16:9 unless asked otherwise — references never " - "bind the generated geometry), the references prepared at their own resolutions, the frame count they " - "imply when it was left open, and the latent geometry every later block keys off." - ) - - @staticmethod - def _check_inputs(components, block_state) -> None: - if (block_state.height is None) != (block_state.width is None): - raise ValueError("`height` and `width` have to be passed together, or neither of them.") - if block_state.height is not None and ( - block_state.height % MINIMAX_H3_CANVAS_MULTIPLE or block_state.width % MINIMAX_H3_CANVAS_MULTIPLE - ): - raise ValueError( - f"`height` and `width` must be multiples of {MINIMAX_H3_CANVAS_MULTIPLE}, got " - f"{block_state.height}x{block_state.width}." - ) - # The duration the request generates is the one of the *aligned* frame count, so that is what the ceiling has - # to hold for: 346 frames would otherwise pass the check and then be rounded up to 362, i.e. 15.083 seconds. - aligned_num_frames = None if block_state.num_frames is None else align_num_frames(block_state.num_frames) - duration = None if aligned_num_frames is None else aligned_num_frames / MINIMAX_H3_FPS - if duration is not None and not MINIMAX_H3_MIN_DURATION <= duration <= MINIMAX_H3_MAX_DURATION: - raise ValueError( - f"MiniMax-H3 generates between {MINIMAX_H3_MIN_DURATION} and {MINIMAX_H3_MAX_DURATION} seconds at " - f"{MINIMAX_H3_FPS} fps, so `num_frames`, rounded up to the next `17 * n + 5` the video VAE can " - f"encode, must be between {int(MINIMAX_H3_MIN_DURATION * MINIMAX_H3_FPS)} and " - f"{int(MINIMAX_H3_MAX_DURATION * MINIMAX_H3_FPS)}, got {block_state.num_frames} (rounded up to " - f"{aligned_num_frames})." - ) - - if not block_state.references: - raise ValueError( - "`ref2va` needs at least one reference; use `MiniMaxH3ModularPipeline` for text-only requests." - ) - kinds = [reference_kind(index, entry) for index, entry in enumerate(block_state.references)] - for kind, limit in ( - ("image", MINIMAX_H3_MAX_REFERENCE_IMAGES), - ("video", MINIMAX_H3_MAX_REFERENCE_VIDEOS), - ("audio", MINIMAX_H3_MAX_REFERENCE_AUDIOS), - ): - if kinds.count(kind) > limit: - raise ValueError(f"MiniMax-H3 accepts at most {limit} {kind} references, got {kinds.count(kind)}.") - if len(kinds) > MINIMAX_H3_MAX_REFERENCES: - raise ValueError( - f"MiniMax-H3 accepts at most {MINIMAX_H3_MAX_REFERENCES} references in total, got {len(kinds)}." - ) - if set(kinds) == {"audio"}: - raise ValueError( - "An audio reference has to be paired with at least one image or video reference and cannot be used " - "on its own." - ) - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="references", - type_hint=list[MiniMaxH3Reference], - required=True, - description=( - "The references to condition on, **in the order the model should read them**: the order labels " - "them in the prompt presentation and lays them out on the shared rotary clock, so a different " - "order is a different request. Every [`MiniMaxH3Reference`] carries exactly one medium, a path or " - "in-memory media — `image` (at most 9), `video` at its own `fps` (at most 3, whose `audio` " - "soundtrack is conditioned on as well), or `audio` at its own `sample_rate` (at most 3) — for at " - "most 12 references in total, and audio references cannot be the only ones. A path is decoded " - "when the reference is built, so these blocks only ever see pixels and samples." - ), - ), - InputParam.template("height", description="Height of the generated video in pixels, a multiple of 32."), - InputParam.template("width", description="Width of the generated video in pixels, a multiple of 32."), - InputParam( - name="num_frames", - type_hint=int, - description=( - "Number of frames to generate, at the fixed 24 fps. Snapped up to the next `17 * n + 5` the video " - "VAE can decode. May be left out, but only when exactly one reference carries audio, in which " - "case the duration is that soundtrack's." - ), - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam("height", type_hint=int, description="Resolved height of the generated video in pixels."), - OutputParam("width", type_hint=int, description="Resolved width of the generated video in pixels."), - OutputParam("num_frames", type_hint=int, description="Resolved number of frames, of the form 17 * n + 5."), - *_latent_geometry_outputs(), - OutputParam( - "prepared_references", - type_hint=list[MiniMaxH3PreparedReference], - description="The references prepared at their own resolutions, in packed order.", - ), - ] - - @staticmethod - def prepare_references( - components, references: list[MiniMaxH3Reference], num_frames: int | None - ) -> tuple[list[MiniMaxH3PreparedReference], int]: - r""" - Resolve the references and, if it was left open, the duration they imply. - - Every reference is prepared at its own resolution: an image is resized to a 2048 pixel short edge, a video is - resampled onto MiniMax-H3's own 24 fps, rescaled onto the 768 pixel canvas of *its own* aspect ratio and - truncated to the generated frame count, and a soundtrack is put on the audio VAE's sample rate and truncated to - the generated duration. None of this touches the target canvas. - - A reference that left its `fps` or its `sample_rate` out is taken to already be at MiniMax-H3's own rate, and - its frames or its samples then flow through untouched. - - A video reference goes through the two passes the reference implementation's `ffmpeg` decode applied, in the - same order: the constant frame rate resample of `resample_reference_frames` and the LANCZOS rescale of - `prepare_reference_frames`. Frames handed over at 24 fps and already at the canvas their own aspect ratio - resolves to therefore reach the VAE untouched, which is the parity-exact route. - - Args: - references (`list[MiniMaxH3Reference]`): - The `references` input of a [`MiniMaxH3Ref2VABlocks`] request. - num_frames (`int`, *optional*): - The requested frame count, or `None` to derive it from the single audio-bearing reference. - - Returns: - `tuple[list[MiniMaxH3PreparedReference], int]`: the prepared references, in packed order, and the frame - count. - """ - resolved = [ - MiniMaxH3PreparedReference(kind=reference_kind(index, entry), has_audio=entry.has_audio) - for index, entry in enumerate(references) - ] - - # The duration may be left open, but then exactly one reference may carry audio, or the request is ambiguous. - if num_frames is None: - audio_bearing = [index for index, reference in enumerate(resolved) if reference.has_audio] - if len(audio_bearing) != 1: - raise ValueError( - "`num_frames` may only be left to the references when exactly one of them carries audio, got " - f"{len(audio_bearing)}." - ) - index = audio_bearing[0] - sample_rate = references[index].sample_rate or components.audio_sampling_rate - duration = references[index].audio.shape[-1] / sample_rate - if not MINIMAX_H3_MIN_DURATION <= duration <= MINIMAX_H3_MAX_DURATION: - raise ValueError( - f"`references[{index}]` is {duration:g} seconds long, outside the " - f"{MINIMAX_H3_MIN_DURATION} to {MINIMAX_H3_MAX_DURATION} seconds MiniMax-H3 generates." - ) - num_frames = align_num_frames(round(duration * MINIMAX_H3_FPS)) - # The duration the request generates is the one of the *aligned* frame count, so that is what the - # ceiling has to hold for: a 14.99 second soundtrack rounds up to 362 frames, i.e. 15.083 seconds. - if num_frames / MINIMAX_H3_FPS > MINIMAX_H3_MAX_DURATION: - raise ValueError( - f"`references[{index}]` is {duration:g} seconds long, which rounds up to {num_frames} frames " - f"(`17 * n + 5`), i.e. {num_frames / MINIMAX_H3_FPS:g} seconds — past the " - f"{MINIMAX_H3_MAX_DURATION} seconds MiniMax-H3 generates. Pass `num_frames` to generate a " - "shorter video from this soundtrack." - ) - num_frames = align_num_frames(num_frames) - - for reference, entry in zip(resolved, references): - if reference.kind == "image": - image = entry.image - if not isinstance(image, Image.Image): - image = Image.fromarray(reference_media_to_uint8(image)) - image = ImageOps.exif_transpose(image).convert("RGB") - height, width = resolve_reference_image_size(*image.size) - reference.image = prepare_reference_image(image, height, width) - elif reference.kind == "video": - frames = resample_reference_frames(reference_media_to_uint8(entry.video), float(entry.fps)) - reference.frames = prepare_reference_frames(frames, num_frames) - if reference.has_audio: - reference.waveform = prepare_reference_waveform( - entry.audio, - entry.sample_rate or components.audio_sampling_rate, - components.audio_sampling_rate, - max_duration=num_frames / MINIMAX_H3_FPS, - ) - return resolved, num_frames - - @torch.no_grad() - def __call__(self, components: MiniMaxH3Ref2VAModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - self._check_inputs(components, block_state) - - if block_state.height is None: - block_state.height, block_state.width = resolve_canvas_size(16, 9) - - requested_num_frames = block_state.num_frames - block_state.prepared_references, block_state.num_frames = self.prepare_references( - components, block_state.references, block_state.num_frames - ) - if requested_num_frames is not None and requested_num_frames != block_state.num_frames: - logger.warning( - f"`num_frames` has to be of the form 17 * n + 5 for the video VAE; rounding {requested_num_frames} up " - f"to {block_state.num_frames}." - ) - - ( - block_state.num_latent_frames, - block_state.latent_height, - block_state.latent_width, - block_state.num_audio_latents, - ) = _latent_geometry(components, block_state.height, block_state.width, block_state.num_frames) - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/minimax_h3/decoders.py b/diffusers/modular_pipelines/minimax_h3/decoders.py deleted file mode 100644 index fc4ee359267b790a8a52c9301e055c7d5369af60..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/minimax_h3/decoders.py +++ /dev/null @@ -1,198 +0,0 @@ -# Copyright 2026 The MiniMax and HuggingFace Teams. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ...configuration_utils import FrozenDict -from ...models import AutoencoderKLMiniMaxH3, AutoencoderKLMiniMaxH3Audio -from ...utils import logging -from ...video_processor import VideoProcessor -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import MiniMaxH3ModularPipeline -from .packing import ( - MINIMAX_H3_PIXEL_MEAN, - MINIMAX_H3_PIXEL_STD, - unpack_audio_tokens, - unpatchify_video_tokens, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -class MiniMaxH3VideoDecodeStep(ModularPipelineBlocks): - model_name = "minimax-h3" - - @property - def description(self) -> str: - return ( - "Unpacks the generated video rows back into latents, denormalizes them and decodes them into video. The " - "spatial tiling of the video VAE covers the canvas exactly, so the decoded frames need no crop back, but " - "the decode itself runs under float16 autocast even though the VAE weights are float32, and the VAE " - "produces ImageNet-normalized RGB that is reverted here." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLMiniMaxH3), - ComponentSpec( - "video_processor", - VideoProcessor, - # The video VAE decodes into ImageNet-normalized RGB over a [0, 1] base range, which this block - # reverts itself, so the processor must not denormalize a second time. - config=FrozenDict({"vae_scale_factor": 16, "do_normalize": False}), - default_creation_method="from_config", - ), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="latents", - type_hint=torch.Tensor, - required=True, - description="The denoised video rows of the packed sequence, conditioning rows first.", - ), - InputParam( - name="num_condition_video_rows", - type_hint=int, - default=0, - description="How many leading video rows are conditioning rows and are dropped here.", - ), - InputParam( - name="num_latent_frames", type_hint=int, required=True, description="Number of video latent frames." - ), - InputParam(name="latent_height", type_hint=int, required=True, description="Height of the video latents."), - InputParam(name="latent_width", type_hint=int, required=True, description="Width of the video latents."), - InputParam.template( - "output_type", description="Output format: 'pil', 'np', 'pt' or 'latent' for the raw latents." - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [OutputParam.template("videos", description="The generated video.")] - - @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - latents = unpatchify_video_tokens( - block_state.latents[block_state.num_condition_video_rows :], - block_state.num_latent_frames, - block_state.latent_height, - block_state.latent_width, - components.vae_latent_channels, - components.patch_size, - ) - latents_mean = torch.tensor(components.vae.config.latents_mean, device=device).view(1, -1, 1, 1, 1) - latents_std = torch.tensor(components.vae.config.latents_std, device=device).view(1, -1, 1, 1, 1) - latents = latents * latents_std + latents_mean - - if block_state.output_type == "latent": - block_state.videos = latents - else: - with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"): - video = components.vae.decode(latents, return_dict=False)[0] - pixel_mean = torch.tensor(MINIMAX_H3_PIXEL_MEAN, device=device).view(1, -1, 1, 1, 1) - pixel_std = torch.tensor(MINIMAX_H3_PIXEL_STD, device=device).view(1, -1, 1, 1, 1) - video = (video.float() * pixel_std + pixel_mean).clamp(0, 1) - block_state.videos = components.video_processor.postprocess_video( - video, output_type=block_state.output_type - ) - - self.set_block_state(state, block_state) - return components, state - - -class MiniMaxH3AudioDecodeStep(ModularPipelineBlocks): - model_name = "minimax-h3" - - @property - def description(self) -> str: - return ( - "Unpacks the generated audio rows back into latents, denormalizes them and decodes them into a stereo " - "waveform. The audio VAE is mono and takes the two stereo channels as two batch items." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("audio_vae", AutoencoderKLMiniMaxH3Audio)] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="audio_latents", - type_hint=torch.Tensor, - required=True, - description="The denoised audio rows of the packed sequence, reference rows first.", - ), - InputParam( - name="num_condition_audio_rows", - type_hint=int, - default=0, - description="How many leading audio rows are reference rows and are dropped here.", - ), - InputParam( - name="num_audio_latents", - type_hint=int, - required=True, - description="Number of audio latents per channel.", - ), - InputParam.template( - "output_type", description="Output format: 'pil', 'np', 'pt' or 'latent' for the raw latents." - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "audio", - type_hint=torch.Tensor, - description="The generated soundtrack, of shape `(1, 2, num_samples)`.", - ), - OutputParam( - "sampling_rate", - type_hint=int, - description="Sample rate of the generated soundtrack in Hz.", - ), - ] - - @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - audio_latents = unpack_audio_tokens( - block_state.audio_latents[block_state.num_condition_audio_rows :], block_state.num_audio_latents - ) - audio_latents_mean = torch.tensor(components.audio_vae.config.latents_mean, device=device).view(1, -1, 1) - audio_latents_std = torch.tensor(components.audio_vae.config.latents_std, device=device).view(1, -1, 1) - audio_latents = audio_latents * audio_latents_std + audio_latents_mean - - if block_state.output_type == "latent": - block_state.audio = audio_latents - else: - audio = components.audio_vae.decode(audio_latents, return_dict=False)[0] - block_state.audio = audio.float().permute(1, 0, 2) - block_state.sampling_rate = components.audio_sampling_rate - - self.set_block_state(state, block_state) - return components, state diff --git a/diffusers/modular_pipelines/minimax_h3/denoise.py b/diffusers/modular_pipelines/minimax_h3/denoise.py deleted file mode 100644 index d149273371bb33b72f20433f6ff9dc4885686652..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/minimax_h3/denoise.py +++ /dev/null @@ -1,325 +0,0 @@ -# Copyright 2026 The MiniMax and HuggingFace Teams. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch - -from ...models import MiniMaxH3Transformer3DModel -from ...schedulers import MiniMaxH3Scheduler -from ...utils import logging -from ..modular_pipeline import ( - BlockState, - LoopSequentialPipelineBlocks, - ModularPipelineBlocks, - PipelineState, -) -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import MiniMaxH3ModularPipeline, MiniMaxH3Ref2VAModularPipeline - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _denoiser_inputs() -> list[InputParam]: - r"""Everything one MiniMax-H3 forward reads, beyond the transformer itself.""" - return [ - InputParam( - name="latents", - type_hint=torch.Tensor, - required=True, - description="The video rows of the packed sequence, conditioning rows first.", - ), - InputParam( - name="audio_latents", - type_hint=torch.Tensor, - required=True, - description="The channel-major audio rows of the packed sequence, reference rows first.", - ), - InputParam.template("prompt_embeds"), - InputParam( - name="row_timestep_plan", - type_hint=list, - required=True, - description="One `(timestep, timestep_indices)` pair per step.", - ), - InputParam( - name="token_tags", type_hint=torch.Tensor, required=True, description="The modality tag of every row." - ), - InputParam( - name="position_ids", - type_hint=torch.Tensor, - required=True, - description="The `(t, h, w)` rotary coordinate of every row.", - ), - InputParam( - name="video_indices", - type_hint=torch.Tensor, - required=True, - description="Sequence positions of the video rows.", - ), - InputParam( - name="audio_indices", - type_hint=torch.Tensor, - required=True, - description="Sequence positions of the audio rows.", - ), - InputParam( - name="text_indices", - type_hint=torch.Tensor, - required=True, - description="Sequence positions of the text rows.", - ), - InputParam.template("attention_kwargs"), - ] - - -def _denoiser_outputs() -> list[OutputParam]: - return [ - OutputParam( - "noise_pred", type_hint=torch.Tensor, description="Predicted velocity of the video rows of the sequence." - ), - OutputParam( - "audio_noise_pred", - type_hint=torch.Tensor, - description="Predicted velocity of the audio rows of the sequence.", - ), - ] - - -def _predict_velocity(transformer: MiniMaxH3Transformer3DModel, block_state: BlockState, i: int): - r"""One MiniMax-H3 forward pass: every row of the packed sequence, at its own noise level, at once.""" - unique_timesteps, timestep_indices = block_state.row_timestep_plan[i] - return transformer( - hidden_states=block_state.latents[None], - audio_hidden_states=block_state.audio_latents[None], - encoder_hidden_states=block_state.prompt_embeds, - timestep=unique_timesteps, - timestep_indices=timestep_indices, - token_tags=block_state.token_tags, - position_ids=block_state.position_ids, - video_indices=block_state.video_indices, - audio_indices=block_state.audio_indices, - text_indices=block_state.text_indices, - attention_kwargs=block_state.attention_kwargs, - return_dict=False, - ) - - -class MiniMaxH3LoopDenoiser(ModularPipelineBlocks): - model_name = "minimax-h3" - - @property - def description(self) -> str: - return ( - "Runs the one MiniMax-H3 forward pass of a denoising iteration, which predicts the velocity of every row " - "of the packed sequence at once. The checkpoint is guidance-distilled, so there is no unconditional pass " - "and no guider." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer", MiniMaxH3Transformer3DModel)] - - @property - def inputs(self) -> list[InputParam]: - return _denoiser_inputs() - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return _denoiser_outputs() - - @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - block_state.noise_pred, block_state.audio_noise_pred = _predict_velocity( - components.transformer, block_state, i - ) - return components, block_state - - -class MiniMaxH3Ref2VALoopDenoiser(ModularPipelineBlocks): - model_name = "minimax-h3-ref2va" - - @property - def description(self) -> str: - return ( - "Runs the one MiniMax-H3 forward pass of a `ref2va` denoising iteration, against the `transformer_ref` " - "partition of the checkpoint." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ComponentSpec("transformer_ref", MiniMaxH3Transformer3DModel)] - - @property - def inputs(self) -> list[InputParam]: - return _denoiser_inputs() - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return _denoiser_outputs() - - @torch.no_grad() - def __call__(self, components: MiniMaxH3Ref2VAModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - block_state.noise_pred, block_state.audio_noise_pred = _predict_velocity( - components.transformer_ref, block_state, i - ) - return components, block_state - - -class MiniMaxH3LoopSchedulerStep(ModularPipelineBlocks): - model_name = "minimax-h3" - - @property - def description(self) -> str: - return ( - "Steps the generated video and audio rows down their own schedule. The conditioning rows are re-imposed " - "by construction: only the generated rows are ever written, so the anchors survive the whole loop." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", MiniMaxH3Scheduler), - ComponentSpec("audio_scheduler", MiniMaxH3Scheduler), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="latents", - type_hint=torch.Tensor, - required=True, - description="The video rows of the packed sequence, conditioning rows first.", - ), - InputParam( - name="audio_latents", - type_hint=torch.Tensor, - required=True, - description="The channel-major audio rows of the packed sequence, reference rows first.", - ), - InputParam( - name="noise_pred", - type_hint=torch.Tensor, - required=True, - description="Predicted velocity of the video rows.", - ), - InputParam( - name="audio_noise_pred", - type_hint=torch.Tensor, - required=True, - description="Predicted velocity of the audio rows.", - ), - InputParam( - name="audio_timesteps", - type_hint=torch.Tensor, - required=True, - description="Timesteps of the audio schedule.", - ), - InputParam( - name="num_condition_video_rows", - type_hint=int, - default=0, - description="How many leading video rows are conditioning rows.", - ), - InputParam( - name="num_condition_audio_rows", - type_hint=int, - default=0, - description="How many leading audio rows are reference rows.", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "latents", - type_hint=torch.Tensor, - description="The video rows of the packed sequence after one step.", - ), - OutputParam( - "audio_latents", - type_hint=torch.Tensor, - description="The audio rows of the packed sequence after one step.", - ), - ] - - @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): - num_condition_video_rows = block_state.num_condition_video_rows - num_condition_audio_rows = block_state.num_condition_audio_rows - - block_state.latents[num_condition_video_rows:] = components.scheduler.step( - block_state.noise_pred[0, num_condition_video_rows:].float(), - t, - block_state.latents[num_condition_video_rows:], - return_dict=False, - )[0] - block_state.audio_latents[num_condition_audio_rows:] = components.audio_scheduler.step( - block_state.audio_noise_pred[0, num_condition_audio_rows:].float(), - block_state.audio_timesteps[i], - block_state.audio_latents[num_condition_audio_rows:], - return_dict=False, - )[0] - return components, block_state - - -class MiniMaxH3DenoiseLoopWrapper(LoopSequentialPipelineBlocks): - model_name = "minimax-h3" - - @property - def description(self) -> str: - return "Iteratively denoises the packed MiniMax-H3 sequence over the two schedules." - - @property - def loop_expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("scheduler", MiniMaxH3Scheduler), - ComponentSpec("audio_scheduler", MiniMaxH3Scheduler), - ] - - @property - def loop_inputs(self) -> list[InputParam]: - return [ - InputParam.template("timesteps", required=True, description="Timesteps of the video schedule."), - ] - - @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - with self.progress_bar(total=len(block_state.timesteps)) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) - progress_bar.update() - self.set_block_state(state, block_state) - return components, state - - -class MiniMaxH3DenoiseStep(MiniMaxH3DenoiseLoopWrapper): - block_classes = [MiniMaxH3LoopDenoiser, MiniMaxH3LoopSchedulerStep] - block_names = ["denoiser", "update"] - - @property - def description(self) -> str: - return "Runs the `t2va` / `fl2va` MiniMax-H3 denoising loop, one forward pass per step." - - -class MiniMaxH3Ref2VADenoiseStep(MiniMaxH3DenoiseLoopWrapper): - model_name = "minimax-h3-ref2va" - block_classes = [MiniMaxH3Ref2VALoopDenoiser, MiniMaxH3LoopSchedulerStep] - block_names = ["denoiser", "update"] - - @property - def description(self) -> str: - return "Runs the `ref2va` MiniMax-H3 denoising loop, one forward pass per step." diff --git a/diffusers/modular_pipelines/minimax_h3/encoders.py b/diffusers/modular_pipelines/minimax_h3/encoders.py deleted file mode 100644 index da2d611eaba0963cfc70e379735f0ea7824af370..0000000000000000000000000000000000000000 --- a/diffusers/modular_pipelines/minimax_h3/encoders.py +++ /dev/null @@ -1,638 +0,0 @@ -# Copyright 2026 The MiniMax and HuggingFace Teams. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import numpy as np -import torch -from transformers import Qwen2TokenizerFast, Qwen3VLForConditionalGeneration, Qwen3VLProcessor - -from ...models import AutoencoderKLMiniMaxH3, AutoencoderKLMiniMaxH3Audio -from ...models.autoencoders.vae import DiagonalGaussianDistribution -from ...schedulers import MiniMaxH3Scheduler -from ...utils import logging -from ..modular_pipeline import ModularPipelineBlocks, PipelineState -from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import MiniMaxH3ModularPipeline, MiniMaxH3Ref2VAModularPipeline -from .packing import ( - MINIMAX_H3_KEYFRAME_ENCODE_SEED, - MINIMAX_H3_KEYFRAME_NOISE_AUG, - MINIMAX_H3_PIXEL_MEAN, - MINIMAX_H3_PIXEL_STD, - MINIMAX_H3_TEXT_ENCODER_LAYER, - MINIMAX_H3_TEXT_TAG, - MINIMAX_H3_VIDEO_TAG, - keyframe_condition_noise, - patchify_video_latents, -) -from .packing_ref2va import ( - MiniMaxH3PreparedReference, - build_ref2va_presentation, - sample_reference_video_frames, - trim_reference_num_frames, -) - - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - - -def _check_prompt(prompt) -> None: - r"""MiniMax-H3 packs one request into one sequence, so a batch of prompts is not a thing.""" - if not isinstance(prompt, str): - raise ValueError( - f"MiniMax-H3 packs one request into one sequence, so `prompt` must be a single string, got {type(prompt)}." - ) - - -def _conditioner_components() -> list[ComponentSpec]: - r"""MiniMax-H3's conditioner: a Qwen3-VL read at its 50th decoder layer, with its language-model head unused.""" - return [ - ComponentSpec("text_encoder", Qwen3VLForConditionalGeneration), - ComponentSpec("tokenizer", Qwen2TokenizerFast), - ComponentSpec("processor", Qwen3VLProcessor), - ] - - -def _conditioner_outputs() -> list[OutputParam]: - return [ - OutputParam.template( - "prompt_embeds", - description=( - "The hidden state MiniMax-H3 conditions on, of shape `(1, num_text_tokens, 5120)`, read after the " - "50th decoder layer of the Qwen3-VL conditioner." - ), - ), - OutputParam( - "text_token_tags", - type_hint=torch.Tensor, - description="The per-row modality tag of every row of `prompt_embeds`; a vision block is tagged as video.", - ), - ] - - -class MiniMaxH3TextEncoderStep(ModularPipelineBlocks): - model_name = "minimax-h3" - - @property - def description(self) -> str: - return ( - "Encodes MiniMax-H3's presentation of a `t2va` / `fl2va` request: the prompt verbatim, preceded by a " - '`": "` label and a vision block per keyframe, with no chat template and no special tokens. ' - "The checkpoint is guidance-distilled, so there is no negative prompt and no unconditional branch." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return _conditioner_components() - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam.template("prompt", description="The prompt to guide generation, a single string."), - InputParam( - name="keyframes", - type_hint=list, - description="The keyframes put onto the target canvas, in packed order (empty or None for `t2va`).", - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return _conditioner_outputs() - - @staticmethod - def encode_prompt( - components, - prompt: str, - images: list | None = None, - device: torch.device | None = None, - dtype: torch.dtype | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: - r""" - Build MiniMax-H3's presentation of a request and encode it. - - The presentation is the verbatim prompt for `t2va`. Every keyframe prepends a `": "` label and a - vision block (`<|vision_start|>`, one `<|image_pad|>` per vision patch, `<|vision_end|>`) — no chat template - and no special tokens. The rows of a vision block are tagged as *video* rather than text, which is what the - transformer's AdaLN modulation keys off. - - Args: - prompt (`str`): The prompt to encode. - images (`list[PIL.Image.Image]`, *optional*): - The keyframes, already prepared onto the target canvas, in packed order. - device (`torch.device`, *optional*): The device to run the conditioner on. - dtype (`torch.dtype`, *optional*): The dtype of the returned embeddings. - - Returns: - `tuple[torch.Tensor, torch.Tensor]`: the `(1, num_text_tokens, 5120)` hidden states and the - `(num_text_tokens,)` per-row modality tags. - """ - device = device or components._execution_device - dtype = dtype or components.transformer.dtype - - num_layers = components.text_encoder.config.text_config.num_hidden_layers - if num_layers <= MINIMAX_H3_TEXT_ENCODER_LAYER: - raise ValueError( - f"MiniMax-H3 conditions on `hidden_states[{MINIMAX_H3_TEXT_ENCODER_LAYER}]` of its Qwen3-VL " - f"conditioner, which needs more than {MINIMAX_H3_TEXT_ENCODER_LAYER} decoder layers, but " - f"`text_encoder` has {num_layers}. The last hidden state of a stack truncated to exactly " - f"{MINIMAX_H3_TEXT_ENCODER_LAYER} layers is post-norm and is not the conditioning MiniMax-H3 expects." - ) - - pixel_values, image_grid_thw = None, None - token_ids, token_tags = [], [] - if images: - vision = components.processor.image_processor(images=images, return_tensors="pt") - pixel_values, image_grid_thw = vision["pixel_values"], vision["image_grid_thw"] - merge_size = components.processor.image_processor.merge_size**2 - for index in range(len(images)): - num_image_tokens = int(image_grid_thw[index].prod()) // merge_size - label_ids = components.tokenizer(f": ", add_special_tokens=False)["input_ids"] - vision_ids = ( - [components.tokenizer.convert_tokens_to_ids("<|vision_start|>")] - + [components.tokenizer.convert_tokens_to_ids("<|image_pad|>")] * num_image_tokens - + [components.tokenizer.convert_tokens_to_ids("<|vision_end|>")] - ) - token_ids += label_ids + vision_ids - token_tags += [MINIMAX_H3_TEXT_TAG] * len(label_ids) + [MINIMAX_H3_VIDEO_TAG] * len(vision_ids) - prompt_ids = components.tokenizer(prompt, add_special_tokens=False)["input_ids"] - token_ids += prompt_ids - token_tags += [MINIMAX_H3_TEXT_TAG] * len(prompt_ids) - - input_ids = torch.tensor([token_ids], dtype=torch.long, device=device) - # Qwen3-VL lays its 3D rotary positions out per modality run, which it reads off the token type ids the - # processor derives from the vision pad ids (`0` text, `1` image, `2` video). - mm_token_type_ids = torch.tensor( - components.processor.create_mm_token_type_ids([token_ids]), dtype=torch.long, device=device - ) - # `text_encoder.model` is a submodule, and a CPU-offload hook — accelerate's or the one the - # `ComponentsManager` attaches — wraps the *top-level* module's `forward` alone, so calling the submodule - # directly would leave the conditioner on the CPU. Fire the hook by hand instead of routing through - # `text_encoder(...)`: MiniMax-H3 reads `hidden_states[50]` and never uses the language-model head, whose - # vocabulary-wide projection over every token is all the top-level forward would add. - hook = getattr(components.text_encoder, "_hf_hook", None) - if hook is not None and hasattr(hook, "pre_forward"): - hook.pre_forward(components.text_encoder) - outputs = components.text_encoder.model( - input_ids=input_ids, - attention_mask=torch.ones_like(input_ids), - mm_token_type_ids=mm_token_type_ids, - pixel_values=None if pixel_values is None else pixel_values.to(device, components.text_encoder.dtype), - image_grid_thw=None if image_grid_thw is None else image_grid_thw.to(device), - use_cache=False, - output_hidden_states=True, - ) - prompt_embeds = outputs.hidden_states[MINIMAX_H3_TEXT_ENCODER_LAYER].to(device=device, dtype=dtype) - return prompt_embeds, torch.tensor(token_tags, dtype=torch.long) - - @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - _check_prompt(block_state.prompt) - - # `encode_prompt` defaults the embedding dtype to the denoiser's; a text encoder block has no denoiser of - # its own — it is meant to run on its own — so it emits the conditioner's dtype, as every other model does. - block_state.prompt_embeds, block_state.text_token_tags = self.encode_prompt( - components, - block_state.prompt, - block_state.keyframes, - device=components._execution_device, - dtype=components.text_encoder.dtype, - ) - - self.set_block_state(state, block_state) - return components, state - - -class MiniMaxH3KeyframeVaeEncoderStep(ModularPipelineBlocks): - model_name = "minimax-h3" - - @property - def description(self) -> str: - return ( - "Encodes the `fl2va` keyframes into packed conditioning rows and noises them to MiniMax-H3's " - "conditioning level. The rows are the anchors of the whole denoising loop: the loop only ever writes the " - "generated rows, so they are never updated again." - ) - - @property - def expected_components(self) -> list[ComponentSpec]: - return [ - ComponentSpec("vae", AutoencoderKLMiniMaxH3), - ComponentSpec("scheduler", MiniMaxH3Scheduler), - ] - - @property - def inputs(self) -> list[InputParam]: - return [ - InputParam( - name="keyframes", - type_hint=list, - required=True, - description="The keyframes put onto the target canvas, in packed order.", - ), - InputParam(name="latent_height", type_hint=int, required=True, description="Height of the video latents."), - InputParam(name="latent_width", type_hint=int, required=True, description="Width of the video latents."), - InputParam.template( - "generator", - description=( - "The generator of the request. The conditioning noise is drawn from it before the target noise " - "of the prepare-latents step." - ), - ), - ] - - @property - def intermediate_outputs(self) -> list[OutputParam]: - return [ - OutputParam( - "condition_latents", - type_hint=torch.Tensor, - description="The noise-augmented video conditioning rows, in packed order.", - ) - ] - - @staticmethod - def encode_keyframes(components, images: list, device: torch.device | None = None) -> torch.Tensor: - r""" - Encode the `fl2va` keyframes into packed conditioning rows. - - The keyframes go through the video VAE's spatial encoder only — they are single frames, so none of its - 17-frame temporal chunking applies — and the posterior is *sampled*, under a generator seeded with 42 - independently of the request seed. The sampled latent is rounded to float16 before being normalized, as in the - reference implementation; both are part of reproducing the released model's conditioning. - - Args: - images (`list[PIL.Image.Image]`): - The keyframes, already prepared onto the target canvas, in packed order. - device (`torch.device`, *optional*): The device to run the VAE on. - - Returns: - `torch.Tensor` of shape `(num_condition_rows, latent_channels * prod(patch_size))`: the float32 - conditioning rows. - """ - device = device or components._execution_device - latents_mean = torch.tensor(components.vae.config.latents_mean).view(1, -1, 1, 1, 1) - latents_std = torch.tensor(components.vae.config.latents_std).view(1, -1, 1, 1, 1) - pixel_mean = torch.tensor(MINIMAX_H3_PIXEL_MEAN, device=device).view(1, -1, 1, 1, 1) - pixel_std = torch.tensor(MINIMAX_H3_PIXEL_STD, device=device).view(1, -1, 1, 1, 1) - - rows = [] - for image in images: - pixels = torch.from_numpy(np.array(image)).to(device).permute(2, 0, 1)[None, :, None] - pixels = (pixels.to(torch.float32).div(255.0) - pixel_mean) / pixel_std - # `vae.encode` chunks along time for videos; a keyframe is one frame and is encoded by the (tiled) - # spatial encoder alone, which is what the released model conditions on. - moments = components.vae._encode_clip(pixels) - posterior = DiagonalGaussianDistribution(moments) - latents = posterior.sample(generator=torch.Generator().manual_seed(MINIMAX_H3_KEYFRAME_ENCODE_SEED)) - # The sampled latent is rounded to float16 before it is normalized: ~11 bits of every conditioning - # latent, so the released model's conditioning cannot be reproduced without it. - latents = latents.to(torch.float16).float().cpu() - rows.append(patchify_video_latents((latents - latents_mean) / latents_std, components.patch_size)) - return torch.cat(rows) - - @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: - block_state = self.get_block_state(state) - device = components._execution_device - - condition_latents = self.encode_keyframes(components, block_state.keyframes, device=device) - noise = keyframe_condition_noise( - ((1, block_state.latent_height, block_state.latent_width),) * len(block_state.keyframes), - components.patch_size, - components.vae_latent_channels, - generator=block_state.generator, - device=device, - ) - block_state.condition_latents = components.scheduler.scale_noise( - condition_latents.to(device), MINIMAX_H3_KEYFRAME_NOISE_AUG, noise - ) - - self.set_block_state(state, block_state) - return components, state - - -class MiniMaxH3Ref2VATextEncoderStep(ModularPipelineBlocks): - model_name = "minimax-h3-ref2va" - - @property - def description(self) -> str: - return ( - "Encodes MiniMax-H3's presentation of a `ref2va` request: a label per reference, numbered per modality " - '(`": "` plus a vision block, `"