From 47666bb87079858c448ffe41a84dc71099b36c75 Mon Sep 17 00:00:00 2001 From: christopher5106 Date: Thu, 24 Sep 2026 13:54:25 +0200 Subject: [PATCH 1/5] [LoRA] accept fal-kontext LoRAs whose global embedder keys lack the base_model.model. prefix Some fal-kontext LoRAs store time_in / vector_in / txt_in / img_in / guidance_in without the `base_model.model.` prefix the block keys use, so `_convert_fal_kontext_lora_to_diffusers` left them in `original_state_dict` and raised "`original_state_dict` should be empty at this point". Map them to their diffusers names before that check. Rebuilt from scenario-labs/diffusers@02abf7eb1 (2026-04-15) without its Kohya Flux.2 hunks, which upstream now covers with `_convert_kohya_flux2_lora_to_diffusers`. --- .../loaders/lora_conversion_utils.py | 22 ++++++ tests/lora/test_lora_conversion_utils.py | 76 +++++++++++++++++++ 2 files changed, 98 insertions(+) create mode 100644 tests/lora/test_lora_conversion_utils.py diff --git a/src/diffusers/loaders/lora_conversion_utils.py b/src/diffusers/loaders/lora_conversion_utils.py index 1b7bcc795d6b..bc111d77dd44 100644 --- a/src/diffusers/loaders/lora_conversion_utils.py +++ b/src/diffusers/loaders/lora_conversion_utils.py @@ -1582,6 +1582,28 @@ def _convert_fal_kontext_lora_to_diffusers(original_state_dict): f"{original_block_prefix}final_layer.linear.{lora_key}.bias" ) + # Some fal-kontext LoRAs carry the global embedder keys (time_in, vector_in, txt_in, img_in, guidance_in) + # without the `base_model.model.` prefix the block keys use. + for lora_key in ["lora_A", "lora_B"]: + for src, dst in [ + (f"time_in.in_layer.{lora_key}.weight", f"time_text_embed.timestep_embedder.linear_1.{lora_key}.weight"), + (f"time_in.out_layer.{lora_key}.weight", f"time_text_embed.timestep_embedder.linear_2.{lora_key}.weight"), + (f"vector_in.in_layer.{lora_key}.weight", f"time_text_embed.text_embedder.linear_1.{lora_key}.weight"), + (f"vector_in.out_layer.{lora_key}.weight", f"time_text_embed.text_embedder.linear_2.{lora_key}.weight"), + (f"txt_in.{lora_key}.weight", f"context_embedder.{lora_key}.weight"), + (f"img_in.{lora_key}.weight", f"x_embedder.{lora_key}.weight"), + ( + f"guidance_in.in_layer.{lora_key}.weight", + f"time_text_embed.guidance_embedder.linear_1.{lora_key}.weight", + ), + ( + f"guidance_in.out_layer.{lora_key}.weight", + f"time_text_embed.guidance_embedder.linear_2.{lora_key}.weight", + ), + ]: + if src in original_state_dict: + converted_state_dict[dst] = original_state_dict.pop(src) + if len(original_state_dict) > 0: raise ValueError(f"`original_state_dict` should be empty at this point but has {original_state_dict.keys()=}.") diff --git a/tests/lora/test_lora_conversion_utils.py b/tests/lora/test_lora_conversion_utils.py new file mode 100644 index 000000000000..2f287ad6c5f3 --- /dev/null +++ b/tests/lora/test_lora_conversion_utils.py @@ -0,0 +1,76 @@ +# coding=utf-8 +# Copyright 2026 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 pytest +import torch + +from diffusers.loaders.lora_conversion_utils import _convert_fal_kontext_lora_to_diffusers + + +def _fal_kontext_state_dict(rank=1, inner_dim=3072, mlp_hidden_dim=12288, num_layers=19, num_single_layers=38): + """Minimal fal-kontext LoRA (block keys only), shaped so the qkv / linear1 splits work.""" + prefix = "base_model.model." + sd = {} + for i in range(num_layers): + for module in ["img_mod.lin", "txt_mod.lin", "img_mlp.0", "img_mlp.2", "txt_mlp.0", "txt_mlp.2"]: + sd[f"{prefix}double_blocks.{i}.{module}.lora_A.weight"] = torch.zeros(rank, 1) + sd[f"{prefix}double_blocks.{i}.{module}.lora_B.weight"] = torch.zeros(1, rank) + for module in ["img_attn.proj", "txt_attn.proj"]: + sd[f"{prefix}double_blocks.{i}.{module}.lora_A.weight"] = torch.zeros(rank, 1) + sd[f"{prefix}double_blocks.{i}.{module}.lora_B.weight"] = torch.zeros(1, rank) + for module in ["img_attn.qkv", "txt_attn.qkv"]: + sd[f"{prefix}double_blocks.{i}.{module}.lora_A.weight"] = torch.zeros(rank, 1) + sd[f"{prefix}double_blocks.{i}.{module}.lora_B.weight"] = torch.zeros(3 * inner_dim, rank) + for i in range(num_single_layers): + sd[f"{prefix}single_blocks.{i}.modulation.lin.lora_A.weight"] = torch.zeros(rank, 1) + sd[f"{prefix}single_blocks.{i}.modulation.lin.lora_B.weight"] = torch.zeros(1, rank) + sd[f"{prefix}single_blocks.{i}.linear1.lora_A.weight"] = torch.zeros(rank, 1) + sd[f"{prefix}single_blocks.{i}.linear1.lora_B.weight"] = torch.zeros(3 * inner_dim + mlp_hidden_dim, rank) + sd[f"{prefix}single_blocks.{i}.linear2.lora_A.weight"] = torch.zeros(rank, 1) + sd[f"{prefix}single_blocks.{i}.linear2.lora_B.weight"] = torch.zeros(1, rank) + sd[f"{prefix}final_layer.linear.lora_A.weight"] = torch.zeros(rank, 1) + sd[f"{prefix}final_layer.linear.lora_B.weight"] = torch.zeros(1, rank) + return sd + + +UNPREFIXED_GLOBAL_KEYS = { + "time_in.in_layer": "time_text_embed.timestep_embedder.linear_1", + "time_in.out_layer": "time_text_embed.timestep_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", + "txt_in": "context_embedder", + "img_in": "x_embedder", + "guidance_in.in_layer": "time_text_embed.guidance_embedder.linear_1", + "guidance_in.out_layer": "time_text_embed.guidance_embedder.linear_2", +} + + +def test_fal_kontext_conversion_blocks_only(): + converted = _convert_fal_kontext_lora_to_diffusers(_fal_kontext_state_dict()) + assert all(k.startswith("transformer.") for k in converted) + assert "transformer.transformer_blocks.0.attn.to_q.lora_B.weight" in converted + assert "transformer.single_transformer_blocks.37.proj_mlp.lora_B.weight" in converted + + +@pytest.mark.parametrize("lora_key", ["lora_A", "lora_B"]) +def test_fal_kontext_conversion_accepts_unprefixed_global_embedders(lora_key): + sd = _fal_kontext_state_dict() + for src in UNPREFIXED_GLOBAL_KEYS: + sd[f"{src}.{lora_key}.weight"] = torch.ones(1, 1) + + converted = _convert_fal_kontext_lora_to_diffusers(sd) + + for src, dst in UNPREFIXED_GLOBAL_KEYS.items(): + assert f"transformer.{dst}.{lora_key}.weight" in converted, src + assert torch.equal(converted[f"transformer.{dst}.{lora_key}.weight"], torch.ones(1, 1)) From 2a27b3e2481a8c318ad33a8b681b08bd6ac879e5 Mon Sep 17 00:00:00 2001 From: christopher5106 Date: Tue, 29 Sep 2026 23:01:50 +0200 Subject: [PATCH 2/5] [LoRA] drop the unit test of the private fal-kontext converter diffusers does not unit-test private converter functions against third-party checkpoint layouts; the fix is reproducible through FluxPipeline.lora_state_dict (see the linked issue). --- tests/lora/test_lora_conversion_utils.py | 76 ------------------------ 1 file changed, 76 deletions(-) delete mode 100644 tests/lora/test_lora_conversion_utils.py diff --git a/tests/lora/test_lora_conversion_utils.py b/tests/lora/test_lora_conversion_utils.py deleted file mode 100644 index 2f287ad6c5f3..000000000000 --- a/tests/lora/test_lora_conversion_utils.py +++ /dev/null @@ -1,76 +0,0 @@ -# coding=utf-8 -# Copyright 2026 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 pytest -import torch - -from diffusers.loaders.lora_conversion_utils import _convert_fal_kontext_lora_to_diffusers - - -def _fal_kontext_state_dict(rank=1, inner_dim=3072, mlp_hidden_dim=12288, num_layers=19, num_single_layers=38): - """Minimal fal-kontext LoRA (block keys only), shaped so the qkv / linear1 splits work.""" - prefix = "base_model.model." - sd = {} - for i in range(num_layers): - for module in ["img_mod.lin", "txt_mod.lin", "img_mlp.0", "img_mlp.2", "txt_mlp.0", "txt_mlp.2"]: - sd[f"{prefix}double_blocks.{i}.{module}.lora_A.weight"] = torch.zeros(rank, 1) - sd[f"{prefix}double_blocks.{i}.{module}.lora_B.weight"] = torch.zeros(1, rank) - for module in ["img_attn.proj", "txt_attn.proj"]: - sd[f"{prefix}double_blocks.{i}.{module}.lora_A.weight"] = torch.zeros(rank, 1) - sd[f"{prefix}double_blocks.{i}.{module}.lora_B.weight"] = torch.zeros(1, rank) - for module in ["img_attn.qkv", "txt_attn.qkv"]: - sd[f"{prefix}double_blocks.{i}.{module}.lora_A.weight"] = torch.zeros(rank, 1) - sd[f"{prefix}double_blocks.{i}.{module}.lora_B.weight"] = torch.zeros(3 * inner_dim, rank) - for i in range(num_single_layers): - sd[f"{prefix}single_blocks.{i}.modulation.lin.lora_A.weight"] = torch.zeros(rank, 1) - sd[f"{prefix}single_blocks.{i}.modulation.lin.lora_B.weight"] = torch.zeros(1, rank) - sd[f"{prefix}single_blocks.{i}.linear1.lora_A.weight"] = torch.zeros(rank, 1) - sd[f"{prefix}single_blocks.{i}.linear1.lora_B.weight"] = torch.zeros(3 * inner_dim + mlp_hidden_dim, rank) - sd[f"{prefix}single_blocks.{i}.linear2.lora_A.weight"] = torch.zeros(rank, 1) - sd[f"{prefix}single_blocks.{i}.linear2.lora_B.weight"] = torch.zeros(1, rank) - sd[f"{prefix}final_layer.linear.lora_A.weight"] = torch.zeros(rank, 1) - sd[f"{prefix}final_layer.linear.lora_B.weight"] = torch.zeros(1, rank) - return sd - - -UNPREFIXED_GLOBAL_KEYS = { - "time_in.in_layer": "time_text_embed.timestep_embedder.linear_1", - "time_in.out_layer": "time_text_embed.timestep_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", - "txt_in": "context_embedder", - "img_in": "x_embedder", - "guidance_in.in_layer": "time_text_embed.guidance_embedder.linear_1", - "guidance_in.out_layer": "time_text_embed.guidance_embedder.linear_2", -} - - -def test_fal_kontext_conversion_blocks_only(): - converted = _convert_fal_kontext_lora_to_diffusers(_fal_kontext_state_dict()) - assert all(k.startswith("transformer.") for k in converted) - assert "transformer.transformer_blocks.0.attn.to_q.lora_B.weight" in converted - assert "transformer.single_transformer_blocks.37.proj_mlp.lora_B.weight" in converted - - -@pytest.mark.parametrize("lora_key", ["lora_A", "lora_B"]) -def test_fal_kontext_conversion_accepts_unprefixed_global_embedders(lora_key): - sd = _fal_kontext_state_dict() - for src in UNPREFIXED_GLOBAL_KEYS: - sd[f"{src}.{lora_key}.weight"] = torch.ones(1, 1) - - converted = _convert_fal_kontext_lora_to_diffusers(sd) - - for src, dst in UNPREFIXED_GLOBAL_KEYS.items(): - assert f"transformer.{dst}.{lora_key}.weight" in converted, src - assert torch.equal(converted[f"transformer.{dst}.{lora_key}.weight"], torch.ones(1, 1)) From a12c38a6893efbfacc6fedc3fd79878fd63c30fa Mon Sep 17 00:00:00 2001 From: Dhruv Nair Date: Wed, 30 Sep 2026 17:38:01 +0530 Subject: [PATCH 3/5] Remove deprecated code paths (#14838) update --- examples/community/README.md | 4 +- examples/community/fresco_v2v.py | 2 +- .../pipeline_animatediff_controlnet.py | 2 +- ..._stable_diffusion_xl_controlnet_adapter.py | 2 +- ...diffusion_xl_controlnet_adapter_inpaint.py | 2 +- ...e_stable_diffusion_xl_instandid_img2img.py | 2 +- .../pipeline_stable_diffusion_xl_instantid.py | 2 +- examples/community/rerender_a_video.py | 2 +- .../stable_diffusion_controlnet_img2img.py | 2 +- .../stable_diffusion_controlnet_inpaint.py | 2 +- .../stable_diffusion_controlnet_reference.py | 2 +- ...table_diffusion_xl_controlnet_reference.py | 2 +- examples/research_projects/anytext/anytext.py | 2 +- .../pipeline_prompt_diffusion.py | 2 +- src/diffusers/__init__.py | 4 - src/diffusers/loaders/__init__.py | 49 +--- src/diffusers/loaders/lora_pipeline.py | 2 +- .../controlnets/controlnet_qwenimage.py | 29 -- src/diffusers/models/embeddings.py | 249 +----------------- .../models/transformers/transformer_chroma.py | 15 +- .../transformers/transformer_hidream_image.py | 17 +- src/diffusers/models/vq_model.py | 17 +- src/diffusers/pipelines/__init__.py | 8 +- .../pipelines/controlnet/__init__.py | 2 - .../pipelines/controlnet/multicontrolnet.py | 12 - .../hidream_image/pipeline_hidream_image.py | 18 +- src/diffusers/pipelines/lumina/__init__.py | 4 +- .../pipelines/lumina/pipeline_lumina.py | 21 -- src/diffusers/pipelines/lumina2/__init__.py | 4 +- .../pipelines/lumina2/pipeline_lumina2.py | 21 -- .../dummy_torch_and_transformers_objects.py | 30 --- 31 files changed, 36 insertions(+), 496 deletions(-) delete mode 100644 src/diffusers/pipelines/controlnet/multicontrolnet.py diff --git a/examples/community/README.md b/examples/community/README.md index f5c3357157e0..9a037db6be18 100644 --- a/examples/community/README.md +++ b/examples/community/README.md @@ -3526,7 +3526,7 @@ from controlnet_aux.midas import MidasDetector from PIL import Image from diffusers import AutoencoderKL, ControlNetModel, MultiAdapter, T2IAdapter -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.utils import load_image from examples.community.pipeline_stable_diffusion_xl_controlnet_adapter import ( StableDiffusionXLControlNetAdapterPipeline, @@ -3591,7 +3591,7 @@ from controlnet_aux.midas import MidasDetector from PIL import Image from diffusers import AutoencoderKL, ControlNetModel, MultiAdapter, T2IAdapter -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.utils import load_image from examples.community.pipeline_stable_diffusion_xl_controlnet_adapter_inpaint import ( StableDiffusionXLControlNetAdapterInpaintPipeline, diff --git a/examples/community/fresco_v2v.py b/examples/community/fresco_v2v.py index 628af232130e..16be186a0b62 100644 --- a/examples/community/fresco_v2v.py +++ b/examples/community/fresco_v2v.py @@ -29,9 +29,9 @@ from diffusers.loaders import StableDiffusionLoraLoaderMixin, TextualInversionLoaderMixin from diffusers.models import AutoencoderKL, ControlNetModel, ImageProjection, UNet2DConditionModel from diffusers.models.attention_processor import AttnProcessor2_0 +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.models.lora import adjust_lora_scale_text_encoder from diffusers.models.unets.unet_2d_condition import UNet2DConditionOutput -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel from diffusers.pipelines.controlnet.pipeline_controlnet_img2img import StableDiffusionControlNetImg2ImgPipeline from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput from diffusers.pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker diff --git a/examples/community/pipeline_animatediff_controlnet.py b/examples/community/pipeline_animatediff_controlnet.py index 5bc53b77324e..a31460533875 100644 --- a/examples/community/pipeline_animatediff_controlnet.py +++ b/examples/community/pipeline_animatediff_controlnet.py @@ -24,10 +24,10 @@ from diffusers.image_processor import PipelineImageInput, VaeImageProcessor from diffusers.loaders import IPAdapterMixin, StableDiffusionLoraLoaderMixin, TextualInversionLoaderMixin from diffusers.models import AutoencoderKL, ControlNetModel, ImageProjection, UNet2DConditionModel, UNetMotionModel +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.models.lora import adjust_lora_scale_text_encoder from diffusers.models.unets.unet_motion_model import MotionAdapter from diffusers.pipelines.animatediff.pipeline_output import AnimateDiffPipelineOutput -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel from diffusers.pipelines.pipeline_utils import DiffusionPipeline, StableDiffusionMixin from diffusers.schedulers import ( DDIMScheduler, diff --git a/examples/community/pipeline_stable_diffusion_xl_controlnet_adapter.py b/examples/community/pipeline_stable_diffusion_xl_controlnet_adapter.py index e38801cd7647..1e297a70bdef 100644 --- a/examples/community/pipeline_stable_diffusion_xl_controlnet_adapter.py +++ b/examples/community/pipeline_stable_diffusion_xl_controlnet_adapter.py @@ -25,8 +25,8 @@ from diffusers.image_processor import PipelineImageInput, VaeImageProcessor from diffusers.loaders import FromSingleFileMixin, StableDiffusionXLLoraLoaderMixin, TextualInversionLoaderMixin from diffusers.models import AutoencoderKL, ControlNetModel, MultiAdapter, T2IAdapter, UNet2DConditionModel +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.models.lora import adjust_lora_scale_text_encoder -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel from diffusers.pipelines.pipeline_utils import DiffusionPipeline, StableDiffusionMixin from diffusers.pipelines.stable_diffusion_xl.pipeline_output import StableDiffusionXLPipelineOutput from diffusers.schedulers import KarrasDiffusionSchedulers diff --git a/examples/community/pipeline_stable_diffusion_xl_controlnet_adapter_inpaint.py b/examples/community/pipeline_stable_diffusion_xl_controlnet_adapter_inpaint.py index 2e05e3380316..8ffb938f7efc 100644 --- a/examples/community/pipeline_stable_diffusion_xl_controlnet_adapter_inpaint.py +++ b/examples/community/pipeline_stable_diffusion_xl_controlnet_adapter_inpaint.py @@ -43,8 +43,8 @@ T2IAdapter, UNet2DConditionModel, ) +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.models.lora import adjust_lora_scale_text_encoder -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel from diffusers.pipelines.pipeline_utils import StableDiffusionMixin from diffusers.pipelines.stable_diffusion_xl.pipeline_output import StableDiffusionXLPipelineOutput from diffusers.schedulers import KarrasDiffusionSchedulers diff --git a/examples/community/pipeline_stable_diffusion_xl_instandid_img2img.py b/examples/community/pipeline_stable_diffusion_xl_instandid_img2img.py index 1710f682d0ed..b7c92e565fee 100644 --- a/examples/community/pipeline_stable_diffusion_xl_instandid_img2img.py +++ b/examples/community/pipeline_stable_diffusion_xl_instandid_img2img.py @@ -25,7 +25,7 @@ from diffusers import StableDiffusionXLControlNetImg2ImgPipeline from diffusers.image_processor import PipelineImageInput from diffusers.models import ControlNetModel -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.pipelines.stable_diffusion_xl import StableDiffusionXLPipelineOutput from diffusers.utils import ( deprecate, diff --git a/examples/community/pipeline_stable_diffusion_xl_instantid.py b/examples/community/pipeline_stable_diffusion_xl_instantid.py index 4dfbcc194dd8..5a27067703ea 100644 --- a/examples/community/pipeline_stable_diffusion_xl_instantid.py +++ b/examples/community/pipeline_stable_diffusion_xl_instantid.py @@ -25,7 +25,7 @@ from diffusers import StableDiffusionXLControlNetPipeline from diffusers.image_processor import PipelineImageInput from diffusers.models import ControlNetModel -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.pipelines.stable_diffusion_xl import StableDiffusionXLPipelineOutput from diffusers.utils import ( deprecate, diff --git a/examples/community/rerender_a_video.py b/examples/community/rerender_a_video.py index 68872e51a792..c9402284bbca 100644 --- a/examples/community/rerender_a_video.py +++ b/examples/community/rerender_a_video.py @@ -26,7 +26,7 @@ from diffusers.image_processor import VaeImageProcessor from diffusers.models import AutoencoderKL, ControlNetModel, UNet2DConditionModel from diffusers.models.attention_processor import Attention, AttnProcessor -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.pipelines.controlnet.pipeline_controlnet_img2img import StableDiffusionControlNetImg2ImgPipeline from diffusers.pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker from diffusers.schedulers import KarrasDiffusionSchedulers diff --git a/examples/community/stable_diffusion_controlnet_img2img.py b/examples/community/stable_diffusion_controlnet_img2img.py index 03c6fe7f6466..ab68ca3786bb 100644 --- a/examples/community/stable_diffusion_controlnet_img2img.py +++ b/examples/community/stable_diffusion_controlnet_img2img.py @@ -9,7 +9,7 @@ from transformers import CLIPImageProcessor, CLIPTextModel, CLIPTokenizer from diffusers import AutoencoderKL, ControlNetModel, UNet2DConditionModel, logging -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.pipelines.pipeline_utils import DiffusionPipeline, StableDiffusionMixin from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput, StableDiffusionSafetyChecker from diffusers.schedulers import KarrasDiffusionSchedulers diff --git a/examples/community/stable_diffusion_controlnet_inpaint.py b/examples/community/stable_diffusion_controlnet_inpaint.py index 9b76faf56a8a..52586e770e33 100644 --- a/examples/community/stable_diffusion_controlnet_inpaint.py +++ b/examples/community/stable_diffusion_controlnet_inpaint.py @@ -10,7 +10,7 @@ from transformers import CLIPImageProcessor, CLIPTextModel, CLIPTokenizer from diffusers import AutoencoderKL, ControlNetModel, UNet2DConditionModel, logging -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.pipelines.pipeline_utils import DiffusionPipeline, StableDiffusionMixin from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput, StableDiffusionSafetyChecker from diffusers.schedulers import KarrasDiffusionSchedulers diff --git a/examples/community/stable_diffusion_controlnet_reference.py b/examples/community/stable_diffusion_controlnet_reference.py index 18c79a0853f9..4c8c00ad8aa9 100644 --- a/examples/community/stable_diffusion_controlnet_reference.py +++ b/examples/community/stable_diffusion_controlnet_reference.py @@ -8,8 +8,8 @@ from diffusers import StableDiffusionControlNetPipeline from diffusers.models import ControlNetModel from diffusers.models.attention import BasicTransformerBlock +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.models.unets.unet_2d_blocks import CrossAttnDownBlock2D, CrossAttnUpBlock2D, DownBlock2D, UpBlock2D -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput from diffusers.utils import logging from diffusers.utils.torch_utils import is_compiled_module, randn_tensor diff --git a/examples/community/stable_diffusion_xl_controlnet_reference.py b/examples/community/stable_diffusion_xl_controlnet_reference.py index a458ee7c6506..2a02ca609599 100644 --- a/examples/community/stable_diffusion_xl_controlnet_reference.py +++ b/examples/community/stable_diffusion_xl_controlnet_reference.py @@ -12,8 +12,8 @@ from diffusers.image_processor import PipelineImageInput from diffusers.models import ControlNetModel from diffusers.models.attention import BasicTransformerBlock +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.models.unets.unet_2d_blocks import CrossAttnDownBlock2D, CrossAttnUpBlock2D, DownBlock2D, UpBlock2D -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel from diffusers.pipelines.stable_diffusion_xl.pipeline_output import StableDiffusionXLPipelineOutput from diffusers.utils import PIL_INTERPOLATION, deprecate, logging, replace_example_docstring from diffusers.utils.torch_utils import is_compiled_module, is_torch_version, randn_tensor diff --git a/examples/research_projects/anytext/anytext.py b/examples/research_projects/anytext/anytext.py index 65fb8775f718..fce1ee4e1bb4 100644 --- a/examples/research_projects/anytext/anytext.py +++ b/examples/research_projects/anytext/anytext.py @@ -53,9 +53,9 @@ TextualInversionLoaderMixin, ) from diffusers.models import AutoencoderKL, ControlNetModel, ImageProjection, UNet2DConditionModel +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.models.lora import adjust_lora_scale_text_encoder from diffusers.models.modeling_utils import ModelMixin -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel from diffusers.pipelines.pipeline_utils import DiffusionPipeline, StableDiffusionMixin from diffusers.pipelines.stable_diffusion.pipeline_output import StableDiffusionPipelineOutput from diffusers.pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker diff --git a/examples/research_projects/promptdiffusion/pipeline_prompt_diffusion.py b/examples/research_projects/promptdiffusion/pipeline_prompt_diffusion.py index 8b23570aea77..f83992725e47 100644 --- a/examples/research_projects/promptdiffusion/pipeline_prompt_diffusion.py +++ b/examples/research_projects/promptdiffusion/pipeline_prompt_diffusion.py @@ -30,8 +30,8 @@ from diffusers.image_processor import PipelineImageInput, VaeImageProcessor from diffusers.loaders import FromSingleFileMixin, StableDiffusionLoraLoaderMixin, TextualInversionLoaderMixin from diffusers.models import AutoencoderKL, ControlNetModel, UNet2DConditionModel +from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel from diffusers.models.lora import adjust_lora_scale_text_encoder -from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel from diffusers.pipelines.pipeline_utils import DiffusionPipeline from diffusers.pipelines.stable_diffusion.pipeline_output import StableDiffusionPipelineOutput from diffusers.pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker diff --git a/src/diffusers/__init__.py b/src/diffusers/__init__.py index 2825e9888c98..942f201fd3fd 100644 --- a/src/diffusers/__init__.py +++ b/src/diffusers/__init__.py @@ -747,9 +747,7 @@ "LTXPipeline", "LucyEditPipeline", "Lumina2Pipeline", - "Lumina2Text2ImgPipeline", "LuminaPipeline", - "LuminaText2ImgPipeline", "MarigoldDepthPipeline", "MarigoldIntrinsicsPipeline", "MarigoldNormalsPipeline", @@ -1604,9 +1602,7 @@ LTXPipeline, LucyEditPipeline, Lumina2Pipeline, - Lumina2Text2ImgPipeline, LuminaPipeline, - LuminaText2ImgPipeline, MarigoldDepthPipeline, MarigoldIntrinsicsPipeline, MarigoldNormalsPipeline, diff --git a/src/diffusers/loaders/__init__.py b/src/diffusers/loaders/__init__.py index 828744386453..5edfa6c01d26 100644 --- a/src/diffusers/loaders/__init__.py +++ b/src/diffusers/loaders/__init__.py @@ -1,56 +1,9 @@ from typing import TYPE_CHECKING -from ..utils import DIFFUSERS_SLOW_IMPORT, _LazyModule, deprecate +from ..utils import DIFFUSERS_SLOW_IMPORT, _LazyModule 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(): diff --git a/src/diffusers/loaders/lora_pipeline.py b/src/diffusers/loaders/lora_pipeline.py index 739ff9d2b3b1..0809066b3dc8 100644 --- a/src/diffusers/loaders/lora_pipeline.py +++ b/src/diffusers/loaders/lora_pipeline.py @@ -3700,7 +3700,7 @@ def unfuse_lora(self, components: list[str] = ["transformer"], **kwargs): class Lumina2LoraLoaderMixin(LoraBaseMixin): r""" - Load LoRA layers into [`Lumina2Transformer2DModel`]. Specific to [`Lumina2Text2ImgPipeline`]. + Load LoRA layers into [`Lumina2Transformer2DModel`]. Specific to [`Lumina2Pipeline`]. """ _lora_loadable_modules = ["transformer"] diff --git a/src/diffusers/models/controlnets/controlnet_qwenimage.py b/src/diffusers/models/controlnets/controlnet_qwenimage.py index f721c51261e1..7ff00c7314e1 100644 --- a/src/diffusers/models/controlnets/controlnet_qwenimage.py +++ b/src/diffusers/models/controlnets/controlnet_qwenimage.py @@ -23,7 +23,6 @@ from ...utils import ( BaseOutput, apply_lora_scale, - deprecate, logging, ) from ..attention import AttentionMixin @@ -138,7 +137,6 @@ def forward( 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: @@ -162,9 +160,6 @@ def forward( 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 @@ -176,17 +171,6 @@ def forward( 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 @@ -273,7 +257,6 @@ def forward( 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: @@ -293,9 +276,6 @@ def forward( 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 @@ -308,15 +288,6 @@ def forward( 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: diff --git a/src/diffusers/models/embeddings.py b/src/diffusers/models/embeddings.py index f3448c07857c..cbebf3de3a50 100644 --- a/src/diffusers/models/embeddings.py +++ b/src/diffusers/models/embeddings.py @@ -85,7 +85,7 @@ def get_3d_sincos_pos_embed( spatial_interpolation_scale: float = 1.0, temporal_interpolation_scale: float = 1.0, device: torch.device | None = None, - output_type: str = "np", + output_type: str = "pt", ) -> torch.Tensor: r""" Creates 3D sinusoidal positional embeddings. @@ -108,14 +108,6 @@ def get_3d_sincos_pos_embed( 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): @@ -152,72 +144,6 @@ def get_3d_sincos_pos_embed( 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, @@ -226,7 +152,7 @@ def get_2d_sincos_pos_embed( interpolation_scale=1.0, base_size=16, device: torch.device | None = None, - output_type: str = "np", + output_type: str = "pt", ): """ Creates 2D sinusoidal positional embeddings. @@ -248,21 +174,6 @@ def get_2d_sincos_pos_embed( 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) @@ -286,7 +197,7 @@ def get_2d_sincos_pos_embed( return pos_embed -def get_2d_sincos_pos_embed_from_grid(embed_dim, grid, output_type="np"): +def get_2d_sincos_pos_embed_from_grid(embed_dim, grid, output_type="pt"): r""" This function generates 2D sinusoidal positional embeddings from a grid. @@ -297,17 +208,6 @@ def get_2d_sincos_pos_embed_from_grid(embed_dim, grid, output_type="np"): 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") @@ -319,14 +219,14 @@ def get_2d_sincos_pos_embed_from_grid(embed_dim, grid, output_type="np"): return emb -def get_1d_sincos_pos_embed_from_grid(embed_dim, pos, output_type="np", flip_sin_to_cos=False, dtype=None): +def get_1d_sincos_pos_embed_from_grid(embed_dim, pos, output_type="pt", 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. + output_type (`str`, *optional*, defaults to `"pt"`): Output type. Only `"pt"` is supported. 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. @@ -334,14 +234,6 @@ def get_1d_sincos_pos_embed_from_grid(embed_dim, pos, output_type="np", flip_sin 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") @@ -368,94 +260,6 @@ def get_1d_sincos_pos_embed_from_grid(embed_dim, pos, output_type="np", flip_sin 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. @@ -973,7 +777,7 @@ def get_3d_rotary_pos_embed_allegro( def get_2d_rotary_pos_embed( - embed_dim, crops_coords, grid_size, use_real=True, device: torch.device | None = None, output_type: str = "np" + embed_dim, crops_coords, grid_size, use_real=True, device: torch.device | None = None, output_type: str = "pt" ): """ RoPE for image tokens with 2d structure. @@ -993,19 +797,6 @@ def get_2d_rotary_pos_embed( 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( @@ -1022,34 +813,6 @@ def get_2d_rotary_pos_embed( 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. diff --git a/src/diffusers/models/transformers/transformer_chroma.py b/src/diffusers/models/transformers/transformer_chroma.py index 8d7d9d5d6a04..92190bb0120d 100644 --- a/src/diffusers/models/transformers/transformer_chroma.py +++ b/src/diffusers/models/transformers/transformer_chroma.py @@ -21,8 +21,7 @@ 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 import apply_lora_scale, logging from ...utils.torch_utils import maybe_allow_in_graph from ..attention import AttentionMixin, FeedForward from ..cache_utils import CacheMixin @@ -216,17 +215,7 @@ def __init__( 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() + processor = FluxAttnProcessor() self.attn = FluxAttention( query_dim=dim, diff --git a/src/diffusers/models/transformers/transformer_hidream_image.py b/src/diffusers/models/transformers/transformer_hidream_image.py index bd69d5de68ca..703230562415 100644 --- a/src/diffusers/models/transformers/transformer_hidream_image.py +++ b/src/diffusers/models/transformers/transformer_hidream_image.py @@ -8,7 +8,7 @@ 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 import apply_lora_scale, logging from ...utils.torch_utils import maybe_adjust_dtype_for_device, maybe_allow_in_graph from ..attention import Attention from ..embeddings import TimestepEmbedding, Timesteps @@ -783,7 +783,6 @@ def forward( 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. @@ -817,20 +816,6 @@ def forward( 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: diff --git a/src/diffusers/models/vq_model.py b/src/diffusers/models/vq_model.py index 635db5310258..52200d9c664d 100644 --- a/src/diffusers/models/vq_model.py +++ b/src/diffusers/models/vq_model.py @@ -11,19 +11,4 @@ # WITHOUT 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) +from .autoencoders.vq_model import VQEncoderOutput, VQModel # noqa: F401 diff --git a/src/diffusers/pipelines/__init__.py b/src/diffusers/pipelines/__init__.py index 32f193a03080..96d10a709d94 100644 --- a/src/diffusers/pipelines/__init__.py +++ b/src/diffusers/pipelines/__init__.py @@ -355,8 +355,8 @@ "JoyImageEditPlusPipeline", "JoyImageEditPlusPipelineOutput", ] - _import_structure["lumina"] = ["LuminaPipeline", "LuminaText2ImgPipeline"] - _import_structure["lumina2"] = ["Lumina2Pipeline", "Lumina2Text2ImgPipeline"] + _import_structure["lumina"] = ["LuminaPipeline"] + _import_structure["lumina2"] = ["Lumina2Pipeline"] _import_structure["lucy"] = ["LucyEditPipeline"] _import_structure["longcat_image"] = ["LongCatImagePipeline", "LongCatImageEditPipeline"] _import_structure["longcat_audio_dit"] = ["LongCatAudioDiTPipeline"] @@ -817,8 +817,8 @@ LTX2VideoDiffusionDecodePipeline, ) from .lucy import LucyEditPipeline - from .lumina import LuminaPipeline, LuminaText2ImgPipeline - from .lumina2 import Lumina2Pipeline, Lumina2Text2ImgPipeline + from .lumina import LuminaPipeline + from .lumina2 import Lumina2Pipeline from .marigold import ( MarigoldDepthPipeline, MarigoldIntrinsicsPipeline, diff --git a/src/diffusers/pipelines/controlnet/__init__.py b/src/diffusers/pipelines/controlnet/__init__.py index 3fb1b8571b43..dcd1926a4f80 100644 --- a/src/diffusers/pipelines/controlnet/__init__.py +++ b/src/diffusers/pipelines/controlnet/__init__.py @@ -21,7 +21,6 @@ _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) else: - _import_structure["multicontrolnet"] = ["MultiControlNetModel"] _import_structure["pipeline_controlnet"] = ["StableDiffusionControlNetPipeline"] _import_structure["pipeline_controlnet_blip_diffusion"] = ["BlipDiffusionControlNetPipeline"] _import_structure["pipeline_controlnet_img2img"] = ["StableDiffusionControlNetImg2ImgPipeline"] @@ -42,7 +41,6 @@ except OptionalDependencyNotAvailable: from ...utils.dummy_torch_and_transformers_objects import * else: - from .multicontrolnet import MultiControlNetModel from .pipeline_controlnet import StableDiffusionControlNetPipeline from .pipeline_controlnet_blip_diffusion import BlipDiffusionControlNetPipeline from .pipeline_controlnet_img2img import StableDiffusionControlNetImg2ImgPipeline diff --git a/src/diffusers/pipelines/controlnet/multicontrolnet.py b/src/diffusers/pipelines/controlnet/multicontrolnet.py deleted file mode 100644 index 6526dd8c9a57..000000000000 --- a/src/diffusers/pipelines/controlnet/multicontrolnet.py +++ /dev/null @@ -1,12 +0,0 @@ -from ...models.controlnets.multicontrolnet import MultiControlNetModel -from ...utils import deprecate, logging - - -logger = logging.get_logger(__name__) - - -class MultiControlNetModel(MultiControlNetModel): - def __init__(self, *args, **kwargs): - deprecation_message = "Importing `MultiControlNetModel` from `diffusers.pipelines.controlnet.multicontrolnet` is deprecated and this will be removed in a future version. Please use `from diffusers.models.controlnets.multicontrolnet import MultiControlNetModel`, instead." - deprecate("diffusers.pipelines.controlnet.multicontrolnet.MultiControlNetModel", "0.34", deprecation_message) - super().__init__(*args, **kwargs) diff --git a/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py b/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py index 959034c89fdd..1bf3ef3699e4 100644 --- a/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py +++ b/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py @@ -30,7 +30,7 @@ from ...loaders import HiDreamImageLoraLoaderMixin from ...models import AutoencoderKL, HiDreamImageTransformer2DModel from ...schedulers import FlowMatchEulerDiscreteScheduler, UniPCMultistepScheduler -from ...utils import deprecate, is_torch_xla_available, logging, replace_example_docstring +from ...utils import is_torch_xla_available, logging, replace_example_docstring from ...utils.torch_utils import randn_tensor from ..pipeline_utils import DiffusionPipeline from .pipeline_output import HiDreamImagePipelineOutput @@ -703,7 +703,6 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 128, - **kwargs, ): r""" Function invoked when calling the pipeline for generation. @@ -807,21 +806,6 @@ def __call__( returning a tuple, the first element is a list with the generated. images. """ - prompt_embeds = kwargs.get("prompt_embeds", None) - negative_prompt_embeds = kwargs.get("negative_prompt_embeds", None) - - if prompt_embeds is not None: - deprecation_message = "The `prompt_embeds` argument is deprecated. Please use `prompt_embeds_t5` and `prompt_embeds_llama3` instead." - deprecate("prompt_embeds", "0.35.0", deprecation_message) - prompt_embeds_t5 = prompt_embeds[0] - prompt_embeds_llama3 = prompt_embeds[1] - - if negative_prompt_embeds is not None: - deprecation_message = "The `negative_prompt_embeds` argument is deprecated. Please use `negative_prompt_embeds_t5` and `negative_prompt_embeds_llama3` instead." - deprecate("negative_prompt_embeds", "0.35.0", deprecation_message) - negative_prompt_embeds_t5 = negative_prompt_embeds[0] - negative_prompt_embeds_llama3 = negative_prompt_embeds[1] - height = height or self.default_sample_size * self.vae_scale_factor width = width or self.default_sample_size * self.vae_scale_factor diff --git a/src/diffusers/pipelines/lumina/__init__.py b/src/diffusers/pipelines/lumina/__init__.py index a19dc7e94641..c9411f9f4783 100644 --- a/src/diffusers/pipelines/lumina/__init__.py +++ b/src/diffusers/pipelines/lumina/__init__.py @@ -22,7 +22,7 @@ _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) else: - _import_structure["pipeline_lumina"] = ["LuminaPipeline", "LuminaText2ImgPipeline"] + _import_structure["pipeline_lumina"] = ["LuminaPipeline"] if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: try: @@ -32,7 +32,7 @@ except OptionalDependencyNotAvailable: from ...utils.dummy_torch_and_transformers_objects import * else: - from .pipeline_lumina import LuminaPipeline, LuminaText2ImgPipeline + from .pipeline_lumina import LuminaPipeline else: import sys diff --git a/src/diffusers/pipelines/lumina/pipeline_lumina.py b/src/diffusers/pipelines/lumina/pipeline_lumina.py index 1cfd9b482d8e..fd6e537a534b 100644 --- a/src/diffusers/pipelines/lumina/pipeline_lumina.py +++ b/src/diffusers/pipelines/lumina/pipeline_lumina.py @@ -30,7 +30,6 @@ from ...schedulers import FlowMatchEulerDiscreteScheduler from ...utils import ( BACKENDS_MAPPING, - deprecate, is_bs4_available, is_ftfy_available, is_torch_xla_available, @@ -936,23 +935,3 @@ def __call__( return (image,) return ImagePipelineOutput(images=image) - - -class LuminaText2ImgPipeline(LuminaPipeline): - def __init__( - self, - transformer: LuminaNextDiT2DModel, - scheduler: FlowMatchEulerDiscreteScheduler, - vae: AutoencoderKL, - text_encoder: GemmaPreTrainedModel, - tokenizer: GemmaTokenizer | GemmaTokenizerFast, - ): - deprecation_message = "`LuminaText2ImgPipeline` has been renamed to `LuminaPipeline` and will be removed in a future version. Please use `LuminaPipeline` instead." - deprecate("diffusers.pipelines.lumina.pipeline_lumina.LuminaText2ImgPipeline", "0.34", deprecation_message) - super().__init__( - transformer=transformer, - scheduler=scheduler, - vae=vae, - text_encoder=text_encoder, - tokenizer=tokenizer, - ) diff --git a/src/diffusers/pipelines/lumina2/__init__.py b/src/diffusers/pipelines/lumina2/__init__.py index b1d6bfeb0d58..300b2f50b5be 100644 --- a/src/diffusers/pipelines/lumina2/__init__.py +++ b/src/diffusers/pipelines/lumina2/__init__.py @@ -22,7 +22,7 @@ _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) else: - _import_structure["pipeline_lumina2"] = ["Lumina2Pipeline", "Lumina2Text2ImgPipeline"] + _import_structure["pipeline_lumina2"] = ["Lumina2Pipeline"] if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: try: @@ -32,7 +32,7 @@ except OptionalDependencyNotAvailable: from ...utils.dummy_torch_and_transformers_objects import * else: - from .pipeline_lumina2 import Lumina2Pipeline, Lumina2Text2ImgPipeline + from .pipeline_lumina2 import Lumina2Pipeline else: import sys diff --git a/src/diffusers/pipelines/lumina2/pipeline_lumina2.py b/src/diffusers/pipelines/lumina2/pipeline_lumina2.py index eb376f9bd9cc..afe0f63e11e4 100644 --- a/src/diffusers/pipelines/lumina2/pipeline_lumina2.py +++ b/src/diffusers/pipelines/lumina2/pipeline_lumina2.py @@ -25,7 +25,6 @@ from ...models.transformers.transformer_lumina2 import Lumina2Transformer2DModel from ...schedulers import FlowMatchEulerDiscreteScheduler from ...utils import ( - deprecate, is_torch_xla_available, logging, replace_example_docstring, @@ -743,23 +742,3 @@ def __call__( return (image,) return ImagePipelineOutput(images=image) - - -class Lumina2Text2ImgPipeline(Lumina2Pipeline): - def __init__( - self, - transformer: Lumina2Transformer2DModel, - scheduler: FlowMatchEulerDiscreteScheduler, - vae: AutoencoderKL, - text_encoder: Gemma2PreTrainedModel, - tokenizer: GemmaTokenizer | GemmaTokenizerFast, - ): - deprecation_message = "`Lumina2Text2ImgPipeline` has been renamed to `Lumina2Pipeline` and will be removed in a future version. Please use `Lumina2Pipeline` instead." - deprecate("diffusers.pipelines.lumina2.pipeline_lumina2.Lumina2Text2ImgPipeline", "0.34", deprecation_message) - super().__init__( - transformer=transformer, - scheduler=scheduler, - vae=vae, - text_encoder=text_encoder, - tokenizer=tokenizer, - ) diff --git a/src/diffusers/utils/dummy_torch_and_transformers_objects.py b/src/diffusers/utils/dummy_torch_and_transformers_objects.py index ed724e7de751..1a5e472c8fa1 100644 --- a/src/diffusers/utils/dummy_torch_and_transformers_objects.py +++ b/src/diffusers/utils/dummy_torch_and_transformers_objects.py @@ -3392,21 +3392,6 @@ def from_pretrained(cls, *args, **kwargs): requires_backends(cls, ["torch", "transformers"]) -class Lumina2Text2ImgPipeline(metaclass=DummyObject): - _backends = ["torch", "transformers"] - - def __init__(self, *args, **kwargs): - requires_backends(self, ["torch", "transformers"]) - - @classmethod - def from_config(cls, *args, **kwargs): - requires_backends(cls, ["torch", "transformers"]) - - @classmethod - def from_pretrained(cls, *args, **kwargs): - requires_backends(cls, ["torch", "transformers"]) - - class LuminaPipeline(metaclass=DummyObject): _backends = ["torch", "transformers"] @@ -3422,21 +3407,6 @@ def from_pretrained(cls, *args, **kwargs): requires_backends(cls, ["torch", "transformers"]) -class LuminaText2ImgPipeline(metaclass=DummyObject): - _backends = ["torch", "transformers"] - - def __init__(self, *args, **kwargs): - requires_backends(self, ["torch", "transformers"]) - - @classmethod - def from_config(cls, *args, **kwargs): - requires_backends(cls, ["torch", "transformers"]) - - @classmethod - def from_pretrained(cls, *args, **kwargs): - requires_backends(cls, ["torch", "transformers"]) - - class MarigoldDepthPipeline(metaclass=DummyObject): _backends = ["torch", "transformers"] From 78405ec9a1c1637c4027136658c96b7e85f915e5 Mon Sep 17 00:00:00 2001 From: yzhautouskay Date: Wed, 30 Sep 2026 17:53:20 +0200 Subject: [PATCH 4/5] [Cosmos3] Fix Transfer SeaCache artifacts with control CFG (#14897) * Fix Cosmos 3 Transfer SeaCache indicators * Clarify SeaCache indicator guidance; tests refactor --------- Co-authored-by: Sayak Paul --- docs/source/en/optimization/cache.md | 15 ++- src/diffusers/hooks/sea_cache.py | 14 ++- .../test_models_transformer_cosmos3.py | 107 +++++++++++++++++- 3 files changed, 121 insertions(+), 15 deletions(-) diff --git a/docs/source/en/optimization/cache.md b/docs/source/en/optimization/cache.md index 9f775ec3b88c..3a771336e4b7 100644 --- a/docs/source/en/optimization/cache.md +++ b/docs/source/en/optimization/cache.md @@ -72,8 +72,14 @@ pipeline.transformer.enable_cache(config) [SeaCache](https://huggingface.co/papers/2602.18993) compares Spectral-Evolution-Aware (SEA) indicators between successive denoising steps. When the accumulated indicator change remains below a threshold, it skips the expensive -transformer block stack and predicts its output from cached residuals. The indicator is computed from the raw vision -latents, including clean conditioning frames for image-to-video generation. +transformer block stack and predicts its output from cached residuals. Build the indicator from the visual latents that +form the generated output. Include clean conditioning frames when they are part of that output trajectory, as in +image-to-video and video-to-video generation. Exclude separate visual hints that condition the generation but are not +part of the output. Text conditioning is excluded because it is not a visual latent. + +Cosmos 3 Transfer packs control hints as separate visual sequences, so its adapter excludes them from the indicator. +Control-CFG branches compare the same output trajectory while retaining their own cached residuals. Control hints still +condition the transformer. The implementation provides built-in adapters for the following models: @@ -87,8 +93,9 @@ Other video transformers can integrate with the generic path when they use `Cach block list, and register the block input/output layout in `TransformerBlockRegistry`. The pipeline must enter a `cache_context` for every transformer call, attach `step_index`, `sigma`, and `num_inference_steps`, and use separate context names for independent trajectories such as conditional and unconditional guidance. Pass a `raw_vision_callback` -that returns the noisy vision latents when no built-in adapter is available. Validate output quality and tune the cache -parameters for each model and scheduler; support and benchmark results do not transfer automatically from Cosmos 3. +that returns the visual latents forming the generated output when no built-in adapter is available. Validate output +quality and tune the cache parameters for each model and scheduler; support and benchmark results do not transfer +automatically from Cosmos 3. ### Cosmos 3 diff --git a/src/diffusers/hooks/sea_cache.py b/src/diffusers/hooks/sea_cache.py index 5cb78db1f6b4..d228b0772e48 100644 --- a/src/diffusers/hooks/sea_cache.py +++ b/src/diffusers/hooks/sea_cache.py @@ -64,8 +64,9 @@ class SeaCacheConfig: power_exp (`float`, defaults to `3.0`): Exponent of the SEA clean-signal power prior. SeaCache uses `3.0` for video features. raw_vision_callback (`Callable`, *optional*): - Advanced model adapter returning raw vision latents with shape `(C, T, H, W)`. When omitted, a built-in - adapter is used if one is available. + Advanced model adapter returning the visual latents forming the generated output, each with shape `(C, T, + H, W)`. Include clean conditioning frames within the output trajectory, but exclude separate visual hints + that are not part of the output. When omitted, a built-in adapter is used if one is available. Example: ```python @@ -326,7 +327,6 @@ def _prepare_cosmos3_raw_vision_metadata( return None raw_vision = [] - has_noisy_vision = False for latent, noisy_frame_indexes in zip(vision_tokens, vision_noisy_frame_indexes): if not isinstance(latent, torch.Tensor) or not isinstance(noisy_frame_indexes, torch.Tensor): return None @@ -340,10 +340,12 @@ def _prepare_cosmos3_raw_vision_metadata( noisy_frame_indexes = noisy_frame_indexes.flatten().to(device=latent.device, dtype=torch.long) if torch.any(noisy_frame_indexes < 0) or torch.any(noisy_frame_indexes >= latent.shape[1]): return None - has_noisy_vision = has_noisy_vision or noisy_frame_indexes.numel() > 0 - raw_vision.append(latent) + # A sequence with noisy frames belongs to the generated output. Keep that sequence whole so clean conditioning + # frames remain in the indicator, but exclude separate clean hints that are not part of the output. + if noisy_frame_indexes.numel() > 0: + raw_vision.append(latent) - return raw_vision if raw_vision and has_noisy_vision else None + return raw_vision or None def _prepare_wan_t2v_raw_vision_metadata( diff --git a/tests/models/transformers/test_models_transformer_cosmos3.py b/tests/models/transformers/test_models_transformer_cosmos3.py index 6b04dd77f8c6..c897d5bfdc69 100644 --- a/tests/models/transformers/test_models_transformer_cosmos3.py +++ b/tests/models/transformers/test_models_transformer_cosmos3.py @@ -108,7 +108,106 @@ def output_shape(self) -> tuple[int, ...]: return (1, 2, 1, 1, 1) -class TestCosmos3OmniTransformerModel(Cosmos3OmniTransformerTesterConfig, ModelTesterMixin): +class TestCosmos3OmniTransformerSeaCache(Cosmos3OmniTransformerTesterConfig, SeaCacheTesterMixin): + cache_input_key = "vision_tokens" + + def test_sea_cache_tracks_output_visual_trajectory(self): + model = self.model_class(**self.get_init_dict()).to(torch_device).eval() + model.enable_cache(SeaCacheConfig(threshold=2.0, cache_end_steps=0)) + target = torch.randn(1, 2, 2, 1, 1, device=torch_device) + control = torch.randn_like(target) + target_with_changed_clean_frame = target.clone() + target_with_changed_clean_frame[:, :, 0] += 10 + inputs = self.get_dummy_inputs() + inputs.update( + sequence_length=6, + position_ids=torch.zeros(3, 6, dtype=torch.long, device=torch_device), + vision_tokens=[control, target], + vision_token_shapes=[(2, 1, 1)] * 2, + vision_sequence_indexes=torch.arange(2, 6, device=torch_device), + vision_mse_loss_indexes=torch.tensor([5], device=torch_device), + vision_noisy_frame_indexes=[ + torch.tensor([], dtype=torch.long, device=torch_device), + torch.tensor([1], device=torch_device), + ], + ) + layer_calls = 0 + + def count_layer_calls(_module, _args, _output): + nonlocal layer_calls + layer_calls += 1 + + model.layers[0].register_forward_hook(count_layer_calls) + decisions = [] + for step, (current_control, current_target) in enumerate( + ( + (control, target), + (control + 100, target), + (control + 100, target_with_changed_clean_frame), + ) + ): + inputs["vision_tokens"] = [current_control, current_target] + with ( + torch.no_grad(), + model.cache_context("cond", step_index=step, sigma=0.9 - step * 0.3, num_inference_steps=3), + ): + model(**inputs) + state = model._diffusers_hook.get_hook(_SEA_CACHE_ROOT_HOOK).state_manager._state_cache["cond"] + decisions.append(state.gate_should_compute) + + assert decisions == [True, False, True] + assert layer_calls == 2 + + @pytest.mark.parametrize("residual_order", [0, 1]) + def test_sea_cache_transfer_branches_share_indicator_with_separate_histories(self, residual_order): + model = self.model_class(**self.get_init_dict()).to(torch_device).eval() + model.enable_cache(SeaCacheConfig(threshold=100.0, residual_order=residual_order, cache_end_steps=0)) + root_hook = model._diffusers_hook.get_hook(_SEA_CACHE_ROOT_HOOK) + target = torch.randn(1, 2, 2, 1, 1, device=torch_device) + control = torch.randn_like(target) + decisions = [] + + for step in range(6): + states = [] + for context, with_control in (("cond", True), ("cond_no_control", False), ("uncond", True)): + inputs = self.get_dummy_inputs() + sequence_length = 6 if with_control else 4 + inputs.update( + sequence_length=sequence_length, + position_ids=torch.zeros(3, sequence_length, dtype=torch.long, device=torch_device), + vision_tokens=[control, target] if with_control else [target], + vision_token_shapes=[(2, 1, 1)] * (2 if with_control else 1), + vision_sequence_indexes=torch.arange(2, sequence_length, device=torch_device), + vision_mse_loss_indexes=torch.tensor([sequence_length - 1], device=torch_device), + vision_noisy_frame_indexes=( + [ + torch.tensor([], dtype=torch.long, device=torch_device), + torch.tensor([1], device=torch_device), + ] + if with_control + else [torch.tensor([1], device=torch_device)] + ), + ) + with ( + torch.no_grad(), + model.cache_context(context, step_index=step, sigma=0.9 - step * 0.1, num_inference_steps=6), + ): + output = model(**inputs) + assert torch.isfinite(output.sample[-1]).all() + states.append(root_hook.state_manager._state_cache[context]) + + assert all(len(state.previous_indicator) == 1 for state in states) + for state in states[1:]: + torch.testing.assert_close(state.previous_indicator[0], states[0].previous_indicator[0]) + assert state.gate_should_compute == states[0].gate_should_compute + assert state.history is not states[0].history + assert states[0].history[-1][2].shape != states[1].history[-1][2].shape + decisions.append(states[0].gate_should_compute) + + assert any(decisions) and not all(decisions) + model._reset_stateful_cache() + assert all(not state.history and state.previous_indicator is None for state in states) + def test_cosmos3_supports_sea_cache_without_changing_state_dict_keys(self): model = self.model_class(**self.get_init_dict()).to(torch_device).eval() state_dict_keys = set(model.state_dict()) @@ -423,6 +522,8 @@ def test_cosmos3_sea_cache_regional_compile_fullgraph_without_recompile(self): assert refreshed.sample[0].shape == self.output_shape assert root_hook.state_manager._state_cache["cond"].history[-1][0] == 2 + +class TestCosmos3OmniTransformerModel(Cosmos3OmniTransformerTesterConfig, ModelTesterMixin): def test_cosmos3_decoder_layer_cache_metadata_tracks_generation_stream(self): metadata = TransformerBlockRegistry.get(Cosmos3VLTextMoTDecoderLayer) @@ -571,10 +672,6 @@ def test_cosmos3_nemotron_rms_norm_multiplies_in_float32(self): torch.testing.assert_close(norm(hidden_states), expected, rtol=0, atol=0) -class TestCosmos3OmniTransformerSeaCache(Cosmos3OmniTransformerTesterConfig, SeaCacheTesterMixin): - cache_input_key = "vision_tokens" - - class TestCosmos3OmniTransformerMemory(Cosmos3OmniTransformerTesterConfig, MemoryTesterMixin): @pytest.mark.skip("The transformer returns one tensor list per generated modality.") def test_layerwise_casting_training(self): From c60830ee365d520ab52b110dda562dd26f7b4d7f Mon Sep 17 00:00:00 2001 From: Steven Liu <59462357+stevhliu@users.noreply.github.com> Date: Wed, 30 Sep 2026 08:59:23 -0700 Subject: [PATCH 5/5] [fix] Add return types (#14874) * add return types * add return blocks * style * modular pipelines * add a test to validate * feedback - skip deprecated stuff --- .github/workflows/pr_modular_tests.yml | 1 + .github/workflows/pr_tests.yml | 1 + .github/workflows/pr_tests_gpu.yml | 1 + Makefile | 5 + src/diffusers/models/activations.py | 8 +- src/diffusers/models/attention.py | 4 +- src/diffusers/models/attention_dispatch.py | 10 +- src/diffusers/models/attention_processor.py | 2 +- .../autoencoders/autoencoder_cosmos3_audio.py | 8 +- .../autoencoder_kl_hunyuanimage.py | 4 +- .../autoencoder_kl_hunyuanimage_refiner.py | 6 +- .../autoencoder_kl_hunyuanvideo15.py | 6 +- .../autoencoders/autoencoder_kl_qwenimage.py | 20 +- .../autoencoder_kl_qwenimage21.py | 24 +-- .../models/autoencoders/autoencoder_kl_wan.py | 24 +-- .../autoencoders/autoencoder_oobleck.py | 12 +- .../models/controlnets/controlnet.py | 2 +- .../models/controlnets/controlnet_hunyuan.py | 13 +- .../models/controlnets/controlnet_union.py | 4 +- .../models/controlnets/controlnet_xs.py | 2 +- .../models/controlnets/controlnet_z_image.py | 14 +- src/diffusers/models/embeddings.py | 58 +++--- src/diffusers/models/lora.py | 2 +- src/diffusers/models/normalization.py | 8 +- src/diffusers/models/resnet.py | 2 +- .../models/transformers/dit_transformer_2d.py | 2 +- .../transformers/dual_transformer_2d.py | 3 +- .../transformers/hunyuan_transformer_2d.py | 6 +- .../transformers/latte_transformer_3d.py | 2 +- .../transformers/pixart_transformer_2d.py | 2 +- .../models/transformers/prior_transformer.py | 2 +- .../models/transformers/sana_transformer.py | 4 +- .../transformers/stable_audio_transformer.py | 2 +- .../transformers/t5_film_transformer.py | 5 +- .../models/transformers/transformer_2d.py | 2 +- .../transformers/transformer_2d_dreamlite.py | 2 +- .../transformers/transformer_allegro.py | 2 +- .../transformers/transformer_anyflow.py | 4 +- .../transformers/transformer_anyflow_far.py | 4 +- .../models/transformers/transformer_bria.py | 4 +- .../transformers/transformer_bria_fibo.py | 6 +- .../models/transformers/transformer_chroma.py | 2 +- .../transformers/transformer_chronoedit.py | 2 +- .../transformers/transformer_cosmos3.py | 4 +- .../transformers/transformer_ernie_image.py | 9 +- .../models/transformers/transformer_helios.py | 6 +- .../transformers/transformer_hidream_image.py | 6 +- .../transformer_hunyuan_video_framepack.py | 6 +- .../transformers/transformer_joyimage.py | 8 +- .../transformer_joyimage_edit_plus.py | 2 +- .../transformers/transformer_kandinsky.py | 20 +- .../transformers/transformer_longcat_image.py | 2 +- .../transformers/transformer_lumina2.py | 4 +- .../models/transformers/transformer_mochi.py | 2 +- .../transformer_nucleusmoe_image.py | 2 +- .../transformers/transformer_omnigen.py | 2 +- .../transformers/transformer_qwenimage.py | 2 +- .../transformers/transformer_sana_video.py | 4 +- .../models/transformers/transformer_sd3.py | 2 +- .../transformers/transformer_skyreels_v2.py | 2 +- .../transformers/transformer_temporal.py | 2 +- .../models/transformers/transformer_wan.py | 2 +- .../transformers/transformer_wan_animate.py | 2 +- .../transformers/transformer_wan_animate_2.py | 4 +- .../transformers/transformer_z_image.py | 14 +- src/diffusers/models/unets/unet_3d_blocks.py | 2 +- src/diffusers/models/unets/unet_kandinsky3.py | 24 ++- .../models/unets/unet_motion_model.py | 4 +- .../models/unets/unet_stable_cascade.py | 16 +- src/diffusers/models/unets/uvit_2d.py | 16 +- .../modular_pipelines/anima/before_denoise.py | 28 ++- .../modular_pipelines/anima/decoders.py | 8 +- .../modular_pipelines/anima/denoise.py | 14 +- .../modular_pipelines/anima/encoders.py | 8 +- .../modular_pipelines/cosmos/after_decode.py | 8 +- .../cosmos/before_denoise.py | 60 ++++-- .../cosmos/before_encoder.py | 4 +- .../modular_pipelines/cosmos/decoders.py | 16 +- .../modular_pipelines/cosmos/denoise.py | 48 +++-- .../modular_pipelines/cosmos/encoders.py | 32 +++- .../cosmos/modular_blocks_cosmos3.py | 4 +- .../ernie_image/before_denoise.py | 12 +- .../modular_pipelines/ernie_image/decoders.py | 4 +- .../modular_pipelines/ernie_image/denoise.py | 16 +- .../modular_pipelines/ernie_image/encoders.py | 8 +- .../modular_pipelines/flux/before_denoise.py | 24 ++- .../modular_pipelines/flux/decoders.py | 3 +- .../modular_pipelines/flux/denoise.py | 12 +- .../modular_pipelines/flux/encoders.py | 16 +- .../modular_pipelines/flux/inputs.py | 16 +- .../modular_pipelines/flux2/before_denoise.py | 24 ++- .../modular_pipelines/flux2/decoders.py | 5 +- .../modular_pipelines/flux2/denoise.py | 14 +- .../modular_pipelines/flux2/encoders.py | 20 +- .../modular_pipelines/flux2/inputs.py | 12 +- .../helios/before_denoise.py | 32 +++- .../modular_pipelines/helios/decoders.py | 3 +- .../modular_pipelines/helios/denoise.py | 40 +++- .../modular_pipelines/helios/encoders.py | 12 +- .../hunyuan_video1_5/before_denoise.py | 16 +- .../hunyuan_video1_5/decoders.py | 3 +- .../hunyuan_video1_5/denoise.py | 16 +- .../hunyuan_video1_5/encoders.py | 12 +- .../ideogram4/before_denoise.py | 16 +- .../modular_pipelines/ideogram4/decoders.py | 4 +- .../modular_pipelines/ideogram4/denoise.py | 20 +- .../modular_pipelines/ideogram4/encoders.py | 8 +- .../modular_pipelines/krea2/before_denoise.py | 24 ++- .../modular_pipelines/krea2/decoders.py | 4 +- .../modular_pipelines/krea2/denoise.py | 20 +- .../modular_pipelines/krea2/encoders.py | 8 +- .../modular_pipelines/ltx/before_denoise.py | 16 +- .../modular_pipelines/ltx/decoders.py | 4 +- .../modular_pipelines/ltx/denoise.py | 24 ++- .../modular_pipelines/ltx/encoders.py | 8 +- .../modular_pipelines/ltx2/before_denoise.py | 25 +-- .../modular_pipelines/ltx2/decoders.py | 9 +- .../modular_pipelines/ltx2/denoise.py | 31 ++- .../modular_pipelines/ltx2/encoders.py | 19 +- .../minimax_h3/before_denoise.py | 32 +++- .../minimax_h3/before_encoder.py | 8 +- .../modular_pipelines/minimax_h3/decoders.py | 12 +- .../modular_pipelines/minimax_h3/denoise.py | 12 +- .../modular_pipelines/minimax_h3/encoders.py | 20 +- .../minimax_music3/before_denoise.py | 4 +- .../minimax_music3/decoders.py | 4 +- .../minimax_music3/denoise.py | 24 ++- .../minimax_music3/encoders.py | 8 +- .../modular_pipelines/modular_pipeline.py | 6 +- .../qwenimage/before_denoise.py | 44 +++-- .../modular_pipelines/qwenimage/decoders.py | 20 +- .../modular_pipelines/qwenimage/denoise.py | 32 +++- .../modular_pipelines/qwenimage/encoders.py | 58 ++++-- .../modular_pipelines/qwenimage/inputs.py | 20 +- .../stable_diffusion_3/before_denoise.py | 16 +- .../stable_diffusion_3/decoders.py | 3 +- .../stable_diffusion_3/denoise.py | 8 +- .../stable_diffusion_3/encoders.py | 12 +- .../stable_diffusion_3/inputs.py | 8 +- .../stable_diffusion_xl/before_denoise.py | 40 +++- .../stable_diffusion_xl/decoders.py | 5 +- .../stable_diffusion_xl/denoise.py | 26 ++- .../stable_diffusion_xl/encoders.py | 16 +- .../modular_pipelines/wan/before_denoise.py | 20 +- .../modular_pipelines/wan/decoders.py | 5 +- .../modular_pipelines/wan/denoise.py | 20 +- .../modular_pipelines/wan/encoders.py | 40 +++- .../wan_animate_2/before_denoise.py | 3 +- .../wan_animate_2/decoders.py | 3 +- .../wan_animate_2/denoise.py | 17 +- .../wan_animate_2/encoders.py | 13 +- .../z_image/before_denoise.py | 24 ++- .../modular_pipelines/z_image/decoders.py | 3 +- .../modular_pipelines/z_image/denoise.py | 14 +- .../modular_pipelines/z_image/encoders.py | 8 +- .../pipelines/ace_step/pipeline_ace_step.py | 2 +- .../animatediff/pipeline_animatediff.py | 2 +- .../pipeline_animatediff_controlnet.py | 2 +- .../animatediff/pipeline_animatediff_sdxl.py | 2 +- .../pipeline_animatediff_sparsectrl.py | 2 +- .../pipeline_animatediff_video2video.py | 2 +- ...line_animatediff_video2video_controlnet.py | 2 +- .../pipelines/anyflow/pipeline_anyflow.py | 2 +- .../pipelines/anyflow/pipeline_anyflow_far.py | 2 +- .../pipelines/audioldm2/pipeline_audioldm2.py | 2 +- src/diffusers/pipelines/bria/pipeline_bria.py | 2 +- .../pipelines/bria_fibo/pipeline_bria_fibo.py | 2 +- .../bria_fibo/pipeline_bria_fibo_edit.py | 2 +- .../pipelines/chroma/pipeline_chroma.py | 2 +- .../chroma/pipeline_chroma_img2img.py | 2 +- .../chroma/pipeline_chroma_inpainting.py | 2 +- .../chronoedit/pipeline_chronoedit.py | 2 +- .../pipeline_consistency_models.py | 2 +- .../controlnet/pipeline_controlnet.py | 2 +- .../pipeline_controlnet_blip_diffusion.py | 2 +- .../controlnet/pipeline_controlnet_img2img.py | 2 +- .../controlnet/pipeline_controlnet_inpaint.py | 2 +- .../pipeline_controlnet_inpaint_sd_xl.py | 2 +- .../controlnet/pipeline_controlnet_sd_xl.py | 2 +- .../pipeline_controlnet_sd_xl_img2img.py | 2 +- ...pipeline_controlnet_union_inpaint_sd_xl.py | 2 +- .../pipeline_controlnet_union_sd_xl.py | 2 +- ...pipeline_controlnet_union_sd_xl_img2img.py | 2 +- .../pipeline_hunyuandit_controlnet.py | 2 +- .../pipeline_stable_diffusion_3_controlnet.py | 2 +- ...table_diffusion_3_controlnet_inpainting.py | 2 +- .../cosmos/pipeline_cosmos2_5_predict.py | 2 +- .../cosmos/pipeline_cosmos2_5_transfer.py | 2 +- .../cosmos/pipeline_cosmos2_text2image.py | 2 +- .../cosmos/pipeline_cosmos2_video2world.py | 2 +- .../pipelines/cosmos/pipeline_cosmos3_omni.py | 2 +- .../cosmos/pipeline_cosmos_text2world.py | 2 +- .../cosmos/pipeline_cosmos_video2world.py | 2 +- .../pipelines/deepfloyd_if/pipeline_if.py | 2 +- .../deepfloyd_if/pipeline_if_img2img.py | 2 +- .../pipeline_if_img2img_superresolution.py | 2 +- .../deepfloyd_if/pipeline_if_inpainting.py | 2 +- .../pipeline_if_inpainting_superresolution.py | 2 +- .../pipeline_if_superresolution.py | 2 +- .../pipeline_diffusion_gemma.py | 2 +- .../pipelines/dreamlite/pipeline_dreamlite.py | 2 +- .../dreamlite/pipeline_dreamlite_mobile.py | 2 +- .../easyanimate/pipeline_easyanimate.py | 2 +- .../pipeline_easyanimate_control.py | 2 +- .../pipeline_easyanimate_inpaint.py | 2 +- .../ernie_image/pipeline_ernie_image.py | 2 +- src/diffusers/pipelines/flux/pipeline_flux.py | 2 +- .../pipelines/flux/pipeline_flux_control.py | 2 +- .../flux/pipeline_flux_control_img2img.py | 2 +- .../flux/pipeline_flux_control_inpaint.py | 2 +- .../flux/pipeline_flux_controlnet.py | 2 +- ...pipeline_flux_controlnet_image_to_image.py | 2 +- .../pipeline_flux_controlnet_inpainting.py | 2 +- .../pipelines/flux/pipeline_flux_fill.py | 2 +- .../pipelines/flux/pipeline_flux_img2img.py | 2 +- .../pipelines/flux/pipeline_flux_inpaint.py | 2 +- .../pipelines/flux/pipeline_flux_kontext.py | 2 +- .../flux/pipeline_flux_kontext_inpaint.py | 2 +- .../flux/pipeline_flux_prior_redux.py | 2 +- .../pipelines/flux2/pipeline_flux2.py | 2 +- .../pipelines/flux2/pipeline_flux2_klein.py | 2 +- .../flux2/pipeline_flux2_klein_inpaint.py | 2 +- .../flux2/pipeline_flux2_klein_kv.py | 2 +- .../pipelines/helios/pipeline_helios.py | 2 +- .../helios/pipeline_helios_pyramid.py | 2 +- .../hidream_image/pipeline_hidream_image.py | 3 +- .../hunyuan_image/pipeline_hunyuanimage.py | 2 +- .../pipeline_hunyuanimage_refiner.py | 2 +- .../pipeline_hunyuan_skyreels_image2video.py | 2 +- .../hunyuan_video/pipeline_hunyuan_video.py | 2 +- .../pipeline_hunyuan_video_framepack.py | 2 +- .../pipeline_hunyuan_video_image2video.py | 2 +- .../pipeline_hunyuan_video1_5.py | 2 +- .../pipeline_hunyuan_video1_5_image2video.py | 2 +- .../hunyuandit/pipeline_hunyuandit.py | 2 +- .../pipelines/ideogram4/pipeline_ideogram4.py | 2 +- .../joyimage/pipeline_joyimage_edit.py | 2 +- .../joyimage/pipeline_joyimage_edit_plus.py | 2 +- .../pipelines/kandinsky/pipeline_kandinsky.py | 2 +- .../kandinsky/pipeline_kandinsky_combined.py | 8 +- .../kandinsky/pipeline_kandinsky_img2img.py | 2 +- .../kandinsky/pipeline_kandinsky_inpaint.py | 2 +- .../kandinsky/pipeline_kandinsky_prior.py | 2 +- .../kandinsky2_2/pipeline_kandinsky2_2.py | 2 +- .../pipeline_kandinsky2_2_combined.py | 8 +- .../pipeline_kandinsky2_2_controlnet.py | 2 +- ...ipeline_kandinsky2_2_controlnet_img2img.py | 2 +- .../pipeline_kandinsky2_2_img2img.py | 2 +- .../pipeline_kandinsky2_2_inpainting.py | 2 +- .../pipeline_kandinsky2_2_prior.py | 2 +- .../pipeline_kandinsky2_2_prior_emb2emb.py | 2 +- .../kandinsky3/pipeline_kandinsky3.py | 2 +- .../kandinsky3/pipeline_kandinsky3_img2img.py | 2 +- .../kandinsky5/pipeline_kandinsky.py | 2 +- .../kandinsky5/pipeline_kandinsky_i2i.py | 2 +- .../kandinsky5/pipeline_kandinsky_i2v.py | 2 +- .../kandinsky5/pipeline_kandinsky_t2i.py | 2 +- .../pipelines/kolors/pipeline_kolors.py | 2 +- .../kolors/pipeline_kolors_img2img.py | 2 +- .../pipelines/krea2/pipeline_krea2.py | 2 +- .../pipeline_latent_consistency_img2img.py | 2 +- .../pipeline_latent_consistency_text2img.py | 2 +- .../pipeline_latent_diffusion.py | 2 +- ...peline_latent_diffusion_superresolution.py | 2 +- .../pipeline_leditspp_stable_diffusion.py | 2 +- .../pipeline_leditspp_stable_diffusion_xl.py | 2 +- .../pipelines/llada2/pipeline_llada2.py | 2 +- .../pipeline_longcat_audio_dit.py | 6 +- .../longcat_image/pipeline_longcat_image.py | 2 +- .../pipeline_longcat_image_edit.py | 2 +- src/diffusers/pipelines/ltx/pipeline_ltx.py | 2 +- .../pipelines/ltx/pipeline_ltx_condition.py | 2 +- .../ltx/pipeline_ltx_i2v_long_multi_prompt.py | 2 +- .../pipelines/ltx/pipeline_ltx_image2video.py | 2 +- .../ltx/pipeline_ltx_latent_upsample.py | 7 +- src/diffusers/pipelines/ltx2/pipeline_ltx2.py | 2 +- .../pipelines/ltx2/pipeline_ltx2_condition.py | 2 +- .../pipelines/ltx2/pipeline_ltx2_dfr.py | 2 +- .../ltx2/pipeline_ltx2_dfr_temporal_refine.py | 2 +- .../ltx2/pipeline_ltx2_diffusion_decode.py | 2 +- .../pipelines/ltx2/pipeline_ltx2_hdr_lora.py | 2 +- .../pipelines/ltx2/pipeline_ltx2_ic_lora.py | 2 +- .../ltx2/pipeline_ltx2_image2video.py | 2 +- .../ltx2/pipeline_ltx2_latent_upsample.py | 2 +- .../pipelines/lucy/pipeline_lucy_edit.py | 2 +- .../marigold/pipeline_marigold_depth.py | 2 +- .../marigold/pipeline_marigold_intrinsics.py | 2 +- .../marigold/pipeline_marigold_normals.py | 2 +- .../pipelines/mochi/pipeline_mochi.py | 2 +- .../motif_video/pipeline_motif_video.py | 2 +- .../pipeline_motif_video_image2video.py | 2 +- .../pipeline_nucleusmoe_image.py | 2 +- .../pipelines/omnigen/pipeline_omnigen.py | 6 +- .../ovis_image/pipeline_ovis_image.py | 2 +- .../pag/pipeline_pag_controlnet_sd.py | 2 +- .../pag/pipeline_pag_controlnet_sd_inpaint.py | 2 +- .../pag/pipeline_pag_controlnet_sd_xl.py | 2 +- .../pipeline_pag_controlnet_sd_xl_img2img.py | 2 +- .../pipelines/pag/pipeline_pag_hunyuandit.py | 2 +- .../pipelines/pag/pipeline_pag_kolors.py | 2 +- .../pipelines/pag/pipeline_pag_sd.py | 2 +- .../pipelines/pag/pipeline_pag_sd_3.py | 2 +- .../pag/pipeline_pag_sd_3_img2img.py | 2 +- .../pag/pipeline_pag_sd_animatediff.py | 2 +- .../pipelines/pag/pipeline_pag_sd_img2img.py | 2 +- .../pipelines/pag/pipeline_pag_sd_inpaint.py | 2 +- .../pipelines/pag/pipeline_pag_sd_xl.py | 2 +- .../pag/pipeline_pag_sd_xl_img2img.py | 2 +- .../pag/pipeline_pag_sd_xl_inpaint.py | 2 +- src/diffusers/pipelines/prx/pipeline_prx.py | 2 +- .../pipelines/prx/pipeline_prx_pixel.py | 2 +- .../pipelines/qwenimage/pipeline_qwenimage.py | 2 +- .../pipeline_qwenimage_controlnet.py | 2 +- .../pipeline_qwenimage_controlnet_inpaint.py | 2 +- .../qwenimage/pipeline_qwenimage_edit.py | 2 +- .../pipeline_qwenimage_edit_inpaint.py | 2 +- .../qwenimage/pipeline_qwenimage_edit_plus.py | 2 +- .../qwenimage/pipeline_qwenimage_img2img.py | 2 +- .../qwenimage/pipeline_qwenimage_inpaint.py | 2 +- .../qwenimage/pipeline_qwenimage_layered.py | 2 +- .../qwenimage21/pipeline_qwenimage21.py | 2 +- .../pipelines/shap_e/pipeline_shap_e.py | 2 +- .../shap_e/pipeline_shap_e_img2img.py | 2 +- .../skyreels_v2/pipeline_skyreels_v2.py | 2 +- .../pipeline_skyreels_v2_diffusion_forcing.py | 2 +- ...eline_skyreels_v2_diffusion_forcing_i2v.py | 2 +- ...eline_skyreels_v2_diffusion_forcing_v2v.py | 2 +- .../skyreels_v2/pipeline_skyreels_v2_i2v.py | 2 +- .../stable_audio/pipeline_stable_audio.py | 2 +- .../stable_audio_3/pipeline_stable_audio_3.py | 2 +- .../pipeline_stable_audio_3_audio2audio.py | 2 +- .../pipeline_stable_audio_3_inpaint.py | 2 +- .../stable_cascade/pipeline_stable_cascade.py | 4 +- .../pipeline_stable_cascade_combined.py | 5 +- .../pipeline_stable_cascade_prior.py | 2 +- .../pipeline_onnx_stable_diffusion.py | 2 +- .../pipeline_onnx_stable_diffusion_img2img.py | 2 +- .../pipeline_onnx_stable_diffusion_inpaint.py | 2 +- .../pipeline_onnx_stable_diffusion_upscale.py | 2 +- .../pipeline_stable_diffusion.py | 2 +- .../pipeline_stable_diffusion_depth2img.py | 2 +- ...peline_stable_diffusion_image_variation.py | 2 +- .../pipeline_stable_diffusion_img2img.py | 2 +- .../pipeline_stable_diffusion_inpaint.py | 2 +- ...eline_stable_diffusion_instruct_pix2pix.py | 2 +- ...ipeline_stable_diffusion_latent_upscale.py | 2 +- .../pipeline_stable_diffusion_upscale.py | 2 +- .../pipeline_stable_unclip.py | 2 +- .../pipeline_stable_unclip_img2img.py | 2 +- .../pipeline_stable_diffusion_3.py | 2 +- .../pipeline_stable_diffusion_3_img2img.py | 2 +- .../pipeline_stable_diffusion_3_inpaint.py | 2 +- .../pipeline_stable_diffusion_xl.py | 2 +- .../pipeline_stable_diffusion_xl_img2img.py | 2 +- .../pipeline_stable_diffusion_xl_inpaint.py | 2 +- ...ne_stable_diffusion_xl_instruct_pix2pix.py | 2 +- .../pipeline_stable_video_diffusion.py | 2 +- .../pipeline_stable_diffusion_adapter.py | 2 +- .../pipeline_stable_diffusion_xl_adapter.py | 2 +- .../pipeline_visualcloze_combined.py | 2 +- .../pipeline_visualcloze_generation.py | 2 +- src/diffusers/pipelines/wan/pipeline_wan.py | 2 +- .../pipelines/wan/pipeline_wan_animate.py | 2 +- .../pipelines/wan/pipeline_wan_i2v.py | 2 +- .../pipelines/wan/pipeline_wan_vace.py | 2 +- .../pipelines/wan/pipeline_wan_video2video.py | 2 +- .../pipelines/z_image/pipeline_z_image.py | 2 +- .../z_image/pipeline_z_image_controlnet.py | 2 +- .../pipeline_z_image_controlnet_inpaint.py | 2 +- .../z_image/pipeline_z_image_img2img.py | 2 +- .../z_image/pipeline_z_image_inpaint.py | 2 +- .../z_image/pipeline_z_image_omni.py | 2 +- utils/check_return_annotations.py | 180 ++++++++++++++++++ 373 files changed, 1700 insertions(+), 809 deletions(-) create mode 100644 utils/check_return_annotations.py diff --git a/.github/workflows/pr_modular_tests.yml b/.github/workflows/pr_modular_tests.yml index f4bf9585bb7d..e6f156235a53 100644 --- a/.github/workflows/pr_modular_tests.yml +++ b/.github/workflows/pr_modular_tests.yml @@ -83,6 +83,7 @@ jobs: python utils/check_dummies.py python utils/check_support_list.py python utils/check_forward_call_docstrings.py + python utils/check_return_annotations.py make deps_table_check_updated - name: Check if failure if: ${{ failure() }} diff --git a/.github/workflows/pr_tests.yml b/.github/workflows/pr_tests.yml index 2034673af942..b5179e82a028 100644 --- a/.github/workflows/pr_tests.yml +++ b/.github/workflows/pr_tests.yml @@ -78,6 +78,7 @@ jobs: python utils/check_dummies.py python utils/check_support_list.py python utils/check_forward_call_docstrings.py + python utils/check_return_annotations.py make deps_table_check_updated - name: Check if failure if: ${{ failure() }} diff --git a/.github/workflows/pr_tests_gpu.yml b/.github/workflows/pr_tests_gpu.yml index 4be060555b8d..2cd321e71160 100644 --- a/.github/workflows/pr_tests_gpu.yml +++ b/.github/workflows/pr_tests_gpu.yml @@ -79,6 +79,7 @@ jobs: python utils/check_dummies.py python utils/check_support_list.py python utils/check_forward_call_docstrings.py + python utils/check_return_annotations.py make deps_table_check_updated - name: Check if failure if: ${{ failure() }} diff --git a/Makefile b/Makefile index 4af58740480e..01dcc65448bc 100644 --- a/Makefile +++ b/Makefile @@ -37,6 +37,7 @@ repo-consistency: python utils/check_repo.py python utils/check_inits.py python utils/check_forward_call_docstrings.py + python utils/check_return_annotations.py # this target runs checks on all files @@ -80,6 +81,10 @@ modular-autodoctrings: check-forward-call-docstrings: python utils/check_forward_call_docstrings.py +# Verify forward() / __call__() have return type annotations +check-return-annotations: + python utils/check_return_annotations.py + # Run tests for the library test: diff --git a/src/diffusers/models/activations.py b/src/diffusers/models/activations.py index 2d1fdb5f7d83..caff731dd599 100644 --- a/src/diffusers/models/activations.py +++ b/src/diffusers/models/activations.py @@ -84,7 +84,7 @@ def gelu(self, gate: torch.Tensor) -> torch.Tensor: 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): + def forward(self, hidden_states) -> torch.Tensor: hidden_states = self.proj(hidden_states) hidden_states = self.gelu(hidden_states) return hidden_states @@ -110,7 +110,7 @@ def gelu(self, gate: torch.Tensor) -> torch.Tensor: return F.gelu(gate.to(dtype=torch.float32)).to(dtype=gate.dtype) return F.gelu(gate) - def forward(self, hidden_states, *args, **kwargs): + def forward(self, hidden_states, *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) @@ -140,7 +140,7 @@ def __init__(self, dim_in: int, dim_out: int, bias: bool = True): self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias) self.activation = nn.SiLU() - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: hidden_states = self.proj(hidden_states) hidden_states, gate = hidden_states.chunk(2, dim=-1) return hidden_states * self.activation(gate) @@ -173,6 +173,6 @@ def __init__(self, dim_in: int, dim_out: int, bias: bool = True, activation: str self.proj = nn.Linear(dim_in, dim_out, bias=bias) self.activation = get_activation(activation) - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: hidden_states = self.proj(hidden_states) return self.activation(hidden_states) diff --git a/src/diffusers/models/attention.py b/src/diffusers/models/attention.py index 65289e4b5f16..504f4afe3af3 100644 --- a/src/diffusers/models/attention.py +++ b/src/diffusers/models/attention.py @@ -1125,7 +1125,7 @@ def __init__( ) self.silu = FP32SiLU() - def forward(self, x): + def forward(self, x) -> torch.Tensor: return self.linear_2(self.silu(self.linear_1(x)) * self.linear_3(x)) @@ -1302,7 +1302,7 @@ def __init__( out_bias=attention_out_bias, ) - def forward(self, hidden_states, encoder_hidden_states, cross_attention_kwargs): + def forward(self, hidden_states, encoder_hidden_states, cross_attention_kwargs) -> torch.Tensor: cross_attention_kwargs = cross_attention_kwargs.copy() if cross_attention_kwargs is not None else {} if self.kv_mapper is not None: diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index 364b7b057e78..237115bac5a7 100644 --- a/src/diffusers/models/attention_dispatch.py +++ b/src/diffusers/models/attention_dispatch.py @@ -2227,7 +2227,7 @@ class SeqAllToAllDim(torch.autograd.Function): """ @staticmethod - def forward(ctx, group, input, scatter_id=2, gather_id=1): + def forward(ctx, group, input, scatter_id=2, gather_id=1) -> torch.Tensor: ctx.group = group ctx.scatter_id = scatter_id ctx.gather_id = gather_id @@ -2408,7 +2408,7 @@ def forward( forward_op, backward_op, _parallel_config: "ParallelConfig" | None = None, - ): + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: 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 @@ -2561,7 +2561,7 @@ def forward( forward_op, backward_op, _parallel_config: "ParallelConfig" | None = None, - ): + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: ulysses_mesh = _parallel_config.context_parallel_config._ulysses_mesh world_size = _parallel_config.context_parallel_config.ulysses_degree group = ulysses_mesh.get_group() @@ -2662,7 +2662,7 @@ def forward( forward_op, backward_op, _parallel_config: "ParallelConfig" | None = None, - ): + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: # Ring attention for arbitrary sequence lengths. if attn_mask is not None: raise ValueError( @@ -2779,7 +2779,7 @@ def forward( backward_op, _parallel_config: "ParallelConfig" | None = None, **kwargs, - ): + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: ulysses_mesh = _parallel_config.context_parallel_config._ulysses_mesh group = ulysses_mesh.get_group() diff --git a/src/diffusers/models/attention_processor.py b/src/diffusers/models/attention_processor.py index 1b923e749663..62484a113f37 100755 --- a/src/diffusers/models/attention_processor.py +++ b/src/diffusers/models/attention_processor.py @@ -985,7 +985,7 @@ def forward( encoder_hidden_states: torch.Tensor | None = None, attention_mask: torch.Tensor | None = None, **kwargs, - ): + ) -> tuple[torch.Tensor, torch.Tensor]: return self.processor( self, hidden_states, diff --git a/src/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py b/src/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py index e5549a47e9f1..a12bcc5113bc 100644 --- a/src/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py +++ b/src/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py @@ -53,7 +53,7 @@ def __init__(self, hidden_dim, logscale=True): self.beta.requires_grad = True self.logscale = logscale - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: shape = hidden_states.shape alpha = self.alpha if not self.logscale else torch.exp(self.alpha) @@ -250,7 +250,7 @@ def __init__(self, dimension: int = 16, dilation: int = 1): self.snake2 = Snake1d(dimension) self.conv2 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=1)) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: """ Forward pass through the residual unit. @@ -300,7 +300,7 @@ def __init__(self, input_dim, output_dim, stride: int = 1, output_padding: int = self.res_unit2 = Cosmos3AudioResidualUnit(output_dim, dilation=3) self.res_unit3 = Cosmos3AudioResidualUnit(output_dim, dilation=9) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.snake1(hidden_state) hidden_state = self.conv_t1(hidden_state) hidden_state = self.res_unit1(hidden_state) @@ -345,7 +345,7 @@ def __init__(self, channels, input_channels, audio_channels, upsampling_ratios, 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): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.conv1(hidden_state) for layer in self.block: diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage.py b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage.py index c1d975ae6bb7..4d46b7c32ec4 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage.py @@ -58,7 +58,7 @@ def __init__(self, in_channels: int, out_channels: int, non_linearity: str = "si else: self.conv_shortcut = None - def forward(self, x): + def forward(self, x) -> torch.Tensor: # Apply shortcut connection residual = x @@ -95,7 +95,7 @@ def __init__(self, in_channels: int): self.to_v = nn.Conv2d(in_channels, in_channels, 1) self.proj = nn.Conv2d(in_channels, in_channels, 1) - def forward(self, x): + def forward(self, x) -> torch.Tensor: identity = x x = self.norm(x) diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage_refiner.py b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage_refiner.py index 9737a822236e..091304660ba8 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage_refiner.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage_refiner.py @@ -86,7 +86,7 @@ def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bi self.gamma = nn.Parameter(torch.ones(shape)) self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - def forward(self, x): + def forward(self, x) -> torch.Tensor: needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( t in str(x.dtype) for t in ("float4_", "float8_") ) @@ -162,7 +162,7 @@ def _dcae_upsample_rearrange(tensor, r1=1, r2=2, r3=2): 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): + def forward(self, x: torch.Tensor) -> torch.Tensor: r1 = 2 if self.add_temporal_upsample else 1 h = self.conv(x) if self.add_temporal_upsample: @@ -205,7 +205,7 @@ def _dcae_downsample_rearrange(tensor, r1=1, r2=2, r3=2): 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): + def forward(self, x: torch.Tensor) -> torch.Tensor: r1 = 2 if self.add_temporal_downsample else 1 h = self.conv(x) if self.add_temporal_downsample: diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanvideo15.py b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanvideo15.py index 9260e1fcbb1d..00e28ee620b6 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanvideo15.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanvideo15.py @@ -86,7 +86,7 @@ def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bi self.gamma = nn.Parameter(torch.ones(shape)) self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - def forward(self, x): + def forward(self, x) -> torch.Tensor: needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( t in str(x.dtype) for t in ("float4_", "float8_") ) @@ -189,7 +189,7 @@ def _dcae_upsample_rearrange(tensor, r1=1, r2=2, r3=2): 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): + def forward(self, x: torch.Tensor) -> torch.Tensor: r1 = 2 if self.add_temporal_upsample else 1 h = self.conv(x) if self.add_temporal_upsample: @@ -240,7 +240,7 @@ def _dcae_downsample_rearrange(tensor, r1=1, r2=2, r3=2): 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): + def forward(self, x: torch.Tensor) -> torch.Tensor: r1 = 2 if self.add_temporal_downsample else 1 h = self.conv(x) if self.add_temporal_downsample: diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage.py b/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage.py index 220520a12e68..fa9f31ac8512 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage.py @@ -72,7 +72,7 @@ def __init__( 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): + def forward(self, x, cache_x=None) -> torch.Tensor: padding = list(self._padding) if cache_x is not None and self._padding[4] > 0: cache_x = cache_x.to(x.device) @@ -104,7 +104,7 @@ def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bi self.gamma = nn.Parameter(torch.ones(shape)) self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - def forward(self, x): + def forward(self, x) -> torch.Tensor: needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( t in str(x.dtype) for t in ("float4_", "float8_") ) @@ -126,7 +126,7 @@ class QwenImageUpsample(nn.Upsample): torch.Tensor: Upsampled tensor with the same data type as the input. """ - def forward(self, x): + def forward(self, x) -> torch.Tensor: return super().forward(x.float()).type_as(x) @@ -171,7 +171,7 @@ def __init__(self, dim: int, mode: str) -> None: else: self.resample = nn.Identity() - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: b, c, t, h, w = x.size() if self.mode == "upsample3d": if feat_cache is not None: @@ -248,7 +248,7 @@ def __init__( 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]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: # Apply shortcut connection h = self.conv_shortcut(x) @@ -308,7 +308,7 @@ def __init__(self, dim): self.to_qkv = nn.Conv2d(dim, dim * 3, 1) self.proj = nn.Conv2d(dim, dim, 1) - def forward(self, x): + def forward(self, x) -> torch.Tensor: identity = x batch_size, channels, time, height, width = x.size() @@ -361,7 +361,7 @@ def __init__(self, dim: int, dropout: float = 0.0, non_linearity: str = "silu", self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: # First residual block x = self.resnets[0](x, feat_cache, feat_idx) @@ -443,7 +443,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: if feat_cache is not None: idx = feat_idx[0] cache_x = x[:, :, -min(CACHE_T, x.shape[2]) :, :, :].clone() @@ -526,7 +526,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: """ Forward pass through the upsampling block. @@ -632,7 +632,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: ## conv1 if feat_cache is not None: idx = feat_idx[0] diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py b/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py index 15575a1c5907..3d99d7a69bf1 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py @@ -169,7 +169,7 @@ def __init__( self._padding = (self.padding[1], self.padding[1], self.padding[0], self.padding[0]) self.padding = (0, 0) - def forward(self, x, cache_x=None): + def forward(self, x, cache_x=None) -> torch.Tensor: padding = list(self._padding) if cache_x is not None: raise ValueError( @@ -206,7 +206,7 @@ def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bi self.gamma = nn.Parameter(torch.ones(shape)) self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - def forward(self, x): + def forward(self, x) -> torch.Tensor: needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( t in str(x.dtype) for t in ("float4_", "float8_") ) @@ -229,7 +229,7 @@ class QwenImage21Upsample(nn.Upsample): torch.Tensor: Upsampled tensor with the same data type as the input. """ - def forward(self, x): + def forward(self, x) -> torch.Tensor: return super().forward(x.float()).type_as(x) @@ -278,7 +278,7 @@ def __init__(self, dim: int, mode: str, upsample_out_dim: int = None) -> None: else: self.resample = nn.Identity() - def forward(self, x, feat_cache=None, feat_idx=None): + def forward(self, x, feat_cache=None, feat_idx=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] b, c, t, h, w = x.size() @@ -355,7 +355,7 @@ def __init__( self.conv2 = QwenImage21CausalConv3d(out_dim, out_dim, 3, padding=1) self.conv_shortcut = QwenImage21CausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity() - def forward(self, x, feat_cache=None, feat_idx=None): + def forward(self, x, feat_cache=None, feat_idx=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] # Apply shortcut connection @@ -418,7 +418,7 @@ def __init__(self, dim): self.to_qkv = nn.Conv2d(dim, dim * 3, 1) self.proj = nn.Conv2d(dim, dim, 1) - def forward(self, x): + def forward(self, x) -> torch.Tensor: identity = x batch_size, channels, time, height, width = x.size() @@ -470,7 +470,7 @@ def __init__(self, dim: int, dropout: float = 0.0, num_layers: int = 1): self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=None): + def forward(self, x, feat_cache=None, feat_idx=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] # First residual block @@ -512,7 +512,7 @@ def __init__(self, in_dim, out_dim, dropout, num_res_blocks, temperal_downsample else: self.downsampler = None - def forward(self, x, feat_cache=None, feat_idx=None): + def forward(self, x, feat_cache=None, feat_idx=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] x_copy = x.clone() @@ -603,7 +603,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=None): + def forward(self, x, feat_cache=None, feat_idx=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] if feat_cache is not None: @@ -700,7 +700,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=False): + def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=False) -> torch.Tensor: if feat_idx is None: feat_idx = [0] """ @@ -775,7 +775,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=None): + def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] """ @@ -889,7 +889,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=False): + def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=False) -> torch.Tensor: if feat_idx is None: feat_idx = [0] ## conv1 diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_wan.py b/src/diffusers/models/autoencoders/autoencoder_kl_wan.py index de8a56edc20e..29bd5738a62b 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_wan.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_wan.py @@ -163,7 +163,7 @@ def __init__( 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): + def forward(self, x, cache_x=None) -> torch.Tensor: padding = list(self._padding) if cache_x is not None and self._padding[4] > 0: cache_x = cache_x.to(x.device) @@ -195,7 +195,7 @@ def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bi self.gamma = nn.Parameter(torch.ones(shape)) self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - def forward(self, x): + def forward(self, x) -> torch.Tensor: needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( t in str(x.dtype) for t in ("float4_", "float8_") ) @@ -217,7 +217,7 @@ class WanUpsample(nn.Upsample): torch.Tensor: Upsampled tensor with the same data type as the input. """ - def forward(self, x): + def forward(self, x) -> torch.Tensor: return super().forward(x.float()).type_as(x) @@ -266,7 +266,7 @@ def __init__(self, dim: int, mode: str, upsample_out_dim: int = None) -> None: else: self.resample = nn.Identity() - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: b, c, t, h, w = x.size() if self.mode == "upsample3d": if feat_cache is not None: @@ -343,7 +343,7 @@ def __init__( 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]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: # Apply shortcut connection h = self.conv_shortcut(x) @@ -403,7 +403,7 @@ def __init__(self, dim): self.to_qkv = nn.Conv2d(dim, dim * 3, 1) self.proj = nn.Conv2d(dim, dim, 1) - def forward(self, x): + def forward(self, x) -> torch.Tensor: identity = x batch_size, channels, time, height, width = x.size() @@ -456,7 +456,7 @@ def __init__(self, dim: int, dropout: float = 0.0, non_linearity: str = "silu", self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: # First residual block x = self.resnets[0](x, feat_cache=feat_cache, feat_idx=feat_idx) @@ -496,7 +496,7 @@ def __init__(self, in_dim, out_dim, dropout, num_res_blocks, temperal_downsample else: self.downsampler = None - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: x_copy = x.clone() for resnet in self.resnets: x = resnet(x, feat_cache=feat_cache, feat_idx=feat_idx) @@ -587,7 +587,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: if feat_cache is not None: idx = feat_idx[0] cache_x = x[:, :, -CACHE_T:, :, :].clone() @@ -684,7 +684,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False): + def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False) -> torch.Tensor: """ Forward pass through the upsampling block. @@ -759,7 +759,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=None): + def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=None) -> torch.Tensor: """ Forward pass through the upsampling block. @@ -876,7 +876,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False): + def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False) -> torch.Tensor: ## conv1 if feat_cache is not None: idx = feat_idx[0] diff --git a/src/diffusers/models/autoencoders/autoencoder_oobleck.py b/src/diffusers/models/autoencoders/autoencoder_oobleck.py index d4251fd9f1a9..639199a8dd36 100644 --- a/src/diffusers/models/autoencoders/autoencoder_oobleck.py +++ b/src/diffusers/models/autoencoders/autoencoder_oobleck.py @@ -41,7 +41,7 @@ def __init__(self, hidden_dim, logscale=True): self.beta.requires_grad = True self.logscale = logscale - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: shape = hidden_states.shape alpha = self.alpha if not self.logscale else torch.exp(self.alpha) @@ -67,7 +67,7 @@ def __init__(self, dimension: int = 16, dilation: int = 1): self.snake2 = Snake1d(dimension) self.conv2 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=1)) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: """ Forward pass through the residual unit. @@ -104,7 +104,7 @@ def __init__(self, input_dim, output_dim, stride: int = 1): nn.Conv1d(input_dim, output_dim, kernel_size=2 * stride, stride=stride, padding=math.ceil(stride / 2)) ) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.res_unit1(hidden_state) hidden_state = self.res_unit2(hidden_state) hidden_state = self.snake1(self.res_unit3(hidden_state)) @@ -133,7 +133,7 @@ def __init__(self, input_dim, output_dim, stride: int = 1): self.res_unit2 = OobleckResidualUnit(output_dim, dilation=3) self.res_unit3 = OobleckResidualUnit(output_dim, dilation=9) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.snake1(hidden_state) hidden_state = self.conv_t1(hidden_state) hidden_state = self.res_unit1(hidden_state) @@ -239,7 +239,7 @@ def __init__(self, encoder_hidden_size, audio_channels, downsampling_ratios, cha 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): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.conv1(hidden_state) for module in self.block: @@ -279,7 +279,7 @@ def __init__(self, channels, input_channels, audio_channels, upsampling_ratios, 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): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.conv1(hidden_state) for layer in self.block: diff --git a/src/diffusers/models/controlnets/controlnet.py b/src/diffusers/models/controlnets/controlnet.py index acd88655c9fe..7dc3e9b18f8d 100644 --- a/src/diffusers/models/controlnets/controlnet.py +++ b/src/diffusers/models/controlnets/controlnet.py @@ -95,7 +95,7 @@ def __init__( nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1) ) - def forward(self, conditioning): + def forward(self, conditioning) -> torch.Tensor: embedding = self.conv_in(conditioning) embedding = F.silu(embedding) diff --git a/src/diffusers/models/controlnets/controlnet_hunyuan.py b/src/diffusers/models/controlnets/controlnet_hunyuan.py index 6ef92d78dd6e..b7cae237e005 100644 --- a/src/diffusers/models/controlnets/controlnet_hunyuan.py +++ b/src/diffusers/models/controlnets/controlnet_hunyuan.py @@ -226,7 +226,7 @@ def forward( style=None, image_rotary_emb=None, return_dict=True, - ): + ) -> HunyuanControlNetOutput | tuple[list[torch.Tensor]]: """ The [`HunyuanDiT2DControlNetModel`] forward method. @@ -257,6 +257,10 @@ def forward( The image rotary embeddings to apply on query and key tensors during attention calculation. return_dict: bool Whether to return a dictionary. + + Returns: + If `return_dict` is True, a [`~models.controlnets.controlnet_hunyuan.HunyuanControlNetOutput`] is returned, + otherwise a `tuple` where the first element is the list of ControlNet block samples. """ height, width = hidden_states.shape[-2:] @@ -339,7 +343,7 @@ def forward( style=None, image_rotary_emb=None, return_dict=True, - ): + ) -> HunyuanControlNetOutput | tuple[list[torch.Tensor]]: """ The [`HunyuanDiT2DControlNetModel`] forward method. @@ -370,6 +374,11 @@ def forward( The image rotary embeddings to apply on query and key tensors during attention calculation. return_dict: bool Whether to return a dictionary. + + Returns: + If `return_dict` is True and only one ControlNet is used, a + [`~models.controlnets.controlnet_hunyuan.HunyuanControlNetOutput`] is returned. Otherwise a `tuple` where + the first element is the list of ControlNet block samples, summed across all ControlNets. """ for i, (image, scale, controlnet) in enumerate(zip(controlnet_cond, conditioning_scale, self.nets)): block_samples = controlnet( diff --git a/src/diffusers/models/controlnets/controlnet_union.py b/src/diffusers/models/controlnets/controlnet_union.py index 8b3ac1c36d85..e24933cace5a 100644 --- a/src/diffusers/models/controlnets/controlnet_union.py +++ b/src/diffusers/models/controlnets/controlnet_union.py @@ -56,7 +56,7 @@ def __init__(self, d_model: int): self.gelu = QuickGELU() self.c_proj = nn.Linear(d_model * 4, d_model) - def forward(self, x: torch.Tensor): + def forward(self, x: torch.Tensor) -> torch.Tensor: x = self.c_fc(x) x = self.gelu(x) x = self.c_proj(x) @@ -76,7 +76,7 @@ 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): + def forward(self, x: torch.Tensor) -> torch.Tensor: x = x + self.attention(self.ln_1(x)) x = x + self.mlp(self.ln_2(x)) return x diff --git a/src/diffusers/models/controlnets/controlnet_xs.py b/src/diffusers/models/controlnets/controlnet_xs.py index a25d5d71a5b1..c3f58ea359a2 100644 --- a/src/diffusers/models/controlnets/controlnet_xs.py +++ b/src/diffusers/models/controlnets/controlnet_xs.py @@ -502,7 +502,7 @@ def from_unet( return model - def forward(self, *args, **kwargs): + def forward(self, *args, **kwargs) -> None: raise ValueError( "A ControlNetXSAdapter cannot be run by itself. Use it together with a UNet2DConditionModel to instantiate a UNetControlNetXSModel." ) diff --git a/src/diffusers/models/controlnets/controlnet_z_image.py b/src/diffusers/models/controlnets/controlnet_z_image.py index a4800b255ef0..904e4bfdf6c5 100644 --- a/src/diffusers/models/controlnets/controlnet_z_image.py +++ b/src/diffusers/models/controlnets/controlnet_z_image.py @@ -62,7 +62,7 @@ def timestep_embedding(t, dim, max_period=10000): embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) return embedding - def forward(self, t): + def forward(self, t) -> torch.Tensor: 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) @@ -166,7 +166,7 @@ def __init__(self, dim: int, hidden_dim: int): def _forward_silu_gating(self, x1, x3): return F.silu(x1) * x3 - def forward(self, x): + def forward(self, x) -> torch.Tensor: return self.w2(self._forward_silu_gating(self.w1(x), self.w3(x))) @@ -238,7 +238,7 @@ def forward( noise_mask: torch.Tensor | None = None, adaln_noisy: torch.Tensor | None = None, adaln_clean: torch.Tensor | None = None, - ): + ) -> torch.Tensor: if self.modulation: seq_len = x.shape[1] @@ -390,7 +390,7 @@ def forward( attn_mask: torch.Tensor, freqs_cis: torch.Tensor, adaln_input: torch.Tensor | None = None, - ): + ) -> torch.Tensor: # Control if self.block_id == 0: c = self.before_proj(c) + x @@ -660,7 +660,7 @@ def forward( conditioning_scale: float = 1.0, patch_size=2, f_patch_size=1, - ): + ) -> dict[int, torch.Tensor]: r""" Args: x (`list` of `torch.Tensor`): @@ -677,6 +677,10 @@ def forward( 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. + + Returns: + `dict[int, torch.Tensor]`: The ControlNet block samples, scaled by `conditioning_scale` and keyed by the + index of the transformer layer each one is added to. """ if ( self.t_scale is None diff --git a/src/diffusers/models/embeddings.py b/src/diffusers/models/embeddings.py index cbebf3de3a50..e1938a99b265 100644 --- a/src/diffusers/models/embeddings.py +++ b/src/diffusers/models/embeddings.py @@ -356,7 +356,7 @@ def cropped_pos_embed(self, height, width): spatial_pos_embed = spatial_pos_embed.reshape(1, -1, spatial_pos_embed.shape[-1]) return spatial_pos_embed - def forward(self, latent): + def forward(self, latent) -> torch.Tensor: if self.pos_embed_max_size is not None: height, width = latent.shape[-2:] else: @@ -408,7 +408,7 @@ def __init__(self, patch_size=2, in_channels=4, embed_dim=768, bias=True): bias=bias, ) - def forward(self, x, freqs_cis): + def forward(self, x, freqs_cis) -> tuple[torch.Tensor, torch.Tensor, list[tuple[int, int]], torch.Tensor]: """ Patchifies and embeds the input tensor(s). @@ -517,7 +517,7 @@ def _get_positional_embeddings( return joint_pos_embedding - def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor): + def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor) -> torch.Tensor: r""" Args: text_embeds (`torch.Tensor`): @@ -1057,7 +1057,7 @@ def __init__( else: self.post_act = get_activation(post_act_fn) - def forward(self, sample, condition=None): + def forward(self, sample, condition=None) -> torch.Tensor: if condition is not None: sample = sample + self.cond_proj(condition) sample = self.linear_1(sample) @@ -1109,7 +1109,7 @@ def __init__( self.weight = self.W del self.W - def forward(self, x): + def forward(self, x) -> torch.Tensor: if self.log: x = torch.log(x) @@ -1143,7 +1143,7 @@ def __init__(self, embed_dim: int, max_seq_length: int = 32): pe[0, :, 1::2] = torch.cos(position * div_term) self.register_buffer("pe", pe) - def forward(self, x): + def forward(self, x) -> torch.Tensor: _, seq_length, _ = x.shape x = x + self.pe[:, :seq_length] return x @@ -1191,7 +1191,7 @@ def __init__( self.height_emb = nn.Embedding(self.height, embed_dim) self.width_emb = nn.Embedding(self.width, embed_dim) - def forward(self, index): + def forward(self, index) -> torch.Tensor: emb = self.emb(index) height_emb = self.height_emb(torch.arange(self.height, device=index.device).view(1, self.height)) @@ -1242,7 +1242,7 @@ def token_drop(self, labels, force_drop_ids=None): labels = torch.where(drop_ids, self.num_classes, labels) return labels - def forward(self, labels: torch.LongTensor, force_drop_ids=None): + def forward(self, labels: torch.LongTensor, force_drop_ids=None) -> torch.Tensor: 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) @@ -1264,7 +1264,7 @@ def __init__( 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): + def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor) -> torch.Tensor: batch_size = text_embeds.shape[0] # image @@ -1290,7 +1290,7 @@ def __init__( 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): + def forward(self, image_embeds: torch.Tensor) -> torch.Tensor: batch_size = image_embeds.shape[0] # image @@ -1308,7 +1308,7 @@ def __init__(self, image_embed_dim=1024, cross_attention_dim=1024): 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): + def forward(self, image_embeds: torch.Tensor) -> torch.Tensor: return self.norm(self.ff(image_embeds)) @@ -1322,7 +1322,7 @@ def __init__(self, image_embed_dim=1024, cross_attention_dim=1024, mult=1, num_t 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): + def forward(self, image_embeds: torch.Tensor) -> torch.Tensor: x = self.ff(image_embeds) x = x.reshape(-1, self.num_tokens, self.cross_attention_dim) return self.norm(x) @@ -1336,7 +1336,7 @@ def __init__(self, num_classes, embedding_dim, class_dropout_prob=0.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): + def forward(self, timestep, class_labels, hidden_dtype=None) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) @@ -1355,7 +1355,7 @@ def __init__(self, embedding_dim, pooled_projection_dim): 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): + def forward(self, timestep, pooled_projection) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D) @@ -1375,7 +1375,7 @@ def __init__(self, embedding_dim, pooled_projection_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): + def forward(self, timestep, guidance, pooled_projection) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D) @@ -1435,7 +1435,7 @@ def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim) self.num_heads = num_heads - def forward(self, x): + def forward(self, x) -> torch.Tensor: 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 @@ -1498,7 +1498,7 @@ def __init__( act_fn="silu_fp32", ) - def forward(self, timestep, encoder_hidden_states, image_meta_size, style, hidden_dtype=None): + def forward(self, timestep, encoder_hidden_states, image_meta_size, style, hidden_dtype=None) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, 256) @@ -1542,7 +1542,7 @@ def __init__(self, hidden_size=4096, cross_attention_dim=2048, frequency_embeddi ), ) - def forward(self, timestep, caption_feat, caption_mask): + def forward(self, timestep, caption_feat, caption_mask) -> torch.Tensor: # timestep embedding: time_freq = self.time_proj(timestep) time_embed = self.timestep_embedder(time_freq.to(dtype=caption_feat.dtype)) @@ -1582,7 +1582,7 @@ def forward( encoder_hidden_states: torch.Tensor, encoder_attention_mask: torch.Tensor, hidden_dtype: torch.dtype | None = None, - ): + ) -> tuple[torch.Tensor, torch.Tensor]: time_proj = self.time_proj(timestep) time_emb = self.timestep_embedder(time_proj.to(dtype=hidden_dtype)) @@ -1601,7 +1601,7 @@ def __init__(self, encoder_dim: int, time_embed_dim: int, num_heads: int = 64): self.proj = nn.Linear(encoder_dim, time_embed_dim) self.norm2 = nn.LayerNorm(time_embed_dim) - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: hidden_states = self.norm1(hidden_states) hidden_states = self.pool(hidden_states) hidden_states = self.proj(hidden_states) @@ -1616,7 +1616,7 @@ def __init__(self, text_embed_dim: int = 768, image_embed_dim: int = 768, time_e 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): + def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor) -> torch.Tensor: # text time_text_embeds = self.text_proj(text_embeds) time_text_embeds = self.text_norm(time_text_embeds) @@ -1633,7 +1633,7 @@ def __init__(self, image_embed_dim: int = 768, time_embed_dim: int = 1536): 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): + def forward(self, image_embeds: torch.Tensor) -> torch.Tensor: # image time_image_embeds = self.image_proj(image_embeds) time_image_embeds = self.image_norm(time_image_embeds) @@ -1663,7 +1663,7 @@ def __init__(self, image_embed_dim: int = 768, time_embed_dim: int = 1536): nn.Conv2d(256, 4, 3, padding=1), ) - def forward(self, image_embeds: torch.Tensor, hint: torch.Tensor): + def forward(self, image_embeds: torch.Tensor, hint: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: # image time_image_embeds = self.image_proj(image_embeds) time_image_embeds = self.image_norm(time_image_embeds) @@ -1684,7 +1684,7 @@ def __init__(self, num_heads, embed_dim, dtype=None): self.num_heads = num_heads self.dim_per_head = embed_dim // self.num_heads - def forward(self, x): + def forward(self, x) -> torch.Tensor: bs, length, width = x.size() def shape(x): @@ -1875,7 +1875,7 @@ def forward( image_masks=None, phrases_embeddings=None, image_embeddings=None, - ): + ) -> torch.Tensor: masks = masks.unsqueeze(-1) # embedding position (it may includes padding as placeholder) @@ -1938,7 +1938,7 @@ def __init__(self, embedding_dim, size_emb_dim, use_additional_conditions: bool 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): + def forward(self, timestep, resolution, aspect_ratio, batch_size, hidden_dtype) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) @@ -1976,7 +1976,7 @@ def __init__(self, in_features, hidden_size, out_features=None, act_fn="gelu_tan 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): + def forward(self, caption) -> torch.Tensor: hidden_states = self.linear_1(caption) hidden_states = self.act_1(hidden_states) hidden_states = self.linear_2(hidden_states) @@ -2007,7 +2007,7 @@ def __init__( FeedForward(embed_dims, embed_dims, activation_fn="gelu", mult=ffn_ratio, bias=False), ) - def forward(self, x, latents, residual): + def forward(self, x, latents, residual) -> torch.Tensor: encoder_hidden_states = self.ln0(x) latents = self.ln1(latents) encoder_hidden_states = torch.cat([encoder_hidden_states, latents], dim=-2) @@ -2346,7 +2346,7 @@ 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]): + def forward(self, image_embeds: list[torch.Tensor]) -> list[torch.Tensor]: projected_image_embeds = [] # currently, we accept `image_embeds` as diff --git a/src/diffusers/models/lora.py b/src/diffusers/models/lora.py index 72e285832737..94ed621940dd 100644 --- a/src/diffusers/models/lora.py +++ b/src/diffusers/models/lora.py @@ -162,7 +162,7 @@ def _unfuse_lora(self): self.w_up = None self.w_down = None - def forward(self, input): + def forward(self, input) -> torch.Tensor: if self.lora_scale is None: self.lora_scale = 1.0 if self.lora_linear_layer is None: diff --git a/src/diffusers/models/normalization.py b/src/diffusers/models/normalization.py index dc872417914e..4b4ba78dc679 100644 --- a/src/diffusers/models/normalization.py +++ b/src/diffusers/models/normalization.py @@ -503,7 +503,7 @@ def __init__(self, dim, eps: float = 1e-5, elementwise_affine: bool = True, bias self.weight = None self.bias = None - def forward(self, input): + def forward(self, input) -> torch.Tensor: return F.layer_norm(input, self.dim, self.weight, self.bias, self.eps) @@ -538,7 +538,7 @@ def __init__(self, dim, eps: float, elementwise_affine: bool = True, bias: bool if bias: self.bias = nn.Parameter(torch.zeros(dim)) - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: # `npu_rms_norm` requires a gamma tensor. When `elementwise_affine=False`, # `self.weight` is `None`, so fall back to the pure PyTorch path. if is_torch_npu_available() and self.weight is not None: @@ -586,7 +586,7 @@ def __init__(self, dim, eps: float, elementwise_affine: bool = True): else: self.weight = None - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: 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) @@ -612,7 +612,7 @@ def __init__(self, dim): self.gamma = nn.Parameter(torch.zeros(1, 1, 1, dim)) self.beta = nn.Parameter(torch.zeros(1, 1, 1, dim)) - def forward(self, x): + def forward(self, x) -> torch.Tensor: 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 diff --git a/src/diffusers/models/resnet.py b/src/diffusers/models/resnet.py index d63e4fd0017b..fc5c89efd200 100644 --- a/src/diffusers/models/resnet.py +++ b/src/diffusers/models/resnet.py @@ -692,7 +692,7 @@ def forward( hidden_states: torch.Tensor, temb: torch.Tensor | None = None, image_only_indicator: torch.Tensor | None = None, - ): + ) -> torch.Tensor: num_frames = image_only_indicator.shape[-1] hidden_states = self.spatial_res_block(hidden_states, temb) diff --git a/src/diffusers/models/transformers/dit_transformer_2d.py b/src/diffusers/models/transformers/dit_transformer_2d.py index 0457acf77108..fc37781ddbd5 100644 --- a/src/diffusers/models/transformers/dit_transformer_2d.py +++ b/src/diffusers/models/transformers/dit_transformer_2d.py @@ -152,7 +152,7 @@ def forward( class_labels: torch.LongTensor | None = None, cross_attention_kwargs: dict[str, Any] = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`DiTTransformer2DModel`] forward method. diff --git a/src/diffusers/models/transformers/dual_transformer_2d.py b/src/diffusers/models/transformers/dual_transformer_2d.py index 778d5128ee23..4d400c18e200 100644 --- a/src/diffusers/models/transformers/dual_transformer_2d.py +++ b/src/diffusers/models/transformers/dual_transformer_2d.py @@ -11,6 +11,7 @@ # WITHOUT WARRANTIES 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 ..modeling_outputs import Transformer2DModelOutput @@ -101,7 +102,7 @@ def forward( attention_mask=None, cross_attention_kwargs=None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ Args: hidden_states ( When discrete, `torch.LongTensor` of shape `(batch size, num latent pixels)`. diff --git a/src/diffusers/models/transformers/hunyuan_transformer_2d.py b/src/diffusers/models/transformers/hunyuan_transformer_2d.py index 83b3797c4fc3..ba0745ce4876 100644 --- a/src/diffusers/models/transformers/hunyuan_transformer_2d.py +++ b/src/diffusers/models/transformers/hunyuan_transformer_2d.py @@ -367,7 +367,7 @@ def forward( image_rotary_emb=None, controlnet_block_samples=None, return_dict=True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`HunyuanDiT2DModel`] forward method. @@ -396,6 +396,10 @@ def forward( A list of tensors that if specified are added to the residuals of transformer blocks. return_dict: bool Whether to return a dictionary. + + Returns: + If `return_dict` is True, a [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a + `tuple` where the first element is the sample tensor. """ height, width = hidden_states.shape[-2:] diff --git a/src/diffusers/models/transformers/latte_transformer_3d.py b/src/diffusers/models/transformers/latte_transformer_3d.py index 01a1e608a927..c953cdabc936 100644 --- a/src/diffusers/models/transformers/latte_transformer_3d.py +++ b/src/diffusers/models/transformers/latte_transformer_3d.py @@ -171,7 +171,7 @@ def forward( encoder_attention_mask: torch.Tensor | None = None, enable_temporal_attentions: bool = True, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`LatteTransformer3DModel`] forward method. diff --git a/src/diffusers/models/transformers/pixart_transformer_2d.py b/src/diffusers/models/transformers/pixart_transformer_2d.py index e5e6178eaf4a..7e08abff85f3 100644 --- a/src/diffusers/models/transformers/pixart_transformer_2d.py +++ b/src/diffusers/models/transformers/pixart_transformer_2d.py @@ -234,7 +234,7 @@ def forward( attention_mask: torch.Tensor | None = None, encoder_attention_mask: torch.Tensor | None = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`PixArtTransformer2DModel`] forward method. diff --git a/src/diffusers/models/transformers/prior_transformer.py b/src/diffusers/models/transformers/prior_transformer.py index f3890446e28e..8847fff3a4ba 100644 --- a/src/diffusers/models/transformers/prior_transformer.py +++ b/src/diffusers/models/transformers/prior_transformer.py @@ -188,7 +188,7 @@ def forward( encoder_hidden_states: torch.Tensor | None = None, attention_mask: torch.BoolTensor | None = None, return_dict: bool = True, - ): + ) -> PriorTransformerOutput | tuple[torch.Tensor]: """ The [`PriorTransformer`] forward method. diff --git a/src/diffusers/models/transformers/sana_transformer.py b/src/diffusers/models/transformers/sana_transformer.py index 1451750d50ef..d05042c08300 100644 --- a/src/diffusers/models/transformers/sana_transformer.py +++ b/src/diffusers/models/transformers/sana_transformer.py @@ -108,7 +108,9 @@ def __init__(self, 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): + def forward( + self, timestep: torch.Tensor, guidance: torch.Tensor = None, hidden_dtype: torch.dtype = None + ) -> tuple[torch.Tensor, torch.Tensor]: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) diff --git a/src/diffusers/models/transformers/stable_audio_transformer.py b/src/diffusers/models/transformers/stable_audio_transformer.py index f4974926ec72..72c42048db71 100644 --- a/src/diffusers/models/transformers/stable_audio_transformer.py +++ b/src/diffusers/models/transformers/stable_audio_transformer.py @@ -48,7 +48,7 @@ def __init__( self.weight = self.W del self.W - def forward(self, x): + def forward(self, x) -> torch.Tensor: if self.log: x = torch.log(x) diff --git a/src/diffusers/models/transformers/t5_film_transformer.py b/src/diffusers/models/transformers/t5_film_transformer.py index 547e72089990..8fb368545d16 100644 --- a/src/diffusers/models/transformers/t5_film_transformer.py +++ b/src/diffusers/models/transformers/t5_film_transformer.py @@ -89,7 +89,7 @@ def encoder_decoder_mask(self, query_input: torch.Tensor, key_input: torch.Tenso 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): + def forward(self, encodings_and_masks, decoder_input_tokens, decoder_noise_time) -> torch.Tensor: """ The [`T5FilmDecoder`] forward method. @@ -101,6 +101,9 @@ def forward(self, encodings_and_masks, decoder_input_tokens, decoder_noise_time) Input tokens for the decoder. decoder_noise_time (`torch.Tensor` of shape `(batch_size,)`): Diffusion timesteps in `[0, 1)` used to condition the decoder. + + Returns: + `torch.Tensor`: The decoded spectrogram of shape `(batch_size, seq_length, input_dims)`. """ batch, _, _ = decoder_input_tokens.shape assert decoder_noise_time.shape == (batch,) diff --git a/src/diffusers/models/transformers/transformer_2d.py b/src/diffusers/models/transformers/transformer_2d.py index 6714383b77ab..50f5a082efec 100644 --- a/src/diffusers/models/transformers/transformer_2d.py +++ b/src/diffusers/models/transformers/transformer_2d.py @@ -332,7 +332,7 @@ def forward( attention_mask: torch.Tensor | None = None, encoder_attention_mask: torch.Tensor | None = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`Transformer2DModel`] forward method. diff --git a/src/diffusers/models/transformers/transformer_2d_dreamlite.py b/src/diffusers/models/transformers/transformer_2d_dreamlite.py index 9d66eeafbd00..370c16235170 100644 --- a/src/diffusers/models/transformers/transformer_2d_dreamlite.py +++ b/src/diffusers/models/transformers/transformer_2d_dreamlite.py @@ -522,7 +522,7 @@ def forward( attention_mask: torch.Tensor | None = None, encoder_attention_mask: torch.Tensor | None = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """Forward pass of :class:`DreamLiteTransformer2DModel`. Args: diff --git a/src/diffusers/models/transformers/transformer_allegro.py b/src/diffusers/models/transformers/transformer_allegro.py index abe82ab578de..57ce1da7fd68 100644 --- a/src/diffusers/models/transformers/transformer_allegro.py +++ b/src/diffusers/models/transformers/transformer_allegro.py @@ -311,7 +311,7 @@ def forward( encoder_attention_mask: torch.Tensor | None = None, image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`AllegroTransformer3DModel`] forward method. diff --git a/src/diffusers/models/transformers/transformer_anyflow.py b/src/diffusers/models/transformers/transformer_anyflow.py index 6b0872ffdb01..1388c63096b0 100644 --- a/src/diffusers/models/transformers/transformer_anyflow.py +++ b/src/diffusers/models/transformers/transformer_anyflow.py @@ -287,7 +287,7 @@ def forward( encoder_hidden_states: torch.Tensor, encoder_hidden_states_image: Optional[torch.Tensor] = None, layout_cfg=None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: if self.deltatime_type == "r": delta_timestep = r_timestep elif self.deltatime_type == "t-r": @@ -384,7 +384,7 @@ def _forward_full_frame(self, num_frames, height, width, device) -> torch.Tensor freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1) return freqs - def forward(self, layout_cfg, device): + def forward(self, layout_cfg, device) -> dict[str, torch.Tensor]: freqs = self._forward_full_frame( num_frames=layout_cfg["total_frames"], height=layout_cfg["full_frame_shape"][0], diff --git a/src/diffusers/models/transformers/transformer_anyflow_far.py b/src/diffusers/models/transformers/transformer_anyflow_far.py index 9ecc16bd04e0..a5fdf84ac829 100644 --- a/src/diffusers/models/transformers/transformer_anyflow_far.py +++ b/src/diffusers/models/transformers/transformer_anyflow_far.py @@ -468,7 +468,7 @@ def forward( encoder_hidden_states_image: Optional[torch.Tensor] = None, far_cfg=None, clean_timestep=None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: if self.deltatime_type == "r": delta_timestep = r_timestep elif self.deltatime_type == "t-r": @@ -749,7 +749,7 @@ def _forward_full_frame(self, num_frames, height, width, device) -> torch.Tensor freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1) return freqs - def forward(self, far_cfg, device, clean_hidden_states=None): + def forward(self, far_cfg, device, clean_hidden_states=None) -> dict[str, torch.Tensor]: full_frame_freqs = self._forward_full_frame( num_frames=far_cfg["total_frames"], height=far_cfg["full_frame_shape"][0], diff --git a/src/diffusers/models/transformers/transformer_bria.py b/src/diffusers/models/transformers/transformer_bria.py index ff4261343ab2..a97b5b772552 100644 --- a/src/diffusers/models/transformers/transformer_bria.py +++ b/src/diffusers/models/transformers/transformer_bria.py @@ -304,7 +304,7 @@ def __init__( self.scale = scale self.time_theta = time_theta - def forward(self, timesteps): + def forward(self, timesteps) -> torch.Tensor: t_emb = get_timestep_embedding( timesteps, self.num_channels, @@ -325,7 +325,7 @@ def __init__(self, embedding_dim, time_theta): ) self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - def forward(self, timestep, dtype): + def forward(self, timestep, dtype) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=dtype)) # (N, D) return timesteps_emb diff --git a/src/diffusers/models/transformers/transformer_bria_fibo.py b/src/diffusers/models/transformers/transformer_bria_fibo.py index 9ec0ea1647a6..3dff8bcf55c7 100644 --- a/src/diffusers/models/transformers/transformer_bria_fibo.py +++ b/src/diffusers/models/transformers/transformer_bria_fibo.py @@ -296,7 +296,7 @@ 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): + def forward(self, caption) -> torch.Tensor: hidden_states = self.linear(caption) return hidden_states @@ -398,7 +398,7 @@ def __init__( self.scale = scale self.time_theta = time_theta - def forward(self, timesteps): + def forward(self, timesteps) -> torch.Tensor: t_emb = get_timestep_embedding( timesteps, self.num_channels, @@ -419,7 +419,7 @@ def __init__(self, embedding_dim, time_theta): ) self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - def forward(self, timestep, dtype): + def forward(self, timestep, dtype) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=dtype)) # (N, D) return timesteps_emb diff --git a/src/diffusers/models/transformers/transformer_chroma.py b/src/diffusers/models/transformers/transformer_chroma.py index 92190bb0120d..81a3dadc0998 100644 --- a/src/diffusers/models/transformers/transformer_chroma.py +++ b/src/diffusers/models/transformers/transformer_chroma.py @@ -190,7 +190,7 @@ def __init__(self, in_dim: int, out_dim: int, hidden_dim: int, n_layers: int = 5 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): + def forward(self, x) -> torch.Tensor: x = self.in_proj(x) for layer, norms in zip(self.layers, self.norms): diff --git a/src/diffusers/models/transformers/transformer_chronoedit.py b/src/diffusers/models/transformers/transformer_chronoedit.py index b39a18a98afb..c676a437cd16 100644 --- a/src/diffusers/models/transformers/transformer_chronoedit.py +++ b/src/diffusers/models/transformers/transformer_chronoedit.py @@ -340,7 +340,7 @@ def forward( encoder_hidden_states: torch.Tensor, encoder_hidden_states_image: torch.Tensor | None = None, timestep_seq_len: int | None = None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: timestep = self.timesteps_proj(timestep) if timestep_seq_len is not None: timestep = timestep.unflatten(0, (-1, timestep_seq_len)) diff --git a/src/diffusers/models/transformers/transformer_cosmos3.py b/src/diffusers/models/transformers/transformer_cosmos3.py index d2cafeca9d27..0295e1c0567b 100644 --- a/src/diffusers/models/transformers/transformer_cosmos3.py +++ b/src/diffusers/models/transformers/transformer_cosmos3.py @@ -144,7 +144,7 @@ def apply_interleaved_mrope(self, freqs, rope_axes_dim): freqs_t[..., idx] = freqs[dim, ..., idx] return freqs_t - def forward(self, position_ids, device, dtype): + def forward(self, position_ids, device, dtype) -> tuple[torch.Tensor, torch.Tensor]: if position_ids.ndim == 2: position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1) # [3,B,N] inv_freq_expanded = ( @@ -188,7 +188,7 @@ def __init__(self, hidden_size: int, intermediate_size: int, hidden_act: str = " 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): + def forward(self, x) -> torch.Tensor: 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)) diff --git a/src/diffusers/models/transformers/transformer_ernie_image.py b/src/diffusers/models/transformers/transformer_ernie_image.py index 0abc5d254bb2..791d4e0bce42 100644 --- a/src/diffusers/models/transformers/transformer_ernie_image.py +++ b/src/diffusers/models/transformers/transformer_ernie_image.py @@ -264,7 +264,7 @@ def forward( rotary_pos_emb, temb: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor], attention_mask: torch.Tensor | None = None, - ): + ) -> torch.Tensor: shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = temb residual = x x = self.adaLN_sa_ln(x) @@ -353,7 +353,7 @@ def forward( text_bth: torch.Tensor, text_lens: torch.Tensor, return_dict: bool = True, - ): + ) -> ErnieImageTransformer2DModelOutput | tuple[torch.Tensor]: """ The [`ErnieImageTransformer2DModel`] forward method. @@ -370,6 +370,11 @@ def forward( 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, a + [`~models.transformers.transformer_ernie_image.ErnieImageTransformer2DModelOutput`] is returned, otherwise + a `tuple` where the first element is the sample tensor. """ device, dtype = hidden_states.device, hidden_states.dtype B, C, H, W = hidden_states.shape diff --git a/src/diffusers/models/transformers/transformer_helios.py b/src/diffusers/models/transformers/transformer_helios.py index b99ab1e3f34f..6733004bf1e5 100644 --- a/src/diffusers/models/transformers/transformer_helios.py +++ b/src/diffusers/models/transformers/transformer_helios.py @@ -87,7 +87,7 @@ def __init__(self, dim: int, eps: float = 1e-6, elementwise_affine: bool = False 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): + def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor, original_context_length: int) -> torch.Tensor: 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) @@ -308,7 +308,7 @@ def forward( timestep: torch.Tensor, encoder_hidden_states: torch.Tensor | None = None, is_return_encoder_hidden_states: bool = True, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]: timestep = self.timesteps_proj(timestep) time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype @@ -354,7 +354,7 @@ def _get_spatial_meshgrid(self, height, width, device_str): return grid_y, grid_x @torch.no_grad() - def forward(self, frame_indices, height, width, device): + def forward(self, frame_indices, height, width, device) -> torch.Tensor: batch_size = frame_indices.shape[0] num_frames = frame_indices.shape[1] diff --git a/src/diffusers/models/transformers/transformer_hidream_image.py b/src/diffusers/models/transformers/transformer_hidream_image.py index 703230562415..afb1f8b8edb9 100644 --- a/src/diffusers/models/transformers/transformer_hidream_image.py +++ b/src/diffusers/models/transformers/transformer_hidream_image.py @@ -295,7 +295,7 @@ def __init__( self._force_inference_output = _force_inference_output - def forward(self, hidden_states): + def forward(self, hidden_states) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]: bsz, seq_len, h = hidden_states.shape ### compute gating score hidden_states = hidden_states.view(-1, h) @@ -362,7 +362,7 @@ def __init__( ) self.num_activated_experts = num_activated_experts - def forward(self, x): + def forward(self, x) -> torch.Tensor: wtype = x.dtype identity = x orig_shape = x.shape @@ -409,7 +409,7 @@ 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): + def forward(self, caption) -> torch.Tensor: hidden_states = self.linear(caption) return hidden_states diff --git a/src/diffusers/models/transformers/transformer_hunyuan_video_framepack.py b/src/diffusers/models/transformers/transformer_hunyuan_video_framepack.py index 9a3dbc00f4ec..4f62a5e85841 100644 --- a/src/diffusers/models/transformers/transformer_hunyuan_video_framepack.py +++ b/src/diffusers/models/transformers/transformer_hunyuan_video_framepack.py @@ -47,7 +47,9 @@ def __init__(self, patch_size: int, patch_size_t: int, rope_dim: list[int], thet self.rope_dim = rope_dim self.theta = theta - def forward(self, frame_indices: torch.Tensor, height: int, width: int, device: torch.device): + def forward( + self, frame_indices: torch.Tensor, height: int, width: int, device: torch.device + ) -> tuple[torch.Tensor, torch.Tensor]: height = height // self.patch_size width = width // self.patch_size grid = torch.meshgrid( @@ -94,7 +96,7 @@ def forward( latents_clean: torch.Tensor | None = None, latents_clean_2x: torch.Tensor | None = None, latents_clean_4x: torch.Tensor | None = None, - ): + ) -> tuple[torch.Tensor | None, torch.Tensor | None, torch.Tensor | None]: if latents_clean is not None: latents_clean = self.proj(latents_clean) latents_clean = latents_clean.flatten(2).transpose(1, 2) diff --git a/src/diffusers/models/transformers/transformer_joyimage.py b/src/diffusers/models/transformers/transformer_joyimage.py index b17ddb05f799..d2fec631f1ac 100644 --- a/src/diffusers/models/transformers/transformer_joyimage.py +++ b/src/diffusers/models/transformers/transformer_joyimage.py @@ -350,7 +350,7 @@ def forward( self, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: timestep = self.timesteps_proj(timestep) time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype @@ -525,7 +525,7 @@ def forward( timestep: torch.Tensor, encoder_hidden_states: torch.Tensor = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`JoyImageEditTransformer3DModel`] forward method. @@ -539,6 +539,10 @@ def forward( 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, a [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a + `tuple` where the first element is the sample tensor. """ # handle multi-item input (b, n, c, t, h, w) is_multi_item = hidden_states.ndim == 6 diff --git a/src/diffusers/models/transformers/transformer_joyimage_edit_plus.py b/src/diffusers/models/transformers/transformer_joyimage_edit_plus.py index 4a13845faad3..a81027b5f40b 100644 --- a/src/diffusers/models/transformers/transformer_joyimage_edit_plus.py +++ b/src/diffusers/models/transformers/transformer_joyimage_edit_plus.py @@ -300,7 +300,7 @@ def forward( self, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: timestep = self.timesteps_proj(timestep) time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype diff --git a/src/diffusers/models/transformers/transformer_kandinsky.py b/src/diffusers/models/transformers/transformer_kandinsky.py index 88ef70d546c8..a908677457c9 100644 --- a/src/diffusers/models/transformers/transformer_kandinsky.py +++ b/src/diffusers/models/transformers/transformer_kandinsky.py @@ -165,7 +165,7 @@ def __init__(self, model_dim, time_dim, max_period=10000.0): self.activation = nn.SiLU() self.out_layer = nn.Linear(time_dim, time_dim, bias=True) - def forward(self, time): + def forward(self, time) -> torch.Tensor: 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))) @@ -178,7 +178,7 @@ def __init__(self, text_dim, model_dim): 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): + def forward(self, text_embed) -> torch.Tensor: text_embed = self.in_layer(text_embed) return self.norm(text_embed).type_as(text_embed) @@ -189,7 +189,7 @@ def __init__(self, visual_dim, model_dim, patch_size): self.patch_size = patch_size self.in_layer = nn.Linear(math.prod(patch_size) * visual_dim, model_dim) - def forward(self, x): + def forward(self, x) -> torch.Tensor: batch_size, duration, height, width, dim = x.shape x = ( x.view( @@ -218,7 +218,7 @@ def __init__(self, dim, max_pos=1024, max_period=10000.0): pos = torch.arange(max_pos, dtype=freq.dtype) self.register_buffer("args", torch.outer(pos, freq), persistent=False) - def forward(self, pos): + def forward(self, pos) -> torch.Tensor: args = self.args[pos] cosine = torch.cos(args) sine = torch.sin(args) @@ -239,7 +239,7 @@ def __init__(self, axes_dims, max_pos=(128, 128, 128), max_period=10000.0): 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)): + def forward(self, shape, pos, scale_factor=(1.0, 1.0, 1.0)) -> torch.Tensor: 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] @@ -268,7 +268,7 @@ def __init__(self, time_dim, model_dim, num_params): self.out_layer.weight.data.zero_() self.out_layer.bias.data.zero_() - def forward(self, x): + def forward(self, x) -> torch.Tensor: return self.out_layer(self.activation(x)) @@ -397,7 +397,7 @@ def __init__(self, dim, ff_dim): self.activation = nn.GELU() self.out_layer = nn.Linear(ff_dim, dim, bias=False) - def forward(self, x): + def forward(self, x) -> torch.Tensor: return self.out_layer(self.activation(self.in_layer(x))) @@ -409,7 +409,7 @@ def __init__(self, model_dim, time_dim, visual_dim, patch_size): 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): + def forward(self, visual_embed, text_embed, time_embed) -> torch.Tensor: shift, scale = torch.chunk(self.modulation(time_embed).unsqueeze(dim=1), 2, dim=-1) visual_embed = ( @@ -449,7 +449,7 @@ def __init__(self, model_dim, time_dim, ff_dim, head_dim): 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): + def forward(self, x, time_embed, rope) -> torch.Tensor: 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) @@ -478,7 +478,7 @@ def __init__(self, model_dim, time_dim, ff_dim, head_dim): 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): + def forward(self, visual_embed, text_embed, time_embed, rope, sparse_params) -> torch.Tensor: self_attn_params, cross_attn_params, ff_params = torch.chunk( self.visual_modulation(time_embed).unsqueeze(dim=1), 3, dim=-1 ) diff --git a/src/diffusers/models/transformers/transformer_longcat_image.py b/src/diffusers/models/transformers/transformer_longcat_image.py index 7b842c42132d..55d4771090e9 100644 --- a/src/diffusers/models/transformers/transformer_longcat_image.py +++ b/src/diffusers/models/transformers/transformer_longcat_image.py @@ -385,7 +385,7 @@ def __init__(self, embedding_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) - def forward(self, timestep, hidden_dtype): + def forward(self, timestep, hidden_dtype) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) diff --git a/src/diffusers/models/transformers/transformer_lumina2.py b/src/diffusers/models/transformers/transformer_lumina2.py index ba822730cb32..6d51b2925ef0 100644 --- a/src/diffusers/models/transformers/transformer_lumina2.py +++ b/src/diffusers/models/transformers/transformer_lumina2.py @@ -260,7 +260,9 @@ def _get_freqs_cis(self, ids: torch.Tensor) -> torch.Tensor: 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): + def forward( + self, hidden_states: torch.Tensor, attention_mask: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, list[int], list[int]]: batch_size, channels, height, width = hidden_states.shape p = self.patch_size post_patch_height, post_patch_width = height // p, width // p diff --git a/src/diffusers/models/transformers/transformer_mochi.py b/src/diffusers/models/transformers/transformer_mochi.py index a1a1f5e9c900..1e544890a4b8 100644 --- a/src/diffusers/models/transformers/transformer_mochi.py +++ b/src/diffusers/models/transformers/transformer_mochi.py @@ -42,7 +42,7 @@ def __init__(self, eps: float): self.eps = eps self.norm = RMSNorm(0, eps, False) - def forward(self, hidden_states, scale=None): + def forward(self, hidden_states, scale=None) -> torch.Tensor: hidden_states_dtype = hidden_states.dtype hidden_states = hidden_states.to(torch.float32) diff --git a/src/diffusers/models/transformers/transformer_nucleusmoe_image.py b/src/diffusers/models/transformers/transformer_nucleusmoe_image.py index f1c0eee949f7..5488c02c579f 100644 --- a/src/diffusers/models/transformers/transformer_nucleusmoe_image.py +++ b/src/diffusers/models/transformers/transformer_nucleusmoe_image.py @@ -127,7 +127,7 @@ def __init__(self, embedding_dim, use_additional_t_cond=False): if use_additional_t_cond: self.addition_t_embedding = nn.Embedding(2, embedding_dim) - def forward(self, timestep, hidden_states, addition_t_cond=None): + def forward(self, timestep, hidden_states, addition_t_cond=None) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_states.dtype)) diff --git a/src/diffusers/models/transformers/transformer_omnigen.py b/src/diffusers/models/transformers/transformer_omnigen.py index f860f5d5ab3e..9b3ee661bb10 100644 --- a/src/diffusers/models/transformers/transformer_omnigen.py +++ b/src/diffusers/models/transformers/transformer_omnigen.py @@ -150,7 +150,7 @@ def __init__( self.long_factor = rope_scaling["long_factor"] self.original_max_position_embeddings = original_max_position_embeddings - def forward(self, hidden_states, position_ids): + def forward(self, hidden_states, position_ids) -> tuple[torch.Tensor, torch.Tensor]: 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) diff --git a/src/diffusers/models/transformers/transformer_qwenimage.py b/src/diffusers/models/transformers/transformer_qwenimage.py index 5a242c8bd5c0..45722ca0d7b0 100644 --- a/src/diffusers/models/transformers/transformer_qwenimage.py +++ b/src/diffusers/models/transformers/transformer_qwenimage.py @@ -214,7 +214,7 @@ def __init__(self, embedding_dim, use_additional_t_cond=False): if use_additional_t_cond: self.addition_t_embedding = nn.Embedding(2, embedding_dim) - def forward(self, timestep, hidden_states, addition_t_cond=None): + def forward(self, timestep, hidden_states, addition_t_cond=None) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_states.dtype)) # (N, D) diff --git a/src/diffusers/models/transformers/transformer_sana_video.py b/src/diffusers/models/transformers/transformer_sana_video.py index db1f08a73a81..84c1fa18a9bb 100644 --- a/src/diffusers/models/transformers/transformer_sana_video.py +++ b/src/diffusers/models/transformers/transformer_sana_video.py @@ -263,7 +263,9 @@ def __init__(self, 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): + def forward( + self, timestep: torch.Tensor, guidance: torch.Tensor = None, hidden_dtype: torch.dtype = None + ) -> tuple[torch.Tensor, torch.Tensor]: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) diff --git a/src/diffusers/models/transformers/transformer_sd3.py b/src/diffusers/models/transformers/transformer_sd3.py index 9a56ca4e226d..eaafe057fd82 100644 --- a/src/diffusers/models/transformers/transformer_sd3.py +++ b/src/diffusers/models/transformers/transformer_sd3.py @@ -59,7 +59,7 @@ def __init__( 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): + def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor) -> 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) diff --git a/src/diffusers/models/transformers/transformer_skyreels_v2.py b/src/diffusers/models/transformers/transformer_skyreels_v2.py index a4a3aa3f6ffa..5687cfcf2964 100644 --- a/src/diffusers/models/transformers/transformer_skyreels_v2.py +++ b/src/diffusers/models/transformers/transformer_skyreels_v2.py @@ -355,7 +355,7 @@ def forward( timestep: torch.Tensor, encoder_hidden_states: torch.Tensor, encoder_hidden_states_image: torch.Tensor | None = None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: timestep = self.timesteps_proj(timestep) time_embedder_dtype = get_parameter_dtype(self.time_embedder) diff --git a/src/diffusers/models/transformers/transformer_temporal.py b/src/diffusers/models/transformers/transformer_temporal.py index 10bad499caf3..5141f07b8235 100644 --- a/src/diffusers/models/transformers/transformer_temporal.py +++ b/src/diffusers/models/transformers/transformer_temporal.py @@ -283,7 +283,7 @@ def forward( encoder_hidden_states: torch.Tensor | None = None, image_only_indicator: torch.Tensor | None = None, return_dict: bool = True, - ): + ) -> TransformerTemporalModelOutput | tuple[torch.Tensor]: """ Args: hidden_states (`torch.Tensor` of shape `(batch size, channel, height, width)`): diff --git a/src/diffusers/models/transformers/transformer_wan.py b/src/diffusers/models/transformers/transformer_wan.py index cf1b4ecc5d78..b9e0b75f46e8 100644 --- a/src/diffusers/models/transformers/transformer_wan.py +++ b/src/diffusers/models/transformers/transformer_wan.py @@ -333,7 +333,7 @@ def forward( encoder_hidden_states: torch.Tensor, encoder_hidden_states_image: torch.Tensor | None = None, timestep_seq_len: int | None = None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: timestep = self.timesteps_proj(timestep) if timestep_seq_len is not None: timestep = timestep.unflatten(0, (-1, timestep_seq_len)) diff --git a/src/diffusers/models/transformers/transformer_wan_animate.py b/src/diffusers/models/transformers/transformer_wan_animate.py index 084c3a2aed7d..3e188ce1ec51 100644 --- a/src/diffusers/models/transformers/transformer_wan_animate.py +++ b/src/diffusers/models/transformers/transformer_wan_animate.py @@ -809,7 +809,7 @@ def forward( encoder_hidden_states: torch.Tensor, encoder_hidden_states_image: torch.Tensor | None = None, timestep_seq_len: int | None = None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: timestep = self.timesteps_proj(timestep) if timestep_seq_len is not None: timestep = timestep.unflatten(0, (-1, timestep_seq_len)) diff --git a/src/diffusers/models/transformers/transformer_wan_animate_2.py b/src/diffusers/models/transformers/transformer_wan_animate_2.py index c19655e6952e..55cfb8548b61 100644 --- a/src/diffusers/models/transformers/transformer_wan_animate_2.py +++ b/src/diffusers/models/transformers/transformer_wan_animate_2.py @@ -544,7 +544,7 @@ def __init__(self, dim, out_dim, patch_size, eps=1e-6): # modulation self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5) - def forward(self, x, e): + def forward(self, x, e) -> torch.Tensor: shift, scale = (self.modulation + e.float().unsqueeze(1)).chunk(2, dim=1) x = self.head((self.norm(x.float()) * (1 + scale) + shift).type_as(x)) return x @@ -562,7 +562,7 @@ def __init__(self, in_dim, out_dim): torch.nn.LayerNorm(out_dim), ) - def forward(self, image_embeds): + def forward(self, image_embeds) -> torch.Tensor: clip_extra_context_tokens = self.proj(image_embeds) return clip_extra_context_tokens diff --git a/src/diffusers/models/transformers/transformer_z_image.py b/src/diffusers/models/transformers/transformer_z_image.py index 4cea745e5ed5..913075ff9ea5 100644 --- a/src/diffusers/models/transformers/transformer_z_image.py +++ b/src/diffusers/models/transformers/transformer_z_image.py @@ -60,7 +60,7 @@ def timestep_embedding(t, dim, max_period=10000): embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) return embedding - def forward(self, t): + def forward(self, t) -> torch.Tensor: 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) @@ -176,7 +176,7 @@ def __init__(self, dim: int, hidden_dim: int): def _forward_silu_gating(self, x1, x3): return F.silu(x1) * x3 - def forward(self, x): + def forward(self, x) -> torch.Tensor: return self.w2(self._forward_silu_gating(self.w1(x), self.w3(x))) @@ -232,7 +232,7 @@ def forward( noise_mask: torch.Tensor | None = None, adaln_noisy: torch.Tensor | None = None, adaln_clean: torch.Tensor | None = None, - ): + ) -> torch.Tensor: if self.modulation: seq_len = x.shape[1] @@ -291,7 +291,7 @@ def __init__(self, hidden_size, out_channels): 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): + def forward(self, x, c=None, noise_mask=None, c_noisy=None, c_clean=None) -> torch.Tensor: seq_len = x.shape[1] if noise_mask is not None: @@ -902,7 +902,7 @@ def forward( image_noise_mask: list[list[int]] | None = None, patch_size: int = 2, f_patch_size: int = 1, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`ZImageTransformer2DModel`] forward method. @@ -930,6 +930,10 @@ def forward( 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. + + Returns: + If `return_dict` is True, a [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a + `tuple` where the first element is the sample tensor. """ assert patch_size in self.all_patch_size and f_patch_size in self.all_f_patch_size omni_mode = isinstance(x[0], list) diff --git a/src/diffusers/models/unets/unet_3d_blocks.py b/src/diffusers/models/unets/unet_3d_blocks.py index e0d7f03bea3a..49ea270e67f4 100644 --- a/src/diffusers/models/unets/unet_3d_blocks.py +++ b/src/diffusers/models/unets/unet_3d_blocks.py @@ -936,7 +936,7 @@ def forward( self, hidden_states: torch.Tensor, image_only_indicator: torch.Tensor, - ): + ) -> torch.Tensor: hidden_states = self.resnets[0]( hidden_states, image_only_indicator=image_only_indicator, diff --git a/src/diffusers/models/unets/unet_kandinsky3.py b/src/diffusers/models/unets/unet_kandinsky3.py index 790d255101a4..571cf0c7c4d1 100644 --- a/src/diffusers/models/unets/unet_kandinsky3.py +++ b/src/diffusers/models/unets/unet_kandinsky3.py @@ -39,7 +39,7 @@ def __init__(self, encoder_hid_dim, cross_attention_dim): 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): + def forward(self, x) -> torch.Tensor: x = self.projection_linear(x) x = self.projection_norm(x) return x @@ -146,7 +146,9 @@ def set_default_attn_processor(self): """ self.set_attn_processor(AttnProcessor()) - def forward(self, sample, timestep, encoder_hidden_states=None, encoder_attention_mask=None, return_dict=True): + def forward( + self, sample, timestep, encoder_hidden_states=None, encoder_attention_mask=None, return_dict=True + ) -> Kandinsky3UNetOutput | tuple[torch.Tensor]: r""" Args: sample (`torch.Tensor`): Input sample. @@ -159,6 +161,10 @@ def forward(self, sample, timestep, encoder_hidden_states=None, encoder_attentio 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, a [`~models.unets.unet_kandinsky3.Kandinsky3UNetOutput`] is returned, otherwise a + `tuple` where the first element is the sample tensor. """ if encoder_attention_mask is not None: encoder_attention_mask = (1 - encoder_attention_mask.to(sample.dtype)) * -10000.0 @@ -260,7 +266,7 @@ def __init__( 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): + def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None) -> torch.Tensor: 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: @@ -328,7 +334,7 @@ def __init__( 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): + def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None) -> torch.Tensor: if self.self_attention: x = self.attentions[0](x, time_embed, image_mask=image_mask) @@ -348,7 +354,7 @@ def __init__(self, groups, normalized_shape, context_dim): self.context_mlp[1].weight.data.zero_() self.context_mlp[1].bias.data.zero_() - def forward(self, x, context): + def forward(self, x, context) -> torch.Tensor: context = self.context_mlp(context) for _ in range(len(x.shape[2:])): @@ -377,7 +383,7 @@ def __init__(self, in_channels, out_channels, time_embed_dim, kernel_size=3, nor else: self.down_sample = nn.Identity() - def forward(self, x, time_embed): + def forward(self, x, time_embed) -> torch.Tensor: x = self.group_norm(x, time_embed) x = self.activation(x) x = self.up_sample(x) @@ -418,7 +424,7 @@ def __init__( else nn.Identity() ) - def forward(self, x, time_embed): + def forward(self, x, time_embed) -> torch.Tensor: out = x for resnet_block in self.resnet_blocks: out = resnet_block(out, time_embed) @@ -441,7 +447,7 @@ def __init__(self, num_channels, context_dim, head_dim=64): out_bias=False, ) - def forward(self, x, context, context_mask=None): + def forward(self, x, context, context_mask=None) -> torch.Tensor: 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) @@ -467,7 +473,7 @@ def __init__(self, num_channels, time_embed_dim, context_dim=None, norm_groups=3 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): + def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None) -> torch.Tensor: 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) diff --git a/src/diffusers/models/unets/unet_motion_model.py b/src/diffusers/models/unets/unet_motion_model.py index faa181d9bfd5..23063bc498f1 100644 --- a/src/diffusers/models/unets/unet_motion_model.py +++ b/src/diffusers/models/unets/unet_motion_model.py @@ -484,7 +484,7 @@ def forward( encoder_attention_mask: torch.Tensor | None = None, cross_attention_kwargs: dict[str, Any] | 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.") @@ -1190,7 +1190,7 @@ def __init__( self.down_blocks = nn.ModuleList(down_blocks) self.up_blocks = nn.ModuleList(up_blocks) - def forward(self, sample): + def forward(self, sample) -> None: r""" Args: sample (`torch.Tensor`): Input sample. diff --git a/src/diffusers/models/unets/unet_stable_cascade.py b/src/diffusers/models/unets/unet_stable_cascade.py index e000fdc51e06..f101de40656e 100644 --- a/src/diffusers/models/unets/unet_stable_cascade.py +++ b/src/diffusers/models/unets/unet_stable_cascade.py @@ -46,7 +46,7 @@ def __init__(self, c, c_timestep, conds=[]): for cname in conds: setattr(self, f"mapper_{cname}", nn.Linear(c_timestep, c * 2)) - def forward(self, x, t): + def forward(self, x, t) -> torch.Tensor: 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): @@ -68,7 +68,7 @@ def __init__(self, c, c_skip=0, kernel_size=3, dropout=0.0): nn.Linear(c * 4, c), ) - def forward(self, x, x_skip=None): + def forward(self, x, x_skip=None) -> torch.Tensor: x_res = x x = self.norm(self.depthwise(x)) if x_skip is not None: @@ -84,7 +84,7 @@ def __init__(self, dim): self.gamma = nn.Parameter(torch.zeros(1, 1, 1, dim)) self.beta = nn.Parameter(torch.zeros(1, 1, 1, dim)) - def forward(self, x): + def forward(self, x) -> torch.Tensor: 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 @@ -99,7 +99,7 @@ def __init__(self, c, c_cond, nhead, self_attn=True, dropout=0.0): 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): + def forward(self, x, kv) -> torch.Tensor: kv = self.kv_mapper(kv) norm_x = self.norm(x) if self.self_attn: @@ -122,7 +122,7 @@ def __init__(self, in_channels, out_channels, mode, enabled=True): 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): + def forward(self, x) -> torch.Tensor: for block in self.blocks: x = block(x) return x @@ -547,7 +547,7 @@ def forward( sca=None, crp=None, return_dict=True, - ): + ) -> StableCascadeUNetOutput | tuple[torch.Tensor]: r""" Args: sample (`torch.Tensor`): The noisy input sample. @@ -569,6 +569,10 @@ def forward( 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. + + Returns: + If `return_dict` is True, a [`~models.unets.unet_stable_cascade.StableCascadeUNetOutput`] is returned, + otherwise a `tuple` where the first element is the sample tensor. """ if pixels is None: pixels = sample.new_zeros(sample.size(0), 3, 8, 8) diff --git a/src/diffusers/models/unets/uvit_2d.py b/src/diffusers/models/unets/uvit_2d.py index 317abe80b1eb..3c4e24642e42 100644 --- a/src/diffusers/models/unets/uvit_2d.py +++ b/src/diffusers/models/unets/uvit_2d.py @@ -148,7 +148,9 @@ def __init__( 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): + def forward( + self, input_ids, encoder_hidden_states, pooled_text_emb, micro_conds, cross_attention_kwargs=None + ) -> torch.Tensor: r""" Args: input_ids (`torch.LongTensor`): @@ -161,6 +163,10 @@ def forward(self, input_ids, encoder_hidden_states, pooled_text_emb, micro_conds 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`. + + Returns: + `torch.Tensor`: The logits over the codebook for each image token, of shape `(batch_size, codebook_size, + height, width)`. """ encoder_hidden_states = self.encoder_proj(encoder_hidden_states) encoder_hidden_states = self.encoder_proj_layer_norm(encoder_hidden_states) @@ -246,7 +252,7 @@ def __init__(self, in_channels, block_out_channels, vocab_size, elementwise_affi 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): + def forward(self, input_ids) -> torch.Tensor: embeddings = self.embeddings(input_ids) embeddings = self.layer_norm(embeddings) embeddings = embeddings.permute(0, 3, 1, 2) @@ -333,7 +339,7 @@ def __init__( else: self.upsample = None - def forward(self, x, pooled_text_emb, encoder_hidden_states, cross_attention_kwargs): + def forward(self, x, pooled_text_emb, encoder_hidden_states, cross_attention_kwargs) -> torch.Tensor: if self.downsample is not None: x = self.downsample(x) @@ -374,7 +380,7 @@ def __init__( 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): + def forward(self, x, cond_embeds) -> torch.Tensor: x_res = x x = self.depthwise(x) @@ -413,7 +419,7 @@ def __init__( 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): + def forward(self, hidden_states) -> torch.Tensor: 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) diff --git a/src/diffusers/modular_pipelines/anima/before_denoise.py b/src/diffusers/modular_pipelines/anima/before_denoise.py index ede832c9d80b..17a94811df6f 100644 --- a/src/diffusers/modular_pipelines/anima/before_denoise.py +++ b/src/diffusers/modular_pipelines/anima/before_denoise.py @@ -214,7 +214,9 @@ def _condition_prompt_embeds( return prompt_embeds.to(dtype=output_dtype, device=device) @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device conditioning_dtype = components.text_conditioner.dtype @@ -300,7 +302,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.batch_size = block_state.prompt_embeds.shape[0] @@ -363,7 +367,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) latent_height, latent_width = block_state.image_latents.shape[-2:] @@ -454,7 +460,9 @@ def prepare_latents( return randn_tensor(shape, generator=generator, device=device, dtype=dtype) @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.height = block_state.height or components.default_height @@ -520,7 +528,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -598,7 +608,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -685,7 +697,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/anima/decoders.py b/src/diffusers/modular_pipelines/anima/decoders.py index f1f4b475a4b8..3be9fc7c2ecb 100644 --- a/src/diffusers/modular_pipelines/anima/decoders.py +++ b/src/diffusers/modular_pipelines/anima/decoders.py @@ -46,7 +46,9 @@ 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: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) latents = block_state.latents.to(components.vae.dtype) @@ -107,7 +109,9 @@ def check_inputs(output_type): raise ValueError(f"Invalid output_type: {output_type}") @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.output_type) diff --git a/src/diffusers/modular_pipelines/anima/denoise.py b/src/diffusers/modular_pipelines/anima/denoise.py index d8146beefe72..86a3569648bc 100644 --- a/src/diffusers/modular_pipelines/anima/denoise.py +++ b/src/diffusers/modular_pipelines/anima/denoise.py @@ -40,7 +40,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[AnimaModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) timestep = t.expand(block_state.latents.shape[0]).to(block_state.dtype) @@ -117,7 +119,7 @@ def inputs(self) -> list[InputParam]: @torch.no_grad() def __call__( self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[AnimaModularPipeline, BlockState]: 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) @@ -151,7 +153,9 @@ 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): + def __call__( + self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[AnimaModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, t, block_state.latents, return_dict=False @@ -181,7 +185,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) num_warmup_steps = len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order diff --git a/src/diffusers/modular_pipelines/anima/encoders.py b/src/diffusers/modular_pipelines/anima/encoders.py index 68950f97be83..726a68280d82 100644 --- a/src/diffusers/modular_pipelines/anima/encoders.py +++ b/src/diffusers/modular_pipelines/anima/encoders.py @@ -235,7 +235,9 @@ def encode_prompt( } @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -379,7 +381,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/cosmos/after_decode.py b/src/diffusers/modular_pipelines/cosmos/after_decode.py index 7f8dd903d615..37c213fcc5ce 100644 --- a/src/diffusers/modular_pipelines/cosmos/after_decode.py +++ b/src/diffusers/modular_pipelines/cosmos/after_decode.py @@ -38,7 +38,9 @@ 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: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, 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: @@ -92,7 +94,9 @@ 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: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) output_path = str(block_state.output_path) fps = int(round(block_state.fps)) diff --git a/src/diffusers/modular_pipelines/cosmos/before_denoise.py b/src/diffusers/modular_pipelines/cosmos/before_denoise.py index 7e9d83fa6316..e6f9dbf05abb 100644 --- a/src/diffusers/modular_pipelines/cosmos/before_denoise.py +++ b/src/diffusers/modular_pipelines/cosmos/before_denoise.py @@ -48,7 +48,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, 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) @@ -124,7 +126,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device sampling_dtype = torch.float32 @@ -225,7 +229,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device sampling_dtype = torch.float32 @@ -324,7 +330,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device sampling_dtype = torch.float32 @@ -464,7 +472,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device has_image_condition = bool(block_state.vision_condition_indexes_for_pack) @@ -556,7 +566,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -652,7 +664,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -737,7 +751,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.cond_position_ids = torch.cat( [ @@ -838,7 +854,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, 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 @@ -927,7 +945,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, 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 @@ -970,7 +990,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device if components.config.use_native_flow_schedule: @@ -1045,7 +1067,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device sampling_dtype = torch.float32 @@ -1142,7 +1166,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device num_hints = len(block_state.control_latents) @@ -1244,7 +1270,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device components.scheduler.set_timesteps(block_state.num_inference_steps, device=device) @@ -1307,7 +1335,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/cosmos/before_encoder.py b/src/diffusers/modular_pipelines/cosmos/before_encoder.py index 2cdf68712cdf..e186559bbcd3 100644 --- a/src/diffusers/modular_pipelines/cosmos/before_encoder.py +++ b/src/diffusers/modular_pipelines/cosmos/before_encoder.py @@ -82,7 +82,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype diff --git a/src/diffusers/modular_pipelines/cosmos/decoders.py b/src/diffusers/modular_pipelines/cosmos/decoders.py index a76e48501d85..688c74629d77 100644 --- a/src/diffusers/modular_pipelines/cosmos/decoders.py +++ b/src/diffusers/modular_pipelines/cosmos/decoders.py @@ -44,7 +44,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("videos")] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -106,7 +108,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, 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.") @@ -171,7 +175,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) latents = block_state.latents vae_dtype = components.vae.dtype @@ -234,7 +240,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/cosmos/denoise.py b/src/diffusers/modular_pipelines/cosmos/denoise.py index 6a369357e96f..294213f48f4b 100644 --- a/src/diffusers/modular_pipelines/cosmos/denoise.py +++ b/src/diffusers/modular_pipelines/cosmos/denoise.py @@ -51,7 +51,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: device = components._execution_device block_state.vision_tokens = [block_state.latents.to(device=device, dtype=components.transformer.dtype)] block_state.vision_timesteps = torch.full( @@ -95,7 +97,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: 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( @@ -141,7 +145,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: 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( @@ -189,7 +195,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: denoiser_input_fields = block_state.denoiser_input_fields loop_input_fields = block_state.as_dict() has_sound = "sound_tokens" in loop_input_fields @@ -297,7 +305,9 @@ 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): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: block_state.latents = components.scheduler.step( block_state.velocity_vision.unsqueeze(0), t, block_state.latents.unsqueeze(0), return_dict=False )[0].squeeze(0) @@ -348,7 +358,9 @@ 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): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: velocity_vision = block_state.velocity_vision.float() latents = block_state.latents.float() @@ -404,7 +416,9 @@ 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): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: 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) @@ -455,7 +469,9 @@ 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): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: 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( @@ -515,7 +531,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) mixed_precision = Cosmos3MixedPrecisionConfig.resolve( components.transformer, @@ -686,7 +704,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: 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] + [ @@ -798,7 +818,9 @@ def _forward(components, static, vision_tokens, vision_timesteps, context_name, return preds_vision[-1] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: # 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 ( @@ -915,7 +937,9 @@ 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): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: block_state.latents = components.scheduler.step( block_state.velocity.unsqueeze(0), t, block_state.latents.unsqueeze(0), return_dict=False )[0].squeeze(0) diff --git a/src/diffusers/modular_pipelines/cosmos/encoders.py b/src/diffusers/modular_pipelines/cosmos/encoders.py index c29b6174edda..735bf7d61f90 100644 --- a/src/diffusers/modular_pipelines/cosmos/encoders.py +++ b/src/diffusers/modular_pipelines/cosmos/encoders.py @@ -162,7 +162,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.num_frames is None: block_state.num_frames = 189 @@ -273,7 +275,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) self._check_inputs(block_state) @@ -412,7 +416,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.num_frames is None: block_state.num_frames = 189 @@ -574,7 +580,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) self._check_inputs(block_state) if block_state.use_system_prompt is None: @@ -668,7 +676,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -772,7 +782,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -941,7 +953,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.vae.dtype @@ -1056,7 +1070,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py b/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py index e0bd89e77ded..ab6ebac2c8cd 100644 --- a/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py +++ b/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py @@ -882,7 +882,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: num_chunks = state.get("num_chunks") state.set("output_chunks", []) state.set("previous_output", None) diff --git a/src/diffusers/modular_pipelines/ernie_image/before_denoise.py b/src/diffusers/modular_pipelines/ernie_image/before_denoise.py index 034230632396..f1301ec3fcd3 100644 --- a/src/diffusers/modular_pipelines/ernie_image/before_denoise.py +++ b/src/diffusers/modular_pipelines/ernie_image/before_denoise.py @@ -118,7 +118,9 @@ def _expand(hiddens: list[torch.Tensor], num_images_per_prompt: int) -> list[tor return [h for h in hiddens for _ in range(num_images_per_prompt)] @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype @@ -177,7 +179,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device num_inference_steps = block_state.num_inference_steps @@ -243,7 +247,9 @@ def _check_inputs(components: ErnieImageModularPipeline, height: int, width: int ) @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype diff --git a/src/diffusers/modular_pipelines/ernie_image/decoders.py b/src/diffusers/modular_pipelines/ernie_image/decoders.py index d7d056b82584..b31c08987eb1 100644 --- a/src/diffusers/modular_pipelines/ernie_image/decoders.py +++ b/src/diffusers/modular_pipelines/ernie_image/decoders.py @@ -73,7 +73,9 @@ 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: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae device = block_state.latents.device diff --git a/src/diffusers/modular_pipelines/ernie_image/denoise.py b/src/diffusers/modular_pipelines/ernie_image/denoise.py index 3a2a2e312486..150947944dd5 100644 --- a/src/diffusers/modular_pipelines/ernie_image/denoise.py +++ b/src/diffusers/modular_pipelines/ernie_image/denoise.py @@ -59,7 +59,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[ErnieImageModularPipeline, BlockState]: 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) @@ -122,7 +124,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[ErnieImageModularPipeline, BlockState]: guider_inputs = { "text_bth": (block_state.text_bth, block_state.negative_text_bth), "text_lens": (block_state.text_lens, block_state.negative_text_lens), @@ -159,7 +163,9 @@ 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): + def __call__( + self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[ErnieImageModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, t, block_state.latents, return_dict=False @@ -208,7 +214,9 @@ 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: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, 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): diff --git a/src/diffusers/modular_pipelines/ernie_image/encoders.py b/src/diffusers/modular_pipelines/ernie_image/encoders.py index 161646d181be..ec016bf520b7 100644 --- a/src/diffusers/modular_pipelines/ernie_image/encoders.py +++ b/src/diffusers/modular_pipelines/ernie_image/encoders.py @@ -121,7 +121,9 @@ def _enhance_prompt( return pe_tokenizer.decode(generated_ids, skip_special_tokens=True).strip() @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -223,7 +225,9 @@ def _encode( return text_hiddens @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/flux/before_denoise.py b/src/diffusers/modular_pipelines/flux/before_denoise.py index 243f9e927d74..8e5a285dbe11 100644 --- a/src/diffusers/modular_pipelines/flux/before_denoise.py +++ b/src/diffusers/modular_pipelines/flux/before_denoise.py @@ -194,7 +194,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -283,7 +285,9 @@ def get_timesteps(scheduler, num_inference_steps, strength, device): return timesteps, num_inference_steps - t_start @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -395,7 +399,9 @@ def prepare_latents( return latents @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, 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 @@ -477,7 +483,9 @@ def check_inputs(image_latents, latents): 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: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(image_latents=block_state.image_latents, latents=block_state.latents) @@ -530,7 +538,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_embeds = block_state.prompt_embeds @@ -582,7 +592,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_embeds = block_state.prompt_embeds diff --git a/src/diffusers/modular_pipelines/flux/decoders.py b/src/diffusers/modular_pipelines/flux/decoders.py index 5fcde5008680..796bff258a6b 100644 --- a/src/diffusers/modular_pipelines/flux/decoders.py +++ b/src/diffusers/modular_pipelines/flux/decoders.py @@ -24,6 +24,7 @@ from ...video_processor import VaeImageProcessor 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 @@ -89,7 +90,7 @@ def intermediate_outputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/flux/denoise.py b/src/diffusers/modular_pipelines/flux/denoise.py index 490ef6d88f57..7f2e20ffcec9 100644 --- a/src/diffusers/modular_pipelines/flux/denoise.py +++ b/src/diffusers/modular_pipelines/flux/denoise.py @@ -92,7 +92,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: FluxModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[FluxModularPipeline, BlockState]: noise_pred = components.transformer( hidden_states=block_state.latents, timestep=t.flatten() / 1000, @@ -174,7 +174,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: FluxModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[FluxModularPipeline, BlockState]: latents = block_state.latents latent_model_input = latents image_latents = block_state.image_latents @@ -219,7 +219,9 @@ 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): + def __call__( + self, components: FluxModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[FluxModularPipeline, BlockState]: # Perform scheduler step using the predicted output latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( @@ -270,7 +272,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/flux/encoders.py b/src/diffusers/modular_pipelines/flux/encoders.py index 5f7e61a535b7..fb7b40024b0a 100644 --- a/src/diffusers/modular_pipelines/flux/encoders.py +++ b/src/diffusers/modular_pipelines/flux/encoders.py @@ -116,7 +116,9 @@ def check_inputs(height, width, vae_scale_factor): 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): + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.resized_image is None and block_state.image is None: @@ -169,7 +171,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam(name="processed_image")] @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState): + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: from ...pipelines.flux.pipeline_flux_kontext import PREFERRED_KONTEXT_RESOLUTIONS block_state = self.get_block_state(state) @@ -260,7 +264,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) image = getattr(block_state, self._image_input_name) @@ -451,7 +457,9 @@ def encode_prompt( return prompt_embeds, pooled_prompt_embeds @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: # Get inputs and intermediates block_state = self.get_block_state(state) self.check_inputs(block_state) diff --git a/src/diffusers/modular_pipelines/flux/inputs.py b/src/diffusers/modular_pipelines/flux/inputs.py index c513d237bee2..24cf107d0128 100644 --- a/src/diffusers/modular_pipelines/flux/inputs.py +++ b/src/diffusers/modular_pipelines/flux/inputs.py @@ -98,7 +98,9 @@ def check_inputs(self, components, block_state): ) @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: # TODO: consider adding negative embeddings? block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -187,7 +189,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam(name="image_width", type_hint=int, description="The width of the image latents"), ] - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs (height/width calculation, patchify, and batch expansion) @@ -246,7 +250,9 @@ def __call__(self, components: FluxModularPipeline, state: PipelineState) -> Pip class FluxKontextAdditionalInputsStep(FluxAdditionalInputsStep): model_name = "flux-kontext" - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs (height/width calculation, patchify, and batch expansion) @@ -334,7 +340,9 @@ def check_inputs(height, width, vae_scale_factor): 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: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) height = block_state.height or components.default_height diff --git a/src/diffusers/modular_pipelines/flux2/before_denoise.py b/src/diffusers/modular_pipelines/flux2/before_denoise.py index 87a6b568a258..5ae79ffaedd9 100644 --- a/src/diffusers/modular_pipelines/flux2/before_denoise.py +++ b/src/diffusers/modular_pipelines/flux2/before_denoise.py @@ -145,7 +145,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -293,7 +295,9 @@ def prepare_latents( return latents @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, 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 @@ -368,7 +372,9 @@ def _prepare_text_ids(x: torch.Tensor, t_coord: torch.Tensor | None = None): return torch.stack(out_ids) - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_embeds = block_state.prompt_embeds @@ -429,7 +435,9 @@ def _prepare_text_ids(x: torch.Tensor, t_coord: torch.Tensor | None = None): return torch.stack(out_ids) - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_embeds = block_state.prompt_embeds @@ -516,7 +524,9 @@ def _pack_latents(latents): return latents @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) image_latents = block_state.image_latents @@ -579,7 +589,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device batch_size = block_state.batch_size * block_state.num_images_per_prompt diff --git a/src/diffusers/modular_pipelines/flux2/decoders.py b/src/diffusers/modular_pipelines/flux2/decoders.py index 81f5ca00dc33..4d8490c45e32 100644 --- a/src/diffusers/modular_pipelines/flux2/decoders.py +++ b/src/diffusers/modular_pipelines/flux2/decoders.py @@ -26,6 +26,7 @@ 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__) # pylint: disable=invalid-name @@ -97,7 +98,7 @@ def _unpack_latents_with_ids(x: torch.Tensor, x_ids: torch.Tensor) -> torch.Tens return torch.stack(x_list, dim=0) @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) latents = block_state.latents @@ -162,7 +163,7 @@ def _unpatchify_latents(latents): return latents @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/flux2/denoise.py b/src/diffusers/modular_pipelines/flux2/denoise.py index 675f14b03c63..fa6877180057 100644 --- a/src/diffusers/modular_pipelines/flux2/denoise.py +++ b/src/diffusers/modular_pipelines/flux2/denoise.py @@ -106,7 +106,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: Flux2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[Flux2ModularPipeline, BlockState]: latents = block_state.latents latent_model_input = latents.to(components.transformer.dtype) img_ids = block_state.latent_ids @@ -195,7 +195,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: Flux2KleinModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[Flux2KleinModularPipeline, BlockState]: latents = block_state.latents latent_model_input = latents.to(components.transformer.dtype) img_ids = block_state.latent_ids @@ -310,7 +310,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: Flux2KleinModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[Flux2KleinModularPipeline, BlockState]: latents = block_state.latents latent_model_input = latents.to(components.transformer.dtype) img_ids = block_state.latent_ids @@ -379,7 +379,9 @@ 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): + def __call__( + self, components: Flux2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Flux2ModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, @@ -430,7 +432,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/flux2/encoders.py b/src/diffusers/modular_pipelines/flux2/encoders.py index 215f33b60ea8..df5e7c218ba7 100644 --- a/src/diffusers/modular_pipelines/flux2/encoders.py +++ b/src/diffusers/modular_pipelines/flux2/encoders.py @@ -152,7 +152,9 @@ def _get_mistral_3_prompt_embeds( return prompt_embeds @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -214,7 +216,9 @@ def check_inputs(block_state): 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: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: import io import requests @@ -353,7 +357,9 @@ def _get_qwen3_prompt_embeds( return prompt_embeds @torch.no_grad() - def __call__(self, components: Flux2KleinModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2KleinModularPipeline, state: PipelineState + ) -> tuple[Flux2KleinModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -495,7 +501,9 @@ def _get_qwen3_prompt_embeds( return prompt_embeds @torch.no_grad() - def __call__(self, components: Flux2KleinModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2KleinModularPipeline, state: PipelineState + ) -> tuple[Flux2KleinModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -585,7 +593,9 @@ def _encode_vae_image(self, vae: AutoencoderKLFlux2, image: torch.Tensor, genera return image_latents @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) condition_images = block_state.condition_images diff --git a/src/diffusers/modular_pipelines/flux2/inputs.py b/src/diffusers/modular_pipelines/flux2/inputs.py index 6bfe6aec97fd..ea868ca795cb 100644 --- a/src/diffusers/modular_pipelines/flux2/inputs.py +++ b/src/diffusers/modular_pipelines/flux2/inputs.py @@ -71,7 +71,9 @@ def intermediate_outputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.batch_size = block_state.prompt_embeds.shape[0] @@ -146,7 +148,9 @@ def intermediate_outputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.batch_size = block_state.prompt_embeds.shape[0] @@ -202,7 +206,9 @@ 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): + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) images = block_state.image diff --git a/src/diffusers/modular_pipelines/helios/before_denoise.py b/src/diffusers/modular_pipelines/helios/before_denoise.py index 593843d48272..bb60ce7f11f3 100644 --- a/src/diffusers/modular_pipelines/helios/before_denoise.py +++ b/src/diffusers/modular_pipelines/helios/before_denoise.py @@ -93,7 +93,9 @@ def check_inputs(self, components, block_state): ) @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -296,7 +298,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return outputs - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) for input_param in self._image_latent_inputs: @@ -400,7 +404,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -504,7 +510,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -614,7 +622,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) batch_size = block_state.batch_size @@ -719,7 +729,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, 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) @@ -757,7 +769,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) history_latents = block_state.history_latents @@ -809,7 +823,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) patch_size = components.transformer.config.patch_size diff --git a/src/diffusers/modular_pipelines/helios/decoders.py b/src/diffusers/modular_pipelines/helios/decoders.py index c448d36136e6..0ab55da4ac63 100644 --- a/src/diffusers/modular_pipelines/helios/decoders.py +++ b/src/diffusers/modular_pipelines/helios/decoders.py @@ -22,6 +22,7 @@ from ...video_processor import VideoProcessor 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 @@ -72,7 +73,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/helios/denoise.py b/src/diffusers/modular_pipelines/helios/denoise.py index 5fcf01a73ffc..32a6f0b11b79 100644 --- a/src/diffusers/modular_pipelines/helios/denoise.py +++ b/src/diffusers/modular_pipelines/helios/denoise.py @@ -132,7 +132,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: keep_first_frame = block_state.keep_first_frame history_sizes = block_state.history_sizes image_latents = block_state.image_latents @@ -218,7 +220,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: keep_first_frame = block_state.keep_first_frame history_sizes = block_state.history_sizes image_latents = block_state.image_latents @@ -257,7 +261,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: device = components._execution_device block_state.latents = randn_tensor( block_state.latent_shape, generator=block_state.generator, device=device, dtype=torch.float32 @@ -291,7 +297,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: device = components._execution_device batch_size, num_channels_latents, num_latent_frames, h_latent, w_latent = block_state.latent_shape @@ -337,7 +345,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: device = components._execution_device components.scheduler.set_timesteps( block_state.num_inference_steps, device=device, sigmas=block_state.sigmas, mu=block_state.mu @@ -392,7 +402,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: latents = block_state.latents timesteps = block_state.timesteps num_inference_steps = block_state.num_inference_steps @@ -511,7 +523,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: device = components._execution_device transformer_dtype = components.transformer.dtype latents = block_state.latents @@ -685,7 +699,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: # e. Collect denoised latents for this chunk block_state.latent_chunks.append(block_state.latents) @@ -733,7 +749,9 @@ def loop_intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.latent_chunks = [] @@ -848,7 +866,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: device = components._execution_device transformer_dtype = components.transformer.dtype latents = block_state.latents diff --git a/src/diffusers/modular_pipelines/helios/encoders.py b/src/diffusers/modular_pipelines/helios/encoders.py index ce11f1b58762..15c55579218f 100644 --- a/src/diffusers/modular_pipelines/helios/encoders.py +++ b/src/diffusers/modular_pipelines/helios/encoders.py @@ -160,7 +160,9 @@ def check_inputs(prompt, negative_prompt): ) @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt = block_state.prompt @@ -248,7 +250,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae @@ -336,7 +340,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/hunyuan_video1_5/before_denoise.py b/src/diffusers/modular_pipelines/hunyuan_video1_5/before_denoise.py index 4c02eb9dd084..f478c25b82e3 100644 --- a/src/diffusers/modular_pipelines/hunyuan_video1_5/before_denoise.py +++ b/src/diffusers/modular_pipelines/hunyuan_video1_5/before_denoise.py @@ -112,7 +112,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, 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) @@ -145,7 +147,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -202,7 +206,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype @@ -297,7 +303,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype diff --git a/src/diffusers/modular_pipelines/hunyuan_video1_5/decoders.py b/src/diffusers/modular_pipelines/hunyuan_video1_5/decoders.py index 630af85c1b10..ce5ecd4c9529 100644 --- a/src/diffusers/modular_pipelines/hunyuan_video1_5/decoders.py +++ b/src/diffusers/modular_pipelines/hunyuan_video1_5/decoders.py @@ -21,6 +21,7 @@ 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__) @@ -59,7 +60,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) latents = block_state.latents.to(components.vae.dtype) / components.vae.config.scaling_factor diff --git a/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py b/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py index 293fad57c93f..223d5e2eee9b 100644 --- a/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py +++ b/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py @@ -49,7 +49,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[HunyuanVideo15ModularPipeline, BlockState]: block_state.latent_model_input = torch.cat( [block_state.latents, block_state.cond_latents_concat, block_state.mask_concat], dim=1 ) @@ -131,7 +133,7 @@ def inputs(self) -> list[InputParam]: @torch.no_grad() def __call__( self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[HunyuanVideo15ModularPipeline, BlockState]: timestep = t.expand(block_state.latent_model_input.shape[0]).to(block_state.latent_model_input.dtype) # Step 1: Collect model inputs @@ -185,7 +187,9 @@ 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): + def __call__( + self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[HunyuanVideo15ModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, t, block_state.latents, return_dict=False @@ -220,7 +224,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( @@ -335,7 +341,7 @@ def inputs(self) -> list[InputParam]: @torch.no_grad() def __call__( self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[HunyuanVideo15ModularPipeline, BlockState]: timestep = t.expand(block_state.latent_model_input.shape[0]).to(block_state.latent_model_input.dtype) # MeanFlow timestep_r (lines 855-862) diff --git a/src/diffusers/modular_pipelines/hunyuan_video1_5/encoders.py b/src/diffusers/modular_pipelines/hunyuan_video1_5/encoders.py index 9d340cc88194..11f511d7e4e6 100644 --- a/src/diffusers/modular_pipelines/hunyuan_video1_5/encoders.py +++ b/src/diffusers/modular_pipelines/hunyuan_video1_5/encoders.py @@ -259,7 +259,9 @@ def encode_prompt( return prompt_embeds, prompt_embeds_mask, prompt_embeds_2, prompt_embeds_mask_2 @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype @@ -363,7 +365,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -424,7 +428,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ideogram4/before_denoise.py b/src/diffusers/modular_pipelines/ideogram4/before_denoise.py index 98be3b141aec..c29ee38e085b 100644 --- a/src/diffusers/modular_pipelines/ideogram4/before_denoise.py +++ b/src/diffusers/modular_pipelines/ideogram4/before_denoise.py @@ -178,7 +178,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_batch = block_state.text_features.shape[0] @@ -256,7 +258,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -351,7 +355,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -520,7 +526,9 @@ def _prepare_ids( return position_ids.to(device), segment_ids.to(device), indicator.to(device) @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ideogram4/decoders.py b/src/diffusers/modular_pipelines/ideogram4/decoders.py index bf5d69270b7c..710734ee7fb2 100644 --- a/src/diffusers/modular_pipelines/ideogram4/decoders.py +++ b/src/diffusers/modular_pipelines/ideogram4/decoders.py @@ -85,7 +85,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("images")] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) z = block_state.latents diff --git a/src/diffusers/modular_pipelines/ideogram4/denoise.py b/src/diffusers/modular_pipelines/ideogram4/denoise.py index 871db69d344c..db1a708c4315 100644 --- a/src/diffusers/modular_pipelines/ideogram4/denoise.py +++ b/src/diffusers/modular_pipelines/ideogram4/denoise.py @@ -56,7 +56,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Ideogram4ModularPipeline, BlockState]: # 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( @@ -150,7 +152,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Ideogram4ModularPipeline, BlockState]: transformer = components.transformer unconditional_transformer = components.unconditional_transformer @@ -200,7 +204,9 @@ 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): + def __call__( + self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Ideogram4ModularPipeline, BlockState]: block_state.latents = components.scheduler.step( block_state.noise_pred, t, block_state.latents, return_dict=False )[0] @@ -280,7 +286,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: @@ -344,7 +352,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) z = block_state.latents diff --git a/src/diffusers/modular_pipelines/ideogram4/encoders.py b/src/diffusers/modular_pipelines/ideogram4/encoders.py index 6e149fa8392e..fa7fc765ea9b 100644 --- a/src/diffusers/modular_pipelines/ideogram4/encoders.py +++ b/src/diffusers/modular_pipelines/ideogram4/encoders.py @@ -133,7 +133,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.prompt_upsampling: @@ -280,7 +282,9 @@ def _get_text_encoder_hidden_states( return [captured[i] for i in QWEN3_VL_ACTIVATION_LAYERS] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/krea2/before_denoise.py b/src/diffusers/modular_pipelines/krea2/before_denoise.py index 63810d30a903..17ed6a3cd376 100644 --- a/src/diffusers/modular_pipelines/krea2/before_denoise.py +++ b/src/diffusers/modular_pipelines/krea2/before_denoise.py @@ -138,7 +138,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_batch, seq_len, num_layers, dim = block_state.prompt_embeds.shape @@ -233,7 +235,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_batch, seq_len, num_layers, dim = block_state.prompt_embeds.shape @@ -317,7 +321,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -415,7 +421,9 @@ 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: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -491,7 +499,9 @@ 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: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -575,7 +585,9 @@ def prepare_position_ids(text_seq_len: int, grid_height: int, grid_width: int, d return torch.cat([text_ids, image_ids], dim=0) @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/krea2/decoders.py b/src/diffusers/modular_pipelines/krea2/decoders.py index fd308b5ef648..2c3fe4c0bf20 100644 --- a/src/diffusers/modular_pipelines/krea2/decoders.py +++ b/src/diffusers/modular_pipelines/krea2/decoders.py @@ -92,7 +92,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("images")] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/krea2/denoise.py b/src/diffusers/modular_pipelines/krea2/denoise.py index 88c6cdca7aba..ccb972ca74c1 100644 --- a/src/diffusers/modular_pipelines/krea2/denoise.py +++ b/src/diffusers/modular_pipelines/krea2/denoise.py @@ -55,7 +55,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Krea2ModularPipeline, BlockState]: 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 @@ -113,7 +115,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Krea2ModularPipeline, BlockState]: transformer = components.transformer latents = block_state.latents.to(transformer.dtype) @@ -190,7 +194,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Krea2ModularPipeline, BlockState]: transformer = components.transformer latents = block_state.latents.to(transformer.dtype) @@ -224,7 +230,9 @@ 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): + def __call__( + self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Krea2ModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, t, block_state.latents, return_dict=False @@ -261,7 +269,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: diff --git a/src/diffusers/modular_pipelines/krea2/encoders.py b/src/diffusers/modular_pipelines/krea2/encoders.py index 7640222e9ad2..a07f75305557 100644 --- a/src/diffusers/modular_pipelines/krea2/encoders.py +++ b/src/diffusers/modular_pipelines/krea2/encoders.py @@ -170,7 +170,9 @@ def _encode_prompt(self, components, prompt, max_sequence_length, device): return hidden_states, attention_mask @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -262,7 +264,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ltx/before_denoise.py b/src/diffusers/modular_pipelines/ltx/before_denoise.py index cd8b3ea82b82..5543742a4d46 100644 --- a/src/diffusers/modular_pipelines/ltx/before_denoise.py +++ b/src/diffusers/modular_pipelines/ltx/before_denoise.py @@ -131,7 +131,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.batch_size = block_state.prompt_embeds.shape[0] @@ -196,7 +198,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -288,7 +292,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -353,7 +359,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ltx/decoders.py b/src/diffusers/modular_pipelines/ltx/decoders.py index 8664dee25bfe..9d5238069bcf 100644 --- a/src/diffusers/modular_pipelines/ltx/decoders.py +++ b/src/diffusers/modular_pipelines/ltx/decoders.py @@ -23,7 +23,7 @@ from ...video_processor import VideoProcessor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import LTXVideoPachifier +from .modular_pipeline import LTXModularPipeline, LTXVideoPachifier logger = logging.get_logger(__name__) @@ -84,7 +84,7 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("videos")] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/ltx/denoise.py b/src/diffusers/modular_pipelines/ltx/denoise.py index b3ed86b51679..44a3c91d471b 100644 --- a/src/diffusers/modular_pipelines/ltx/denoise.py +++ b/src/diffusers/modular_pipelines/ltx/denoise.py @@ -49,7 +49,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTXModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) return components, block_state @@ -115,7 +117,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[LTXModularPipeline, BlockState]: 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 @@ -171,7 +173,9 @@ def description(self) -> str: ) @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTXModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, @@ -211,7 +215,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( @@ -275,7 +281,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTXModularPipeline, BlockState]: 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 @@ -342,7 +350,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[LTXModularPipeline, BlockState]: 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 @@ -411,7 +419,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTXModularPipeline, BlockState]: 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 diff --git a/src/diffusers/modular_pipelines/ltx/encoders.py b/src/diffusers/modular_pipelines/ltx/encoders.py index 55405ad0aefe..f8e1851eb5bc 100644 --- a/src/diffusers/modular_pipelines/ltx/encoders.py +++ b/src/diffusers/modular_pipelines/ltx/encoders.py @@ -148,7 +148,9 @@ def encode_prompt( return prompt_embeds, prompt_attention_mask, negative_prompt_embeds, negative_prompt_attention_mask @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -235,7 +237,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ltx2/before_denoise.py b/src/diffusers/modular_pipelines/ltx2/before_denoise.py index 81ffc28188ea..1878c01305fb 100644 --- a/src/diffusers/modular_pipelines/ltx2/before_denoise.py +++ b/src/diffusers/modular_pipelines/ltx2/before_denoise.py @@ -30,6 +30,7 @@ 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 LTX2ModularPipeline logger = logging.get_logger(__name__) @@ -318,7 +319,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) # `repeat_interleave` keeps each prompt's copies contiguous, matching how the latents are laid out @@ -376,7 +377,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -484,7 +485,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -580,7 +581,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -682,7 +683,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -779,7 +780,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -924,7 +925,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1225,7 +1226,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1452,7 +1453,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1523,7 +1524,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1632,7 +1633,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1739,7 +1740,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ltx2/decoders.py b/src/diffusers/modular_pipelines/ltx2/decoders.py index fc957a3f9925..cbeb1a66826d 100644 --- a/src/diffusers/modular_pipelines/ltx2/decoders.py +++ b/src/diffusers/modular_pipelines/ltx2/decoders.py @@ -30,6 +30,7 @@ from ...video_processor import VideoProcessor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import LTX2ModularPipeline logger = logging.get_logger(__name__) @@ -118,7 +119,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.latents = block_state.latents[:, : block_state.base_token_count] self.set_block_state(state, block_state) @@ -168,7 +169,7 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("videos")] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) decoder = components.diffusion_decoder @@ -258,7 +259,7 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("videos")] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae @@ -362,7 +363,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) audio_vae = components.audio_vae diff --git a/src/diffusers/modular_pipelines/ltx2/denoise.py b/src/diffusers/modular_pipelines/ltx2/denoise.py index b1c4657d4d04..2fb4951377a1 100644 --- a/src/diffusers/modular_pipelines/ltx2/denoise.py +++ b/src/diffusers/modular_pipelines/ltx2/denoise.py @@ -29,6 +29,7 @@ PipelineState, ) from ..modular_pipeline_utils import ComponentSpec, InputParam +from .modular_pipeline import LTX2ModularPipeline # Velocity-space helpers, mirrored from `diffusers.pipelines.ltx2.pipeline_ltx2.LTX2Pipeline` and redefined here @@ -94,7 +95,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) block_state.audio_latent_model_input = block_state.audio_latents.to(block_state.dtype) timestep = t.expand(block_state.latents.shape[0]) @@ -130,7 +133,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) block_state.audio_latent_model_input = block_state.audio_latents.to(block_state.dtype) timestep = t.expand(block_state.latents.shape[0]) @@ -172,7 +177,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) block_state.audio_latent_model_input = block_state.audio_latents.to(block_state.dtype) timestep = t.expand(block_state.latents.shape[0]) @@ -337,7 +344,9 @@ def inputs(self) -> list[InputParam]: return inputs @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: 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 @@ -461,7 +470,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: noise_pred_video = convert_x0_to_velocity( block_state.latents, block_state.noise_pred_video, i, components.scheduler ) @@ -514,7 +525,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: spatial_patch = components.transformer_spatial_patch_size temporal_patch = components.transformer_temporal_patch_size latent_num_frames = (block_state.num_frames - 1) // components.vae_temporal_compression_ratio + 1 @@ -587,7 +600,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: # Conditioning strengths run from 0 (always use the denoised sample) to 1 (always use the condition), with # intermediate values specifying how strongly to follow the condition. Applied in x0 space, not velocity # space (which is what the transformer outputs). @@ -632,7 +647,7 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/ltx2/encoders.py b/src/diffusers/modular_pipelines/ltx2/encoders.py index b261597c0f68..ae3d29b5cfe3 100644 --- a/src/diffusers/modular_pipelines/ltx2/encoders.py +++ b/src/diffusers/modular_pipelines/ltx2/encoders.py @@ -50,6 +50,7 @@ from ...video_processor import VideoProcessor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import LTX2ModularPipeline logger = logging.get_logger(__name__) @@ -221,7 +222,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not block_state.enable_prompt_enhancement: @@ -318,7 +319,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not block_state.enable_prompt_enhancement: @@ -427,7 +428,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not block_state.enable_prompt_enhancement: @@ -543,7 +544,7 @@ def check_inputs(block_state): 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, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -631,7 +632,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) padding_side = components.tokenizer.padding_side @@ -721,7 +722,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if getattr(components, "duration_head", None) is None: @@ -887,7 +888,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1006,7 +1007,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1228,7 +1229,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py b/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py index 247b9e88d761..94395a7c3ebe 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py +++ b/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py @@ -155,7 +155,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.keyframe_anchors = () @@ -371,7 +373,9 @@ def build_packed_sequence( return position_ids, token_tags, video_indices, audio_indices, text_indices, num_condition_rows, 0 @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -725,7 +729,9 @@ def build_ref2va_packed_sequence( ) @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -845,7 +851,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device patch_size = components.patch_size @@ -944,7 +952,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device patch_size = components.patch_size @@ -1012,7 +1022,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.latents = torch.cat([block_state.condition_rows, block_state.latents]) @@ -1088,7 +1100,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1222,7 +1236,9 @@ def build_row_timesteps( return torch.unique(row_timesteps, sorted=True, return_inverse=True) @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/minimax_h3/before_encoder.py b/src/diffusers/modular_pipelines/minimax_h3/before_encoder.py index ebae839fda65..dea178af4cb7 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/before_encoder.py +++ b/src/diffusers/modular_pipelines/minimax_h3/before_encoder.py @@ -112,7 +112,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) keyframes = [keyframe for keyframe in (block_state.image, block_state.last_image) if keyframe is not None] @@ -383,7 +385,9 @@ def _normalize_audio_condition( return torchaudio.transforms.Resample(sample_rate, target_sample_rate)(waveform) @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) # 1. Validate the request. diff --git a/src/diffusers/modular_pipelines/minimax_h3/decoders.py b/src/diffusers/modular_pipelines/minimax_h3/decoders.py index e5b624cdb6c4..bdc7adf7c5e7 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/decoders.py +++ b/src/diffusers/modular_pipelines/minimax_h3/decoders.py @@ -96,7 +96,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) patch_t, patch_h, patch_w = components.patch_size channels = components.vae_latent_channels @@ -169,7 +171,9 @@ 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: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -234,7 +238,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/minimax_h3/denoise.py b/src/diffusers/modular_pipelines/minimax_h3/denoise.py index 2f2ce59bfda6..efcf36c1130b 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/denoise.py +++ b/src/diffusers/modular_pipelines/minimax_h3/denoise.py @@ -108,7 +108,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: MiniMaxH3ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[MiniMaxH3ModularPipeline, BlockState]: transformer = getattr(components, self.transformer_name) unique_timesteps, timestep_indices = block_state.row_timestep_plan[i] # The layout tags its outputs `denoiser_input_fields`, and their names are the transformer's own parameter @@ -218,7 +220,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: MiniMaxH3ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[MiniMaxH3ModularPipeline, BlockState]: num_condition_video_rows = block_state.num_condition_video_rows num_condition_audio_rows = block_state.num_condition_audio_rows @@ -258,7 +262,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, 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): diff --git a/src/diffusers/modular_pipelines/minimax_h3/encoders.py b/src/diffusers/modular_pipelines/minimax_h3/encoders.py index 8ed4d548e9d1..80e84895f6e0 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/encoders.py +++ b/src/diffusers/modular_pipelines/minimax_h3/encoders.py @@ -181,7 +181,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not isinstance(block_state.prompt, str): raise ValueError( @@ -258,7 +260,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not isinstance(block_state.prompt, str): raise ValueError( @@ -354,7 +358,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -606,7 +612,9 @@ def emit(segment: tuple[list[int], list[int]]) -> None: return token_ids, token_tags @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not isinstance(block_state.prompt, str): raise ValueError( @@ -705,7 +713,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/minimax_music3/before_denoise.py b/src/diffusers/modular_pipelines/minimax_music3/before_denoise.py index be58527d681f..31f8411b5357 100644 --- a/src/diffusers/modular_pipelines/minimax_music3/before_denoise.py +++ b/src/diffusers/modular_pipelines/minimax_music3/before_denoise.py @@ -61,7 +61,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxMusic3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxMusic3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) num_frames = block_state.frame_hiddens.shape[1] diff --git a/src/diffusers/modular_pipelines/minimax_music3/decoders.py b/src/diffusers/modular_pipelines/minimax_music3/decoders.py index 2472542f1563..c74773862e59 100644 --- a/src/diffusers/modular_pipelines/minimax_music3/decoders.py +++ b/src/diffusers/modular_pipelines/minimax_music3/decoders.py @@ -73,7 +73,9 @@ def check_inputs(block_state): raise ValueError(f"Invalid output_type: {block_state.output_type}") @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxMusic3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxMusic3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) diff --git a/src/diffusers/modular_pipelines/minimax_music3/denoise.py b/src/diffusers/modular_pipelines/minimax_music3/denoise.py index d65155af07c4..0a709b8bf85b 100644 --- a/src/diffusers/modular_pipelines/minimax_music3/denoise.py +++ b/src/diffusers/modular_pipelines/minimax_music3/denoise.py @@ -74,7 +74,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int + ) -> tuple[MiniMaxMusic3ModularPipeline, BlockState]: device = components._execution_device chunk_start = block_state.chunk_starts[k] @@ -111,7 +113,9 @@ def inputs(self) -> list[InputParam]: return [InputParam.template("generator")] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int + ) -> tuple[MiniMaxMusic3ModularPipeline, BlockState]: device = components._execution_device latents = randn_tensor( @@ -148,7 +152,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int + ) -> tuple[MiniMaxMusic3ModularPipeline, BlockState]: device = components._execution_device sigmas = np.linspace(1.0, 1.0 / block_state.num_inference_steps, block_state.num_inference_steps) @@ -194,7 +200,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int + ) -> tuple[MiniMaxMusic3ModularPipeline, BlockState]: latents = block_state.latents timesteps = block_state.timesteps overlap = block_state.overlap @@ -246,7 +254,9 @@ def description(self) -> str: ) @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int + ) -> tuple[MiniMaxMusic3ModularPipeline, BlockState]: latents = block_state.latents if block_state.overlap > 0: latents[..., : block_state.overlap] = block_state.previous_latent[..., : block_state.overlap] @@ -292,7 +302,9 @@ def loop_intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxMusic3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxMusic3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.latent_chunks = [] diff --git a/src/diffusers/modular_pipelines/minimax_music3/encoders.py b/src/diffusers/modular_pipelines/minimax_music3/encoders.py index 6fae35be3ffe..fafd8cccc0ea 100644 --- a/src/diffusers/modular_pipelines/minimax_music3/encoders.py +++ b/src/diffusers/modular_pipelines/minimax_music3/encoders.py @@ -206,7 +206,9 @@ def check_inputs(block_state): raise ValueError(f"`lyrics` must be a non-empty string, got {block_state.lyrics!r}") @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxMusic3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxMusic3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -285,7 +287,9 @@ def check_inputs(block_state): raise ValueError(f"`audio_duration` must be positive, got {block_state.audio_duration}") @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxMusic3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxMusic3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) diff --git a/src/diffusers/modular_pipelines/modular_pipeline.py b/src/diffusers/modular_pipelines/modular_pipeline.py index e9e5463c1e72..0cca16877623 100644 --- a/src/diffusers/modular_pipelines/modular_pipeline.py +++ b/src/diffusers/modular_pipelines/modular_pipeline.py @@ -777,7 +777,7 @@ def select_block(self, **kwargs) -> str | None: raise NotImplementedError(f"Subclass {self.__class__.__name__} must implement the `select_block` method.") @torch.no_grad() - def __call__(self, pipeline, state: PipelineState) -> PipelineState: + def __call__(self, pipeline, state: PipelineState) -> tuple["ModularPipeline", PipelineState]: trigger_kwargs = {name: state.get(name) for name in self.block_trigger_inputs if name is not None} block_name = self.select_block(**trigger_kwargs) @@ -1149,7 +1149,7 @@ def outputs(self) -> list[str]: return self.intermediate_outputs @torch.no_grad() - def __call__(self, pipeline, state: PipelineState) -> PipelineState: + def __call__(self, pipeline, state: PipelineState) -> tuple["ModularPipeline", PipelineState]: for block_name, block in self.sub_blocks.items(): try: pipeline, state = block(pipeline, state) @@ -1533,7 +1533,7 @@ def loop_step(self, components, state: PipelineState, **kwargs): raise return components, state - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple["ModularPipeline", PipelineState]: raise NotImplementedError("`__call__` method needs to be implemented by the subclass") @property diff --git a/src/diffusers/modular_pipelines/qwenimage/before_denoise.py b/src/diffusers/modular_pipelines/qwenimage/before_denoise.py index b928bf7fce9e..e4720b16244d 100644 --- a/src/diffusers/modular_pipelines/qwenimage/before_denoise.py +++ b/src/diffusers/modular_pipelines/qwenimage/before_denoise.py @@ -196,7 +196,9 @@ def check_inputs(height, width, vae_scale_factor): raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -315,7 +317,9 @@ def check_inputs(height, width, vae_scale_factor): raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -433,7 +437,9 @@ def check_inputs(image_latents, latents): raise ValueError(f"`image_latents` must have 3 dimensions (patchified), but got {image_latents.ndim}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -515,7 +521,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -601,7 +609,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -683,7 +693,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -780,7 +790,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -875,7 +887,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.img_shapes = [ @@ -960,7 +974,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # for edit, image size can be different from the target size (height/width) @@ -1072,7 +1088,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae_scale_factor = components.vae_scale_factor @@ -1182,7 +1200,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1281,7 +1299,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) controlnet = unwrap_module(components.controlnet) diff --git a/src/diffusers/modular_pipelines/qwenimage/decoders.py b/src/diffusers/modular_pipelines/qwenimage/decoders.py index e4ccb6b8e047..9ba2365416b2 100644 --- a/src/diffusers/modular_pipelines/qwenimage/decoders.py +++ b/src/diffusers/modular_pipelines/qwenimage/decoders.py @@ -90,7 +90,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae_scale_factor = components.vae_scale_factor @@ -158,7 +160,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Unpack: (B, seq, C*4) -> (B, C, layers+1, H, W) @@ -225,7 +227,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("images", note="tensor output of the vae decoder.")] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # YiYi Notes: remove support for output_type = "latents', we can just skip decode/encode step in modular @@ -307,7 +311,7 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("images")] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) latents = block_state.latents @@ -409,7 +413,9 @@ def check_inputs(output_type): raise ValueError(f"Invalid output_type: {output_type}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.output_type) @@ -492,7 +498,9 @@ def check_inputs(output_type, mask_overlay_kwargs): raise ValueError("only support output_type 'pil' for mask overlay") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.output_type, block_state.mask_overlay_kwargs) diff --git a/src/diffusers/modular_pipelines/qwenimage/denoise.py b/src/diffusers/modular_pipelines/qwenimage/denoise.py index de8ea05c5047..7f271782f82b 100644 --- a/src/diffusers/modular_pipelines/qwenimage/denoise.py +++ b/src/diffusers/modular_pipelines/qwenimage/denoise.py @@ -57,7 +57,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: # one timestep block_state.timestep = t.expand(block_state.latents.shape[0]).to(block_state.latents.dtype) block_state.latent_model_input = block_state.latents @@ -88,7 +90,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: # one timestep block_state.latent_model_input = torch.cat([block_state.latents, block_state.image_latents], dim=1) @@ -138,7 +142,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[QwenImageModularPipeline, BlockState]: # cond_scale for the timestep (controlnet input) if isinstance(block_state.controlnet_keep[i], list): block_state.cond_scale = [ @@ -205,7 +211,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: guider_inputs = { "encoder_hidden_states": ( getattr(block_state, "prompt_embeds", None), @@ -290,7 +298,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: guider_inputs = { "encoder_hidden_states": ( getattr(block_state, "prompt_embeds", None), @@ -366,7 +376,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, @@ -419,7 +431,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: block_state.init_latents_proper = block_state.image_latents if i < len(block_state.timesteps) - 1: block_state.noise_timestep = block_state.timesteps[i + 1] @@ -466,7 +480,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/qwenimage/encoders.py b/src/diffusers/modular_pipelines/qwenimage/encoders.py index 5dade5716a49..1c414bb1f3c5 100644 --- a/src/diffusers/modular_pipelines/qwenimage/encoders.py +++ b/src/diffusers/modular_pipelines/qwenimage/encoders.py @@ -325,7 +325,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) images = block_state.image @@ -414,7 +416,9 @@ def check_inputs(resolution: int): raise ValueError(f"Resolution must be 1024 or 640 but is {resolution}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(resolution=block_state.resolution) @@ -505,7 +509,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) images = block_state.image @@ -619,7 +625,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -743,7 +751,9 @@ def check_inputs(prompt, negative_prompt, max_sequence_length): raise ValueError(f"`max_sequence_length` cannot be greater than 1024 but is {max_sequence_length}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -874,7 +884,9 @@ def check_inputs(prompt, negative_prompt): raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.prompt, block_state.negative_prompt) @@ -1002,7 +1014,9 @@ def check_inputs(prompt, negative_prompt): raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.prompt, block_state.negative_prompt) @@ -1132,7 +1146,9 @@ def check_inputs(height, width, vae_scale_factor): raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -1228,7 +1244,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) width, height = block_state.resized_image[0].size @@ -1312,7 +1330,9 @@ def check_inputs(height, width, vae_scale_factor): raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -1387,7 +1407,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) width, height = block_state.resized_image[0].size @@ -1459,7 +1481,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) image = block_state.resized_image @@ -1563,7 +1587,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [self._output] # default is "image_latents" @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1670,7 +1696,9 @@ def check_inputs(height, width, vae_scale_factor): raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.height, block_state.width, components.vae_scale_factor) @@ -1769,7 +1797,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Permute: (B, C, 1, H, W) -> (B, 1, C, H, W) diff --git a/src/diffusers/modular_pipelines/qwenimage/inputs.py b/src/diffusers/modular_pipelines/qwenimage/inputs.py index 38a49e07345f..2d35d833c60f 100644 --- a/src/diffusers/modular_pipelines/qwenimage/inputs.py +++ b/src/diffusers/modular_pipelines/qwenimage/inputs.py @@ -205,7 +205,9 @@ def check_inputs( ): raise ValueError("`negative_prompt_embeds_mask` must have the same batch size as `prompt_embeds`") - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -411,7 +413,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return outputs - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs @@ -626,7 +630,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return outputs - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs @@ -852,7 +858,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return outputs - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs @@ -969,7 +977,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) if isinstance(components.controlnet, QwenImageMultiControlNetModel): diff --git a/src/diffusers/modular_pipelines/stable_diffusion_3/before_denoise.py b/src/diffusers/modular_pipelines/stable_diffusion_3/before_denoise.py index 5007faa12f67..09f5a5f4f54c 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_3/before_denoise.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_3/before_denoise.py @@ -190,7 +190,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -285,7 +287,9 @@ def get_timesteps(scheduler, num_inference_steps, strength): return timesteps, num_inference_steps - t_start @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -372,7 +376,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device batch_size = block_state.batch_size * block_state.num_images_per_prompt @@ -446,7 +452,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) latent_timestep = block_state.timesteps[:1].repeat(block_state.latents.shape[0]) block_state.initial_noise = block_state.latents diff --git a/src/diffusers/modular_pipelines/stable_diffusion_3/decoders.py b/src/diffusers/modular_pipelines/stable_diffusion_3/decoders.py index b1a8df1c7fa7..f2665ed06e47 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_3/decoders.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_3/decoders.py @@ -21,6 +21,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import StableDiffusion3ModularPipeline logger = logging.get_logger(__name__) @@ -62,7 +63,7 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("images", type_hint=list[PIL.Image.Image] | torch.Tensor)] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py b/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py index 33bd98095d8a..cde6ace66245 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py @@ -102,7 +102,7 @@ def __call__( block_state: BlockState, i: int, t: torch.Tensor, - ) -> PipelineState: + ) -> tuple[StableDiffusion3ModularPipeline, BlockState]: do_cfg = block_state.negative_prompt_embeds is not None guider_inputs = { @@ -174,7 +174,7 @@ def __call__( block_state: BlockState, i: int, t: torch.Tensor, - ): + ) -> tuple[StableDiffusion3ModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, @@ -207,7 +207,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, 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, diff --git a/src/diffusers/modular_pipelines/stable_diffusion_3/encoders.py b/src/diffusers/modular_pipelines/stable_diffusion_3/encoders.py index bef2a0f812ec..ea163b3b9cf9 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_3/encoders.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_3/encoders.py @@ -364,7 +364,9 @@ def check_inputs(height, width, vae_scale_factor, patch_size): raise ValueError(f"Width must be divisible by {vae_scale_factor * patch_size} but is {width}") @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState): + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.image is None: @@ -432,7 +434,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) image = getattr(block_state, self._image_input_name) @@ -526,7 +530,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device diff --git a/src/diffusers/modular_pipelines/stable_diffusion_3/inputs.py b/src/diffusers/modular_pipelines/stable_diffusion_3/inputs.py index d7e88b571612..8dd997b07b57 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_3/inputs.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_3/inputs.py @@ -187,7 +187,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.batch_size = block_state.prompt_embeds.shape[0] @@ -282,7 +284,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) for input_name in self._image_latent_inputs: diff --git a/src/diffusers/modular_pipelines/stable_diffusion_xl/before_denoise.py b/src/diffusers/modular_pipelines/stable_diffusion_xl/before_denoise.py index 92c74219bd06..f88256fbd9f4 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_xl/before_denoise.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_xl/before_denoise.py @@ -334,7 +334,9 @@ def check_inputs(self, components, block_state): ) @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -482,7 +484,9 @@ def get_timesteps(components, num_inference_steps, strength, device, denoising_s return timesteps, num_inference_steps @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -567,7 +571,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -823,7 +829,9 @@ def prepare_mask_latents( return mask, masked_image_latents @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.dtype = block_state.dtype if block_state.dtype is not None else components.vae.dtype @@ -926,7 +934,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.dtype = block_state.dtype if block_state.dtype is not None else components.vae.dtype @@ -1027,7 +1037,9 @@ def prepare_latents(comp, batch_size, num_channels_latents, height, width, dtype return latents @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.dtype is None: @@ -1217,7 +1229,9 @@ def get_guidance_scale_embedding( return emb @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -1395,7 +1409,9 @@ def get_guidance_scale_embedding( return emb @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -1560,7 +1576,9 @@ def prepare_control_image( return image @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) # (1) prepare controlnet inputs @@ -1789,7 +1807,9 @@ def prepare_control_image( return image @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) controlnet = unwrap_module(components.controlnet) diff --git a/src/diffusers/modular_pipelines/stable_diffusion_xl/decoders.py b/src/diffusers/modular_pipelines/stable_diffusion_xl/decoders.py index b4f15df8b411..ea6fdfa1df98 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_xl/decoders.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_xl/decoders.py @@ -27,6 +27,7 @@ PipelineState, ) from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import StableDiffusionXLModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -84,7 +85,7 @@ def upcast_vae(components): components.vae.to(dtype=torch.float32) @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not block_state.output_type == "latent": @@ -182,7 +183,7 @@ def inputs(self) -> list[tuple[str, Any]]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.padding_mask_crop is not None and block_state.crops_coords is not None: diff --git a/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py b/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py index ec344fe0ad37..16a8b236ce2e 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py @@ -66,7 +66,9 @@ def inputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: block_state.scaled_latents = components.scheduler.scale_model_input(block_state.latents, t) return components, block_state @@ -131,7 +133,9 @@ def check_inputs(components, block_state): ) @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: self.check_inputs(components, block_state) block_state.scaled_latents = components.scheduler.scale_model_input(block_state.latents, t) @@ -198,7 +202,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int - ) -> PipelineState: + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: # Map the keys we'll see on each `guider_state_batch` (e.g. guider_state_batch.prompt_embeds) # to the corresponding (cond, uncond) fields on block_state. (e.g. block_state.prompt_embeds, block_state.negative_prompt_embeds) guider_inputs = { @@ -351,7 +355,9 @@ def prepare_extra_kwargs(func, exclude_kwargs=[], **kwargs): return extra_kwargs @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: extra_controlnet_kwargs = self.prepare_extra_kwargs( components.controlnet.forward, **block_state.controlnet_kwargs ) @@ -508,7 +514,9 @@ def prepare_extra_kwargs(func, exclude_kwargs=[], **kwargs): return extra_kwargs @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: # Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline block_state.extra_step_kwargs = self.prepare_extra_kwargs( components.scheduler.step, generator=block_state.generator, eta=block_state.eta @@ -603,7 +611,9 @@ def check_inputs(self, components, block_state): raise ValueError(f"noise is required for this step {self.__class__.__name__}") @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: self.check_inputs(components, block_state) # Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline @@ -684,7 +694,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.disable_guidance = True if components.unet.config.time_cond_proj_dim is not None else False diff --git a/src/diffusers/modular_pipelines/stable_diffusion_xl/encoders.py b/src/diffusers/modular_pipelines/stable_diffusion_xl/encoders.py index 26e5524309f1..381badb2204d 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_xl/encoders.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_xl/encoders.py @@ -188,7 +188,9 @@ def prepare_ip_adapter_image_embeds( return ip_adapter_image_embeds @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.prepare_unconditional_embeds = components.guider.num_conditions > 1 @@ -534,7 +536,9 @@ def encode_prompt( return prompt_embeds, negative_prompt_embeds, pooled_prompt_embeds, negative_pooled_prompt_embeds @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: # Get inputs and intermediates block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -655,7 +659,9 @@ def _encode_vae_image(self, components, image: torch.Tensor, generator: torch.Ge return image_latents @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.preprocess_kwargs = block_state.preprocess_kwargs or {} block_state.device = components._execution_device @@ -825,7 +831,9 @@ def prepare_mask_latents( return mask, masked_image_latents @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.dtype = block_state.dtype if block_state.dtype is not None else components.vae.dtype diff --git a/src/diffusers/modular_pipelines/wan/before_denoise.py b/src/diffusers/modular_pipelines/wan/before_denoise.py index 1d90c20d8124..0e7aec5364fb 100644 --- a/src/diffusers/modular_pipelines/wan/before_denoise.py +++ b/src/diffusers/modular_pipelines/wan/before_denoise.py @@ -245,7 +245,9 @@ def check_inputs(self, components, block_state): ) @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -355,7 +357,9 @@ def inputs(self) -> list[InputParam]: return inputs - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs (height/width calculation, patchify, and batch expansion) @@ -436,7 +440,9 @@ def check_inputs(block_state): "Generating multiple videos per prompt is not yet supported. This may be supported in the future." ) - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -469,7 +475,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -568,7 +576,9 @@ def prepare_latents( return latents @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) diff --git a/src/diffusers/modular_pipelines/wan/decoders.py b/src/diffusers/modular_pipelines/wan/decoders.py index 529c9291c250..f820d86339ef 100644 --- a/src/diffusers/modular_pipelines/wan/decoders.py +++ b/src/diffusers/modular_pipelines/wan/decoders.py @@ -24,6 +24,7 @@ from ...video_processor import VideoProcessor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import WanModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -54,7 +55,7 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.latents = block_state.latents[:, :, block_state.num_reference_images :] @@ -107,7 +108,7 @@ def intermediate_outputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae_dtype = components.vae.dtype diff --git a/src/diffusers/modular_pipelines/wan/denoise.py b/src/diffusers/modular_pipelines/wan/denoise.py index 4beb0425a5fc..3036b33868d1 100644 --- a/src/diffusers/modular_pipelines/wan/denoise.py +++ b/src/diffusers/modular_pipelines/wan/denoise.py @@ -64,7 +64,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[WanModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) return components, block_state @@ -104,7 +106,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[WanModularPipeline, BlockState]: block_state.latent_model_input = torch.cat( [block_state.latents, block_state.image_condition_latents], dim=1 ).to(block_state.dtype) @@ -181,7 +185,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[WanModularPipeline, BlockState]: components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) # The guider splits model inputs into separate batches for conditional/unconditional predictions. @@ -315,7 +319,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[WanModularPipeline, BlockState]: boundary_timestep = components.config.boundary_ratio * components.num_train_timesteps if t >= boundary_timestep: block_state.current_model = components.transformer @@ -391,7 +395,9 @@ def description(self) -> str: ) @torch.no_grad() - def __call__(self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[WanModularPipeline, BlockState]: # Perform scheduler step using the predicted output latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( @@ -441,7 +447,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/wan/encoders.py b/src/diffusers/modular_pipelines/wan/encoders.py index 3bebe341a511..6eaf42ce4f6a 100644 --- a/src/diffusers/modular_pipelines/wan/encoders.py +++ b/src/diffusers/modular_pipelines/wan/encoders.py @@ -273,7 +273,9 @@ def encode_prompt( return prompt_embeds, negative_prompt_embeds @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: # Get inputs and intermediates block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -319,7 +321,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("resized_image", type_hint=PIL.Image.Image), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) max_area = block_state.height * block_state.width @@ -356,7 +360,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("resized_last_image", type_hint=PIL.Image.Image), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) height = block_state.resized_image.height @@ -403,7 +409,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("image_embeds", type_hint=torch.Tensor, description="The image embeddings"), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -448,7 +456,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("image_embeds", type_hint=torch.Tensor, description="The image embeddings"), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -521,7 +531,9 @@ def check_inputs(components, block_state): f"`num_frames` has to be greater than 0, and (num_frames - 1) must be divisible by {components.vae_scale_factor_temporal}, but got {block_state.num_frames}." ) - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -822,7 +834,9 @@ def prepare_masks(components, mask, reference_images): return torch.stack(mask_list) @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -897,7 +911,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("image_condition_latents", type_hint=torch.Tensor | None), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) batch_size, _, _, latent_height, latent_width = block_state.first_frame_latents.shape @@ -976,7 +992,9 @@ def check_inputs(components, block_state): f"`num_frames` has to be greater than 0, and (num_frames - 1) must be divisible by {components.vae_scale_factor_temporal}, but got {block_state.num_frames}." ) - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -1046,7 +1064,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("image_condition_latents", type_hint=torch.Tensor | None), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) batch_size, _, _, latent_height, latent_width = block_state.first_last_frame_latents.shape diff --git a/src/diffusers/modular_pipelines/wan_animate_2/before_denoise.py b/src/diffusers/modular_pipelines/wan_animate_2/before_denoise.py index 0ad038e8ccaa..7cf13afea769 100644 --- a/src/diffusers/modular_pipelines/wan_animate_2/before_denoise.py +++ b/src/diffusers/modular_pipelines/wan_animate_2/before_denoise.py @@ -20,6 +20,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import InputParam, OutputParam +from .modular_pipeline import WanAnimate2ModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -82,7 +83,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) latent_height, latent_width = block_state.reference_image_latents.shape[-2:] diff --git a/src/diffusers/modular_pipelines/wan_animate_2/decoders.py b/src/diffusers/modular_pipelines/wan_animate_2/decoders.py index 5317ee85e67a..cc9e6fb4f6e2 100644 --- a/src/diffusers/modular_pipelines/wan_animate_2/decoders.py +++ b/src/diffusers/modular_pipelines/wan_animate_2/decoders.py @@ -20,6 +20,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import WanAnimate2ModularPipeline from .video_processor import WanAnimate2VideoProcessor @@ -87,7 +88,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) video = torch.cat(block_state.segment_frames, dim=2)[:, :, : block_state.real_frame_len] diff --git a/src/diffusers/modular_pipelines/wan_animate_2/denoise.py b/src/diffusers/modular_pipelines/wan_animate_2/denoise.py index d96b8f814239..1e18b3a09eeb 100644 --- a/src/diffusers/modular_pipelines/wan_animate_2/denoise.py +++ b/src/diffusers/modular_pipelines/wan_animate_2/denoise.py @@ -28,6 +28,7 @@ from ..modular_pipeline import BlockState, LoopSequentialPipelineBlocks, ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam from .encoders import encode_vae, get_i2v_mask +from .modular_pipeline import WanAnimate2ModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -115,7 +116,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: device = components._execution_device latent_height, latent_width = block_state.reference_image_latents.shape[-2:] @@ -190,7 +191,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: # `block_state.out_frames` is seeded by the loop wrapper and written by the decode step of the # previous iteration. device = components._execution_device @@ -270,7 +271,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: device = components._execution_device block_state.latents = randn_tensor( @@ -319,7 +320,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: device = components._execution_device components.scheduler.set_timesteps(block_state.num_inference_steps, device=device) @@ -400,7 +401,7 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: device = components._execution_device transformer_dtype = components.transformer.dtype @@ -527,7 +528,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: transformer_dtype = components.transformer.dtype guider_inputs = { @@ -673,7 +674,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: latents = block_state.latents.to(torch.float32) # The first latent frame is the reference image's slot, not video content. out_frames = decode_vae(components.vae, latents[:, 1:]) @@ -730,7 +731,7 @@ def loop_intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Seed the loop-carried state: `segment_frames` collects each segment's decoded frames (the decode step diff --git a/src/diffusers/modular_pipelines/wan_animate_2/encoders.py b/src/diffusers/modular_pipelines/wan_animate_2/encoders.py index 21b70f636f7d..ffba4c672691 100644 --- a/src/diffusers/modular_pipelines/wan_animate_2/encoders.py +++ b/src/diffusers/modular_pipelines/wan_animate_2/encoders.py @@ -25,6 +25,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import WanAnimate2ModularPipeline from .video_processor import WanAnimate2VideoProcessor @@ -169,7 +170,7 @@ def check_inputs(block_state): raise ValueError(f"`prompt` has to be of type `str` but is {type(block_state.prompt)}") @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -269,7 +270,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -387,7 +388,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -473,7 +474,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -523,7 +524,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -578,7 +579,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/z_image/before_denoise.py b/src/diffusers/modular_pipelines/z_image/before_denoise.py index 5216529d460f..e96d80b4bfca 100644 --- a/src/diffusers/modular_pipelines/z_image/before_denoise.py +++ b/src/diffusers/modular_pipelines/z_image/before_denoise.py @@ -258,7 +258,9 @@ def check_inputs(self, components, block_state): ) @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -366,7 +368,9 @@ def inputs(self) -> list[InputParam]: return inputs - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs (height/width calculation, patchify, and batch expansion) @@ -467,7 +471,9 @@ def prepare_latents( return latents @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -524,7 +530,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -580,7 +588,9 @@ def check_inputs(self, components, block_state): raise ValueError(f"Strength must be between 0.0 and 1.0, but got {block_state.strength}") @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -613,7 +623,9 @@ def inputs(self) -> list[InputParam]: InputParam("timesteps", required=True), ] - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) latent_timestep = block_state.timesteps[:1].repeat(block_state.latents.shape[0]) diff --git a/src/diffusers/modular_pipelines/z_image/decoders.py b/src/diffusers/modular_pipelines/z_image/decoders.py index 353253102376..21e307d53df4 100644 --- a/src/diffusers/modular_pipelines/z_image/decoders.py +++ b/src/diffusers/modular_pipelines/z_image/decoders.py @@ -24,6 +24,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import ZImageModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -74,7 +75,7 @@ def intermediate_outputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae_dtype = components.vae.dtype diff --git a/src/diffusers/modular_pipelines/z_image/denoise.py b/src/diffusers/modular_pipelines/z_image/denoise.py index 863df312389a..899800a5019a 100644 --- a/src/diffusers/modular_pipelines/z_image/denoise.py +++ b/src/diffusers/modular_pipelines/z_image/denoise.py @@ -63,7 +63,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: ZImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[ZImageModularPipeline, BlockState]: latents = block_state.latents.unsqueeze(2).to( block_state.dtype ) # [batch_size, num_channels, 1, height, width] @@ -152,7 +154,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: ZImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[ZImageModularPipeline, BlockState]: components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) # The guider splits model inputs into separate batches for conditional/unconditional predictions. @@ -219,7 +221,9 @@ def description(self) -> str: ) @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: ZImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[ZImageModularPipeline, BlockState]: # Perform scheduler step using the predicted output latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( @@ -269,7 +273,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/z_image/encoders.py b/src/diffusers/modular_pipelines/z_image/encoders.py index 06deb8236893..c3ab717078e1 100644 --- a/src/diffusers/modular_pipelines/z_image/encoders.py +++ b/src/diffusers/modular_pipelines/z_image/encoders.py @@ -244,7 +244,9 @@ def encode_prompt( return prompt_embeds, negative_prompt_embeds @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: # Get inputs and intermediates block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -316,7 +318,9 @@ def check_inputs(components, block_state): f"`height` and `width` have to be divisible by {components.vae_scale_factor_spatial} but are {block_state.height} and {block_state.width}." ) - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) diff --git a/src/diffusers/pipelines/ace_step/pipeline_ace_step.py b/src/diffusers/pipelines/ace_step/pipeline_ace_step.py index b11e08208eb1..8f8e2e1be95b 100644 --- a/src/diffusers/pipelines/ace_step/pipeline_ace_step.py +++ b/src/diffusers/pipelines/ace_step/pipeline_ace_step.py @@ -831,7 +831,7 @@ def __call__( cfg_interval_end: float = 1.0, timesteps: Optional[List[float]] = None, attention_kwargs: Optional[dict] = None, - ): + ) -> AudioPipelineOutput | tuple: r""" The call function to the pipeline for music generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff.py index b9e5b40b65ff..fe97aedf627a 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff.py @@ -595,7 +595,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], decode_chunk_size: int = 16, **kwargs, - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff_controlnet.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff_controlnet.py index a9630cc3c00f..4cc3136a21b9 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff_controlnet.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff_controlnet.py @@ -748,7 +748,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], decode_chunk_size: int = 16, - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff_sdxl.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff_sdxl.py index 70c6a5dc5cd6..4bb6bee294c1 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff_sdxl.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff_sdxl.py @@ -905,7 +905,7 @@ def __call__( clip_skip: int | None = None, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff_sparsectrl.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff_sparsectrl.py index fcf260b47f3a..6be516c0938f 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff_sparsectrl.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff_sparsectrl.py @@ -738,7 +738,7 @@ def __call__( clip_skip: int | None = None, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video.py index b8aa82ab9d2f..e3f6966a159a 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video.py @@ -770,7 +770,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], decode_chunk_size: int = 16, - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video_controlnet.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video_controlnet.py index 7c649b501f32..35821b09cd52 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video_controlnet.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video_controlnet.py @@ -940,7 +940,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], decode_chunk_size: int = 16, - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/anyflow/pipeline_anyflow.py b/src/diffusers/pipelines/anyflow/pipeline_anyflow.py index 61240b40e6b5..152183bbb311 100644 --- a/src/diffusers/pipelines/anyflow/pipeline_anyflow.py +++ b/src/diffusers/pipelines/anyflow/pipeline_anyflow.py @@ -406,7 +406,7 @@ def __call__( callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 512, use_mean_velocity: bool = True, - ): + ) -> AnyFlowPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/anyflow/pipeline_anyflow_far.py b/src/diffusers/pipelines/anyflow/pipeline_anyflow_far.py index 63ac8fa4d6bd..829bc4a236d8 100644 --- a/src/diffusers/pipelines/anyflow/pipeline_anyflow_far.py +++ b/src/diffusers/pipelines/anyflow/pipeline_anyflow_far.py @@ -476,7 +476,7 @@ def __call__( use_mean_velocity: bool = True, use_kv_cache: bool = True, chunk_partition: Optional[List[int]] = None, - ): + ) -> AnyFlowPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/audioldm2/pipeline_audioldm2.py b/src/diffusers/pipelines/audioldm2/pipeline_audioldm2.py index 110ddd1bfef3..b9e64c18c468 100644 --- a/src/diffusers/pipelines/audioldm2/pipeline_audioldm2.py +++ b/src/diffusers/pipelines/audioldm2/pipeline_audioldm2.py @@ -864,7 +864,7 @@ def __call__( callback_steps: int | None = 1, cross_attention_kwargs: dict[str, Any] | None = None, output_type: str | None = "np", - ): + ) -> AudioPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/bria/pipeline_bria.py b/src/diffusers/pipelines/bria/pipeline_bria.py index 9b80278af21e..a3f262ce935e 100644 --- a/src/diffusers/pipelines/bria/pipeline_bria.py +++ b/src/diffusers/pipelines/bria/pipeline_bria.py @@ -469,7 +469,7 @@ def __call__( max_sequence_length: int = 128, clip_value: None | float = None, normalize: bool = False, - ): + ) -> BriaPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo.py b/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo.py index 398758294bcc..96f713851cf4 100644 --- a/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo.py +++ b/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo.py @@ -453,7 +453,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 3000, do_patching=False, - ): + ) -> BriaFiboPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo_edit.py b/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo_edit.py index 28858322fb96..3b7ae9b1c206 100644 --- a/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo_edit.py +++ b/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo_edit.py @@ -621,7 +621,7 @@ def __call__( max_sequence_length: int = 3000, do_patching=False, _auto_resize: bool = True, - ): + ) -> BriaFiboPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/chroma/pipeline_chroma.py b/src/diffusers/pipelines/chroma/pipeline_chroma.py index 2f60d21091dd..37562f4b6a52 100644 --- a/src/diffusers/pipelines/chroma/pipeline_chroma.py +++ b/src/diffusers/pipelines/chroma/pipeline_chroma.py @@ -612,7 +612,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ChromaPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py b/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py index f1a3cc6c7b24..da6615e71c7f 100644 --- a/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py +++ b/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py @@ -673,7 +673,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ChromaPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py b/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py index 4900ef14e3f1..fe95a103cd4f 100644 --- a/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py +++ b/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py @@ -795,7 +795,7 @@ def __call__( max_sequence_length: int = 256, prompt_attention_mask: torch.Tensor | None = None, negative_prompt_attention_mask: torch.Tensor | None = None, - ): + ) -> ChromaPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py b/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py index 3fefa951c9b8..95a3ad3c20e3 100644 --- a/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py +++ b/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py @@ -493,7 +493,7 @@ def __call__( max_sequence_length: int = 512, enable_temporal_reasoning: bool = False, num_temporal_reasoning_steps: int = 0, - ): + ) -> ChronoEditPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/consistency_models/pipeline_consistency_models.py b/src/diffusers/pipelines/consistency_models/pipeline_consistency_models.py index b2b18b52e824..d8fb22f1c262 100644 --- a/src/diffusers/pipelines/consistency_models/pipeline_consistency_models.py +++ b/src/diffusers/pipelines/consistency_models/pipeline_consistency_models.py @@ -182,7 +182,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, - ): + ) -> ImagePipelineOutput | tuple: r""" Args: batch_size (`int`, *optional*, defaults to 1): diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet.py index 7ef287ae8154..bc95c1efb940 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet.py @@ -936,7 +936,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_blip_diffusion.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_blip_diffusion.py index 3b91229814cc..1e0dfd3fb94f 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_blip_diffusion.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_blip_diffusion.py @@ -255,7 +255,7 @@ def __call__( prompt_reps: int = 20, output_type: str | None = "pil", return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_img2img.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_img2img.py index b40dda940959..15aecfa4d4b4 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_img2img.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_img2img.py @@ -934,7 +934,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint.py index c981dd14abd6..7f21b964417d 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint.py @@ -1025,7 +1025,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint_sd_xl.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint_sd_xl.py index 86905fb39b29..0503974c101e 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint_sd_xl.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint_sd_xl.py @@ -1210,7 +1210,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl.py index f0174cb6ce98..3c75d15ea6e9 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl.py @@ -1039,7 +1039,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl_img2img.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl_img2img.py index 92366e465673..ba24ed70db0e 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl_img2img.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl_img2img.py @@ -1119,7 +1119,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_inpaint_sd_xl.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_inpaint_sd_xl.py index 3dd12ec72989..7e7aeddcbbac 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_inpaint_sd_xl.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_inpaint_sd_xl.py @@ -1190,7 +1190,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl.py index 4802279b81f4..2177f30b3e72 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl.py @@ -1016,7 +1016,7 @@ def __call__( clip_skip: int | None = None, callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl_img2img.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl_img2img.py index 62c9c06f46a4..0cbc923fce10 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl_img2img.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl_img2img.py @@ -1110,7 +1110,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet_hunyuandit/pipeline_hunyuandit_controlnet.py b/src/diffusers/pipelines/controlnet_hunyuandit/pipeline_hunyuandit_controlnet.py index ba241bf4feb6..bea400c19f86 100644 --- a/src/diffusers/pipelines/controlnet_hunyuandit/pipeline_hunyuandit_controlnet.py +++ b/src/diffusers/pipelines/controlnet_hunyuandit/pipeline_hunyuandit_controlnet.py @@ -662,7 +662,7 @@ def __call__( target_size: tuple[int, int] | None = None, crops_coords_top_left: tuple[int, int] = (0, 0), use_resolution_binning: bool = True, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation with HunyuanDiT. diff --git a/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet.py b/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet.py index 4530a424adb4..5a7ee29a1c23 100644 --- a/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet.py +++ b/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet.py @@ -852,7 +852,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet_inpainting.py b/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet_inpainting.py index d2890d55811c..60b5828cb8ac 100644 --- a/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet_inpainting.py +++ b/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet_inpainting.py @@ -1020,7 +1020,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_predict.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_predict.py index c2c5e6d2c824..74c708f25444 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_predict.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_predict.py @@ -566,7 +566,7 @@ def __call__( max_sequence_length: int = 512, conditional_frame_timestep: float = 0.0001, num_latent_conditional_frames: int = 2, - ): + ) -> CosmosPipelineOutput | tuple: r""" The call function to the pipeline for generation. Supports three modes: diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_transfer.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_transfer.py index e38d926bbd28..d8d084946a98 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_transfer.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_transfer.py @@ -595,7 +595,7 @@ def __call__( conditional_frame_timestep: float = 0.1, num_ar_conditional_frames: Optional[int] = 1, num_ar_latent_conditional_frames: Optional[int] = None, - ): + ) -> CosmosPipelineOutput | tuple: r""" `controls` drive the conditioning through ControlNet. Controls are assumed to be pre-processed, e.g. edge maps are pre-computed. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_text2image.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_text2image.py index 8c6de18b3a9a..e29fde090c22 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_text2image.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_text2image.py @@ -434,7 +434,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> CosmosImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_video2world.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_video2world.py index 2a708e1118e0..01b8d7b73680 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_video2world.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_video2world.py @@ -507,7 +507,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, sigma_conditioning: float = 0.0001, - ): + ) -> CosmosPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py index 02dc70b29cfc..edd8f44409c2 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py @@ -1350,7 +1350,7 @@ def __call__( mixed_precision_first_steps: int | None = None, mixed_precision_last_steps: int | None = None, mixed_precision_reasoner_policy: str | None = None, - ) -> Cosmos3OmniPipelineOutput: + ) -> Cosmos3OmniPipelineOutput | tuple: r""" Run the Cosmos 3 omni pipeline end-to-end: encode the (optional) conditioning image/video, denoise vision and (optional) sound latents jointly, and decode them back into a video and audio waveform. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos_text2world.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos_text2world.py index 61d9ec8f0574..55eec93d67a9 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos_text2world.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos_text2world.py @@ -420,7 +420,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> CosmosPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos_video2world.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos_video2world.py index bf7e28584967..7298218a4770 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos_video2world.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos_video2world.py @@ -536,7 +536,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> CosmosPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if.py index b8c70fc6528c..aaf177b31edc 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if.py @@ -566,7 +566,7 @@ def __call__( callback_steps: int = 1, clean_caption: bool = True, cross_attention_kwargs: dict[str, Any] | None = None, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img.py index 3dadc63f4952..6bc077aef03d 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img.py @@ -685,7 +685,7 @@ def __call__( callback_steps: int = 1, clean_caption: bool = True, cross_attention_kwargs: dict[str, Any] | None = None, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img_superresolution.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img_superresolution.py index 4839a0860462..ed28d7b59fa4 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img_superresolution.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img_superresolution.py @@ -770,7 +770,7 @@ def __call__( cross_attention_kwargs: dict[str, Any] | None = None, noise_level: int = 250, clean_caption: bool = True, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting.py index 03a9d6f7c5e8..1ed581f56a0e 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting.py @@ -783,7 +783,7 @@ def __call__( callback_steps: int = 1, clean_caption: bool = True, cross_attention_kwargs: dict[str, Any] | None = None, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting_superresolution.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting_superresolution.py index 841382ad9c63..b7a1736c61a6 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting_superresolution.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting_superresolution.py @@ -864,7 +864,7 @@ def __call__( cross_attention_kwargs: dict[str, Any] | None = None, noise_level: int = 0, clean_caption: bool = True, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_superresolution.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_superresolution.py index 52ebebb6f9b4..45f1ee6b1610 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_superresolution.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_superresolution.py @@ -635,7 +635,7 @@ def __call__( cross_attention_kwargs: dict[str, Any] | None = None, noise_level: int = 250, clean_caption: bool = True, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py b/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py index 5d608d7c49fb..207142dccb67 100644 --- a/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py +++ b/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py @@ -184,7 +184,7 @@ def __call__( | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] | None = None, - ) -> DiffusionGemmaPipelineOutput | tuple[torch.LongTensor, list[str] | None]: + ) -> DiffusionGemmaPipelineOutput | tuple: """ Generate text with block diffusion. diff --git a/src/diffusers/pipelines/dreamlite/pipeline_dreamlite.py b/src/diffusers/pipelines/dreamlite/pipeline_dreamlite.py index e9a0e3c2a767..85b303b32dec 100644 --- a/src/diffusers/pipelines/dreamlite/pipeline_dreamlite.py +++ b/src/diffusers/pipelines/dreamlite/pipeline_dreamlite.py @@ -403,7 +403,7 @@ def __call__( return_dict: bool = True, max_sequence_length: int = 200, text_pad_embedding: Optional[torch.Tensor] = None, - ): + ) -> DreamLitePipelineOutput | tuple: r"""Run the DreamLite pipeline. Args: diff --git a/src/diffusers/pipelines/dreamlite/pipeline_dreamlite_mobile.py b/src/diffusers/pipelines/dreamlite/pipeline_dreamlite_mobile.py index ca9e6b7b4c40..339b58c9cbf9 100644 --- a/src/diffusers/pipelines/dreamlite/pipeline_dreamlite_mobile.py +++ b/src/diffusers/pipelines/dreamlite/pipeline_dreamlite_mobile.py @@ -398,7 +398,7 @@ def __call__( return_dict: bool = True, max_sequence_length: int = 200, text_pad_embedding: Optional[torch.Tensor] = None, - ): + ) -> DreamLitePipelineOutput | tuple: r"""Run the distilled DreamLite Mobile pipeline. Args: diff --git a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate.py b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate.py index 72e19a8cce1f..428a20039778 100755 --- a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate.py +++ b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate.py @@ -546,7 +546,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], guidance_rescale: float = 0.0, - ): + ) -> EasyAnimatePipelineOutput | tuple: r""" Generates images or video using the EasyAnimate pipeline based on the provided prompts. diff --git a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_control.py b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_control.py index 4ad3a48b70ec..9f773e18b5c3 100755 --- a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_control.py +++ b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_control.py @@ -695,7 +695,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], guidance_rescale: float = 0.0, timesteps: list[int] | None = None, - ): + ) -> EasyAnimatePipelineOutput | tuple: r""" Generates images or video using the EasyAnimate pipeline based on the provided prompts. diff --git a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_inpaint.py b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_inpaint.py index 69bb332944d6..768c4730fcbb 100755 --- a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_inpaint.py +++ b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_inpaint.py @@ -815,7 +815,7 @@ def __call__( strength: float = 1.0, noise_aug_strength: float = 0.0563, timesteps: list[int] | None = None, - ): + ) -> EasyAnimatePipelineOutput | tuple: r""" The call function to the pipeline for generation with HunyuanDiT. diff --git a/src/diffusers/pipelines/ernie_image/pipeline_ernie_image.py b/src/diffusers/pipelines/ernie_image/pipeline_ernie_image.py index 11fce6a204bf..a0097cae660f 100644 --- a/src/diffusers/pipelines/ernie_image/pipeline_ernie_image.py +++ b/src/diffusers/pipelines/ernie_image/pipeline_ernie_image.py @@ -222,7 +222,7 @@ def __call__( callback_on_step_end: Optional[Callable[[int, int, dict], None]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], use_pe: bool = True, # é»˜č®¤ä½æē”ØPEčæ›č”Œę”¹å†™ - ): + ) -> ErnieImagePipelineOutput | tuple: """ Generate images from text prompts. diff --git a/src/diffusers/pipelines/flux/pipeline_flux.py b/src/diffusers/pipelines/flux/pipeline_flux.py index eb831a7975ba..2e580388874c 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux.py +++ b/src/diffusers/pipelines/flux/pipeline_flux.py @@ -628,7 +628,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_control.py b/src/diffusers/pipelines/flux/pipeline_flux_control.py index d483161b33b2..b50861cb437f 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_control.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_control.py @@ -603,7 +603,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py b/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py index 56bec7a637ce..15d876caf6ca 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py @@ -658,7 +658,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py b/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py index 5a8c7ff2900e..6fa0c1dff9f6 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py @@ -777,7 +777,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py b/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py index 177a1e5e4ef2..ab2cebb1e765 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py @@ -710,7 +710,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_controlnet_image_to_image.py b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_image_to_image.py index 70e4df7acacd..482fca3b59c3 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_controlnet_image_to_image.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_image_to_image.py @@ -663,7 +663,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py index 02a2e93420bf..26a434807694 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py @@ -770,7 +770,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_fill.py b/src/diffusers/pipelines/flux/pipeline_flux_fill.py index 929f7530bb86..55d3464cfb4e 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_fill.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_fill.py @@ -723,7 +723,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_img2img.py b/src/diffusers/pipelines/flux/pipeline_flux_img2img.py index 81bb499ac4ae..3288bba94772 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_img2img.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_img2img.py @@ -708,7 +708,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_inpaint.py b/src/diffusers/pipelines/flux/pipeline_flux_inpaint.py index 466fba8ac7e6..15d9bb9868fd 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_inpaint.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_inpaint.py @@ -810,7 +810,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_kontext.py b/src/diffusers/pipelines/flux/pipeline_flux_kontext.py index e5fc95e5a1c1..3579c3abd8d1 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_kontext.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_kontext.py @@ -726,7 +726,7 @@ def __call__( max_sequence_length: int = 512, max_area: int = 1024**2, _auto_resize: bool = True, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py b/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py index 020f9761e121..05bf8200f588 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py @@ -919,7 +919,7 @@ def __call__( max_sequence_length: int = 512, max_area: int = 1024**2, _auto_resize: bool = True, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_prior_redux.py b/src/diffusers/pipelines/flux/pipeline_flux_prior_redux.py index fb39ad2a1583..e74d6d4e13b6 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_prior_redux.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_prior_redux.py @@ -387,7 +387,7 @@ def __call__( prompt_embeds_scale: float | list[float] | None = 1.0, pooled_prompt_embeds_scale: float | list[float] | None = 1.0, return_dict: bool = True, - ): + ) -> FluxPriorReduxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux2/pipeline_flux2.py b/src/diffusers/pipelines/flux2/pipeline_flux2.py index aa73c53fb58d..7bb716772e3a 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2.py @@ -766,7 +766,7 @@ def __call__( max_sequence_length: int = 512, text_encoder_out_layers: tuple[int] = (10, 20, 30), caption_upsample_temperature: float = None, - ): + ) -> Flux2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py b/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py index 92005750e551..50fd51e283f6 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py @@ -634,7 +634,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, text_encoder_out_layers: tuple[int] = (9, 18, 27), - ): + ) -> Flux2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py index 0f9051a99b12..3150c1144f67 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py @@ -849,7 +849,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, text_encoder_out_layers: tuple[int, ...] = (9, 18, 27), - ): + ) -> Flux2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for inpainting. diff --git a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py index 82a33a84568d..711246db71d5 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py @@ -628,7 +628,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, text_encoder_out_layers: tuple[int] = (9, 18, 27), - ): + ) -> Flux2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/helios/pipeline_helios.py b/src/diffusers/pipelines/helios/pipeline_helios.py index 90ac654bc77c..c80158ce8625 100644 --- a/src/diffusers/pipelines/helios/pipeline_helios.py +++ b/src/diffusers/pipelines/helios/pipeline_helios.py @@ -483,7 +483,7 @@ def __call__( num_latent_frames_per_chunk: int = 9, keep_first_frame: bool = True, is_skip_first_chunk: bool = False, - ): + ) -> HeliosPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py b/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py index c187e436a857..c32108d613ed 100644 --- a/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py +++ b/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py @@ -552,7 +552,7 @@ def __call__( zero_steps: int | None = 1, # ------------ DMD ------------ is_amplify_first_chunk: bool = False, - ): + ) -> HeliosPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py b/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py index 1bf3ef3699e4..e4cb9f7c12ea 100644 --- a/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py +++ b/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py @@ -703,7 +703,8 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 128, - ): + **kwargs, + ) -> HiDreamImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py index 50239e9afa22..4faa73d40247 100644 --- a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py +++ b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py @@ -528,7 +528,7 @@ def __call__( attention_kwargs: dict[str, Any] | None = None, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> HunyuanImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py index efdb5505e604..4b926b82c474 100644 --- a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py +++ b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py @@ -457,7 +457,7 @@ def __call__( attention_kwargs: dict[str, Any] | None = None, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> HunyuanImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py index bd54f2563b52..22e8346b8a14 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py @@ -510,7 +510,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], prompt_template: dict[str, Any] = DEFAULT_PROMPT_TEMPLATE, max_sequence_length: int = 256, - ): + ) -> HunyuanVideoPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py index 9e7c198c19cc..1edd4f6db9f6 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py @@ -466,7 +466,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], prompt_template: dict[str, Any] = DEFAULT_PROMPT_TEMPLATE, max_sequence_length: int = 256, - ): + ) -> HunyuanVideoPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py index 349481492ac0..0208559bb56e 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py @@ -619,7 +619,7 @@ def __call__( prompt_template: dict[str, Any] = DEFAULT_PROMPT_TEMPLATE, max_sequence_length: int = 256, sampling_type: FramepackSamplingType = FramepackSamplingType.INVERTED_ANTI_DRIFTING, - ): + ) -> HunyuanVideoFramepackPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py index 13eb35386001..929118cc0af8 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py @@ -651,7 +651,7 @@ def __call__( prompt_template: dict[str, Any] = DEFAULT_PROMPT_TEMPLATE, max_sequence_length: int = 256, image_embed_interleave: int | None = None, - ): + ) -> HunyuanVideoPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py index 7232ebbee5b8..cba5f0aa5c68 100644 --- a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py +++ b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py @@ -564,7 +564,7 @@ def __call__( output_type: str | None = "np", return_dict: bool = True, attention_kwargs: dict[str, Any] | None = None, - ): + ) -> HunyuanVideo15PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py index 71a36a1c51cd..b171f83c3668 100644 --- a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py +++ b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py @@ -671,7 +671,7 @@ def __call__( output_type: str | None = "np", return_dict: bool = True, attention_kwargs: dict[str, Any] | None = None, - ): + ) -> HunyuanVideo15PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuandit/pipeline_hunyuandit.py b/src/diffusers/pipelines/hunyuandit/pipeline_hunyuandit.py index 5d656a3c370a..1e9a78130ec5 100644 --- a/src/diffusers/pipelines/hunyuandit/pipeline_hunyuandit.py +++ b/src/diffusers/pipelines/hunyuandit/pipeline_hunyuandit.py @@ -596,7 +596,7 @@ def __call__( target_size: tuple[int, int] | None = None, crops_coords_top_left: tuple[int, int] = (0, 0), use_resolution_binning: bool = True, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation with HunyuanDiT. diff --git a/src/diffusers/pipelines/ideogram4/pipeline_ideogram4.py b/src/diffusers/pipelines/ideogram4/pipeline_ideogram4.py index 7577e0463ca7..8aa76660be15 100644 --- a/src/diffusers/pipelines/ideogram4/pipeline_ideogram4.py +++ b/src/diffusers/pipelines/ideogram4/pipeline_ideogram4.py @@ -501,7 +501,7 @@ def __call__( attention_kwargs: dict[str, Any] | None = None, callback_on_step_end: Callable[["Ideogram4Pipeline", int, int, dict[str, Any]], dict[str, Any]] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ) -> Ideogram4PipelineOutput | tuple[Any]: + ) -> Ideogram4PipelineOutput | tuple: r""" Run text-to-image generation. diff --git a/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit.py b/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit.py index 2eff843d5225..29b2269ef018 100644 --- a/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit.py +++ b/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit.py @@ -629,7 +629,7 @@ def __call__( callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 4096, enable_denormalization: bool = True, - ): + ) -> JoyImageEditPipelineOutput | tuple: r""" Generate an edited image conditioned on a reference image and a text prompt. diff --git a/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit_plus.py b/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit_plus.py index ac8e01278e3e..0f4d1c207325 100644 --- a/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit_plus.py +++ b/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit_plus.py @@ -465,7 +465,7 @@ def __call__( | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 4096, - ): + ) -> JoyImageEditPlusPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky.py b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky.py index aa6006ddc082..5274453551ca 100644 --- a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky.py +++ b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky.py @@ -252,7 +252,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_combined.py b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_combined.py index 3a0c3c07e8ca..feb552e26afb 100644 --- a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_combined.py +++ b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_combined.py @@ -28,7 +28,7 @@ from ...utils import ( replace_example_docstring, ) -from ..pipeline_utils import DiffusionPipeline +from ..pipeline_utils import DiffusionPipeline, ImagePipelineOutput from .pipeline_kandinsky import KandinskyPipeline from .pipeline_kandinsky_img2img import KandinskyImg2ImgPipeline from .pipeline_kandinsky_inpaint import KandinskyInpaintPipeline @@ -231,7 +231,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. @@ -452,7 +452,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. @@ -693,7 +693,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_img2img.py b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_img2img.py index f374741bd9fb..740dad385c79 100644 --- a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_img2img.py +++ b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_img2img.py @@ -314,7 +314,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_inpaint.py b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_inpaint.py index 9bc52ce871f5..61d0c371ea89 100644 --- a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_inpaint.py +++ b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_inpaint.py @@ -419,7 +419,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_prior.py b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_prior.py index 5708286b7d5a..c21589d15043 100644 --- a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_prior.py +++ b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_prior.py @@ -415,7 +415,7 @@ def __call__( guidance_scale: float = 4.0, output_type: str | None = "pt", return_dict: bool = True, - ): + ) -> KandinskyPriorPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2.py index 668df4ac5980..2ef2f245c841 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2.py @@ -145,7 +145,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_combined.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_combined.py index d5616958a563..3fc6e68b7c66 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_combined.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_combined.py @@ -21,7 +21,7 @@ from ...models import PriorTransformer, UNet2DConditionModel, VQModel from ...schedulers import DDPMScheduler, UnCLIPScheduler from ...utils import deprecate, logging, replace_example_docstring -from ..pipeline_utils import DiffusionPipeline +from ..pipeline_utils import DiffusionPipeline, ImagePipelineOutput from .pipeline_kandinsky2_2 import KandinskyV22Pipeline from .pipeline_kandinsky2_2_img2img import KandinskyV22Img2ImgPipeline from .pipeline_kandinsky2_2_inpainting import KandinskyV22InpaintPipeline @@ -222,7 +222,7 @@ def __call__( prior_callback_on_step_end_tensor_inputs: list[str] = ["latents"], callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. @@ -468,7 +468,7 @@ def __call__( prior_callback_on_step_end_tensor_inputs: list[str] = ["latents"], callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. @@ -723,7 +723,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet.py index a9b7be516114..4a8c19784f2d 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet.py @@ -174,7 +174,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet_img2img.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet_img2img.py index f77a40595194..99a7ee527068 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet_img2img.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet_img2img.py @@ -215,7 +215,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_img2img.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_img2img.py index dc17f49bbfe0..649c1d7c678b 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_img2img.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_img2img.py @@ -198,7 +198,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_inpainting.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_inpainting.py index f258dfc07094..e4ba5558f63e 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_inpainting.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_inpainting.py @@ -318,7 +318,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior.py index 8095f79280d4..c26afa61ff6f 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior.py @@ -387,7 +387,7 @@ def __call__( return_dict: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> KandinskyPriorPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior_emb2emb.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior_emb2emb.py index 72f1d8556ec5..9ed4ef589789 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior_emb2emb.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior_emb2emb.py @@ -410,7 +410,7 @@ def __call__( guidance_scale: float = 4.0, output_type: str | None = "pt", # pt only return_dict: bool = True, - ): + ) -> KandinskyPriorPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3.py b/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3.py index ca8f124c74cf..ba26777f4366 100644 --- a/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3.py +++ b/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3.py @@ -353,7 +353,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3_img2img.py b/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3_img2img.py index beb4caafb6d3..9ce1da92f459 100644 --- a/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3_img2img.py +++ b/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3_img2img.py @@ -418,7 +418,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py index 7a5dc2cd1ac1..1ad08729e8eb 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py @@ -702,7 +702,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> KandinskyPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py index 1784b4e42972..9fa3378d0143 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py @@ -589,7 +589,7 @@ def __call__( callback_on_step_end: Callable[[int, int, None], PipelineCallback | MultiPipelineCallbacks] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> KandinskyImagePipelineOutput | tuple: r""" The call function to the pipeline for image-to-image generation. diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py index d39478547d8e..34099f191891 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py @@ -771,7 +771,7 @@ def __call__( callback_on_step_end: Callable[[int, int, None], PipelineCallback | MultiPipelineCallbacks] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> KandinskyPipelineOutput | tuple: r""" The call function to the pipeline for image-to-video generation. diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py index d86fff668771..46002e086a28 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py @@ -555,7 +555,7 @@ def __call__( callback_on_step_end: Callable[[int, int, None], PipelineCallback | MultiPipelineCallbacks] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> KandinskyImagePipelineOutput | tuple: r""" The call function to the pipeline for text-to-image generation. diff --git a/src/diffusers/pipelines/kolors/pipeline_kolors.py b/src/diffusers/pipelines/kolors/pipeline_kolors.py index 1e11faf8b9b6..a4d92e278d70 100644 --- a/src/diffusers/pipelines/kolors/pipeline_kolors.py +++ b/src/diffusers/pipelines/kolors/pipeline_kolors.py @@ -679,7 +679,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> KolorsPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kolors/pipeline_kolors_img2img.py b/src/diffusers/pipelines/kolors/pipeline_kolors_img2img.py index d9b519267216..39ed6e37ffaf 100644 --- a/src/diffusers/pipelines/kolors/pipeline_kolors_img2img.py +++ b/src/diffusers/pipelines/kolors/pipeline_kolors_img2img.py @@ -814,7 +814,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> KolorsPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/krea2/pipeline_krea2.py b/src/diffusers/pipelines/krea2/pipeline_krea2.py index 51d33cb48619..dbc9a19e74be 100644 --- a/src/diffusers/pipelines/krea2/pipeline_krea2.py +++ b/src/diffusers/pipelines/krea2/pipeline_krea2.py @@ -466,7 +466,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], attention_kwargs: dict[str, Any] | None = None, max_sequence_length: int = 512, - ): + ) -> Krea2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_img2img.py b/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_img2img.py index 424a2c46e06b..6a0fdb96147e 100644 --- a/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_img2img.py +++ b/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_img2img.py @@ -730,7 +730,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_text2img.py b/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_text2img.py index 60f59ec7f9d3..947421628577 100644 --- a/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_text2img.py +++ b/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_text2img.py @@ -661,7 +661,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion.py b/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion.py index 2b0e85d393a7..63d6dc6c7117 100644 --- a/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion.py +++ b/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion.py @@ -85,7 +85,7 @@ def __call__( output_type: str | None = "pil", return_dict: bool = True, **kwargs, - ) -> tuple | ImagePipelineOutput: + ) -> ImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion_superresolution.py b/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion_superresolution.py index c44d49944ea3..13f28e3ee8c7 100644 --- a/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion_superresolution.py +++ b/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion_superresolution.py @@ -78,7 +78,7 @@ def __call__( generator: torch.Generator | list[torch.Generator] | None = None, output_type: str | None = "pil", return_dict: bool = True, - ) -> tuple | ImagePipelineOutput: + ) -> ImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion.py b/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion.py index c6cdb127309b..a4dd6ddd1dd5 100644 --- a/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion.py +++ b/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion.py @@ -745,7 +745,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> LEditsPPDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for editing. The [`~pipelines.ledits_pp.LEditsPPPipelineStableDiffusion.invert`] method has to be called beforehand. Edits will diff --git a/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion_xl.py b/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion_xl.py index 6e97b1ea2ed4..039bc276aa6c 100644 --- a/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion_xl.py +++ b/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion_xl.py @@ -813,7 +813,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> LEditsPPDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for editing. The [`~pipelines.ledits_pp.LEditsPPPipelineStableDiffusionXL.invert`] method has to be called beforehand. Edits diff --git a/src/diffusers/pipelines/llada2/pipeline_llada2.py b/src/diffusers/pipelines/llada2/pipeline_llada2.py index 06b4875f18a9..a9b8f351144e 100644 --- a/src/diffusers/pipelines/llada2/pipeline_llada2.py +++ b/src/diffusers/pipelines/llada2/pipeline_llada2.py @@ -271,7 +271,7 @@ def __call__( | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] | None = None, - ) -> LLaDA2PipelineOutput | tuple[torch.LongTensor, list[str] | None]: + ) -> LLaDA2PipelineOutput | tuple: """ Generate text with block-wise refinement. diff --git a/src/diffusers/pipelines/longcat_audio_dit/pipeline_longcat_audio_dit.py b/src/diffusers/pipelines/longcat_audio_dit/pipeline_longcat_audio_dit.py index e6478535b373..f48d3376916c 100644 --- a/src/diffusers/pipelines/longcat_audio_dit/pipeline_longcat_audio_dit.py +++ b/src/diffusers/pipelines/longcat_audio_dit/pipeline_longcat_audio_dit.py @@ -231,7 +231,7 @@ def __call__( return_dict: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> AudioPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. @@ -254,6 +254,10 @@ def __call__( Tensor inputs passed to `callback_on_step_end`. Examples: + + Returns: + [`~pipelines.AudioPipelineOutput`] or `tuple`: [`~pipelines.AudioPipelineOutput`] if `return_dict` is True, + otherwise a `tuple`. When returning a tuple, the first element is the generated audio waveform. """ if prompt is None: prompt = [] diff --git a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py index 41ca3eb54f83..eed7b40c48fc 100644 --- a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py +++ b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py @@ -490,7 +490,7 @@ def __call__( enable_cfg_renorm: bool | None = True, cfg_renorm_min: float | None = 0.0, enable_prompt_rewrite: bool | None = True, - ): + ) -> LongCatImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py index 9f35bb685d9f..8f107bd456c4 100644 --- a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py +++ b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py @@ -546,7 +546,7 @@ def __call__( output_type: str | None = "pil", return_dict: bool = True, joint_attention_kwargs: dict[str, Any] | None = None, - ): + ) -> LongCatImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx.py b/src/diffusers/pipelines/ltx/pipeline_ltx.py index ce9177547c52..e66f8dc18cf4 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx.py @@ -561,7 +561,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 128, - ): + ) -> LTXPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py b/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py index 28d296695998..dc2dde5aad93 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py @@ -881,7 +881,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> LTXPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py b/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py index 838d5afc5c5a..f057f8d7908e 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py @@ -972,7 +972,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 128, - ): + ) -> LTXPipelineOutput | tuple: r""" Generate an image-to-video sequence via temporal sliding windows and multi-prompt scheduling. diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py b/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py index 81ecfce50efa..11e65270cfb6 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py @@ -623,7 +623,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 128, - ): + ) -> LTXPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx_latent_upsample.py b/src/diffusers/pipelines/ltx/pipeline_ltx_latent_upsample.py index 6f325237a248..c6eff670908f 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx_latent_upsample.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx_latent_upsample.py @@ -199,7 +199,7 @@ def __call__( generator: torch.Generator | list[torch.Generator] | None = None, output_type: str | None = "pil", return_dict: bool = True, - ): + ) -> LTXPipelineOutput | tuple: r""" Function invoked when calling the pipeline for latent upsampling. @@ -227,6 +227,11 @@ def __call__( The output format of the generated video. Choose between `PIL.Image`, `np.array`, or `latent`. return_dict (`bool`, *optional*, defaults to `True`): Whether or not to return a [`~pipelines.ltx.LTXPipelineOutput`] instead of a plain tuple. + + Returns: + [`~pipelines.ltx.LTXPipelineOutput`] or `tuple`: [`~pipelines.ltx.LTXPipelineOutput`] if `return_dict` is + True, otherwise a `tuple`. When returning a tuple, the first element is the upsampled video (or the latents + if `output_type="latent"`). """ self.check_inputs( video=video, diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2.py index 22948a7ecf3a..4e5ced0b4ec8 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2.py @@ -970,7 +970,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py index bd2ee3ec6708..e1bff845302c 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py @@ -1390,7 +1390,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py index 0cda9c6079b6..936af16d8805 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py @@ -1563,7 +1563,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2DFRPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py index bebf3739956a..bdf9c8e67677 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py @@ -1377,7 +1377,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2DFRPipelineOutput | tuple: r""" Run one temporal refine round. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_diffusion_decode.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_diffusion_decode.py index 2f1137830a57..28d740991a54 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_diffusion_decode.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_diffusion_decode.py @@ -84,7 +84,7 @@ def __call__( output_type: str = "pil", return_dict: bool = True, denormalize: bool = True, - ): + ) -> LTX2VideoDecodeOutput | tuple: r""" Args: latents (`torch.Tensor`): diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py index 91173bc6e161..d9d95cdc0a0b 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py @@ -1080,7 +1080,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2PipelineOutput | tuple: r""" Run HDR IC-LoRA video generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py index dc92b6eb965a..2924a086c721 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py @@ -1822,7 +1822,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py index c7c81d26cb45..36b7effdfd6c 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py @@ -1026,7 +1026,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_latent_upsample.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_latent_upsample.py index 8aa72a425e1c..09d842927d12 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_latent_upsample.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_latent_upsample.py @@ -280,7 +280,7 @@ def __call__( generator: torch.Generator | list[torch.Generator] | None = None, output_type: str | None = "pil", return_dict: bool = True, - ): + ) -> LTXPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py b/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py index 1bd7ab4ca675..b0d2f91736ea 100644 --- a/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py +++ b/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py @@ -472,7 +472,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> LucyPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/marigold/pipeline_marigold_depth.py b/src/diffusers/pipelines/marigold/pipeline_marigold_depth.py index a81d1c51742c..f8eaf683e0c0 100644 --- a/src/diffusers/pipelines/marigold/pipeline_marigold_depth.py +++ b/src/diffusers/pipelines/marigold/pipeline_marigold_depth.py @@ -363,7 +363,7 @@ def __call__( output_uncertainty: bool = False, output_latent: bool = False, return_dict: bool = True, - ): + ) -> MarigoldDepthOutput | tuple: """ Function invoked when calling the pipeline. diff --git a/src/diffusers/pipelines/marigold/pipeline_marigold_intrinsics.py b/src/diffusers/pipelines/marigold/pipeline_marigold_intrinsics.py index 9488d8f5c9b8..e0fe93d561b2 100644 --- a/src/diffusers/pipelines/marigold/pipeline_marigold_intrinsics.py +++ b/src/diffusers/pipelines/marigold/pipeline_marigold_intrinsics.py @@ -375,7 +375,7 @@ def __call__( output_uncertainty: bool = False, output_latent: bool = False, return_dict: bool = True, - ): + ) -> MarigoldIntrinsicsOutput | tuple: """ Function invoked when calling the pipeline. diff --git a/src/diffusers/pipelines/marigold/pipeline_marigold_normals.py b/src/diffusers/pipelines/marigold/pipeline_marigold_normals.py index 3f94ce441232..14a46be58ea6 100644 --- a/src/diffusers/pipelines/marigold/pipeline_marigold_normals.py +++ b/src/diffusers/pipelines/marigold/pipeline_marigold_normals.py @@ -348,7 +348,7 @@ def __call__( output_uncertainty: bool = False, output_latent: bool = False, return_dict: bool = True, - ): + ) -> MarigoldNormalsOutput | tuple: """ Function invoked when calling the pipeline. diff --git a/src/diffusers/pipelines/mochi/pipeline_mochi.py b/src/diffusers/pipelines/mochi/pipeline_mochi.py index c146d2d1e564..cbb8f3b2c6c7 100644 --- a/src/diffusers/pipelines/mochi/pipeline_mochi.py +++ b/src/diffusers/pipelines/mochi/pipeline_mochi.py @@ -466,7 +466,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> MochiPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/motif_video/pipeline_motif_video.py b/src/diffusers/pipelines/motif_video/pipeline_motif_video.py index 8ad37932e970..8bd2eb5bb5c9 100644 --- a/src/diffusers/pipelines/motif_video/pipeline_motif_video.py +++ b/src/diffusers/pipelines/motif_video/pipeline_motif_video.py @@ -516,7 +516,7 @@ def __call__( callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 512, vae_batch_size: int | None = None, - ): + ) -> MotifVideoPipelineOutput | tuple: r""" The call function to the pipeline for text-to-video generation. diff --git a/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py b/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py index 1b32ba74f24b..d30aebc610b8 100644 --- a/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py +++ b/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py @@ -644,7 +644,7 @@ def __call__( ] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> MotifVideoPipelineOutput | tuple: r""" The call function to the pipeline for image-to-video generation. diff --git a/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py b/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py index f50f11c8c152..71b7a82de4ea 100644 --- a/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py +++ b/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py @@ -401,7 +401,7 @@ def __call__( attention_kwargs: dict[str, Any] | None = None, callback_on_step_end: Callable[[int, int, dict], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> NucleusMoEImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/omnigen/pipeline_omnigen.py b/src/diffusers/pipelines/omnigen/pipeline_omnigen.py index 6564b2a672a0..69d369c0ab34 100644 --- a/src/diffusers/pipelines/omnigen/pipeline_omnigen.py +++ b/src/diffusers/pipelines/omnigen/pipeline_omnigen.py @@ -293,7 +293,7 @@ def __call__( return_dict: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> ImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. @@ -356,6 +356,10 @@ def __call__( Returns: [`~pipelines.ImagePipelineOutput`] or `tuple`: If `return_dict` is `True`, [`~pipelines.ImagePipelineOutput`] is returned, otherwise a `tuple` is returned where the first element is a list with the generated images. + + Returns: + [`~pipelines.ImagePipelineOutput`] or `tuple`: [`~pipelines.ImagePipelineOutput`] if `return_dict` is True, + otherwise a `tuple`. When returning a tuple, the first element is a list with the generated images. """ height = height or self.default_sample_size * self.vae_scale_factor diff --git a/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py b/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py index b22f2f0cec2d..841f4e1471e2 100644 --- a/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py +++ b/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py @@ -476,7 +476,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> OvisImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd.py b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd.py index 3a88272f24f4..7b6abb32acc5 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd.py @@ -893,7 +893,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_inpaint.py b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_inpaint.py index 98221e4c30ba..3cdd6bcdc6e9 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_inpaint.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_inpaint.py @@ -1003,7 +1003,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl.py b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl.py index 72060682c196..27c3c2f4ded9 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl.py @@ -1043,7 +1043,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl_img2img.py b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl_img2img.py index 4f2b8b2044fc..133acd82168d 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl_img2img.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl_img2img.py @@ -1122,7 +1122,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_hunyuandit.py b/src/diffusers/pipelines/pag/pipeline_pag_hunyuandit.py index a443a19bd952..6c96e16f51cb 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_hunyuandit.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_hunyuandit.py @@ -612,7 +612,7 @@ def __call__( use_resolution_binning: bool = True, pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation with HunyuanDiT. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_kolors.py b/src/diffusers/pipelines/pag/pipeline_pag_kolors.py index 4f138d91d9c6..9c588c2c168a 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_kolors.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_kolors.py @@ -699,7 +699,7 @@ def __call__( pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, max_sequence_length: int = 256, - ): + ) -> KolorsPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd.py b/src/diffusers/pipelines/pag/pipeline_pag_sd.py index b12597460f65..5548ec8acdd2 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd.py @@ -771,7 +771,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_3.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_3.py index d86adccc2ccf..aebdd2495c3b 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_3.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_3.py @@ -712,7 +712,7 @@ def __call__( max_sequence_length: int = 256, pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_3_img2img.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_3_img2img.py index 24f3d828bd81..def0cc3c3b03 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_3_img2img.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_3_img2img.py @@ -765,7 +765,7 @@ def __call__( max_sequence_length: int = 256, pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_animatediff.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_animatediff.py index 2baeda5649ad..4303faa08671 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_animatediff.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_animatediff.py @@ -600,7 +600,7 @@ def __call__( decode_chunk_size: int = 16, pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_img2img.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_img2img.py index de6dfbc585fa..dacd2c066a24 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_img2img.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_img2img.py @@ -806,7 +806,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_inpaint.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_inpaint.py index 426419f12f73..2b10667c5cdb 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_inpaint.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_inpaint.py @@ -941,7 +941,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl.py index ca9c6b5aadd9..dd873edc63a6 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl.py @@ -873,7 +873,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_img2img.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_img2img.py index 31fdf19cbade..4566225a240f 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_img2img.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_img2img.py @@ -1029,7 +1029,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_inpaint.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_inpaint.py index 77933867631c..3d7d3500e669 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_inpaint.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_inpaint.py @@ -1125,7 +1125,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/prx/pipeline_prx.py b/src/diffusers/pipelines/prx/pipeline_prx.py index f4ec214313e3..b07fc13d2f0e 100644 --- a/src/diffusers/pipelines/prx/pipeline_prx.py +++ b/src/diffusers/pipelines/prx/pipeline_prx.py @@ -576,7 +576,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], tokenizer_max_length: int | None = None, skip_text_cleaning: bool = False, - ): + ) -> PRXPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/prx/pipeline_prx_pixel.py b/src/diffusers/pipelines/prx/pipeline_prx_pixel.py index 22a4d8dd4b18..4aa6e41aac5d 100644 --- a/src/diffusers/pipelines/prx/pipeline_prx_pixel.py +++ b/src/diffusers/pipelines/prx/pipeline_prx_pixel.py @@ -426,7 +426,7 @@ def __call__( use_resolution_binning: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> PRXPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py index 1da0518a4f65..34bfb8637f68 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py @@ -430,7 +430,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py index f946fdf27d00..ff619a781585 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py @@ -537,7 +537,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py index 97f510a6dbf4..8c09df5574c1 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py @@ -603,7 +603,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py index 85abb815cf23..1861141810f6 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py @@ -527,7 +527,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py index 57d1fdaaf99f..6413205424eb 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py @@ -664,7 +664,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py index 84d1b60152b1..8b17f17c0e3e 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py @@ -549,7 +549,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py index 9b9af83737e5..9beaa769fae2 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py @@ -506,7 +506,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py index 3d5f0040932a..a7b0e4d9912b 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py @@ -619,7 +619,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py index 7e06a7d36ffd..e7053794e9b6 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py @@ -562,7 +562,7 @@ def __call__( resolution: int = 640, cfg_normalize: bool = False, use_en_prompt: bool = False, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py b/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py index 786b09b4e3cd..6aecf130d64f 100644 --- a/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py +++ b/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py @@ -526,7 +526,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], output_resolution: int = 1024, use_kv_cache: bool = True, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/shap_e/pipeline_shap_e.py b/src/diffusers/pipelines/shap_e/pipeline_shap_e.py index eea83aff9e10..56e8f57dab35 100644 --- a/src/diffusers/pipelines/shap_e/pipeline_shap_e.py +++ b/src/diffusers/pipelines/shap_e/pipeline_shap_e.py @@ -200,7 +200,7 @@ def __call__( frame_size: int = 64, output_type: str | None = "pil", # pil, np, latent, mesh return_dict: bool = True, - ): + ) -> ShapEPipelineOutput | tuple: """ The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/shap_e/pipeline_shap_e_img2img.py b/src/diffusers/pipelines/shap_e/pipeline_shap_e_img2img.py index f59fd298c684..403d5fa8a743 100644 --- a/src/diffusers/pipelines/shap_e/pipeline_shap_e_img2img.py +++ b/src/diffusers/pipelines/shap_e/pipeline_shap_e_img2img.py @@ -182,7 +182,7 @@ def __call__( frame_size: int = 64, output_type: str | None = "pil", # pil, np, latent, mesh return_dict: bool = True, - ): + ) -> ShapEPipelineOutput | tuple: """ The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py index 0c9e6add9937..b4a4242b0036 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py @@ -396,7 +396,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> SkyReelsV2PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py index 31b75bfb336f..73b1541f9969 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py @@ -623,7 +623,7 @@ def __call__( ar_step: int = 0, causal_block_size: int | None = None, fps: int = 24, - ): + ) -> SkyReelsV2PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py index 576681b1b957..c19997155d66 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py @@ -672,7 +672,7 @@ def __call__( ar_step: int = 0, causal_block_size: int | None = None, fps: int = 24, - ): + ) -> SkyReelsV2PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py index df6076263238..c18a80fdae0e 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py @@ -710,7 +710,7 @@ def __call__( ar_step: int = 0, causal_block_size: int | None = None, fps: int = 24, - ): + ) -> SkyReelsV2PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py index b1f70b60b22a..977073a60a1e 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py @@ -501,7 +501,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> SkyReelsV2PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_audio/pipeline_stable_audio.py b/src/diffusers/pipelines/stable_audio/pipeline_stable_audio.py index 475f4032edab..b5f7cbfdf93f 100644 --- a/src/diffusers/pipelines/stable_audio/pipeline_stable_audio.py +++ b/src/diffusers/pipelines/stable_audio/pipeline_stable_audio.py @@ -483,7 +483,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int | None = 1, output_type: str | None = "pt", - ): + ) -> AudioPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3.py b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3.py index a1ef3cf3cca7..11004b31b70b 100644 --- a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3.py +++ b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3.py @@ -418,7 +418,7 @@ def __call__( callback_on_step_end: Optional[Callable[[int, int, dict], dict]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], output_type: str = "pt", - ) -> Union[AudioPipelineOutput, tuple]: + ) -> AudioPipelineOutput | tuple: r""" Generate audio from a text prompt. diff --git a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_audio2audio.py b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_audio2audio.py index 50cd4775d5f1..b84388adb12a 100644 --- a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_audio2audio.py +++ b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_audio2audio.py @@ -442,7 +442,7 @@ def __call__( callback_on_step_end: Optional[Callable[[int, int, dict], dict]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], output_type: str = "pt", - ) -> Union[AudioPipelineOutput, tuple]: + ) -> AudioPipelineOutput | tuple: r""" Generate an audio variation conditioned on a text prompt and a reference waveform. diff --git a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_inpaint.py b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_inpaint.py index 687546947eb5..14bafdadbab3 100644 --- a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_inpaint.py +++ b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_inpaint.py @@ -475,7 +475,7 @@ def __call__( callback_on_step_end: Optional[Callable[[int, int, dict], dict]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], output_type: str = "pt", - ) -> Union[AudioPipelineOutput, tuple]: + ) -> AudioPipelineOutput | tuple: r""" Generate inpainted audio conditioned on a text prompt and reference. diff --git a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade.py b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade.py index 0961d1c46e94..c7af35aa955e 100644 --- a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade.py +++ b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade.py @@ -14,6 +14,8 @@ from typing import Callable +import numpy as np +import PIL.Image import torch from transformers import CLIPTextModelWithProjection, CLIPTokenizer @@ -321,7 +323,7 @@ def __call__( return_dict: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> ImagePipelineOutput | list[PIL.Image.Image] | np.ndarray | torch.Tensor: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_combined.py b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_combined.py index 71bfc3f7bab6..a1d2ca6c56b3 100644 --- a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_combined.py +++ b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_combined.py @@ -13,6 +13,7 @@ # limitations under the License. from typing import Callable +import numpy as np import PIL import torch from transformers import CLIPImageProcessor, CLIPTextModelWithProjection, CLIPTokenizer, CLIPVisionModelWithProjection @@ -21,7 +22,7 @@ from ...schedulers import DDPMWuerstchenScheduler from ...utils import is_torch_version, replace_example_docstring from ..deprecated.wuerstchen.modeling_paella_vq_model import PaellaVQModel -from ..pipeline_utils import DeprecatedPipelineMixin, DiffusionPipeline +from ..pipeline_utils import DeprecatedPipelineMixin, DiffusionPipeline, ImagePipelineOutput from .pipeline_stable_cascade import StableCascadeDecoderPipeline from .pipeline_stable_cascade_prior import StableCascadePriorPipeline @@ -181,7 +182,7 @@ def __call__( prior_callback_on_step_end_tensor_inputs: list[str] = ["latents"], callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> ImagePipelineOutput | list[PIL.Image.Image] | np.ndarray | torch.Tensor: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_prior.py b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_prior.py index fb58094f964b..b2fc2799ff31 100644 --- a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_prior.py +++ b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_prior.py @@ -396,7 +396,7 @@ def __call__( return_dict: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> StableCascadePriorPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion.py b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion.py index d7776c7b5196..3c47b5e45b62 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion.py @@ -282,7 +282,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, np.ndarray], None] | None = None, callback_steps: int = 1, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_img2img.py b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_img2img.py index 88b84cff804d..dd27a0de50d2 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_img2img.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_img2img.py @@ -330,7 +330,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, np.ndarray], None] | None = None, callback_steps: int = 1, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_inpaint.py b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_inpaint.py index cf04bdf4da7b..1eb0ea87a8a1 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_inpaint.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_inpaint.py @@ -339,7 +339,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, np.ndarray], None] | None = None, callback_steps: int = 1, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_upscale.py b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_upscale.py index 8494a253f54f..548fdd9e1734 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_upscale.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_upscale.py @@ -366,7 +366,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, np.ndarray], None] | None = None, callback_steps: int | None = 1, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py index d28bb2a9fe59..36b350bd8258 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py @@ -803,7 +803,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_depth2img.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_depth2img.py index 977de5d7fb39..5ffa48daf6d2 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_depth2img.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_depth2img.py @@ -653,7 +653,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_image_variation.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_image_variation.py index 15b8daf334ed..914af27b6a97 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_image_variation.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_image_variation.py @@ -272,7 +272,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_img2img.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_img2img.py index 719be9258341..54a6005d17fe 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_img2img.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_img2img.py @@ -881,7 +881,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_inpaint.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_inpaint.py index 96794eaa297a..dc1f74c74b3d 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_inpaint.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_inpaint.py @@ -908,7 +908,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_instruct_pix2pix.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_instruct_pix2pix.py index 7a24e6008351..c0df751cf054 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_instruct_pix2pix.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_instruct_pix2pix.py @@ -192,7 +192,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], cross_attention_kwargs: dict[str, Any] | None = None, **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_latent_upscale.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_latent_upscale.py index 1920f033c126..c057434b8f8b 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_latent_upscale.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_latent_upscale.py @@ -411,7 +411,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, - ): + ) -> ImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_upscale.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_upscale.py index 1a0a7412e5d7..402f3f67844e 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_upscale.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_upscale.py @@ -553,7 +553,7 @@ def __call__( callback_steps: int = 1, cross_attention_kwargs: dict[str, Any] | None = None, clip_skip: int = None, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip.py index d2c9bf0c4162..1d6feec05846 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip.py @@ -670,7 +670,7 @@ def __call__( prior_guidance_scale: float = 4.0, prior_latents: torch.Tensor | None = None, clip_skip: int | None = None, - ): + ) -> ImagePipelineOutput | tuple: """ The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip_img2img.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip_img2img.py index 059ae1e6fd4d..c560b411620e 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip_img2img.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip_img2img.py @@ -646,7 +646,7 @@ def __call__( noise_level: int = 0, image_embeds: torch.Tensor | None = None, clip_skip: int | None = None, - ): + ) -> ImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3.py b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3.py index 5c05b469660f..9509adde741b 100644 --- a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3.py +++ b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3.py @@ -805,7 +805,7 @@ def __call__( skip_layer_guidance_stop: float = 0.2, skip_layer_guidance_start: float = 0.01, mu: float | None = None, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_img2img.py b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_img2img.py index c0ab805a4ef4..54ae68d19fb4 100644 --- a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_img2img.py +++ b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_img2img.py @@ -860,7 +860,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, mu: float | None = None, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_inpaint.py b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_inpaint.py index 321e9f8dd80e..9f623ec02091 100644 --- a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_inpaint.py +++ b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_inpaint.py @@ -955,7 +955,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, mu: float | None = None, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py index c116a49d81c6..d170e3e805af 100644 --- a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py +++ b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py @@ -859,7 +859,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_img2img.py b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_img2img.py index aedd131aae3c..7ea71e76844b 100644 --- a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_img2img.py +++ b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_img2img.py @@ -1014,7 +1014,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_inpaint.py b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_inpaint.py index 407b1a856216..15b978fab8ab 100644 --- a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_inpaint.py +++ b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_inpaint.py @@ -1124,7 +1124,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_instruct_pix2pix.py b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_instruct_pix2pix.py index bcd337414bac..10afb1b89ee0 100644 --- a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_instruct_pix2pix.py +++ b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_instruct_pix2pix.py @@ -624,7 +624,7 @@ def __call__( original_size: tuple[int, int] = None, crops_coords_top_left: tuple[int, int] = (0, 0), target_size: tuple[int, int] = None, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_video_diffusion/pipeline_stable_video_diffusion.py b/src/diffusers/pipelines/stable_video_diffusion/pipeline_stable_video_diffusion.py index 007d2b8da0cb..eb05eea7837f 100644 --- a/src/diffusers/pipelines/stable_video_diffusion/pipeline_stable_video_diffusion.py +++ b/src/diffusers/pipelines/stable_video_diffusion/pipeline_stable_video_diffusion.py @@ -406,7 +406,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], return_dict: bool = True, - ): + ) -> StableVideoDiffusionPipelineOutput | list[list[PIL.Image.Image]] | np.ndarray | torch.Tensor: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_adapter.py b/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_adapter.py index ffb877cfd0f6..3e111fa752ba 100644 --- a/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_adapter.py +++ b/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_adapter.py @@ -712,7 +712,7 @@ def __call__( cross_attention_kwargs: dict[str, Any] | None = None, adapter_conditioning_scale: float | list[float] = 1.0, clip_skip: int | None = None, - ): + ) -> StableDiffusionAdapterPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_xl_adapter.py b/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_xl_adapter.py index 1e7966192650..56c68ae8c9d0 100644 --- a/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_xl_adapter.py +++ b/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_xl_adapter.py @@ -896,7 +896,7 @@ def __call__( adapter_conditioning_scale: float | list[float] = 1.0, adapter_conditioning_factor: float = 1.0, clip_skip: int | None = None, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_combined.py b/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_combined.py index 2d881e22a783..6b5317e723c9 100644 --- a/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_combined.py +++ b/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_combined.py @@ -273,7 +273,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, upsampling_strength: float = 1.0, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the VisualCloze pipeline for generation. diff --git a/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py b/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py index b34f4c3faeab..cd6e61cff1df 100644 --- a/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py +++ b/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py @@ -677,7 +677,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the VisualCloze pipeline for generation. diff --git a/src/diffusers/pipelines/wan/pipeline_wan.py b/src/diffusers/pipelines/wan/pipeline_wan.py index b33a2a7af3db..452911cf899d 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan.py +++ b/src/diffusers/pipelines/wan/pipeline_wan.py @@ -402,7 +402,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> WanPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/wan/pipeline_wan_animate.py b/src/diffusers/pipelines/wan/pipeline_wan_animate.py index a6b340c2d19f..e96729826e70 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_animate.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_animate.py @@ -791,7 +791,7 @@ def __call__( callback_on_step_end: Callable[[int, int, None], PipelineCallback | MultiPipelineCallbacks] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> WanPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/wan/pipeline_wan_i2v.py b/src/diffusers/pipelines/wan/pipeline_wan_i2v.py index a98f0324e0f0..2d1f7b94750e 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_i2v.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_i2v.py @@ -533,7 +533,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> WanPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/wan/pipeline_wan_vace.py b/src/diffusers/pipelines/wan/pipeline_wan_vace.py index 9186304b5953..f7689496c968 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_vace.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_vace.py @@ -714,7 +714,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> WanPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/wan/pipeline_wan_video2video.py b/src/diffusers/pipelines/wan/pipeline_wan_video2video.py index cfb26f5cb3b1..b192147acb64 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_video2video.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_video2video.py @@ -502,7 +502,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> WanPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image.py b/src/diffusers/pipelines/z_image/pipeline_z_image.py index 3e2055c6257f..b97bdda170bf 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image.py @@ -318,7 +318,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet.py b/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet.py index 81373ffb56ff..6ba698a30db3 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet.py @@ -410,7 +410,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet_inpaint.py b/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet_inpaint.py index 178e74dea4fa..cb72c807700c 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet_inpaint.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet_inpaint.py @@ -419,7 +419,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image_img2img.py b/src/diffusers/pipelines/z_image/pipeline_z_image_img2img.py index b5c7740bb0c1..efa6f50962db 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image_img2img.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image_img2img.py @@ -392,7 +392,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for image-to-image generation. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image_inpaint.py b/src/diffusers/pipelines/z_image/pipeline_z_image_inpaint.py index 132c22c0cff3..62a53ea261d1 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image_inpaint.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image_inpaint.py @@ -560,7 +560,7 @@ def __call__( callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for inpainting. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image_omni.py b/src/diffusers/pipelines/z_image/pipeline_z_image_omni.py index 50776ceaf34d..48ecc6e93163 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image_omni.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image_omni.py @@ -383,7 +383,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/utils/check_return_annotations.py b/utils/check_return_annotations.py new file mode 100644 index 000000000000..be8b1dc61ffc --- /dev/null +++ b/utils/check_return_annotations.py @@ -0,0 +1,180 @@ +# coding=utf-8 +# Copyright 2026 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. +""" +Check that these methods have a return type annotation: + +* `forward()` on every class in `src/diffusers/models` +* `__call__()` on every pipeline in `src/diffusers/pipelines` +* `__call__()` on every modular pipeline block in `src/diffusers/modular_pipelines` + +A class counts as a pipeline if it inherits from `DiffusionPipeline`, either directly or through another class. A class +counts as a modular pipeline block if it inherits from `ModularPipelineBlocks` in the same way. + +Deprecated code is skipped: + +* anything in a folder named `deprecated`, such as `pipelines/deprecated` +* pipelines that inherit from `DeprecatedPipelineMixin` +* classes and methods whose `# Copied from` comment points to deprecated code, because they can't change unless the + deprecated code changes too + +A method is only checked on the class where it's written, not on classes that inherit it. Any annotation passes, +including `-> None`. + +Run from the repository root: + + python utils/check_return_annotations.py +""" + +from __future__ import annotations + +import ast +import sys +from collections import defaultdict +from pathlib import Path + + +REPO_ROOT = Path(__file__).resolve().parents[1] +SRC_DIR = REPO_ROOT / "src" / "diffusers" +MODELS_DIR = SRC_DIR / "models" +PIPELINES_DIR = SRC_DIR / "pipelines" +MODULAR_DIR = SRC_DIR / "modular_pipelines" + +PIPELINE_BASE = "DiffusionPipeline" +DEPRECATED_PIPELINE_BASE = "DeprecatedPipelineMixin" +BLOCKS_BASE = "ModularPipelineBlocks" + + +def _base_names(class_def: ast.ClassDef) -> list[str]: + """Return the names of the classes this class inherits from. For a name like `nn.Module`, keep only `Module`.""" + names = [] + for base in class_def.bases: + if isinstance(base, ast.Name): + names.append(base.id) + elif isinstance(base, ast.Attribute): + names.append(base.attr) + return names + + +def _find_method(class_def: ast.ClassDef, method_name: str) -> ast.FunctionDef | ast.AsyncFunctionDef | None: + for node in class_def.body: + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) and node.name == method_name: + return node + return None + + +def _parse_classes(paths: list[Path]) -> list[tuple[Path, ast.ClassDef, list[str]]]: + """Return every class in `paths`, along with its file and the lines of that file.""" + classes = [] + for path in paths: + try: + source = path.read_text(encoding="utf-8") + tree = ast.parse(source) + except (SyntaxError, UnicodeDecodeError): + continue + lines = source.splitlines() + classes.extend((path, node, lines) for node in ast.walk(tree) if isinstance(node, ast.ClassDef)) + return classes + + +def _is_deprecated_path(path: Path) -> bool: + """Return whether the file is inside a folder named `deprecated`.""" + return "deprecated" in path.relative_to(SRC_DIR).parts[:-1] + + +def _copied_from_deprecated(lines: list[str], node: ast.ClassDef | ast.FunctionDef | ast.AsyncFunctionDef) -> bool: + """Return whether the `# Copied from` comment right above a class or method points to deprecated code.""" + first_line = min([decorator.lineno for decorator in node.decorator_list] + [node.lineno]) + if first_line < 2: + return False + comment = lines[first_line - 2].strip() + return comment.startswith("# Copied from") and ".deprecated." in comment + + +def _subclass_checker(classes: list[tuple[Path, ast.ClassDef, list[str]]]): + """ + Return a function `is_subclass(name, base)` that tells whether the class `name` inherits from the class `base`, + either directly or through other classes. + + Classes are matched by name only, so two classes with the same name in different files are treated as one class. + """ + bases_by_name: dict[str, set[str]] = defaultdict(set) + for _, class_def, _ in classes: + bases_by_name[class_def.name].update(_base_names(class_def)) + + cache: dict[tuple[str, str], bool] = {} + + def is_subclass(name: str, base: str, _seen: frozenset[str] = frozenset()) -> bool: + if name == base: + return True + if (name, base) in cache: + return cache[(name, base)] + if name in _seen: # stop if this class was already visited, so a loop in the class names can't run forever + return False + result = any(is_subclass(parent, base, _seen | {name}) for parent in bases_by_name.get(name, ())) + cache[(name, base)] = result + return result + + return is_subclass + + +def _is_under(path: Path, directory: Path) -> bool: + return directory in path.parents + + +def main() -> int: + classes = _parse_classes(sorted(SRC_DIR.rglob("*.py"))) + is_subclass = _subclass_checker(classes) + + errors = [] + for path, class_def, lines in classes: + if _is_deprecated_path(path): + continue + if _is_under(path, MODELS_DIR): + method_name = "forward" + elif _is_under(path, PIPELINES_DIR): + if not is_subclass(class_def.name, PIPELINE_BASE) or is_subclass(class_def.name, DEPRECATED_PIPELINE_BASE): + continue + method_name = "__call__" + elif _is_under(path, MODULAR_DIR): + if not is_subclass(class_def.name, BLOCKS_BASE): + continue + method_name = "__call__" + else: + continue + + method = _find_method(class_def, method_name) + if method is None or method.returns is not None: + continue + if _copied_from_deprecated(lines, class_def) or _copied_from_deprecated(lines, method): + continue + rel = path.relative_to(REPO_ROOT).as_posix() + errors.append(f"{rel}:{method.lineno}: {class_def.name}.{method_name} has no return type annotation") + + if errors: + print("\n".join(errors)) + sys.stdout.flush() # print the list before the summary, even when both end up in the same log + if len(errors) == 1: + summary = "Found 1 method without a return type annotation. Add one to the method above." + else: + summary = f"Found {len(errors)} methods without a return type annotation. Add one to each method above." + print(f"\n{summary}", file=sys.stderr) + return 1 + + print("All forward/__call__ methods have return type annotations.") + return 0 + + +if __name__ == "__main__": + sys.exit(main())