diff --git a/README.md b/README.md index d280be824a89f87d3a7278dd8ea587f932cd6778..b30f73aa4a00385e1fc52d033cc5d5a78d2a7edd 100644 --- a/README.md +++ b/README.md @@ -17,20 +17,19 @@ Joint video **and** soundtrack out of a single denoising pass, at **bfloat16 wit This Space is the denoising half: the 61.73 GiB transformer and the two autoencoders. The 62.14 GiB Qwen3-VL conditioner runs in -[`minimax-h3-conditioner`](https://huggingface.co/spaces/diffusers-internal-dev/minimax-h3-conditioner), which this -Space calls over the gradio API for every request. +[`qwen3vl-conditioner`](https://huggingface.co/spaces/multimodalart/qwen3vl-conditioner), which this +Space calls over the gradio API for every request. The weights are the public +[`MiniMaxAI/MiniMax-H3`](https://huggingface.co/MiniMaxAI/MiniMax-H3) diffusers checkpoint. ## 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 — the existing demos -([`minimax-h3`](https://huggingface.co/spaces/diffusers-internal-dev/minimax-h3), -[`-fp8`](https://huggingface.co/spaces/diffusers-internal-dev/minimax-h3-fp8)) run NVFP4 and float8 weights for that -reason alone. Cut the `MiniMaxH3Blocks` sequence at its `text_encoder` step and both halves fit unquantized: +Space is therefore impossible, which is why quantized demos of it run NVFP4 or float8 weights. Cut the +`MiniMaxH3Blocks` sequence at its `text_encoder` step and both halves fit unquantized: | Space | Subfolders | Download | Resident | |---|---|---|---| -| [`minimax-h3-conditioner`](https://huggingface.co/spaces/diffusers-internal-dev/minimax-h3-conditioner) | `text_encoder/` + `tokenizer/` + `processor/` | 66.7 GB | 62.15 GiB bf16 | +| [`qwen3vl-conditioner`](https://huggingface.co/spaces/multimodalart/qwen3vl-conditioner) | `text_encoder/` + `tokenizer/` + `processor/` | 66.7 GB | 62.15 GiB bf16 | | this one | `transformer/` + `vae/` + `audio_vae/` | 77.3 GB | 61.73 GiB bf16 + 10.43 GiB float32 | Besides the quality argument, unquantized weights are the ones AoTI can export; an NVFP4 checkpoint cannot be @@ -59,19 +58,22 @@ Space pays, where the whole cost is the traffic auto-offload has to move. ## How the split is expressed -`MiniMaxH3Blocks` is a `SequentialPipelineBlocks` of eight steps: +`MiniMaxH3Blocks` is a `SequentialPipelineBlocks` whose branches are picked per request — and per `workflow=` — from +the inputs: ``` -setup -> text_encoder -> vae_encoder -> prepare_layout -> prepare_latents -> set_timesteps -> denoise -> decode +before_encode -> text_encoder -> vae_encoder -> denoise -> after_denoise -> decode ``` +where `denoise` is itself `prepare_layout -> prepare_latents -> set_timesteps -> denoise`. + `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` / `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 = MiniMaxH3GeneratorBlocks().init_pipeline("diffusers-internal-dev/MiniMax-H3") +pipe = MiniMaxH3GeneratorBlocks().init_pipeline("MiniMaxAI/MiniMax-H3") pipe.load_components(dtype=torch.bfloat16) state = pipe(prompt_embeds=..., text_token_tags=..., height=768, width=1344, num_frames=124, num_inference_steps=30) ``` @@ -80,10 +82,12 @@ The wire format is exactly those two tensors — `(1, num_text_tokens, 5120)` bf carried as one safetensors file with the resolved `height` / `width` / `num_frames` in its metadata header. A text-only request is 246 KB of it; one 768x1344 keyframe adds 1016 vision rows and takes it to 10.7 MB. -The `setup` step runs on **both** halves. It owns no component (PIL and arithmetic) and it resolves the canvas, the -`17 * n + 5` frame count and the keyframes placed onto that canvas — which the conditioner needs to build its vision -blocks and this Space needs to encode with the video VAE. It is deterministic, and the conditioner returns the plan -it resolved so this Space pins the same canvas rather than re-deriving it. +The keyframe `resize` step runs on **both** halves. It owns no pretrained component (PIL and arithmetic) and it puts +the keyframes onto the target canvas — which the conditioner needs to build its vision blocks and this Space needs to +encode with the video VAE. It is deterministic, and the conditioner returns the plan it resolved so this Space pins +the same canvas rather than re-deriving it. Two things that step no longer does, and that both halves therefore do +themselves: EXIF-transposing a keyframe into upright RGB, and snapping `num_frames` to `17 * n + 5` — the frame count +is resolved by the layout step, which lives on this side of the cut. ## Nothing is paid for with GPU time @@ -147,18 +151,25 @@ one-time `PIPE.to("cuda")` is inside the first row's 339 s and does not reappear | Variable | Default | Meaning | |---|---|---| -| `H3_CONDITIONER` | `diffusers-internal-dev/minimax-h3-conditioner` | The Space this one asks for embeddings. | +| `H3_MODEL_REPO` | `MiniMaxAI/MiniMax-H3` | The diffusers-layout checkpoint. Public. | +| `H3_CONDITIONER` | `multimodalart/qwen3vl-conditioner` | The Space this one asks for embeddings. | | `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. | | `H3_GPU_DURATION` | `900` | Seconds per request; the pool applies a 1.5 duration factor. | | `H3_GPU_SIZE` | `xlarge` | ZeroGPU allocation size. `large` does not fit. | -## Required secret +## Secrets -`HF_TOKEN` — `diffusers-internal-dev/MiniMax-H3` is private, and so is the conditioner Space this one calls. +The weights and the conditioner Space are both public, so neither needs a token. `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. \ No newline at end of file +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, so block names and the shape of the split are exactly what a new head +is liable to move. \ No newline at end of file diff --git a/app.py b/app.py index 75dfa838191ba0463abcb0545115116361d5ab2f..cbeb31e2cbbda1870c9fe90d541a2a16a83fda6e 100644 --- a/app.py +++ b/app.py @@ -13,7 +13,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 @@ -85,8 +85,9 @@ def load_models() -> str | None: """Load the denoising half. At **startup**, but *not* onto the card. `MiniMaxH3GeneratorBlocks` declares `transformer`, `vae`, `audio_vae`, `scheduler`, `audio_scheduler` and - `video_processor`, so `load_components` fetches exactly those subfolders out of the shared - `modular_model_index.json` — `text_encoder/` and `transformer_ref/` are never touched. + `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 `transformer_ref/` are never touched. Both autoencoders carry `_keep_in_fp32_modules` over every module, so the `dtype` below is refused for them and they stay float32: a bfloat16 audio VAE decodes the soundtrack roughly 20 dB too quiet. @@ -105,11 +106,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 @@ -121,7 +117,9 @@ def load_models() -> str | None: blocks = MiniMaxH3GeneratorBlocks() print(f"[gen] 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) pipe.transformer.set_attention_backend(ATTENTION) # Still startup, still free: an AoTI package carries no weights and opens its compiled archive lazily inside @@ -182,8 +180,14 @@ def conditioner(): return CLIENT -def encode_remote(prompt, image_path, last_image_path, canvas, num_frames): - """Ask the conditioner Space for `prompt_embeds` + `text_token_tags`. Off this Space's GPU time entirely.""" +def encode_remote(prompt, image_path, last_image_path, canvas, num_frames, rewrite_prompt=False): + """Ask the conditioner Space for `prompt_embeds` + `text_token_tags`. Off this Space's GPU time entirely. + + `rewrite_prompt` is the conditioner's prompt upsampling: it rewrites the request into MiniMax-H3's trained format + with its own Qwen3-VL and encodes *that*, 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 it costs this + Space's `get_duration` nothing. + """ from gradio_client import handle_file from safetensors import safe_open @@ -193,6 +197,7 @@ def encode_remote(prompt, image_path, last_image_path, canvas, num_frames): last_image_path=handle_file(last_image_path) if last_image_path else None, canvas=canvas, num_frames=num_frames, + rewrite_prompt=bool(rewrite_prompt), api_name="/encode", ) with safe_open(path, framework="pt") as handle: @@ -250,7 +255,8 @@ def _generate(prompt_embeds, text_token_tags, image, last_image, height, width, return state.get("videos")[0], state.get("audio")[0].cpu(), state.get("sampling_rate") -def generate(prompt, image_path=None, last_image_path=None, canvas=DEFAULT_CANVAS, duration=5, steps=28, seed=42, progress=gr.Progress(track_tqdm=True)): +def generate(prompt, image_path=None, last_image_path=None, canvas=DEFAULT_CANVAS, duration=5, steps=28, seed=42, upsample=False, progress=gr.Progress(track_tqdm=True)): + """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: @@ -258,27 +264,34 @@ def generate(prompt, image_path=None, last_image_path=None, canvas=DEFAULT_CANVA if not prompt or not prompt.strip(): raise gr.Error("MiniMax-H3 always takes a prompt, keyframes or not.") - from PIL import Image + from PIL import Image, ImageOps from diffusers.utils import encode_video num_frames = snap_frames(duration) - progress(0.0, desc=f"Conditioning on {CONDITIONER_SPACE} ...") + progress(0.0, desc=f"Upsampling the prompt on {CONDITIONER_SPACE} ..." if upsample else f"Conditioning on {CONDITIONER_SPACE} ...") conditioned = time.time() prompt_embeds, text_token_tags, metadata, plan = encode_remote( - prompt, image_path, last_image_path, canvas, num_frames + prompt, image_path, last_image_path, canvas, num_frames, rewrite_prompt=upsample ) 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 "" + + # EXIF-transposed and in RGB before the blocks see them, which the resize step that replaced the old setup step no + # longer does itself. Both halves have to prepare a keyframe the same way or the conditioning latents encoded here + # would not be of the image the conditioner looked at. + def keyframe(path): + return ImageOps.exif_transpose(Image.open(path)).convert("RGB") if path else None progress(0.1, desc=f"Denoising {steps} steps at {width}x{height}, {num_frames} frames ...") started = time.time() frames, audio, sampling_rate = _generate( prompt_embeds, text_token_tags, - Image.open(image_path) if image_path else None, - Image.open(last_image_path) if last_image_path else None, + keyframe(image_path), + keyframe(last_image_path), height, width, num_frames, @@ -294,11 +307,12 @@ def generate(prompt, image_path=None, last_image_path=None, canvas=DEFAULT_CANVA report = ( f"`{width}x{height}`, {num_frames} frames ({num_frames / FPS:.3f} s), {int(steps)} steps · " - f"conditioner {condition_seconds:.0f}s ({plan['num_text_tokens']} tokens) · " + f"conditioner {condition_seconds:.0f}s ({plan['num_text_tokens']} tokens" + f"{', upsampled' if refined else ''}) · " f"denoise + decode {generate_seconds:.0f}s ({generate_seconds / int(steps):.1f} s/step) · seed {int(seed)}" ) print(f"[gen] {report}", flush=True) - return path, report + return path, report, refined @@ -378,10 +392,22 @@ with gr.Blocks(title="MiniMax-H3") as demo: duration = gr.Slider(label="Duration (s)", minimum=2, maximum=MAX_UI_DURATION, step=1, value=5) 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(): video = gr.Video(label="Video + soundtrack") report = gr.Markdown(visible=False) + 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.", + ) image.upload(_fit_keyframe, [image, canvas], [image, canvas]) @@ -394,16 +420,19 @@ with gr.Blocks(title="MiniMax-H3") as demo: ["A slow seamless camera move from the first view to the last", "examples/first.png", "examples/last.png", "1344x768 · 16:9 full"], ], inputs=[prompt, image, last_image, canvas], - outputs=[video, report], + outputs=[video, report, upsampled], fn=generate, cache_examples=True, cache_mode="lazy", ) + # `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. Same on the way out: the + # video and the report stay first and the upsampled prompt is appended last. run.click( generate, - [prompt, image, last_image, canvas, duration, steps, seed], - [video, report], + [prompt, image, last_image, canvas, duration, steps, seed, upsample], + [video, report, upsampled], api_name="generate", ) 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, `"