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/docs/source/en/api/pipelines/krea2.md b/docs/source/en/api/pipelines/krea2.md index 4c7107425d86..19c8a1cb5184 100644 --- a/docs/source/en/api/pipelines/krea2.md +++ b/docs/source/en/api/pipelines/krea2.md @@ -70,6 +70,17 @@ image = pipe( image.save("krea2_turbo.png") ``` +## Loading single-file checkpoints + +```python +import torch +from diffusers import Krea2Pipeline, Krea2Transformer2DModel + +transformer = Krea2Transformer2DModel.from_single_file( + "https://huggingface.co/krea/Krea-2-Turbo/blob/main/turbo.safetensors", dtype=torch.bfloat16 +) +pipe = Krea2Pipeline.from_pretrained("krea/Krea-2-Turbo", transformer=transformer, dtype=torch.bfloat16).to("cuda") +``` ## Krea2Pipeline diff --git a/docs/source/en/api/utilities.md b/docs/source/en/api/utilities.md index 69e69742249f..91ffb1a41ae7 100644 --- a/docs/source/en/api/utilities.md +++ b/docs/source/en/api/utilities.md @@ -50,6 +50,10 @@ Utility and helper functions for working with 🤗 Diffusers. [[autodoc]] utils.torch_utils.randn_tensor +## TorchDeviceBackend + +[[autodoc]] utils.torch_utils.TorchDeviceBackend + ## apply_layerwise_casting [[autodoc]] hooks.layerwise_casting.apply_layerwise_casting diff --git a/docs/source/en/modular_diffusers/modular_pipeline.md b/docs/source/en/modular_diffusers/modular_pipeline.md index 35a4c7c919db..e234a8aa4aae 100644 --- a/docs/source/en/modular_diffusers/modular_pipeline.md +++ b/docs/source/en/modular_diffusers/modular_pipeline.md @@ -471,7 +471,7 @@ pipe.save_pretrained("local/path", repo_id="my-username/flux2-custom-transformer Pass `overwrite_modular_index=False` to keep the loading specs in `modular_model_index.json` as they are. A saved component whose loading spec is empty is still filled in with the destination, since there is nothing to preserve. -Note that moving the files any other way (uploading with `hf upload`, downloading a repository with `hf download --local-dir`) doesn't rewrite the index, so the copy still points to the old location; update the index manually in that case. +Moving the files any other way doesn't rewrite the index. A copy downloaded with `hf download --local-dir` still works: when a pipeline is loaded from a local directory, every component whose files are present in that directory is loaded from it instead of the recorded repository. A copy uploaded with `hf upload` keeps pointing at the old location, so update the index manually in that case. A modular repository can also include custom pipeline blocks as Python code. This allows you to share specialized blocks that aren't native to Diffusers. For example, [diffusers/Florence2-image-Annotator](https://huggingface.co/diffusers/Florence2-image-Annotator) contains custom blocks alongside the loading configuration: 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/docs/source/en/training/distributed_inference.md b/docs/source/en/training/distributed_inference.md index 733b67081983..818488b4cfce 100644 --- a/docs/source/en/training/distributed_inference.md +++ b/docs/source/en/training/distributed_inference.md @@ -436,43 +436,49 @@ pipeline = DiffusionPipeline.from_pretrained( [Tensor parallelism](https://huggingface.co/spaces/nanotron/ultrascale-playbook?section=tensor_parallelism) shards the weight matrices of a model across devices. Each device holds a column-wise (`"colwise"`) or row-wise (`"rowwise"`) slice of each layer, computes a partial result, and an `AllReduce`/`AllGather` at the layer boundary reconstructs the full output. Unlike context parallelism, it reduces the per-device *weight* memory, which is useful for models that do not fit on a single device. -Pass a [`TensorParallelConfig`] to [`~ModelMixin.enable_parallelism`]. `tp_degree` is the number of devices to shard across and must divide the model's number of attention heads. The model must define a `_tp_plan` (a flat mapping of module-name globs to a `"colwise"`/`"rowwise"` style). +Pass a [`TensorParallelConfig`] to the `parallel_config` argument of the model's [`~ModelMixin.from_pretrained`]. `tp_degree` is the number of devices to shard across and must divide the model's number of attention heads. The model must define a `_tp_plan` (a flat mapping of module-name globs to a `"colwise"`/`"rowwise"` style). + +Loading this way shards the checkpoint *while reading it*: each rank reads only its own slice of each sharded weight and places it straight onto its own device. Nothing full-size is ever materialized, so per-rank memory falls as `tp_degree` rises. + +Compared to loading the full model and then calling [`~ModelMixin.enable_parallelism`], it loads faster and uses less CPU memory per rank, with the gap growing as `tp_degree` rises. Numbers below are for [black-forest-labs/FLUX.2-dev](https://huggingface.co/black-forest-labs/FLUX.2-dev) (transformer only, 32B params, bf16) on 4x A10G (23GB). + +| tp_degree | method | load time | peak CPU/rank | +|---|---|---|---| +| 4 | `from_pretrained(parallel_config=...)` | 12.5s | 6.8GB | +| 4 | `from_pretrained` + `enable_parallelism` | 30.4s | 64.1GB | ```py import torch from torch import distributed as dist -from diffusers import DiffusionPipeline, TensorParallelConfig +from diffusers import DiffusionPipeline, Flux2Transformer2DModel, TensorParallelConfig -def setup_distributed(): - if not dist.is_initialized(): - dist.init_process_group(backend="nccl") - rank = dist.get_rank() +def main(): + dist.init_process_group(backend="nccl") + rank, world_size = dist.get_rank(), dist.get_world_size() device = torch.device(f"cuda:{rank}") torch.cuda.set_device(device) - return device -def main(): - device = setup_distributed() - world_size = dist.get_world_size() + # Each rank reads only its own shard of every planned weight, straight onto `cuda:rank`. + transformer = Flux2Transformer2DModel.from_pretrained( + "black-forest-labs/FLUX.2-dev", + subfolder="transformer", + torch_dtype=torch.bfloat16, + parallel_config=TensorParallelConfig(tp_degree=world_size), + ) pipeline = DiffusionPipeline.from_pretrained( - "black-forest-labs/FLUX.2-dev", torch_dtype=torch.bfloat16 - ) # weights stay on CPU - - # Shard the transformer first, then move only each rank's slice onto the accelerator. - pipeline.transformer.enable_parallelism(config=TensorParallelConfig(tp_degree=world_size)) - pipeline.transformer.to(device) - - # Move the remaining, non-sharded components onto the accelerator individually. + "black-forest-labs/FLUX.2-dev", transformer=transformer, torch_dtype=torch.bfloat16 + ) + # The transformer is already on its device; move the remaining components individually. Do not call + # `pipeline.to(device)` — that would move every rank's shards onto the same device. pipeline.text_encoder.to(device) pipeline.vae.to(device) generator = torch.Generator().manual_seed(42) image = pipeline(prompt="a cat holding a sign that says hello", generator=generator).images[0] - if dist.get_rank() == 0: + if rank == 0: image.save("output.png") - if dist.is_initialized(): - dist.destroy_process_group() + dist.destroy_process_group() if __name__ == "__main__": main() @@ -484,6 +490,15 @@ torchrun --nproc-per-node 4 tensor_parallel_flux.py `tp_degree` is taken from `world_size` above, so `--nproc-per-node 4` shards the transformer across 4 devices. +> [!CAUTION] +> Loading with a tensor-parallel `parallel_config` isn't supported yet with `device_map`, `quantization_config`, `low_cpu_mem_usage=False`, `use_flashpack=True`, or non-safetensors weights; each raises rather than quietly falling back to loading the full checkpoint. +> +> Combining tensor parallelism with quantization, offloading, or LoRA adapters isn't supported yet either, so those raise however the model is sharded. +> +> To shard a model that is already in memory, call [`~ModelMixin.enable_parallelism`] with the same config instead — that loads everything first and reshards it, so it costs full checkpoint memory on every rank. + +Saving a tensor-parallel model isn't supported yet, and [`~ModelMixin.save_pretrained`] raises on one. Save the model before sharding it. + ### Writing a tensor parallelism plan Tensor parallelism only works on models that define a `_tp_plan`, a flat class attribute mapping module-name globs to a sharding style. Writing one is mostly a matter of pairing each projection that *expands* the hidden dimension with the projection that *contracts* it back. @@ -536,9 +551,11 @@ Anything absent from the plan stays replicated on every rank, which is the right #### Constraints and verification -- `tp_degree` must divide `config.num_attention_heads`. This is validated in [`~ModelMixin.enable_parallelism`]. +- `tp_degree` must divide `config.num_attention_heads`. - Every packed block must *individually* be divisible by `tp_degree`, not just their sum. +Both are validated by [`~ModelMixin.from_pretrained`] and [`~ModelMixin.enable_parallelism`] before any weight is loaded or sharded. + Validate a new plan numerically rather than by eye: generate with a fixed seed on a single device, then again under tensor parallelism, and compare the outputs. A misplaced `"colwise"`/`"rowwise"` usually still runs and produces a plausible but wrong image. > [!TIP] diff --git a/src/diffusers/hooks/group_offloading.py b/src/diffusers/hooks/group_offloading.py index e6c8d8732113..2c3ac131e4d8 100644 --- a/src/diffusers/hooks/group_offloading.py +++ b/src/diffusers/hooks/group_offloading.py @@ -23,6 +23,7 @@ import torch from ..utils import get_logger, is_accelerate_available, is_torchao_available +from ..utils.torch_utils import TorchDeviceBackend from ._common import _GO_LC_SUPPORTED_PYTORCH_LAYERS from .hooks import HookRegistry, ModelHook @@ -166,11 +167,7 @@ def __init__( else: self.cpu_param_dict = self._init_cpu_param_dict() - self._torch_accelerator_module = ( - getattr(torch, torch.accelerator.current_accelerator().type) - if hasattr(torch, "accelerator") - else torch.cuda - ) + self._torch_accelerator_module = TorchDeviceBackend(self.onload_device) @staticmethod def _to_cpu(tensor, low_cpu_mem_usage): @@ -671,12 +668,13 @@ def apply_group_offloading( stream = None if use_stream: - if torch.cuda.is_available(): - stream = torch.cuda.Stream() - elif hasattr(torch, "xpu") and torch.xpu.is_available(): - stream = torch.Stream() - else: - raise ValueError("Using streams for data transfer requires a CUDA device, or an Intel XPU device.") + backend = TorchDeviceBackend(onload_device) + if onload_device.type == "cpu" or not hasattr(backend, "Stream"): + raise ValueError( + "Using streams for data transfer requires an onload device whose backend implements streams, " + f"got `{onload_device.type}`. Pass `use_stream=False`." + ) + stream = backend.Stream() if not use_stream and record_stream: raise ValueError("`record_stream` cannot be True when `use_stream=False`.") 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/src/diffusers/hooks/tensor_parallel.py b/src/diffusers/hooks/tensor_parallel.py index b90a5761d043..ce8abd844ab5 100644 --- a/src/diffusers/hooks/tensor_parallel.py +++ b/src/diffusers/hooks/tensor_parallel.py @@ -12,15 +12,17 @@ # See the License for the specific language governing permissions and # limitations under the License. +from typing import NamedTuple + import torch from ..models._modeling_parallel import TensorParallelConfig -from ..utils import get_logger +from ..utils import get_logger, is_peft_available logger = get_logger(__name__) # pylint: disable=invalid-name -_SUPPORTED_TP_DEVICES = ("cuda", "neuron") +_SUPPORTED_TP_DEVICES = ("cuda", "neuron", "tpu") class PackedColwiseParallel: @@ -48,6 +50,32 @@ def __init__(self, blocks: "list[int] | None" = None): self.blocks = blocks +class ReplicatedInputRowwiseParallel: + """Row-wise sharding for a Linear whose input arrives replicated instead of column-sharded. + + Plain `"rowwise"` is the second half of a colwise/rowwise pair, so it expects its input to already be `Shard(-1)` + — which it is when the preceding Linear was colwise-sharded. A Linear that instead reads a replicated activation, + such as a modulation projection off the shared timestep embedding, needs its input sharded on the way in (a local + narrow, no collective) and its partial output all-reduced on the way out. + + Weight and bias shard exactly as for plain `"rowwise"`: the weight over its input columns, the bias replicated and + added after the all-reduce. Use this to shard a large standalone projection whose output must keep the full + feature dimension, where colwise sharding would need an extra all-gather to rebuild it. + """ + + +def _packed_blocks(style: "PackedColwiseParallel | PackedRowwiseParallel", module, path: str) -> "list[int]": + """The blocks of a packed style: its own `blocks`, else the `_tp_packed_*_blocks` attribute on the Linear.""" + attr = "_tp_packed_col_blocks" if isinstance(style, PackedColwiseParallel) else "_tp_packed_row_blocks" + blocks = style.blocks if style.blocks is not None else getattr(module, attr, None) + if not blocks: + raise ValueError( + f"'{path}' uses {type(style).__name__} but has no blocks: pass `blocks` to it, or set `{attr}` on the " + f"Linear in the model's `__init__`." + ) + return blocks + + def _blocks_to_block_sizes(total_size: int, blocks: "list[int]") -> "list[int]": """Convert proportional block counts to absolute sizes. @@ -65,6 +93,110 @@ def _blocks_to_block_sizes(total_size: int, blocks: "list[int]") -> "list[int]": return [b * unit for b in blocks] +class TPShardSpec(NamedTuple): + """How one parameter is laid out across the tensor-parallel ranks. + + `dim` is the dimension sharded across ranks, or `None` when the parameter is replicated on every rank (a rowwise + bias, which is added after the all-reduce). `block_sizes` partitions `dim` into independently sharded blocks; a + plain `"colwise"` / `"rowwise"` style has a single block covering the whole dimension, and packed styles have one + per fused projection. + """ + + dim: "int | None" + block_sizes: "list[int] | None" + + +def _local_shard(tensor, dim: int, block_sizes: "list[int]", tp_mesh) -> torch.Tensor: + """Extract this rank's slice of `tensor` along `dim`. + + `tensor` may be a `torch.Tensor` or a safetensors `PySafeSlice`, so the same arithmetic serves both resharding a + weight already in memory and reading only one rank's slice off disk. Note that `PySafeSlice` exposes `get_shape()` + rather than `.ndim`, and that a `dim`-1 slice comes back strided, hence the final `contiguous()` — + `DTensor.from_local` needs a contiguous local tensor. + """ + rank = tp_mesh.get_local_rank() + tp_size = tp_mesh.size() + ndim = tensor.dim() if isinstance(tensor, torch.Tensor) else len(tensor.get_shape()) + + parts, offset = [], 0 + for block_size in block_sizes: + chunk = block_size // tp_size + index = [slice(None)] * ndim + index[dim] = slice(offset + rank * chunk, offset + (rank + 1) * chunk) + parts.append(tensor[tuple(index)]) + offset += block_size + + local = parts[0] if len(parts) == 1 else torch.cat(parts, dim=dim) + return local.contiguous() + + +def resolve_tp_shard_specs(model: torch.nn.Module, tp_plan: dict, tp_degree: int) -> "dict[str, TPShardSpec]": + """Map every `_tp_plan`-covered parameter name to its `TPShardSpec`. + + Parameters absent from the result are untouched by tensor parallelism. Both `weight` and `bias` of each planned + module are covered. + + Raises if any sharded block is not divisible by `tp_degree`, so an uneven split is rejected before any weight is + read or sharded. Left to `Shard`, it would give the trailing ranks a smaller slice and break the paired + colwise/rowwise matmul. + + The plan is expanded by `_resolve_tp_plan` so there is a single implementation of the glob rules; the `id(module) + -> name` map recovers qualified names from the submodules it returns. Going back through `_resolve_tp_plan` also + handles a model that reuses one block instance in two places, which re-expanding the globs here would get wrong. + + Safe to call on a meta model: only shapes, `bias is not None`, and the packed-block attributes set in the module's + `__init__` are read. + """ + names = {id(module): name for name, module in model.named_modules()} + specs: dict[str, TPShardSpec] = {} + + for block, relative_plan in _resolve_tp_plan(model, tp_plan): + prefix = names[id(block)] + for relative_path, style in relative_plan.items(): + submodule = block + for atom in relative_path.split("."): + submodule = getattr(submodule, atom) + path = f"{prefix}.{relative_path}" if prefix else relative_path + + # `_tp_packed_*_blocks` hold absolute sizes rather than proportions; that works because + # they sum to the full dimension, so `_blocks_to_block_sizes` computes `unit == 1`. + if style == "colwise": + weight_spec = TPShardSpec(0, [submodule.weight.shape[0]]) + bias_spec = weight_spec + elif style == "rowwise" or isinstance(style, ReplicatedInputRowwiseParallel): + # Both place the weight the same way; they differ only in the forward input/output hooks. + weight_spec = TPShardSpec(1, [submodule.weight.shape[1]]) + bias_spec = TPShardSpec(None, None) + elif isinstance(style, PackedColwiseParallel): + blocks = _packed_blocks(style, submodule, path) + weight_spec = TPShardSpec(0, _blocks_to_block_sizes(submodule.weight.shape[0], blocks)) + bias_spec = weight_spec + elif isinstance(style, PackedRowwiseParallel): + blocks = _packed_blocks(style, submodule, path) + weight_spec = TPShardSpec(1, _blocks_to_block_sizes(submodule.weight.shape[1], blocks)) + bias_spec = TPShardSpec(None, None) + else: + raise ValueError( + f"Unsupported tensor-parallel style '{style}' for '{path}'. " + f"Expected 'colwise', 'rowwise', PackedColwiseParallel, PackedRowwiseParallel, or " + f"ReplicatedInputRowwiseParallel." + ) + + specs[f"{path}.weight"] = weight_spec + if submodule.bias is not None: + specs[f"{path}.bias"] = bias_spec + + for name, spec in specs.items(): + for block_size in spec.block_sizes or []: + if block_size % tp_degree != 0: + raise ValueError( + f"Cannot shard '{name}' across {tp_degree} tensor-parallel ranks: a block of size {block_size} " + f"along dim {spec.dim} is not divisible by {tp_degree}." + ) + + return specs + + def _resolve_tp_plan(model: torch.nn.Module, tp_plan: dict) -> list: """Group a flat `_tp_plan` into per-block `(submodule, {relative_path: style})` plans. @@ -107,127 +239,73 @@ def _resolve_tp_plan(model: torch.nn.Module, tp_plan: dict) -> list: return [grouped[key] for key in order] +def _shard_packed_param(param, dim: int, blocks: "list[int]", device_mesh, src_data_rank) -> torch.nn.Parameter: + """Shard a packed `param` along `dim`, splitting each fused block across the ranks separately. + + The parameter is replicated before slicing: the broadcast from `src_data_rank` is what makes one rank's weights + authoritative when the model was randomly initialized rather than loaded from a checkpoint, in which case every + rank starts with different values. + """ + from torch.distributed.tensor import DTensor, Replicate, Shard, distribute_tensor + + full = distribute_tensor(param, device_mesh, [Replicate()], src_data_rank=src_data_rank).to_local() + local = _local_shard(full, dim, _blocks_to_block_sizes(full.shape[dim], blocks), device_mesh) + return torch.nn.Parameter( + DTensor.from_local(local, device_mesh, [Shard(dim)], run_check=False), requires_grad=param.requires_grad + ) + + def _styles(relative_plan: dict) -> dict: """Map a `{relative_path: style}` plan to `parallelize_module` style instances. - Values may be plain strings (`"colwise"` / `"rowwise"`) or `PackedColwiseParallel` / `PackedRowwiseParallel` marker - instances. Returns `{relative_path: ColwiseParallel() | RowwiseParallel() | }`, each subclassed to - reject a sharded dim that is not divisible by the TP degree. + Values may be plain strings (`"colwise"` / `"rowwise"`) or `PackedColwiseParallel` / `PackedRowwiseParallel` / + `ReplicatedInputRowwiseParallel` marker instances. Returns `{relative_path: ColwiseParallel() | RowwiseParallel() | + }`. Divisibility by the TP degree is checked earlier, by `resolve_tp_shard_specs`. """ - import torch.nn as nn - from torch.distributed.tensor import DTensor, Replicate, Shard, distribute_tensor + from torch.distributed.tensor import Replicate, distribute_tensor from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel def _make_packed_col(marker: PackedColwiseParallel) -> ColwiseParallel: - _blocks = marker.blocks - class _PackedColwiseImpl(ColwiseParallel): def _partition_linear_fn(self, name, module, device_mesh): - blocks = _blocks if _blocks is not None else getattr(module, "_tp_packed_col_blocks") - rank = device_mesh.get_local_rank() - tp_size = device_mesh.size() + blocks = _packed_blocks(marker, module, name) # Both weight (`[out, in]`) and bias (`[out]`) are sharded row-wise (dim 0) with the same per-block # slicing so each rank's bias rows line up with its weight rows for the packed layout. for param_name, param in module.named_parameters(): - full = distribute_tensor( - param, device_mesh, [Replicate()], src_data_rank=self.src_data_rank - ).to_local() - block_sizes = _blocks_to_block_sizes(full.shape[0], blocks) - parts, offset = [], 0 - for bs in block_sizes: - if bs % tp_size != 0: - raise ValueError( - f"Cannot shard packed block of size {bs} across {tp_size} tensor-parallel ranks: " - f"{bs} is not divisible by {tp_size}." - ) - chunk = bs // tp_size - parts.append(full[offset + rank * chunk : offset + (rank + 1) * chunk].contiguous()) - offset += bs - local = torch.cat(parts, dim=0) - dist_param = nn.Parameter( - DTensor.from_local(local, device_mesh, [Shard(0)], run_check=False), - requires_grad=param.requires_grad, - ) - module.register_parameter(param_name, dist_param) + sharded = _shard_packed_param(param, 0, blocks, device_mesh, self.src_data_rank) + module.register_parameter(param_name, sharded) return _PackedColwiseImpl() def _make_packed_row(marker: PackedRowwiseParallel) -> RowwiseParallel: - _blocks = marker.blocks - class _PackedRowwiseImpl(RowwiseParallel): def _partition_linear_fn(self, name, module, device_mesh): - blocks = _blocks if _blocks is not None else getattr(module, "_tp_packed_row_blocks") - rank = device_mesh.get_local_rank() - tp_size = device_mesh.size() + blocks = _packed_blocks(marker, module, name) + # Only the weight (`[out, in]`) is sharded, column-wise (dim 1); the bias is added after the + # all-reduce, so every rank keeps it whole. for param_name, param in module.named_parameters(): if param_name == "weight": - full = distribute_tensor( - param, device_mesh, [Replicate()], src_data_rank=self.src_data_rank - ).to_local() - block_sizes = _blocks_to_block_sizes(full.shape[1], blocks) - parts, offset = [], 0 - for bs in block_sizes: - if bs % tp_size != 0: - raise ValueError( - f"Cannot shard packed block of size {bs} across {tp_size} tensor-parallel ranks: " - f"{bs} is not divisible by {tp_size}." - ) - chunk = bs // tp_size - parts.append(full[:, offset + rank * chunk : offset + (rank + 1) * chunk].contiguous()) - offset += bs - local = torch.cat(parts, dim=1) - dist_param = nn.Parameter( - DTensor.from_local(local, device_mesh, [Shard(1)], run_check=False), - requires_grad=param.requires_grad, - ) + sharded = _shard_packed_param(param, 1, blocks, device_mesh, self.src_data_rank) else: - dist_param = nn.Parameter( + sharded = torch.nn.Parameter( distribute_tensor(param, device_mesh, [Replicate()], src_data_rank=self.src_data_rank), requires_grad=param.requires_grad, ) - module.register_parameter(param_name, dist_param) + module.register_parameter(param_name, sharded) return _PackedRowwiseImpl() - # `distribute_tensor` accepts an indivisible shard dim and just gives the trailing ranks a smaller (or empty) - # slice, so an uneven split does not raise here — it surfaces much later as a shape or numerics error, because - # the attention head split and the paired colwise/rowwise Linear both assume equal shards. Reject it up front, - # matching what the packed styles above and the Neuron pre-shard path already do. - def _make_checked_col(path: str) -> ColwiseParallel: - class _CheckedColwiseImpl(ColwiseParallel): - def _partition_linear_fn(self, name, module, device_mesh): - tp_size = device_mesh.size() - out_features = module.weight.shape[0] - if out_features % tp_size != 0: - raise ValueError( - f"Cannot colwise-shard '{path}' weight rows ({out_features}) across {tp_size} " - f"tensor-parallel ranks: not divisible by {tp_size}." - ) - super()._partition_linear_fn(name, module, device_mesh) - - return _CheckedColwiseImpl() - - def _make_checked_row(path: str) -> RowwiseParallel: - class _CheckedRowwiseImpl(RowwiseParallel): - def _partition_linear_fn(self, name, module, device_mesh): - tp_size = device_mesh.size() - in_features = module.weight.shape[1] - if in_features % tp_size != 0: - raise ValueError( - f"Cannot rowwise-shard '{path}' weight columns ({in_features}) across {tp_size} " - f"tensor-parallel ranks: not divisible by {tp_size}." - ) - super()._partition_linear_fn(name, module, device_mesh) - - return _CheckedRowwiseImpl() - resolved = {} for path, style in relative_plan.items(): if style == "colwise": - resolved[path] = _make_checked_col(path) + resolved[path] = ColwiseParallel() elif style == "rowwise": - resolved[path] = _make_checked_row(path) + resolved[path] = RowwiseParallel() + elif isinstance(style, ReplicatedInputRowwiseParallel): + # `input_layouts=Replicate()` makes `prepare_input` narrow the replicated activation down to this rank's + # columns rather than trusting it to already be `Shard(-1)`; the default would read a full-width tensor + # as if it were one rank's shard. + resolved[path] = RowwiseParallel(input_layouts=Replicate()) elif isinstance(style, PackedColwiseParallel): resolved[path] = _make_packed_col(style) elif isinstance(style, PackedRowwiseParallel): @@ -235,17 +313,247 @@ def _partition_linear_fn(self, name, module, device_mesh): else: raise ValueError( f"Unsupported tensor-parallel style '{style}' for '{path}'. " - f"Expected 'colwise', 'rowwise', PackedColwiseParallel, or PackedRowwiseParallel." + f"Expected 'colwise', 'rowwise', PackedColwiseParallel, PackedRowwiseParallel, or " + f"ReplicatedInputRowwiseParallel." ) return resolved +def _hooks_only_styles(relative_plan: dict) -> dict: + """Map a `{relative_path: style}` plan to styles that partition nothing. + + Used when the caller has already placed every planned parameter as a sharded `DTensor`. `parallelize_module` then + runs only to register the forward input/output hooks; `_partition_linear_fn` must not re-partition. Packed and + plain styles share hook behaviour, so both collapse onto the two styles here. + + Note this is not purely additive: `distribute_module` still replicates any *remaining* plain parameter of the + targeted module into a `Replicate()` DTensor via a broadcast. Callers should therefore place every planned + parameter themselves, and must ensure none is left on `meta` — the broadcast would be issued on a meta tensor. + """ + from torch.distributed.tensor import Replicate + from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel + + class _NoPartitionColwise(ColwiseParallel): + def _partition_linear_fn(self, name, module, device_mesh): + pass # weight already Shard(0) + + class _NoPartitionRowwise(RowwiseParallel): + def _partition_linear_fn(self, name, module, device_mesh): + pass # weight already Shard(1) + + resolved = {} + for path, style in relative_plan.items(): + if style == "colwise" or isinstance(style, PackedColwiseParallel): + resolved[path] = _NoPartitionColwise() + elif style == "rowwise" or isinstance(style, PackedRowwiseParallel): + resolved[path] = _NoPartitionRowwise() + elif isinstance(style, ReplicatedInputRowwiseParallel): + resolved[path] = _NoPartitionRowwise(input_layouts=Replicate()) + else: + raise ValueError( + f"Unsupported tensor-parallel style '{style}' for '{path}'. " + f"Expected 'colwise', 'rowwise', PackedColwiseParallel, PackedRowwiseParallel, or " + f"ReplicatedInputRowwiseParallel." + ) + return resolved + + +def _pre_shard_and_parallelize( + model: torch.nn.Module, + tp_mesh: "torch.distributed.device_mesh.DeviceMesh", + groups: list, + specs: "dict[str, TPShardSpec]", + device: torch.device, +) -> None: + """Slice every planned parameter on CPU, place only this rank's shard on `device`, then register the hooks. + + Unlike the default path this does not broadcast from a single rank, so every rank must already hold the same + weights. Model weights must be on CPU when this is called. + """ + import torch.nn as nn + from torch.distributed.tensor import DTensor, Replicate, Shard + from torch.distributed.tensor.parallel import parallelize_module + + for name, spec in specs.items(): + path, _, param_name = name.rpartition(".") + module = model.get_submodule(path) + param = getattr(module, param_name) + + if spec.dim is None: + # A rowwise bias is added after the all-reduce, so every rank needs the whole vector. + local, placement = param.data, Replicate() + else: + local, placement = _local_shard(param.data, spec.dim, spec.block_sizes, tp_mesh), Shard(spec.dim) + + module.register_parameter( + param_name, + nn.Parameter( + DTensor.from_local(local.to(device), tp_mesh, [placement]), + requires_grad=param.requires_grad, + ), + ) + + # `parallelize_module` is now a no-op for weight distribution (they are already DTensors) but still registers the + # input/output hooks required for the forward pass. + for block, relative_plan in groups: + parallelize_module(block, tp_mesh, _hooks_only_styles(relative_plan)) + + +def _tp_degree(tp_config) -> int: + """The TP degree, read the same way `_resolve_parallel_config` will, for checks run before the mesh is built.""" + return tp_config.mesh.size() if tp_config.mesh is not None else tp_config.tp_degree + + +def _is_tensor_parallel(module: torch.nn.Module) -> bool: + """Whether `module` has been sharded with tensor parallelism.""" + parallel_config = getattr(module, "_parallel_config", None) + return parallel_config is not None and parallel_config.tensor_parallel_config is not None + + +def _raise_if_tensor_parallel(module: torch.nn.Module, action: str, reason: str, error_cls=ValueError) -> None: + """Reject an operation that cannot run on a model already sharded with tensor parallelism. + + `action` completes "so it cannot ...", and `reason` says why or what to do instead. + """ + if _is_tensor_parallel(module): + raise error_cls( + f"'{module.__class__.__name__}' is sharded with tensor parallelism, so it cannot {action}. {reason}" + ) + + +def _check_tp_weights_format(model_files: list) -> None: + """Reject non-safetensors weights for a tensor-parallel load, which needs to read one slice of each tensor. + + Separate from `_check_tp_supported` because it needs the resolved checkpoint files, which `from_pretrained` only + knows later. + """ + non_safetensors = [f for f in model_files if not str(f).endswith(".safetensors")] + if non_safetensors: + raise ValueError( + f"A tensor-parallel `parallel_config` requires safetensors weights, so that each rank can " + f"read only its own slice of each tensor. Got {non_safetensors}." + ) + + +def _check_tp_supported( + model_name: str, + tp_plan: "dict | None", + num_heads: "int | None", + tp_config, + *, + model: "torch.nn.Module | None" = None, + device_map=None, + hf_quantizer=None, + low_cpu_mem_usage: bool = True, + use_flashpack: bool = False, + use_safetensors: bool = True, +) -> None: + """Reject a model class, `tp_degree`, loading option, or model state that tensor parallelism cannot shard. + + Both entry points run it before anything is loaded or sharded. `from_pretrained` passes its loading options: + sharding on load needs a meta-initialized model and lazily sliceable safetensors files, and rather than silently + falling back to loading the full checkpoint — which would quietly give up the memory saving — each unsupported + option raises. `enable_parallelism` passes the `model` already in memory instead, whose parameters must be plain + parameters owned by the model itself: a quantized, offloaded, or adapter-wrapped model is rejected up front rather + than failing deep inside `parallelize_module` — or, worse, sharding successfully and producing wrong numbers. + """ + if tp_plan is None: + raise ValueError( + f"`_tp_plan` must be set on the model class to use tensor parallelism. '{model_name}' does not define one." + ) + + tp_degree = _tp_degree(tp_config) + if num_heads is not None and num_heads % tp_degree != 0: + raise ValueError(f"`tp_degree` ({tp_degree}) must divide the number of attention heads ({num_heads}).") + + if device_map is not None: + raise ValueError( + "`device_map` cannot be combined with a tensor-parallel `parallel_config`: tensor parallelism " + "already places each rank's shard on that rank's device. Drop `device_map`." + ) + if hf_quantizer is not None: + raise ValueError( + "`quantization_config` cannot be combined with a tensor-parallel `parallel_config`: quantized " + "parameters are packed into a quantizer-specific layout that cannot be sharded into `DTensor`s. " + "Load the model unquantized to shard it." + ) + if not low_cpu_mem_usage: + raise ValueError( + "`low_cpu_mem_usage=False` cannot be combined with a tensor-parallel `parallel_config`: " + "streaming each rank's shard requires the model to be initialized on the meta device." + ) + if use_flashpack: + raise ValueError( + "`use_flashpack=True` cannot be combined with a tensor-parallel `parallel_config`; FlashPack " + "checkpoints cannot be sliced per rank." + ) + if not use_safetensors: + raise ValueError( + "`use_safetensors=False` cannot be combined with a tensor-parallel `parallel_config`: each rank reads " + "only its own slice of each tensor, which needs safetensors weights." + ) + + if model is not None: + if getattr(model, "hf_quantizer", None) is not None or getattr(model, "is_quantized", False): + raise ValueError( + f"'{model.__class__.__name__}' is quantized, which cannot be combined with tensor parallelism: its " + "parameters are packed into a quantizer-specific layout that cannot be sharded into `DTensor`s. Load " + "the model unquantized to shard it." + ) + + from .group_offloading import _is_group_offload_enabled + + if _is_group_offload_enabled(model): + raise ValueError( + f"'{model.__class__.__name__}' has group offloading enabled, which cannot be combined with tensor " + "parallelism: both decide where a parameter lives. Tensor parallelism already keeps only one shard of " + "each weight per rank, so offloading is not needed on top of it." + ) + + # `device_map` dispatch and accelerate's CPU offloading both leave an `_hf_hook` on every module they placed, + # and the weights they offloaded are `meta` tensors that `DTensor.from_local` cannot shard. + if getattr(model, "hf_device_map", None) is not None or any( + hasattr(module, "_hf_hook") for module in model.modules() + ): + raise ValueError( + f"'{model.__class__.__name__}' is placed by accelerate — through `device_map` or CPU offloading — " + "which cannot be combined with tensor parallelism: tensor parallelism already places each rank's " + "shard on that rank's device. Load the model without `device_map` and without offloading to shard it." + ) + + if is_peft_available(): + from peft.tuners.tuners_utils import BaseTunerLayer + + if any(isinstance(module, BaseTunerLayer) for module in model.modules()): + raise ValueError( + f"'{model.__class__.__name__}' has adapter (LoRA) layers injected, which cannot be combined with " + "tensor parallelism: `_tp_plan` covers the base `Linear` layers only, so the adapter weights " + "would stay unsharded and the result would be wrong. Unload the adapter before sharding." + ) + + def apply_tensor_parallel( model: torch.nn.Module, config: TensorParallelConfig, tp_plan: dict, + weights_already_sharded: bool = False, ) -> None: - """Apply tensor parallel on a model from its flat `_tp_plan`.""" + """Apply tensor parallel on a model from its flat `_tp_plan`. + + Args: + model (`torch.nn.Module`): + The model to shard in place. + config (`TensorParallelConfig`): + The tensor-parallel config. Its device mesh must already be set up with `config.setup(...)`. + tp_plan (`dict`): + A flat mapping of module-name globs to a `"colwise"`/`"rowwise"` style (or a packed variant), usually the + model's `_tp_plan`. + weights_already_sharded (`bool`, defaults to `False`): + Whether the planned parameters are already `DTensor` shards, as they are after a streaming + `from_pretrained` load. If `True`, only the forward hooks are registered. This is passed explicitly rather + than detected, because a planned parameter missing from the checkpoint would still be a meta tensor and + would make detection say "not sharded" for a model that is in fact half-sharded. + """ tp_mesh = config._mesh if tp_mesh is None: raise ValueError("`config._mesh` is None. Call `config.setup(rank, world_size, device)` before applying TP.") @@ -257,17 +565,31 @@ def apply_tensor_parallel( f"or from the active accelerator when the mesh is built from `tp_degree`." ) - backend = "neuron" if tp_mesh.device_type == "neuron" else "default" + backend = tp_mesh.device_type if tp_mesh.device_type in ("neuron", "tpu") else "default" groups = _resolve_tp_plan(model, tp_plan) logger.debug(f"Applying tensor parallel (backend={backend}) over {len(groups)} module group(s) on mesh {tp_mesh}.") + from torch.distributed.tensor.parallel import parallelize_module + + if weights_already_sharded: + for submodule, relative_plan in groups: + parallelize_module(submodule, tp_mesh, _hooks_only_styles(relative_plan)) + return + + # Also validates every planned shard, so an uneven split raises before any module is sharded. + specs = resolve_tp_shard_specs(model, tp_plan, tp_mesh.size()) + if backend == "neuron": from .tensor_parallel_neuron import _apply_tp_neuron - _apply_tp_neuron(model, tp_mesh, groups) + _apply_tp_neuron(model, tp_mesh, groups, specs) return - from torch.distributed.tensor.parallel import parallelize_module + if backend == "tpu": + # `parallelize_module` would materialize every full weight on each chip before scattering it, which runs out + # of HBM on models larger than one chip. Pre-sharding on CPU keeps only this rank's slice on the device. + _pre_shard_and_parallelize(model, tp_mesh, groups, specs, config._device) + return for submodule, relative_plan in groups: parallelize_module(submodule, tp_mesh, _styles(relative_plan)) diff --git a/src/diffusers/hooks/tensor_parallel_neuron.py b/src/diffusers/hooks/tensor_parallel_neuron.py index 6b8f219a17ff..a87f28b1edb1 100644 --- a/src/diffusers/hooks/tensor_parallel_neuron.py +++ b/src/diffusers/hooks/tensor_parallel_neuron.py @@ -17,162 +17,32 @@ The difference from the generic path is a workaround for a Neuron NRT bug: consecutive `reduce_scatter` collectives for large weight tensors (≥ 5120×5120) can fail when all layers are distributed in a single `parallelize_module` call. The fix is to pre-shard each weight locally on CPU via `DTensor.from_local` *before* calling `parallelize_module`; the -latter then sees already-placed DTensors, skips the collective for weights, but still registers the required +latter then sees already-placed DTensors and skips the collective for weights, while still registering the required input/output hooks for the forward pass. + +Only needed for a model that is already in memory. `from_pretrained` with a tensor-parallel `parallel_config` streams +each rank's slice straight off disk into its DTensor, which issues no weight collectives at all and so cannot hit the +bug in the first place. """ import torch -import torch.distributed as dist import torch.nn as nn - -def _neuron_styles(relative_plan: dict) -> dict: - """Map a `{relative_path: style}` plan to no-op-partition styles for Neuron. - - Weights (and biases) are pre-sharded in `_pre_shard_and_tp`, so `parallelize_module` runs only to register the - forward hooks; `_partition_linear_fn` must not re-partition. Packed and plain styles share hook behavior, so both - collapse onto the two no-op styles. - """ - from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel - - from .tensor_parallel import PackedColwiseParallel, PackedRowwiseParallel - - class _NeuronColwise(ColwiseParallel): - def _partition_linear_fn(self, name, module, device_mesh): - pass # weight already Shard(0) via DTensor.from_local; parallelize_module runs only for the hooks - - class _NeuronRowwise(RowwiseParallel): - def _partition_linear_fn(self, name, module, device_mesh): - pass # weight already Shard(1) via DTensor.from_local; parallelize_module runs only for the hooks - - resolved = {} - for path, style in relative_plan.items(): - if style == "colwise" or isinstance(style, PackedColwiseParallel): - resolved[path] = _NeuronColwise() - elif style == "rowwise" or isinstance(style, PackedRowwiseParallel): - resolved[path] = _NeuronRowwise() - else: - raise ValueError( - f"Unsupported tensor-parallel style '{style}' for '{path}'. " - f"Expected 'colwise', 'rowwise', PackedColwiseParallel, or PackedRowwiseParallel." - ) - return resolved - - -def _pre_shard_and_tp( - module: nn.Module, - tp_mesh: "torch.distributed.device_mesh.DeviceMesh", - original_plan: dict, - rank: int, - tp_size: int, -) -> None: - """Pre-shard Linear weights via `DTensor.from_local`, then call `parallelize_module`. - - Workaround for a Neuron NRT bug where consecutive `reduce_scatter` calls for large weight tensors (≥ 5120×5120) - fail when all layers are distributed in a single `parallelize_module` call. Pre-sharding each weight on CPU means - it is already an on-device DTensor when `parallelize_module` runs (via `_neuron_styles`), so the collective is - skipped while the forward hooks are still registered. - """ - from torch.distributed.tensor import DTensor, Replicate, Shard - from torch.distributed.tensor.parallel import parallelize_module - - from .tensor_parallel import PackedColwiseParallel, PackedRowwiseParallel, _blocks_to_block_sizes - - device = torch.neuron.current_device() - - for path, orig_style in original_plan.items(): - # Resolve nested attribute path (e.g. "attn.to_q" or "attn.to_out.0") - submod = module - for part in path.split("."): - submod = getattr(submod, part) - - if not hasattr(submod, "weight"): - raise ValueError(f"`_tp_plan` entry '{path}' does not resolve to a module with a `weight` parameter.") - - w = submod.weight.data # CPU at this point - b = submod.bias.data if submod.bias is not None else None - if isinstance(orig_style, PackedColwiseParallel): - blocks = orig_style.blocks if orig_style.blocks is not None else getattr(submod, "_tp_packed_col_blocks") - block_sizes = _blocks_to_block_sizes(w.shape[0], blocks) - parts, bias_parts, offset = [], [], 0 - for bs in block_sizes: - if bs % tp_size != 0: - raise ValueError( - f"Cannot shard packed block of size {bs} across {tp_size} tensor-parallel ranks: " - f"{bs} is not divisible by {tp_size}." - ) - chunk = bs // tp_size - sl = slice(offset + rank * chunk, offset + (rank + 1) * chunk) - parts.append(w[sl, :].contiguous()) - if b is not None: - bias_parts.append(b[sl].contiguous()) - offset += bs - shard = torch.cat(parts, dim=0).to(device) - submod.weight = nn.Parameter(DTensor.from_local(shard, tp_mesh, [Shard(0)])) - if b is not None: - bias_shard = torch.cat(bias_parts, dim=0).to(device) - submod.bias = nn.Parameter(DTensor.from_local(bias_shard, tp_mesh, [Shard(0)])) - elif isinstance(orig_style, PackedRowwiseParallel): - blocks = orig_style.blocks if orig_style.blocks is not None else getattr(submod, "_tp_packed_row_blocks") - block_sizes = _blocks_to_block_sizes(w.shape[1], blocks) - parts, offset = [], 0 - for bs in block_sizes: - if bs % tp_size != 0: - raise ValueError( - f"Cannot shard packed block of size {bs} across {tp_size} tensor-parallel ranks: " - f"{bs} is not divisible by {tp_size}." - ) - chunk = bs // tp_size - parts.append(w[:, offset + rank * chunk : offset + (rank + 1) * chunk].contiguous()) - offset += bs - shard = torch.cat(parts, dim=1).to(device) - submod.weight = nn.Parameter(DTensor.from_local(shard, tp_mesh, [Shard(1)])) - if b is not None: # rowwise bias is added post-reduction → keep it replicated - submod.bias = nn.Parameter(DTensor.from_local(b.to(device), tp_mesh, [Replicate()])) - elif orig_style == "colwise": - if w.shape[0] % tp_size != 0: - raise ValueError( - f"Cannot colwise-shard '{path}' weight rows ({w.shape[0]}) across {tp_size} " - f"tensor-parallel ranks: not divisible by {tp_size}." - ) - rows = w.shape[0] // tp_size - sl = slice(rank * rows, (rank + 1) * rows) - submod.weight = nn.Parameter(DTensor.from_local(w[sl, :].contiguous().to(device), tp_mesh, [Shard(0)])) - if b is not None: - submod.bias = nn.Parameter(DTensor.from_local(b[sl].contiguous().to(device), tp_mesh, [Shard(0)])) - elif orig_style == "rowwise": - if w.shape[1] % tp_size != 0: - raise ValueError( - f"Cannot rowwise-shard '{path}' weight columns ({w.shape[1]}) across {tp_size} " - f"tensor-parallel ranks: not divisible by {tp_size}." - ) - cols = w.shape[1] // tp_size - shard = w[:, rank * cols : (rank + 1) * cols].contiguous().to(device) - submod.weight = nn.Parameter(DTensor.from_local(shard, tp_mesh, [Shard(1)])) - if b is not None: # rowwise bias is added post-reduction → keep it replicated - submod.bias = nn.Parameter(DTensor.from_local(b.to(device), tp_mesh, [Replicate()])) - - # parallelize_module is now a no-op for weight distribution (already DTensors) - # but still registers the input/output hooks required for the forward pass. - parallelize_module(module, tp_mesh, _neuron_styles(original_plan)) +from .tensor_parallel import TPShardSpec, _pre_shard_and_parallelize def _apply_tp_neuron( model: nn.Module, tp_mesh: "torch.distributed.device_mesh.DeviceMesh", groups: list, + specs: "dict[str, TPShardSpec]", ) -> None: - """Apply tensor parallelism on Neuron from resolved `_tp_plan` groups. + """Pre-shard the planned parameters via `DTensor.from_local`, then register the forward hooks. - `groups` is produced by `diffusers.hooks.tensor_parallel._resolve_tp_plan` — the same source of truth used by the - generic path, so the two backends shard identical layers. For each `(block, relative_plan)` group this pre-shards - the weights via `DTensor.from_local` (Neuron NRT consecutive-reduce-scatter workaround), then calls - `parallelize_module` to register the forward hooks. + `groups` and `specs` both come from the model's `_tp_plan` via `diffusers.hooks.tensor_parallel._resolve_tp_plan` / + `resolve_tp_shard_specs`, the same source of truth the generic path uses, so the two backends shard identical + layers. Model weights must be on CPU when this is called. """ - rank = dist.get_rank() - tp_size = tp_mesh.size() - - for block, relative_plan in groups: - _pre_shard_and_tp(block, tp_mesh, relative_plan, rank, tp_size) + _pre_shard_and_parallelize(model, tp_mesh, groups, specs, torch.neuron.current_device()) diff --git a/src/diffusers/loaders/lora_pipeline.py b/src/diffusers/loaders/lora_pipeline.py index 0809066b3dc8..1003aa57c420 100644 --- a/src/diffusers/loaders/lora_pipeline.py +++ b/src/diffusers/loaders/lora_pipeline.py @@ -30,6 +30,7 @@ logging, require_peft_backend, ) +from ..utils.torch_utils import get_device from .lora_base import ( # noqa LORA_WEIGHT_NAME, LORA_WEIGHT_NAME_SAFE, @@ -109,7 +110,7 @@ def _maybe_dequantize_weight_for_expanded_lora(model, module): if module.weight.device.type == "cpu": weight_on_cpu = True - device = torch.accelerator.current_accelerator().type if hasattr(torch, "accelerator") else "cuda" + device = get_device() if is_bnb_4bit_quantized or is_bnb_8bit_quantized: module_weight = dequantize_bnb_weight( module.weight.to(device) if weight_on_cpu else module.weight, diff --git a/src/diffusers/loaders/peft.py b/src/diffusers/loaders/peft.py index b0494207f48e..5ab34f4f0a3d 100644 --- a/src/diffusers/loaders/peft.py +++ b/src/diffusers/loaders/peft.py @@ -153,6 +153,16 @@ def load_lora_adapter( from peft.tuners.tuners_utils import BaseTunerLayer from ..hooks.group_offloading import _maybe_remove_and_reapply_group_offloading + from ..hooks.tensor_parallel import _raise_if_tensor_parallel + + # `_tp_plan` covers the base `Linear` layers only, so the injected adapter weights would stay unsharded + # and the sharded base layer would be added to a full-sized adapter output. + _raise_if_tensor_parallel( + self, + "have a LoRA adapter loaded", + "The adapter layers are not covered by the model's `_tp_plan`. Load the adapter before sharding the " + "model.", + ) cache_dir = kwargs.pop("cache_dir", None) force_download = kwargs.pop("force_download", False) diff --git a/src/diffusers/loaders/single_file_model.py b/src/diffusers/loaders/single_file_model.py index cd49ddef69f2..1556227673f4 100644 --- a/src/diffusers/loaders/single_file_model.py +++ b/src/diffusers/loaders/single_file_model.py @@ -42,6 +42,7 @@ convert_flux_transformer_checkpoint_to_diffusers, convert_hidream_transformer_to_diffusers, convert_hunyuan_video_transformer_to_diffusers, + convert_krea2_transformer_checkpoint_to_diffusers, convert_ldm_unet_checkpoint, convert_ldm_vae_checkpoint, convert_ltx2_audio_vae_to_diffusers, @@ -50,6 +51,7 @@ convert_ltx_transformer_checkpoint_to_diffusers, convert_ltx_vae_checkpoint_to_diffusers, convert_lumina2_to_diffusers, + convert_minimax_h3_transformer_checkpoint_to_diffusers, convert_mochi_transformer_checkpoint_to_diffusers, convert_qwen_image21_transformer_checkpoint_to_diffusers, convert_sana_transformer_to_diffusers, @@ -198,6 +200,10 @@ "checkpoint_mapping_fn": convert_qwen_image21_transformer_checkpoint_to_diffusers, "default_subfolder": "transformer", }, + "Krea2Transformer2DModel": { + "checkpoint_mapping_fn": convert_krea2_transformer_checkpoint_to_diffusers, + "default_subfolder": "transformer", + }, "Flux2Transformer2DModel": { "checkpoint_mapping_fn": convert_flux2_transformer_checkpoint_to_diffusers, "default_subfolder": "transformer", @@ -225,6 +231,10 @@ "checkpoint_mapping_fn": lambda checkpoint, **kwargs: checkpoint, "default_subfolder": "transformer", }, + "MiniMaxH3Transformer3DModel": { + "checkpoint_mapping_fn": convert_minimax_h3_transformer_checkpoint_to_diffusers, + "default_subfolder": "transformer", + }, } diff --git a/src/diffusers/loaders/single_file_utils.py b/src/diffusers/loaders/single_file_utils.py index faa6317ea9b4..3e160be02b4b 100644 --- a/src/diffusers/loaders/single_file_utils.py +++ b/src/diffusers/loaders/single_file_utils.py @@ -127,6 +127,7 @@ ], "z-image-turbo-controlnet": "control_all_x_embedder.2-1.weight", "z-image-turbo-controlnet-2.x": "control_layers.14.adaLN_modulation.0.weight", + "minimax-h3": "token_refiner.blocks.0.attn.qkv_proj.weight", "sana": [ "blocks.0.cross_attn.q_linear.weight", "blocks.0.cross_attn.q_linear.bias", @@ -159,6 +160,7 @@ "audio_vae.per_channel_statistics.mean-of-means", ], "qwen-image-2.1": ["model.diffusion_model.txt_in.text_norm.weight", "txt_in.text_norm.weight"], + "krea2": ["model.diffusion_model.txtfusion.projector.weight", "txtfusion.projector.weight"], } DIFFUSERS_DEFAULT_PIPELINE_PATHS = { @@ -244,6 +246,8 @@ "z-image-turbo-controlnet-2.1": {"pretrained_model_name_or_path": "hlky/Z-Image-Turbo-Fun-Controlnet-Union-2.1"}, "ltx2-dev": {"pretrained_model_name_or_path": "Lightricks/LTX-2"}, "qwen-image-2.1": {"pretrained_model_name_or_path": "Qwen/Qwen-Image-2.1"}, + "krea2": {"pretrained_model_name_or_path": "krea/Krea-2-Raw"}, + "minimax-h3": {"pretrained_model_name_or_path": "MiniMaxAI/MiniMax-H3"}, } # Use to configure model sample size when original config is provided @@ -789,6 +793,9 @@ def infer_diffusers_model_type(checkpoint): elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["qwen-image-2.1"]): model_type = "qwen-image-2.1" + elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["krea2"]): + model_type = "krea2" + elif CHECKPOINT_KEY_NAMES["wan_vae"] in checkpoint: # All Wan models use the same VAE so we can use the same default model repo to fetch the config model_type = "wan-t2v-14B" @@ -829,6 +836,9 @@ def infer_diffusers_model_type(checkpoint): elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["ltx2"]): model_type = "ltx2-dev" + elif CHECKPOINT_KEY_NAMES["minimax-h3"] in checkpoint: + model_type = "minimax-h3" + else: model_type = "v1" @@ -4229,6 +4239,91 @@ def convert_ernie_image_transformer_checkpoint_to_diffusers(checkpoint, **kwargs return checkpoint +def convert_minimax_h3_transformer_checkpoint_to_diffusers(checkpoint, config, qkv_layout="stacked", **kwargs): + if qkv_layout not in ("stacked", "interleaved"): + raise ValueError( + f'`qkv_layout` must be "stacked" or "interleaved", got {qkv_layout!r}. "stacked" is the `[q; k; v]` row ' + "order of every published single-file MiniMax-H3 checkpoint (Comfy-Org/MiniMax-H3 and the GGUFs derived " + 'from it). "interleaved" is the per-head `[q k v]` row order of the MiniMaxAI/MiniMax-H3 shards, for a ' + "file merged from those shards by hand. The two cannot be told apart from the checkpoint itself." + ) + if "adaln_t_table" in checkpoint: + raise ValueError( + "This is a pruned MiniMax-H3 checkpoint: it replaces `time_embedder` with `adaln_t_table` of shape " + f"{tuple(checkpoint['adaln_t_table'].shape)}, which `MiniMaxH3Transformer3DModel` does not support. " + "Use the unpruned checkpoint from https://huggingface.co/MiniMaxAI/MiniMax-H3." + ) + + MINIMAX_H3_KEYS_RENAME_DICT = { + "token_refiner.blocks.": "token_refiner.refiner_blocks.", + "time_embedder.proj_in.": "time_embedder.linear_1.", + "time_embedder.proj_out.": "time_embedder.linear_2.", + "video_patch_proj.": "proj_in.", + "audio_patch_proj.": "audio_proj_in.", + "condition_proj.": "context_embedder.", + "final_layer.norm.": "norm_out.norm.", + "final_layer.adaln_proj.linear.": "norm_out.linear.", + "final_layer.video_out.": "proj_out.", + "final_layer.audio_out.": "audio_proj_out.", + ".attn.q_norm.": ".attn.norm_q.", + ".attn.k_norm.": ".attn.norm_k.", + ".attn.out_proj.": ".attn.to_out.0.", + ".mlp.fc1.": ".ff.net.0.proj.", + ".mlp.fc2.": ".ff.net.2.", + } + + def convert_minimax_h3_fused_attention(key: str, state_dict: dict[str, object]) -> None: + # Published single files store the reference model's post-load `[q; k; v]` stack; the MiniMaxAI/MiniMax-H3 + # shards interleave the rows per head, `[head0: q k v, head1: q k v, ...]`. Keys and shapes are identical, so + # the caller has to say which one it is. + fused_qkv_weight = state_dict.pop(key) + if qkv_layout == "interleaved": + fused_qkv_weight = fused_qkv_weight.unflatten( + 0, (config["num_attention_heads"], 3, config["attention_head_dim"]) + ) + to_q_weight, to_k_weight, to_v_weight = [weight.flatten(0, 1) for weight in fused_qkv_weight.unbind(dim=1)] + else: + to_q_weight, to_k_weight, to_v_weight = torch.chunk(fused_qkv_weight, 3, dim=0) + state_dict[key.replace(".attn.qkv_proj.weight", ".attn.to_q.weight")] = to_q_weight + state_dict[key.replace(".attn.qkv_proj.weight", ".attn.to_k.weight")] = to_k_weight + state_dict[key.replace(".attn.qkv_proj.weight", ".attn.to_v.weight")] = to_v_weight + + def convert_minimax_h3_gated_ff(key: str, state_dict: dict[str, object]) -> None: + # The checkpoint fuses `[gate; value]`, `SwiGLU` reads `[value; gate]`. + gate, value = torch.chunk(state_dict[key], 2, dim=0) + state_dict[key] = torch.cat([value, gate], dim=0) + + TRANSFORMER_SPECIAL_KEYS_REMAP = { + ".attn.qkv_proj.weight": convert_minimax_h3_fused_attention, + ".ff.net.0.proj.weight": convert_minimax_h3_gated_ff, + } + + def update_state_dict(state_dict: dict[str, object], old_key: str, new_key: str) -> None: + state_dict[new_key] = state_dict.pop(old_key) + + converted_state_dict = {key: checkpoint.pop(key) for key in list(checkpoint.keys())} + + # `MiniMaxH3RotaryPosEmbed` recomputes this buffer from the config. + converted_state_dict.pop("rope.inv_freq", None) + + for key in list(converted_state_dict.keys()): + new_key = key[:] + if new_key.startswith("blocks."): + new_key = new_key.replace("blocks.", "transformer_blocks.", 1) + for replace_key, rename_key in MINIMAX_H3_KEYS_RENAME_DICT.items(): + new_key = new_key.replace(replace_key, rename_key) + + update_state_dict(converted_state_dict, key, new_key) + + for key in list(converted_state_dict.keys()): + for special_key, handler_fn_inplace in TRANSFORMER_SPECIAL_KEYS_REMAP.items(): + if special_key not in key: + continue + handler_fn_inplace(key, converted_state_dict) + + return converted_state_dict + + def convert_qwen_image21_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): converted_state_dict = {} @@ -4245,3 +4340,51 @@ def convert_qwen_image21_transformer_checkpoint_to_diffusers(checkpoint, **kwarg converted_state_dict[new_key] = value return converted_state_dict + + +def convert_krea2_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): + prefix_rename_dict = { + "first.": "img_in.", + "tmlp.0.": "time_embed.linear_1.", + "tmlp.2.": "time_embed.linear_2.", + "tproj.1.": "time_mod_proj.", + "txtmlp.0.scale": "txt_in.norm.weight", + "txtmlp.1.": "txt_in.linear_1.", + "txtmlp.3.": "txt_in.linear_2.", + "txtfusion.": "text_fusion.", + "blocks.": "transformer_blocks.", + "last.linear.": "final_layer.linear.", + "last.norm.scale": "final_layer.norm.weight", + "last.modulation.lin": "final_layer.scale_shift_table", + } + block_rename_dict = { + ".attn.wq.": ".attn.to_q.", + ".attn.wk.": ".attn.to_k.", + ".attn.wv.": ".attn.to_v.", + ".attn.wo.": ".attn.to_out.0.", + ".attn.gate.": ".attn.to_gate.", + ".attn.qknorm.qnorm.scale": ".attn.norm_q.weight", + ".attn.qknorm.knorm.scale": ".attn.norm_k.weight", + ".mlp.": ".ff.", + ".prenorm.scale": ".norm1.weight", + ".postnorm.scale": ".norm2.weight", + ".mod.lin": ".scale_shift_table", + } + + converted_state_dict = {} + for key in list(checkpoint.keys()): + new_key = key.replace("model.diffusion_model.", "") + for old, new in prefix_rename_dict.items(): + if new_key.startswith(old): + new_key = new + new_key[len(old) :] + break + for old, new in block_rename_dict.items(): + new_key = new_key.replace(old, new) + + value = checkpoint.pop(key) + # The original checkpoint stores each block's six modulation vectors flattened into one. + if new_key.startswith("transformer_blocks.") and new_key.endswith(".scale_shift_table"): + value = value.reshape(6, -1) + converted_state_dict[new_key] = value + + return converted_state_dict diff --git a/src/diffusers/models/_modeling_parallel.py b/src/diffusers/models/_modeling_parallel.py index 86627284e078..0f1da9f3b353 100644 --- a/src/diffusers/models/_modeling_parallel.py +++ b/src/diffusers/models/_modeling_parallel.py @@ -161,7 +161,7 @@ class TensorParallelConfig: Tensor parallelism shards weight matrices (column-wise and row-wise) across devices. Each device computes a partial result; an AllReduce/AllGather at layer boundaries reconstructs the full output. Uses `torch.distributed.tensor.parallelize_module` with `ColwiseParallel` / `RowwiseParallel` sharding styles. Supported - device types are `"cuda"` and `"neuron"`. + device types are `"cuda"`, `"neuron"` and `"tpu"`. Args: tp_degree (`int`, defaults to `1`): @@ -186,6 +186,8 @@ def __post_init__(self): raise ValueError("`tp_degree` must be >= 1.") def setup(self, rank: int, world_size: int, device: torch.device, mesh: torch.distributed.device_mesh.DeviceMesh): + if mesh.size() > world_size: + raise ValueError(f"Tensor parallel degree ({mesh.size()}) cannot exceed the world size ({world_size}).") self._rank = rank self._world_size = world_size self._device = device @@ -198,6 +200,11 @@ class ParallelConfig: """ Configuration for applying different parallelisms. + Both may be set at once. The two are then applied over one device mesh with a dimension each — `("ring", "ulysses", + "tp")`, built by `enable_parallelism` — so their collectives stay in separate process groups. This is what lets a + model too large for one device (TP shards the weights) also run a sequence too long for one device's attention + scratchpad (CP splits the sequence): TP alone cannot divide the sequence, and CP alone cannot divide the weights. + Args: context_parallel_config (`ContextParallelConfig`, *optional*): Configuration for context parallelism. @@ -214,12 +221,17 @@ class ParallelConfig: _mesh: torch.distributed.device_mesh.DeviceMesh = None def __post_init__(self): - if self.context_parallel_config is not None and self.tensor_parallel_config is not None: + if self.context_parallel_config is None and self.tensor_parallel_config is None: raise ValueError( - "Combining context parallelism and tensor parallelism in a single `ParallelConfig` is not supported. " - "Please specify only one of `context_parallel_config` or `tensor_parallel_config`." + "A `ParallelConfig` must specify at least one of `context_parallel_config` or " + "`tensor_parallel_config`." ) + @property + def _is_combined(self) -> bool: + """Whether both context and tensor parallelism are requested, i.e. they must share one mesh.""" + return self.context_parallel_config is not None and self.tensor_parallel_config is not None + def setup( self, rank: int, @@ -232,10 +244,32 @@ def setup( self._world_size = world_size self._device = device self._mesh = mesh + + # Context and tensor parallelism compose because they cut the model along different axes: CP splits the + # sequence (and, under Ulysses, trades sequence for heads inside attention) while TP shards the Linear + # weights. Composing them means giving each its own mesh *dimension*, so that every collective one issues + # stays inside its own process group: TP's all-reduce must not reach a rank holding a different sequence + # chunk, and CP's sequence all-gather must not reach a rank holding a different weight shard. + cp_mesh = tp_mesh = mesh + if self._is_combined: + dim_names = mesh.mesh_dim_names if mesh is not None else None + missing = [d for d in ("ring", "ulysses", "tp") if dim_names is None or d not in dim_names] + if missing: + raise ValueError( + f"Combining context and tensor parallelism requires a device mesh with 'ring', 'ulysses' and " + f"'tp' dimensions, but {missing} {'is' if len(missing) == 1 else 'are'} missing (got " + f"{dim_names}). `enable_parallelism` builds this mesh from `ring_degree`/`ulysses_degree`/" + f"`tp_degree`; if you build it yourself, set `mesh=` on one of the two configs." + ) + # `ContextParallelConfig.setup` slices ("ring", "ulysses") off whatever mesh it is handed, so it takes + # the full mesh. TP instead needs its own 1-D submesh: `parallelize_module` shards a weight over every + # rank of the mesh it is given, so handing it the 3-D mesh would shard over the CP ranks as well. + tp_mesh = mesh["tp"] + if self.context_parallel_config is not None: - self.context_parallel_config.setup(rank, world_size, device, mesh) + self.context_parallel_config.setup(rank, world_size, device, cp_mesh) if self.tensor_parallel_config is not None: - self.tensor_parallel_config.setup(rank, world_size, device, mesh) + self.tensor_parallel_config.setup(rank, world_size, device, tp_mesh) @dataclass(frozen=True) 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..ac6af47ab493 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py @@ -31,6 +31,7 @@ CACHE_T = 2 +# Copied from diffusers.models.autoencoders.autoencoder_kl_wan.AvgDown3D with AvgDown3D->QwenImage21AvgDown3D class QwenImage21AvgDown3D(nn.Module): def __init__( self, @@ -46,7 +47,6 @@ def __init__( f"`in_channels` ({in_channels}) times the downsampling factor ({factor}) must be divisible by " f"`out_channels` ({out_channels})." ) - self.in_channels = in_channels self.out_channels = out_channels self.factor_t = factor_t @@ -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..dcb280ff21a2 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_wan.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_wan.py @@ -40,14 +40,18 @@ def __init__( factor_s=1, ): super().__init__() + factor = factor_t * factor_s * factor_s + if in_channels * factor % out_channels != 0: + raise ValueError( + f"`in_channels` ({in_channels}) times the downsampling factor ({factor}) must be divisible by " + f"`out_channels` ({out_channels})." + ) self.in_channels = in_channels self.out_channels = out_channels self.factor_t = factor_t self.factor_s = factor_s - self.factor = self.factor_t * self.factor_s * self.factor_s - - assert in_channels * self.factor % out_channels == 0 - self.group_size = in_channels * self.factor // out_channels + self.factor = factor + self.group_size = in_channels * factor // out_channels def forward(self, x: torch.Tensor) -> torch.Tensor: pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t @@ -163,7 +167,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 +199,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 +221,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 +270,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 +347,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 +407,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 +460,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 +500,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 +591,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 +688,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 +763,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 +880,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/model_loading_utils.py b/src/diffusers/models/model_loading_utils.py index be2e01e4ce18..b6d9d7997d81 100644 --- a/src/diffusers/models/model_loading_utils.py +++ b/src/diffusers/models/model_loading_utils.py @@ -381,6 +381,99 @@ def _load_shard_file( return offload_index, state_dict_index, mismatched_keys, error_msgs +def _load_shard_file_tp( + shard_file, + model, + model_state_dict, + tp_shard_specs, + tp_config, + dtype=None, + keep_in_fp32_modules=None, + unexpected_keys=None, + ignore_mismatched_sizes=False, +): + """Load one safetensors shard, reading only this rank's slice of each tensor-parallel parameter. + + The counterpart of `_load_shard_file` for a model being sharded by `_tp_plan`, with the same return contract so it + can be swapped in as `load_fn`. Parameters covered by `tp_shard_specs` are sliced while still on disk and placed as + `DTensor`s; everything else is read whole and replicated on every rank, exactly as tensor parallelism requires. + + Slicing before the dtype cast is the point of the whole exercise: `load_model_dict_into_meta` casts the full tensor + first, which would materialize it in full on every rank. + """ + from safetensors import safe_open + from torch.distributed.tensor import DTensor, Replicate, Shard + + from ..hooks.tensor_parallel import _local_shard + + tp_mesh = tp_config._mesh + # `TensorParallelConfig._device` is derived from the default accelerator, which is not meaningful on + # Neuron; resolve it the way the Neuron pre-shard backend does. + if tp_mesh.device_type == "neuron": + device = torch.neuron.current_device() + else: + device = tp_config._device + + mismatched_keys = [] + + # The slices are lazy views over the file, so every read has to happen inside this block. + with safe_open(shard_file, framework="pt", device="cpu") as f: + for key in f.keys(): + if key not in model_state_dict: + unexpected_keys.append(key) + continue + + checkpoint_slice = f.get_slice(key) + expected_shape = model_state_dict[key].shape + if tuple(checkpoint_slice.get_shape()) != tuple(expected_shape): + # Checkpoints always hold full tensors, so the comparison is against the unsharded shape. + if not ignore_mismatched_sizes: + raise ValueError( + f"Cannot load {key} because it has shape {tuple(checkpoint_slice.get_shape())} in the " + f"checkpoint but shape {tuple(expected_shape)} in {model.__class__.__name__}. Pass " + "`ignore_mismatched_sizes=True` to skip it and keep the randomly initialized weight." + ) + mismatched_keys.append((key, tuple(checkpoint_slice.get_shape()), tuple(expected_shape))) + continue + + spec = tp_shard_specs.get(key) + if spec is None or spec.dim is None: + param = checkpoint_slice[...] + else: + param = _local_shard(checkpoint_slice, spec.dim, spec.block_sizes, tp_mesh) + + # Mirror `load_model_dict_into_meta`: only floating point weights are cast, and modules held + # in fp32 override the requested dtype. + if dtype is not None and torch.is_floating_point(param): + if keep_in_fp32_modules is not None and any( + module_to_keep_in_fp32 in key.split(".") for module_to_keep_in_fp32 in keep_in_fp32_modules + ): + param = param.to(torch.float32) + else: + param = param.to(dtype) + + if spec is None: + set_module_tensor_to_device(model, key, device, value=param) + continue + + path, _, param_name = key.rpartition(".") + module = model.get_submodule(path) + # A rowwise bias is added after the all-reduce, so it stays replicated. It still has to be a + # DTensor: a plain tensor next to a sharded weight fails the `addmm` dispatch. + placement = Replicate() if spec.dim is None else Shard(spec.dim) + module.register_parameter( + param_name, + torch.nn.Parameter( + DTensor.from_local(param.to(device), tp_mesh, [placement], run_check=False), + requires_grad=getattr(module, param_name).requires_grad, + ), + ) + + # `offload_index` / `state_dict_index` are always None here: offloading and tensor parallelism are + # rejected as a combination by `from_pretrained`. + return None, None, mismatched_keys, [] + + def _load_shard_files_with_threadpool( shard_files, model, @@ -443,28 +536,6 @@ def _load_shard_files_with_threadpool( return offload_index, state_dict_index, mismatched_keys, error_msgs -def _find_mismatched_keys( - state_dict, - model_state_dict, - loaded_keys, - ignore_mismatched_sizes, -): - mismatched_keys = [] - if ignore_mismatched_sizes: - for checkpoint_key in loaded_keys: - model_key = checkpoint_key - # If the checkpoint is sharded, we may not have the key here. - if checkpoint_key not in state_dict: - continue - - if model_key in model_state_dict and state_dict[checkpoint_key].shape != model_state_dict[model_key].shape: - mismatched_keys.append( - (checkpoint_key, state_dict[checkpoint_key].shape, model_state_dict[model_key].shape) - ) - del state_dict[checkpoint_key] - return mismatched_keys - - def _load_state_dict_into_model( model_to_load, state_dict: OrderedDict, assign_to_params_buffers: bool = False ) -> list[str]: diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index 425f2f29235e..262e5a8ae699 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -77,6 +77,7 @@ _fetch_index_file, _fetch_index_file_legacy, _load_shard_file, + _load_shard_file_tp, _load_shard_files_with_threadpool, load_state_dict, ) @@ -574,6 +575,14 @@ def enable_group_offload( "2. Or, run a forward pass with tiling disabled (can still use small dummy inputs)." ) logger.warning(msg) + from ..hooks.tensor_parallel import _raise_if_tensor_parallel + + _raise_if_tensor_parallel( + self, + "be group-offloaded", + "Both decide where a parameter lives, and tensor parallelism already keeps only one shard of each weight " + "per rank, so offloading is not needed on top of it.", + ) if not self._supports_group_offloading: raise ValueError( f"{self.__class__.__name__} does not support group offloading. Please make sure to set the boolean attribute " @@ -726,6 +735,15 @@ def save_pretrained( logger.error(f"Provided path ({save_directory}) should be a directory, not a file") return + from ..hooks.tensor_parallel import _raise_if_tensor_parallel + + _raise_if_tensor_parallel( + self, + "be saved yet", + "Its parameters are sharded across ranks. Save the model before sharding it instead.", + error_cls=NotImplementedError, + ) + hf_quantizer = getattr(self, "hf_quantizer", None) if hf_quantizer is not None: quantization_serializable = ( @@ -1037,7 +1055,9 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None use_safetensors = kwargs.pop("use_safetensors", None) quantization_config = kwargs.pop("quantization_config", None) disable_mmap = kwargs.pop("disable_mmap", False) - parallel_config: ParallelConfig | ContextParallelConfig | None = kwargs.pop("parallel_config", None) + parallel_config: ParallelConfig | ContextParallelConfig | TensorParallelConfig | None = kwargs.pop( + "parallel_config", None + ) use_flashpack = kwargs.pop("use_flashpack", False) flashpack_kwargs = kwargs.pop("flashpack_kwargs", {}) @@ -1212,6 +1232,35 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None else: keep_in_fp32_modules = [] + # A tensor-parallel `parallel_config` makes `from_pretrained` shard while it reads, so each rank only + # ever materializes its own slice. Validate the combination before any file is fetched. + tp_config = None + if parallel_config is not None: + tp_config = ( + parallel_config + if isinstance(parallel_config, TensorParallelConfig) + else parallel_config.tensor_parallel_config + ) + if tp_config is not None and tp_config.tp_degree == 1 and tp_config.mesh is None: + # Nothing to shard, so take the ordinary loader rather than building 1-rank DTensors. + tp_config = None + if tp_config is not None: + from ..hooks.tensor_parallel import _check_tp_supported + + # Before the checkpoint files are resolved, so that e.g. `use_flashpack` fails with the real reason + # instead of a missing-file error. The weights-format check lives where the resolved file list is known. + _check_tp_supported( + cls.__name__, + cls._tp_plan, + config.get("num_attention_heads"), + tp_config, + device_map=device_map, + low_cpu_mem_usage=low_cpu_mem_usage, + use_flashpack=use_flashpack, + use_safetensors=use_safetensors, + hf_quantizer=hf_quantizer, + ) + is_sharded = False resolved_model_file = None @@ -1305,6 +1354,12 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None if not isinstance(resolved_model_file, list): resolved_model_file = [resolved_model_file] + if tp_config is not None: + from ..hooks.tensor_parallel import _check_tp_weights_format + + # As soon as the resolved files are known, before the model is built. + _check_tp_weights_format(resolved_model_file) + # set dtype to instantiate the model under: # 1. If torch_dtype is not None, we use that dtype # 2. If torch_dtype is float8, we don't use _set_default_torch_dtype and we downcast after loading the model @@ -1324,6 +1379,22 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None with ContextManagers(init_contexts): model = cls.from_config(config, **unused_kwargs) + # Resolve the tensor-parallel mesh before any weights are read, so each rank can stream only its own + # slice of every planned parameter straight into a DTensor instead of materializing the full + # checkpoint and resharding it afterwards. + tp_shard_specs = None + if tp_config is not None: + from ..hooks.tensor_parallel import resolve_tp_shard_specs + + parallel_config = model._resolve_parallel_config(parallel_config) + tp_config = parallel_config.tensor_parallel_config + tp_shard_specs = resolve_tp_shard_specs(model, cls._tp_plan, tp_config._mesh.size()) + # Each rank opens every shard file but only reads its own slices, so threading the files buys + # nothing and would have several threads calling `register_parameter` on the same modules. + if is_parallel_loading_enabled: + logger.debug("Disabling parallel loading: a tensor-parallel load reads the shard files sequentially.") + is_parallel_loading_enabled = False + if use_flashpack: if is_flashpack_available(): import flashpack @@ -1366,7 +1437,7 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None torch.set_default_dtype(dtype_orig) state_dict = None - if not is_sharded: + if not is_sharded and tp_shard_specs is None: # Time to load the checkpoint state_dict = load_state_dict(resolved_model_file[0], disable_mmap=disable_mmap) # We only fix it for non sharded checkpoints as we don't need it yet for sharded one. @@ -1374,6 +1445,13 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None if is_sharded: loaded_keys = sharded_metadata["all_checkpoint_keys"] + elif tp_shard_specs is not None: + # Read the key names out of the safetensors header without materializing any tensor, and leave + # `state_dict` as None so `_load_pretrained_model` keeps reading from the file itself. + from safetensors import safe_open + + with safe_open(resolved_model_file[0], framework="pt") as f: + loaded_keys = list(f.keys()) else: loaded_keys = list(state_dict.keys()) @@ -1421,6 +1499,8 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None keep_in_fp32_modules=keep_in_fp32_modules, is_parallel_loading_enabled=is_parallel_loading_enabled, disable_mmap=disable_mmap, + tp_shard_specs=tp_shard_specs, + tp_config=tp_config, ) loading_info = { "missing_keys": missing_keys, @@ -1461,7 +1541,27 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None model.eval() if parallel_config is not None: - model.enable_parallelism(config=parallel_config) + if tp_shard_specs is not None: + # The weights are already sharded, so this only registers the forward hooks. `_parallel_config` + # was recorded by `_resolve_parallel_config` before loading. + from torch.distributed.tensor import DTensor + + from ..hooks.tensor_parallel import apply_tensor_parallel + + # Non-persistent buffers are absent from both the state dict and the checkpoint, and + # `init_empty_weights` leaves them as real CPU tensors, so move them across explicitly. + tp_device = ( + torch.neuron.current_device() if tp_config._mesh.device_type == "neuron" else tp_config._device + ) + for name, buffer in model.named_buffers(): + if buffer.device != tp_device and not isinstance(buffer, DTensor): + module_path, _, buffer_name = name.rpartition(".") + module = model.get_submodule(module_path) if module_path else model + module._buffers[buffer_name] = buffer.to(tp_device) + + apply_tensor_parallel(model, tp_config, cls._tp_plan, weights_already_sharded=True) + else: + model.enable_parallelism(config=parallel_config) if output_loading_info: return model, loading_info @@ -1607,6 +1707,93 @@ def compile_repeated_blocks(self, *args, **kwargs): f"Regional compilation failed because {repeated_blocks} classes are not found in the model. " ) + def _resolve_parallel_config( + self, config: ParallelConfig | ContextParallelConfig | TensorParallelConfig + ) -> ParallelConfig: + """Normalize `config`, build its device mesh, and record it on the model. + + Split out of `enable_parallelism` because `from_pretrained` needs the mesh *before* it reads any weights, in + order to stream each rank's shard straight into place. Whichever of the two runs first builds the mesh exactly + once. + """ + if not torch.distributed.is_available() or not torch.distributed.is_initialized(): + raise RuntimeError( + "torch.distributed must be available and initialized before applying a `parallel_config`." + ) + + if isinstance(config, ContextParallelConfig): + config = ParallelConfig(context_parallel_config=config) + elif isinstance(config, TensorParallelConfig): + config = ParallelConfig(tensor_parallel_config=config) + + rank = torch.distributed.get_rank() + world_size = torch.distributed.get_world_size() + device_type = torch._C._get_accelerator().type + device_module = torch.get_device_module(device_type) + device = torch.device(device_type, rank % device_module.device_count()) + + mesh = None + cp_config = config.context_parallel_config + tp_config = config.tensor_parallel_config + if cp_config is not None and tp_config is not None: + # One mesh, one dimension per parallelism, so `ParallelConfig.setup` can hand each config its own + # submesh. "tp" goes last, which makes TP ranks adjacent: its all-reduce fires twice per block, more + # often than the CP collectives, so it is the one that wants the closest devices. The CP dimensions are + # then strided, which is fine for them — on Neuron only `all_to_all` (Ulysses) constrains its replica + # groups to be block-aligned; `all_gather` (ring) and `all_reduce` (TP) accept any grouping. + mesh = ( + cp_config.mesh + or tp_config.mesh + or torch.distributed.device_mesh.init_device_mesh( + device_type=device_type, + mesh_shape=(cp_config.ring_degree, cp_config.ulysses_degree, tp_config.tp_degree), + mesh_dim_names=("ring", "ulysses", "tp"), + ) + ) + elif cp_config is not None: + mesh = cp_config.mesh or torch.distributed.device_mesh.init_device_mesh( + device_type=device_type, + mesh_shape=cp_config.mesh_shape, + mesh_dim_names=cp_config.mesh_dim_names, + ) + elif tp_config is not None: + mesh = tp_config.mesh or torch.distributed.device_mesh.init_device_mesh( + device_type=device_type, + mesh_shape=(tp_config.tp_degree,), + mesh_dim_names=("tp",), + ) + + # `config.setup()` records the mesh resolved above onto the config; see `ParallelConfig.setup`. + config.setup(rank, world_size, device, mesh=mesh) + + # Validate the combination up front — after `setup`, which resolves `_tp_degree` from the mesh, but before + # anything is recorded on the model or applied to it. CP hooks are applied before TP, so a check left to the + # TP branch would raise on a model already carrying CP hooks: half-parallelised, and not usable. + if cp_config is not None and tp_config is not None: + tp_degree = tp_config._tp_degree + requested = cp_config.ring_degree * cp_config.ulysses_degree * tp_degree + if requested > world_size: + raise ValueError( + f"Combining context and tensor parallelism needs `ring_degree` ({cp_config.ring_degree}) * " + f"`ulysses_degree` ({cp_config.ulysses_degree}) * `tp_degree` ({tp_degree}) = {requested} " + f"devices, which exceeds the world size ({world_size})." + ) + num_heads = getattr(self.config, "num_attention_heads", None) + divisor = tp_degree * cp_config.ulysses_degree + if num_heads is not None and not cp_config.ulysses_anything and num_heads % divisor != 0: + # Ulysses trades sequence for heads inside attention (`SeqAllToAllDim` scatters over dim 2), and it + # only ever sees the heads TP left on this rank, so the count has to survive both splits. + raise ValueError( + f"Combining tensor parallelism (`tp_degree`={tp_degree}) with Ulysses context parallelism " + f"(`ulysses_degree`={cp_config.ulysses_degree}) requires the number of attention heads " + f"({num_heads}) to be divisible by their product ({divisor}): TP shards the heads first, and " + f"Ulysses splits what is left on each rank. Pass `ulysses_anything=True` to pad the head " + f"dimension instead, or pick degrees whose product divides {num_heads}." + ) + + self._parallel_config = config + return config + def enable_parallelism( self, *, @@ -1617,26 +1804,36 @@ def enable_parallelism( "`enable_parallelism` is an experimental feature. The API may change in the future and breaking changes may be introduced at any time without warning." ) - if not torch.distributed.is_available() and not torch.distributed.is_initialized(): - raise RuntimeError( - "torch.distributed must be available and initialized before calling `enable_parallelism`." - ) - from ..hooks.context_parallel import apply_context_parallel from .attention import AttentionModuleMixin from .attention_dispatch import AttentionBackendName, _AttentionBackendRegistry from .attention_processor import Attention, MochiAttention - if isinstance(config, ContextParallelConfig): - config = ParallelConfig(context_parallel_config=config) - elif isinstance(config, TensorParallelConfig): - config = ParallelConfig(tensor_parallel_config=config) + if self._parallel_config is not None: + raise RuntimeError( + f"Parallelism is already applied to this {self.__class__.__name__}. `enable_parallelism` cannot be " + "called twice, and it must not be called on a model loaded with `from_pretrained(..., " + "parallel_config=...)` — that already sharded the weights while reading the checkpoint." + ) - rank = torch.distributed.get_rank() - world_size = torch.distributed.get_world_size() - device_type = torch._C._get_accelerator().type - device_module = torch.get_device_module(device_type) - device = torch.device(device_type, rank % device_module.device_count()) + tp_config = ( + config if isinstance(config, TensorParallelConfig) else getattr(config, "tensor_parallel_config", None) + ) + if tp_config is not None: + from ..hooks.tensor_parallel import _check_tp_supported, _tp_degree, resolve_tp_shard_specs + + # Before `_resolve_parallel_config`, which records the config on the model: a model that fails these + # checks is left untouched. + _check_tp_supported( + self.__class__.__name__, + self._tp_plan, + getattr(self.config, "num_attention_heads", None), + tp_config, + model=self, + ) + resolve_tp_shard_specs(self, self._tp_plan, _tp_degree(tp_config)) + + config = self._resolve_parallel_config(config) attention_classes = (Attention, MochiAttention, AttentionModuleMixin) @@ -1668,26 +1865,6 @@ def enable_parallelism( # iterate over all modules after checking the first processor break - mesh = None - if config.context_parallel_config is not None: - cp_config = config.context_parallel_config - mesh = cp_config.mesh or torch.distributed.device_mesh.init_device_mesh( - device_type=device_type, - mesh_shape=cp_config.mesh_shape, - mesh_dim_names=cp_config.mesh_dim_names, - ) - elif config.tensor_parallel_config is not None: - tp_config = config.tensor_parallel_config - mesh = tp_config.mesh or torch.distributed.device_mesh.init_device_mesh( - device_type=device_type, - mesh_shape=(tp_config.tp_degree,), - mesh_dim_names=("tp",), - ) - - # `config.setup()` records the mesh resolved above onto the config; see `ParallelConfig.setup`. - config.setup(rank, world_size, device, mesh=mesh) - self._parallel_config = config - # Only context parallelism needs the config inside attention: it replaces the attention computation itself # (Ulysses all-to-all / ring). Tensor parallelism only shards `Linear` weights, so each rank runs the ordinary # attention op over its own heads and the processors must stay unaware of it. @@ -1708,16 +1885,6 @@ def enable_parallelism( apply_context_parallel(self, config.context_parallel_config, cp_plan) if config.tensor_parallel_config is not None: - if self._tp_plan is None: - raise ValueError( - "`_tp_plan` must be set on the model class to use tensor parallelism. " - f"'{self.__class__.__name__}' does not define one." - ) - tp_degree = config.tensor_parallel_config._tp_degree - num_heads = getattr(self.config, "num_attention_heads", None) - if num_heads is not None and num_heads % tp_degree != 0: - raise ValueError(f"`tp_degree` ({tp_degree}) must divide the number of attention heads ({num_heads}).") - from ..hooks.tensor_parallel import apply_tensor_parallel apply_tensor_parallel(self, config.tensor_parallel_config, self._tp_plan) @@ -1741,6 +1908,8 @@ def _load_pretrained_model( offload_folder: str | os.PathLike | None = None, is_parallel_loading_enabled: bool | None = False, disable_mmap: bool = False, + tp_shard_specs: dict | None = None, + tp_config: TensorParallelConfig | None = None, ): model_state_dict = model.state_dict() expected_keys = list(model_state_dict.keys()) @@ -1757,6 +1926,17 @@ def _load_pretrained_model( mismatched_keys = [] error_msgs = [] + if tp_shard_specs is not None: + # `_hooks_only_styles` lets `parallelize_module` broadcast any planned parameter it finds still + # plain, and for a key the checkpoint does not carry that broadcast would be issued on a `meta` + # tensor. + missing_planned_keys = sorted(set(tp_shard_specs) & set(missing_keys)) + if missing_planned_keys: + raise ValueError( + f"Cannot shard {cls.__name__} across tensor-parallel ranks because its `_tp_plan` covers " + f"parameters that the checkpoint does not contain: {missing_planned_keys}." + ) + # Deal with offload if device_map is not None and "disk" in device_map.values(): if offload_folder is None: @@ -1792,25 +1972,37 @@ def _load_pretrained_model( resolved_model_file = [state_dict] # Prepare the loading function sharing the attributes shared between them. - load_fn = functools.partial( - _load_shard_files_with_threadpool if is_parallel_loading_enabled else _load_shard_file, - model=model, - model_state_dict=model_state_dict, - device_map=device_map, - dtype=dtype, - hf_quantizer=hf_quantizer, - keep_in_fp32_modules=keep_in_fp32_modules, - loaded_keys=loaded_keys, - unexpected_keys=unexpected_keys, - offload_index=offload_index, - offload_folder=offload_folder, - state_dict_index=state_dict_index, - state_dict_folder=state_dict_folder, - ignore_mismatched_sizes=ignore_mismatched_sizes, - low_cpu_mem_usage=low_cpu_mem_usage, - disable_mmap=disable_mmap, - ) - + if tp_shard_specs is not None: + load_fn = functools.partial( + _load_shard_file_tp, + model=model, + model_state_dict=model_state_dict, + tp_shard_specs=tp_shard_specs, + tp_config=tp_config, + dtype=dtype, + keep_in_fp32_modules=keep_in_fp32_modules, + unexpected_keys=unexpected_keys, + ignore_mismatched_sizes=ignore_mismatched_sizes, + ) + else: + load_fn = functools.partial( + _load_shard_files_with_threadpool if is_parallel_loading_enabled else _load_shard_file, + model=model, + model_state_dict=model_state_dict, + device_map=device_map, + dtype=dtype, + hf_quantizer=hf_quantizer, + keep_in_fp32_modules=keep_in_fp32_modules, + loaded_keys=loaded_keys, + unexpected_keys=unexpected_keys, + offload_index=offload_index, + offload_folder=offload_folder, + state_dict_index=state_dict_index, + state_dict_folder=state_dict_folder, + ignore_mismatched_sizes=ignore_mismatched_sizes, + low_cpu_mem_usage=low_cpu_mem_usage, + disable_mmap=disable_mmap, + ) if is_parallel_loading_enabled: offload_index, state_dict_index, _mismatched_keys, _error_msgs = load_fn(resolved_model_file) error_msgs += _error_msgs 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_krea2.py b/src/diffusers/models/transformers/transformer_krea2.py index 55d275e5dca7..980d942831e6 100644 --- a/src/diffusers/models/transformers/transformer_krea2.py +++ b/src/diffusers/models/transformers/transformer_krea2.py @@ -21,7 +21,7 @@ import torch.nn.functional as F from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin +from ...loaders import FromOriginalModelMixin, PeftAdapterMixin from ...utils import apply_lora_scale, logging from ...utils.torch_utils import maybe_adjust_dtype_for_device from ..attention import AttentionMixin, AttentionModuleMixin @@ -336,7 +336,7 @@ def forward(self, ids: torch.Tensor) -> torch.Tensor: return freqs_cos, freqs_sin -class Krea2Transformer2DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixin): +class Krea2Transformer2DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixin, FromOriginalModelMixin): r""" The single-stream MMDiT flow-matching backbone used by the Krea 2 pipeline. 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_minimax_h3.py b/src/diffusers/models/transformers/transformer_minimax_h3.py index f49cdaca2eb6..0f9e38bdb965 100644 --- a/src/diffusers/models/transformers/transformer_minimax_h3.py +++ b/src/diffusers/models/transformers/transformer_minimax_h3.py @@ -19,7 +19,8 @@ import torch.nn as nn from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin +from ...hooks.tensor_parallel import PackedColwiseParallel, ReplicatedInputRowwiseParallel +from ...loaders import FromOriginalModelMixin, PeftAdapterMixin from ...utils import BaseOutput, apply_lora_scale, logging from .._modeling_parallel import ContextParallelInput, ContextParallelOutput from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward @@ -179,9 +180,11 @@ def __call__( key = attn.to_k(hidden_states) value = attn.to_v(hidden_states) - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) + # Reshape by a fixed `head_dim` and let `-1` absorb the head count. Under tensor parallelism each rank + # holds a column-sharded slice (`attn.heads // tp_degree` heads); this keeps the processor TP-agnostic. + query = query.unflatten(-1, (-1, attn.head_dim)) + key = key.unflatten(-1, (-1, attn.head_dim)) + value = value.unflatten(-1, (-1, attn.head_dim)) query = attn.norm_q(query) key = attn.norm_k(key) @@ -373,7 +376,9 @@ def forward( return hidden_states -class MiniMaxH3Transformer3DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixin, CacheMixin): +class MiniMaxH3Transformer3DModel( + ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin +): r""" A Transformer model for joint video + audio generation, introduced in MiniMax-H3. @@ -449,6 +454,45 @@ class MiniMaxH3Transformer3DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftA "audio_proj_out", "rope", ] + # Tensor-parallel plan: how each block's Linears shard across the TP mesh. Q/K/V and the attention output + # are unfused, so they are plain "colwise"/"rowwise" (torch's ColwiseParallel / RowwiseParallel). The one + # packed projection is the SwiGLU input `ff.net.0.proj`, a single Linear producing `[value; gate]` in equal + # halves, hence PackedColwiseParallel([1, 1]) so each half is sharded independently. + # + # `adaln_proj.linear` is the one projection that has to keep its full output width: the six modulation + # parameters it produces scale and shift the full hidden dim of the packed sequence, which is already all-reduced + # by the time they are applied. Sharding it colwise would need an all-gather to rebuild that width, so it is + # sharded `ReplicatedInputRowwiseParallel` instead — over its `time_embed_dim` input, reading the replicated + # `temb` and all-reducing the result. It is worth the collective: at `6 * hidden_size * MINIMAX_H3_MODALITY_NUM` + # outputs per block it is ~40% of the denoiser's weights, and leaving it replicated puts the per-rank floor above + # a single device's memory at any TP degree. The all-reduce itself is over `num_timesteps * 6 * hidden_size * + # MINIMAX_H3_MODALITY_NUM` elements, i.e. hundreds of KB, once per block per step. + # + # Intentionally absent, i.e. replicated on every rank: the RMSNorms (`norm1`, `norm2`, + # `token_refiner.final_norm`, `norm_out.norm`); the QK-norms, which apply over `head_dim` after the heads are + # already split; `norm_out.linear`, which indexes the full hidden dim as `adaln_proj` does but is a single + # `2 * hidden_size` projection rather than one per block, so sharding it would buy nothing; and the patch/text + # embedders and the two output heads. + # + # `attn.to_qkv` is deliberately not listed: it only exists after `fuse_projections()`, and the plan is + # resolved by attribute lookup, so an unconditional entry would break the ordinary unfused model. + _tp_plan = { + # denoiser block stack + "transformer_blocks.*.attn.to_q": "colwise", + "transformer_blocks.*.attn.to_k": "colwise", + "transformer_blocks.*.attn.to_v": "colwise", + "transformer_blocks.*.attn.to_out.0": "rowwise", + "transformer_blocks.*.ff.net.0.proj": PackedColwiseParallel([1, 1]), + "transformer_blocks.*.ff.net.2": "rowwise", + "transformer_blocks.*.adaln_proj.linear": ReplicatedInputRowwiseParallel(), + # the token-refiner blocks are the same attention + SwiGLU FFN, minus AdaLN and rotary + "token_refiner.refiner_blocks.*.attn.to_q": "colwise", + "token_refiner.refiner_blocks.*.attn.to_k": "colwise", + "token_refiner.refiner_blocks.*.attn.to_v": "colwise", + "token_refiner.refiner_blocks.*.attn.to_out.0": "rowwise", + "token_refiner.refiner_blocks.*.ff.net.0.proj": PackedColwiseParallel([1, 1]), + "token_refiner.refiner_blocks.*.ff.net.2": "rowwise", + } # Context parallelism shards the packed sequence, so the split cannot happen on the inputs of `forward`: the rows # of the three modalities are scattered into the packed buffer with sequence-wide indices, which only address the # full sequence. The split therefore happens once the buffer is built, at the first block, and everything that is 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/components_manager.py b/src/diffusers/modular_pipelines/components_manager.py index 87b43bc3a630..4d467e53ea45 100644 --- a/src/diffusers/modular_pipelines/components_manager.py +++ b/src/diffusers/modular_pipelines/components_manager.py @@ -28,7 +28,7 @@ is_accelerate_available, logging, ) -from ..utils.torch_utils import get_device +from ..utils.torch_utils import TorchDeviceBackend, empty_device_cache, get_device if is_accelerate_available(): @@ -179,13 +179,7 @@ def __call__(self, hooks, model_id, model, execution_device): except AttributeError: raise AttributeError(f"Do not know how to compute memory footprint of `{model.__class__.__name__}.") - device_type = execution_device.type - device_module = getattr(torch, device_type, torch.cuda) - try: - mem_on_device = device_module.mem_get_info(execution_device.index)[0] - except AttributeError: - raise AttributeError(f"Do not know how to obtain obtain memory info for {str(device_module)}.") - + mem_on_device = TorchDeviceBackend(execution_device).mem_get_info()[0] mem_on_device = mem_on_device - self.memory_reserve_margin if current_module_size < mem_on_device: return [] @@ -513,10 +507,7 @@ def remove(self, component_id: str = None): import gc gc.collect() - if torch.cuda.is_available(): - torch.cuda.empty_cache() - if torch.xpu.is_available(): - torch.xpu.empty_cache() + empty_device_cache() # YiYi TODO: rename to search_components for now, may remove this method def search_components( @@ -743,12 +734,8 @@ def enable_auto_cpu_offload( if not isinstance(device, torch.device): device = torch.device(device) - device_type = device.type - device_module = getattr(torch, device_type, torch.cuda) - if not hasattr(device_module, "mem_get_info"): - raise NotImplementedError( - f"`enable_auto_cpu_offload() relies on the `mem_get_info()` method. It's not implemented for {str(device.type)}." - ) + # Fail here rather than on the first forward: the strategy cannot run without a free-memory query. + TorchDeviceBackend(device).mem_get_info() if device.index is None: device = torch.device(f"{device.type}:{0}") 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..bb1c1cd54b4f 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 @@ -1214,6 +1228,13 @@ def build_row_timesteps( Returns: `tuple[torch.Tensor, torch.Tensor]`: the distinct timesteps, sorted, and the index of every row into them. """ + # The row plan is built on CPU and the caller moves the finished pair to the denoiser's device: `unique` + # has a data-dependent output shape, which an accelerator would rather not trace. The layout step hands the + # index tensors over already on that device, so bring them back for the scatter below — indexing a CPU + # tensor with an accelerator one does not work. + video_indices = video_indices.cpu() + audio_indices = audio_indices.cpu() + sequence_length = int(video_indices.numel() + audio_indices.numel() + num_text_tokens) row_timesteps = torch.full((sequence_length,), video_timestep, dtype=torch.float32) row_timesteps[video_indices[:num_condition_video_rows]] = condition_video_timestep @@ -1222,7 +1243,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..edd22fa573e3 100644 --- a/src/diffusers/modular_pipelines/modular_pipeline.py +++ b/src/diffusers/modular_pipelines/modular_pipeline.py @@ -30,13 +30,23 @@ from typing_extensions import Self from ..configuration_utils import ConfigMixin, FrozenDict +from ..models.auto_model import AutoModel +from ..models.modeling_utils import ModelMixin from ..pipelines.pipeline_loading_utils import ( LOADABLE_CLASSES, _fetch_class_library_tuple, _unwrap_model, + filter_model_files, simple_get_class_obj, ) -from ..utils import PushToHubMixin, deprecate, is_accelerate_available, logging +from ..utils import ( + TRANSFORMERS_COMPONENT_AUX_FILES, + PushToHubMixin, + deprecate, + is_accelerate_available, + is_transformers_available, + logging, +) from ..utils.dynamic_modules_utils import get_class_from_dynamic_module, resolve_trust_remote_code from ..utils.hub_utils import _resolve_revision, load_or_create_model_card, populate_model_card from ..utils.torch_utils import empty_device_cache, is_compiled_module @@ -59,12 +69,46 @@ ) +# classes whose components are loaded from weight files; a component without a type hint is loaded with `AutoModel` +_MODEL_CLASSES = (ModelMixin, AutoModel) +if is_transformers_available(): + from transformers import PreTrainedModel + + _MODEL_CLASSES = (*_MODEL_CLASSES, PreTrainedModel) + if is_accelerate_available(): import accelerate logger = logging.get_logger(__name__) # pylint: disable=invalid-name +def _is_local_component( + pretrained_model_name_or_path: str | os.PathLike | None, component_spec: ComponentSpec +) -> bool: + """ + Whether the component's files are in `pretrained_model_name_or_path`, a local pipeline directory: weight files for + a model, the config file its class saves for a diffusers component without weights (schedulers, guiders, ...), one + of `TRANSFORMERS_COMPONENT_AUX_FILES` for a transformers one (tokenizers, processors, ...). + """ + if pretrained_model_name_or_path is None: + return False + component_dir = os.path.join(pretrained_model_name_or_path, component_spec.subfolder or "") + if not os.path.isdir(component_dir): + return False + filenames = os.listdir(component_dir) + + class_obj = component_spec.type_hint + is_model = class_obj is None or issubclass(class_obj, _MODEL_CLASSES) + + if is_model: + return len(filter_model_files(filenames)) > 0 + + if issubclass(class_obj, ConfigMixin): + return class_obj.config_name in filenames + + return any(filename in filenames for filename in TRANSFORMERS_COMPONENT_AUX_FILES) + + # map regular pipeline to modular pipeline class name @@ -777,7 +821,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 +1193,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 +1577,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 @@ -1765,6 +1809,11 @@ def __init__( library, class_name, component_spec_dict = value component_spec = self._dict_to_component_spec(name, component_spec_dict) component_spec.default_creation_method = "from_pretrained" + # a local copy of the repo (e.g. `hf download --local-dir`) keeps the original index, which + # points at the Hub; load the components whose files are present locally from the copy + if _is_local_component(pretrained_model_name_or_path, component_spec): + component_spec.pretrained_model_name_or_path = str(pretrained_model_name_or_path) + component_spec.revision = None self._component_specs[name] = component_spec elif name in self._config_specs: 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_sana.py b/src/diffusers/pipelines/pag/pipeline_pag_sana.py index f621255325a0..2b8dd5f14d24 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sana.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sana.py @@ -35,7 +35,7 @@ logging, replace_example_docstring, ) -from ...utils.torch_utils import get_device, is_torch_version, randn_tensor +from ...utils.torch_utils import randn_tensor from ..pipeline_utils import DiffusionPipeline, ImagePipelineOutput from ..pixart_alpha.pipeline_pixart_alpha import ( ASPECT_RATIO_512_BIN, @@ -892,15 +892,9 @@ def __call__( image = latents else: latents = latents.to(self.vae.dtype) - torch_accelerator_module = getattr(torch, get_device(), torch.cuda) - oom_error = ( - torch.OutOfMemoryError - if is_torch_version(">=", "2.5.0") - else torch_accelerator_module.OutOfMemoryError - ) try: image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0] - except oom_error as e: + except torch.OutOfMemoryError as e: warnings.warn( f"{e}. \n" f"Try to use VAE tiling for large images. For example: \n" 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/pipeline_loading_utils.py b/src/diffusers/pipelines/pipeline_loading_utils.py index 69bce1a1c533..6958f49c8ddd 100644 --- a/src/diffusers/pipelines/pipeline_loading_utils.py +++ b/src/diffusers/pipelines/pipeline_loading_utils.py @@ -65,19 +65,6 @@ TRANSFORMERS_DUMMY_MODULES_FOLDER = "transformers.utils" CONNECTED_PIPES_KEYS = ["prior"] -# Auxiliary (non-weight) files a transformers component saves next to its weights. Repos with a flat, -# transformers-style layout host a component's files at the repo root instead of in a subfolder, where the -# folder-based allow patterns of `DiffusionPipeline.download` would miss them. Root-hosted weights and -# `config.json` are matched by their own patterns, so only these auxiliary filenames need listing. -# Currently the set needed by DiffusionGemma — extend as new flat-layout pipelines require it. -TRANSFORMERS_COMPONENT_AUX_FILES = [ - "chat_template.jinja", - "generation_config.json", - "processor_config.json", - "tokenizer.json", - "tokenizer_config.json", -] - logger = logging.get_logger(__name__) LOADABLE_CLASSES = { diff --git a/src/diffusers/pipelines/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index 82e6c4c2aff4..750fed48cbd0 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -59,6 +59,7 @@ from ..utils import ( CONFIG_NAME, DEPRECATED_REVISION_ARGS, + TRANSFORMERS_COMPONENT_AUX_FILES, BaseOutput, PushToHubMixin, _get_detailed_type, @@ -92,7 +93,6 @@ CONNECTED_PIPES_KEYS, CUSTOM_PIPELINE_FILE_NAME, LOADABLE_CLASSES, - TRANSFORMERS_COMPONENT_AUX_FILES, _fetch_class_library_tuple, _get_custom_components_and_folders, _get_custom_pipeline_class, @@ -1187,6 +1187,7 @@ def enable_model_cpu_offload(self, gpu_id: int | None = None, device: torch.devi automatically detect the available accelerator and use. """ self._maybe_raise_error_if_group_offload_active(raise_error=True) + self._maybe_raise_error_if_tensor_parallel_active() is_pipeline_device_mapped = self._is_pipeline_device_mapped() if is_pipeline_device_mapped: @@ -1287,7 +1288,7 @@ def maybe_free_model_hooks(self): return # make sure the model is in the same state as before calling it - self.enable_model_cpu_offload(device=getattr(self, "_offload_device", "cuda")) + self.enable_model_cpu_offload(device=getattr(self, "_offload_device", get_device())) def enable_sequential_cpu_offload(self, gpu_id: int | None = None, device: torch.device | str = None): r""" @@ -1305,6 +1306,7 @@ def enable_sequential_cpu_offload(self, gpu_id: int | None = None, device: torch automatically detect the available accelerator and use. """ self._maybe_raise_error_if_group_offload_active(raise_error=True) + self._maybe_raise_error_if_tensor_parallel_active() if is_accelerate_available() and is_accelerate_version(">=", "0.14.0"): from accelerate import cpu_offload @@ -2242,6 +2244,23 @@ def _maybe_raise_error_if_group_offload_active( return True return False + def _maybe_raise_error_if_tensor_parallel_active(self) -> None: + """Raise if any component is sharded with tensor parallelism, which CPU offloading cannot be applied on top of. + + A tensor-parallel component's parameters are `DTensor` shards tied to that rank's device and process group; + moving them to CPU and back, as the offload hooks do, is not supported. + """ + from ..hooks.tensor_parallel import _raise_if_tensor_parallel + + for component in self.components.values(): + if isinstance(component, torch.nn.Module): + _raise_if_tensor_parallel( + component, + "be CPU-offloaded (model or sequential)", + "Tensor parallelism already keeps only one shard of each weight per rank, so offloading is not " + "needed on top of it.", + ) + def _is_pipeline_device_mapped(self): # We support passing `device_map="cuda"`, for example. This is helpful, in case # users want to pass `device_map="cpu"` when initializing a pipeline. This explicit declaration is desirable 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/sana/pipeline_sana.py b/src/diffusers/pipelines/sana/pipeline_sana.py index 553e45a628d9..f8a4fbe508e3 100644 --- a/src/diffusers/pipelines/sana/pipeline_sana.py +++ b/src/diffusers/pipelines/sana/pipeline_sana.py @@ -38,7 +38,7 @@ scale_lora_layers, unscale_lora_layers, ) -from ...utils.torch_utils import get_device, is_torch_version, randn_tensor +from ...utils.torch_utils import randn_tensor from ..pipeline_utils import DiffusionPipeline from ..pixart_alpha.pipeline_pixart_alpha import ( ASPECT_RATIO_512_BIN, @@ -957,15 +957,9 @@ def __call__( image = latents else: latents = latents.to(self.vae.dtype) - torch_accelerator_module = getattr(torch, get_device(), torch.cuda) - oom_error = ( - torch.OutOfMemoryError - if is_torch_version(">=", "2.5.0") - else torch_accelerator_module.OutOfMemoryError - ) try: image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0] - except oom_error as e: + except torch.OutOfMemoryError as e: warnings.warn( f"{e}. \n" f"Try to use VAE tiling for large images. For example: \n" diff --git a/src/diffusers/pipelines/sana/pipeline_sana_controlnet.py b/src/diffusers/pipelines/sana/pipeline_sana_controlnet.py index de1910f68192..5205ede5904a 100644 --- a/src/diffusers/pipelines/sana/pipeline_sana_controlnet.py +++ b/src/diffusers/pipelines/sana/pipeline_sana_controlnet.py @@ -38,7 +38,7 @@ scale_lora_layers, unscale_lora_layers, ) -from ...utils.torch_utils import get_device, is_torch_version, randn_tensor +from ...utils.torch_utils import randn_tensor from ..pipeline_utils import DiffusionPipeline from ..pixart_alpha.pipeline_pixart_alpha import ( ASPECT_RATIO_512_BIN, @@ -1053,15 +1053,9 @@ def __call__( image = latents else: latents = latents.to(self.vae.dtype) - torch_accelerator_module = getattr(torch, get_device(), torch.cuda) - oom_error = ( - torch.OutOfMemoryError - if is_torch_version(">=", "2.5.0") - else torch_accelerator_module.OutOfMemoryError - ) try: image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0] - except oom_error as e: + except torch.OutOfMemoryError as e: warnings.warn( f"{e}. \n" f"Try to use VAE tiling for large images. For example: \n" diff --git a/src/diffusers/pipelines/sana/pipeline_sana_sprint.py b/src/diffusers/pipelines/sana/pipeline_sana_sprint.py index 812441d8e462..9be1c57b07e7 100644 --- a/src/diffusers/pipelines/sana/pipeline_sana_sprint.py +++ b/src/diffusers/pipelines/sana/pipeline_sana_sprint.py @@ -38,7 +38,7 @@ scale_lora_layers, unscale_lora_layers, ) -from ...utils.torch_utils import get_device, is_torch_version, randn_tensor +from ...utils.torch_utils import randn_tensor from ..pipeline_utils import DiffusionPipeline from ..pixart_alpha.pipeline_pixart_alpha import ASPECT_RATIO_1024_BIN from .pipeline_output import SanaPipelineOutput @@ -839,15 +839,9 @@ def __call__( image = latents else: latents = latents.to(self.vae.dtype) - torch_accelerator_module = getattr(torch, get_device(), torch.cuda) - oom_error = ( - torch.OutOfMemoryError - if is_torch_version(">=", "2.5.0") - else torch_accelerator_module.OutOfMemoryError - ) try: image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0] - except oom_error as e: + except torch.OutOfMemoryError as e: warnings.warn( f"{e}. \n" f"Try to use VAE tiling for large images. For example: \n" diff --git a/src/diffusers/pipelines/sana/pipeline_sana_sprint_img2img.py b/src/diffusers/pipelines/sana/pipeline_sana_sprint_img2img.py index e149e0c597b2..9f97ff835512 100644 --- a/src/diffusers/pipelines/sana/pipeline_sana_sprint_img2img.py +++ b/src/diffusers/pipelines/sana/pipeline_sana_sprint_img2img.py @@ -39,7 +39,7 @@ scale_lora_layers, unscale_lora_layers, ) -from ...utils.torch_utils import get_device, is_torch_version, randn_tensor +from ...utils.torch_utils import randn_tensor from ..pipeline_utils import DiffusionPipeline from ..pixart_alpha.pipeline_pixart_alpha import ASPECT_RATIO_1024_BIN from .pipeline_output import SanaPipelineOutput @@ -929,15 +929,9 @@ def __call__( image = latents else: latents = latents.to(self.vae.dtype) - torch_accelerator_module = getattr(torch, get_device(), torch.cuda) - oom_error = ( - torch.OutOfMemoryError - if is_torch_version(">=", "2.5.0") - else torch_accelerator_module.OutOfMemoryError - ) try: image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0] - except oom_error as e: + except torch.OutOfMemoryError as e: warnings.warn( f"{e}. \n" f"Try to use VAE tiling for large images. For example: \n" diff --git a/src/diffusers/pipelines/sana_video/pipeline_sana_video.py b/src/diffusers/pipelines/sana_video/pipeline_sana_video.py index 7ae85639e358..8cf51d3bb0e8 100644 --- a/src/diffusers/pipelines/sana_video/pipeline_sana_video.py +++ b/src/diffusers/pipelines/sana_video/pipeline_sana_video.py @@ -37,7 +37,7 @@ scale_lora_layers, unscale_lora_layers, ) -from ...utils.torch_utils import get_device, is_torch_version, randn_tensor +from ...utils.torch_utils import randn_tensor from ...video_processor import VideoProcessor from ..pipeline_utils import DiffusionPipeline from .pipeline_output import SanaVideoPipelineOutput @@ -990,12 +990,6 @@ def __call__( video = latents else: latents = latents.to(self.vae.dtype) - torch_accelerator_module = getattr(torch, get_device(), torch.cuda) - oom_error = ( - torch.OutOfMemoryError - if is_torch_version(">=", "2.5.0") - else torch_accelerator_module.OutOfMemoryError - ) if isinstance(self.vae, AutoencoderKLLTX2Video): latents_mean = self.vae.latents_mean latents_std = self.vae.latents_std @@ -1014,7 +1008,7 @@ def __call__( latents = latents / latents_std + latents_mean try: video = self.vae.decode(latents, return_dict=False)[0] - except oom_error as e: + except torch.OutOfMemoryError as e: warnings.warn( f"{e}. \n" f"Try to use VAE tiling for large images. For example: \n" diff --git a/src/diffusers/pipelines/sana_video/pipeline_sana_video_i2v.py b/src/diffusers/pipelines/sana_video/pipeline_sana_video_i2v.py index 81df1d0759da..cc772af22665 100644 --- a/src/diffusers/pipelines/sana_video/pipeline_sana_video_i2v.py +++ b/src/diffusers/pipelines/sana_video/pipeline_sana_video_i2v.py @@ -39,7 +39,7 @@ scale_lora_layers, unscale_lora_layers, ) -from ...utils.torch_utils import get_device, is_torch_version, randn_tensor +from ...utils.torch_utils import randn_tensor from ...video_processor import VideoProcessor from ..pipeline_utils import DiffusionPipeline from .pipeline_output import SanaVideoPipelineOutput @@ -1043,12 +1043,6 @@ def __call__( video = latents else: latents = latents.to(self.vae.dtype) - torch_accelerator_module = getattr(torch, get_device(), torch.cuda) - oom_error = ( - torch.OutOfMemoryError - if is_torch_version(">=", "2.5.0") - else torch_accelerator_module.OutOfMemoryError - ) if isinstance(self.vae, AutoencoderKLLTX2Video): latents_mean = self.vae.latents_mean latents_std = self.vae.latents_std @@ -1067,7 +1061,7 @@ def __call__( latents = latents / latents_std + latents_mean try: video = self.vae.decode(latents, return_dict=False)[0] - except oom_error as e: + except torch.OutOfMemoryError as e: warnings.warn( f"{e}. \n" f"Try to use VAE tiling for large images. For example: \n" 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/src/diffusers/quantizers/gguf/gguf_quantizer.py b/src/diffusers/quantizers/gguf/gguf_quantizer.py index 42ea3982c912..e2bee2e4023d 100644 --- a/src/diffusers/quantizers/gguf/gguf_quantizer.py +++ b/src/diffusers/quantizers/gguf/gguf_quantizer.py @@ -18,6 +18,7 @@ is_torch_available, logging, ) +from ...utils.torch_utils import get_device if is_torch_available() and is_gguf_available(): @@ -177,12 +178,7 @@ def _dequantize(self, model): logger.info( "Model was found to be on CPU (could happen as a result of `enable_model_cpu_offload()`). So, moving it to accelerator. After dequantization, will move the model back to CPU again to preserve the previous device." ) - device = ( - torch.accelerator.current_accelerator() - if hasattr(torch, "accelerator") - else torch.cuda.current_device() - ) - model.to(device) + model.to(get_device()) model = _dequantize_gguf_and_restore_linear(model, self.modules_to_not_convert) if is_model_on_cpu: diff --git a/src/diffusers/training_utils.py b/src/diffusers/training_utils.py index 68c0c5254234..384a73e51296 100644 --- a/src/diffusers/training_utils.py +++ b/src/diffusers/training_utils.py @@ -37,6 +37,7 @@ is_torchvision_available, is_transformers_available, ) +from .utils.torch_utils import empty_device_cache if is_transformers_available(): @@ -412,15 +413,7 @@ def free_memory(): Runs garbage collection. Then clears the cache of the available accelerator. """ gc.collect() - - if torch.cuda.is_available(): - torch.cuda.empty_cache() - elif torch.backends.mps.is_available(): - torch.mps.empty_cache() - elif is_torch_npu_available(): - torch_npu.npu.empty_cache() - elif hasattr(torch, "xpu") and torch.xpu.is_available(): - torch.xpu.empty_cache() + empty_device_cache() @contextmanager diff --git a/src/diffusers/utils/__init__.py b/src/diffusers/utils/__init__.py index 5c63a4bc7661..b3051dfcb9d1 100644 --- a/src/diffusers/utils/__init__.py +++ b/src/diffusers/utils/__init__.py @@ -36,6 +36,7 @@ SAFE_WEIGHTS_INDEX_NAME, SAFETENSORS_FILE_EXTENSION, SAFETENSORS_WEIGHTS_NAME, + TRANSFORMERS_COMPONENT_AUX_FILES, USE_PEFT_BACKEND, WEIGHTS_INDEX_NAME, WEIGHTS_NAME, diff --git a/src/diffusers/utils/constants.py b/src/diffusers/utils/constants.py index fcf0e4518800..d17587c65bdc 100644 --- a/src/diffusers/utils/constants.py +++ b/src/diffusers/utils/constants.py @@ -37,6 +37,17 @@ FLASHPACK_FILE_EXTENSION = "flashpack" GGUF_FILE_EXTENSION = "gguf" ONNX_EXTERNAL_WEIGHTS_NAME = "weights.pb" +# Auxiliary (non-weight) files a transformers component saves next to its weights, or as its only files for tokenizers +# and processors. `DiffusionPipeline.download` uses them to fetch components hosted at the root of a flat, +# transformers-style repo, and `ModularPipeline` to tell that such a component is present in a local directory. +TRANSFORMERS_COMPONENT_AUX_FILES = [ + "chat_template.jinja", + "generation_config.json", + "preprocessor_config.json", + "processor_config.json", + "tokenizer.json", + "tokenizer_config.json", +] HUGGINGFACE_CO_RESOLVE_ENDPOINT = os.environ.get("HF_ENDPOINT", "https://huggingface.co") DIFFUSERS_DYNAMIC_MODULE_NAME = "diffusers_modules" HF_MODULES_CACHE = os.getenv("HF_MODULES_CACHE", os.path.join(HF_HOME, "modules")) diff --git a/src/diffusers/utils/torch_utils.py b/src/diffusers/utils/torch_utils.py index b3292d5cf0d2..801044c58106 100644 --- a/src/diffusers/utils/torch_utils.py +++ b/src/diffusers/utils/torch_utils.py @@ -25,9 +25,7 @@ from . import logging from .import_utils import ( is_torch_available, - is_torch_mlu_available, is_torch_neuronx_available, - is_torch_npu_available, is_torch_version, ) @@ -40,71 +38,6 @@ import torch from torch.fft import fftn, fftshift, ifftn, ifftshift - BACKEND_SUPPORTS_TRAINING = { - "cuda": True, - "xpu": True, - "cpu": True, - "mps": False, - "neuron": False, - "default": True, - } - BACKEND_EMPTY_CACHE = { - "cuda": torch.cuda.empty_cache, - "xpu": torch.xpu.empty_cache, - "cpu": None, - "mps": torch.mps.empty_cache, - "neuron": None, - "default": None, - } - BACKEND_DEVICE_COUNT = { - "cuda": torch.cuda.device_count, - "xpu": torch.xpu.device_count, - "cpu": lambda: 0, - "mps": lambda: 0, - "neuron": lambda: getattr(getattr(torch, "neuron", None), "device_count", lambda: 0)(), - "default": 0, - } - BACKEND_MANUAL_SEED = { - "cuda": torch.cuda.manual_seed, - "xpu": torch.xpu.manual_seed, - "cpu": torch.manual_seed, - "mps": torch.mps.manual_seed, - "neuron": torch.manual_seed, - "default": torch.manual_seed, - } - BACKEND_RESET_PEAK_MEMORY_STATS = { - "cuda": torch.cuda.reset_peak_memory_stats, - "xpu": getattr(torch.xpu, "reset_peak_memory_stats", None), - "cpu": None, - "mps": None, - "neuron": None, - "default": None, - } - BACKEND_RESET_MAX_MEMORY_ALLOCATED = { - "cuda": torch.cuda.reset_max_memory_allocated, - "xpu": getattr(torch.xpu, "reset_peak_memory_stats", None), - "cpu": None, - "mps": None, - "neuron": None, - "default": None, - } - BACKEND_MAX_MEMORY_ALLOCATED = { - "cuda": torch.cuda.max_memory_allocated, - "xpu": getattr(torch.xpu, "max_memory_allocated", None), - "cpu": 0, - "mps": 0, - "neuron": 0, - "default": 0, - } - BACKEND_SYNCHRONIZE = { - "cuda": torch.cuda.synchronize, - "xpu": getattr(torch.xpu, "synchronize", None), - "cpu": None, - "mps": None, - "neuron": getattr(getattr(torch, "neuron", None), "synchronize", None), - "default": None, - } - _FP64_UNSUPPORTED_DEVICES = frozenset({"mps", "npu", "neuron"}) _INT64_UNSUPPORTED_DEVICES = frozenset({"mps", "npu", "neuron"}) _DTYPE_DOWNCAST = {torch.float64: torch.float32, torch.int64: torch.int32} @@ -120,62 +53,6 @@ def maybe_allow_in_graph(cls): return cls -# This dispatches a defined function according to the accelerator from the function definitions. -def _device_agnostic_dispatch(device: str, dispatch_table: dict[str, callable], *args, **kwargs): - if device not in dispatch_table: - return dispatch_table["default"](*args, **kwargs) - - fn = dispatch_table[device] - - # Some device agnostic functions return values. Need to guard against 'None' instead at - # user level - if not callable(fn): - return fn - - return fn(*args, **kwargs) - - -# These are callables which automatically dispatch the function specific to the accelerator -def backend_manual_seed(device: str, seed: int): - return _device_agnostic_dispatch(device, BACKEND_MANUAL_SEED, seed) - - -def backend_synchronize(device: str): - return _device_agnostic_dispatch(device, BACKEND_SYNCHRONIZE) - - -def backend_empty_cache(device: str): - return _device_agnostic_dispatch(device, BACKEND_EMPTY_CACHE) - - -def backend_device_count(device: str): - return _device_agnostic_dispatch(device, BACKEND_DEVICE_COUNT) - - -def backend_reset_peak_memory_stats(device: str): - return _device_agnostic_dispatch(device, BACKEND_RESET_PEAK_MEMORY_STATS) - - -def backend_reset_max_memory_allocated(device: str): - return _device_agnostic_dispatch(device, BACKEND_RESET_MAX_MEMORY_ALLOCATED) - - -def backend_max_memory_allocated(device: str): - return _device_agnostic_dispatch(device, BACKEND_MAX_MEMORY_ALLOCATED) - - -# These are callables which return boolean behaviour flags and can be used to specify some -# device agnostic alternative where the feature is unsupported. -def backend_supports_training(device: str): - if not is_torch_available(): - return False - - if device not in BACKEND_SUPPORTS_TRAINING: - device = "default" - - return BACKEND_SUPPORTS_TRAINING[device] - - def maybe_adjust_dtype_for_device(dtype: "torch.dtype", device: "torch.device") -> "torch.dtype": unsupported = _DTYPE_UNSUPPORTED_DEVICES.get(dtype) return _DTYPE_DOWNCAST[dtype] if unsupported and device.type in unsupported else dtype @@ -348,38 +225,136 @@ def get_torch_cuda_device_capability(): return None -@functools.lru_cache -def get_device(): - if torch.cuda.is_available(): - return "cuda" - elif is_torch_npu_available(): - return "npu" - elif hasattr(torch, "xpu") and torch.xpu.is_available(): - return "xpu" - elif torch.backends.mps.is_available(): - return "mps" - elif is_torch_mlu_available(): - return "mlu" - elif is_torch_neuronx_available() and hasattr(torch, "neuron") and torch.neuron.is_available(): - return "neuron" - else: +class TorchDeviceBackend: + """ + A proxy for the `torch.` namespace (`torch.cuda`, `torch.xpu`, `torch.mps`, ...) of one device. Attributes + the class does not define are the module's own (`synchronize`, `device_count`, `Stream`, `current_stream`, ...); + the methods defined here override the operations whose availability differs between backends and need a fallback: + cache clearing, seeding and memory queries. With no `device`, detects the host accelerator through + `torch.accelerator`. Raises if torch has no module for the backend rather than silently falling back to + `torch.cuda`. + """ + + def __init__(self, device: str | torch.device | None = None): + self.device = torch.device(self._detect_device_type() if device is None else device) + self.module = torch.get_device_module(self.device.type) + + @staticmethod + @functools.lru_cache + def _detect_device_type() -> str: + if torch.accelerator.is_available(): + return torch.accelerator.current_accelerator().type + + # Neuron is XLA-based and never registers as a torch accelerator. + if is_torch_neuronx_available() and hasattr(torch, "neuron") and torch.neuron.is_available(): + return "neuron" return "cpu" + def empty_cache(self) -> None: + # Backends without a caching allocator (cpu, neuron) have nothing to clear. + empty_cache = getattr(self.module, "empty_cache", None) + if empty_cache is not None: + empty_cache() + + def manual_seed(self, seed: int) -> None: + # `torch.manual_seed` seeds every device, so it is the correct fallback for backends without their own. + manual_seed = getattr(self.module, "manual_seed", None) + if manual_seed is None: + torch.manual_seed(seed) + return + manual_seed(seed) + + def __getattr__(self, name: str): + # Proxy: anything not overridden here is `torch.`'s own attribute. + if name == "module": + raise AttributeError(name) + return getattr(self.module, name) + + def _accelerator_serves(self, min_torch_version: str) -> bool: + # `torch.accelerator` only serves the process accelerator, and its memory API arrived in 2.9 (statistics) and + # 2.10 (`get_memory_info`). + current = torch.accelerator.current_accelerator() + return is_torch_version(">=", min_torch_version) and current is not None and current.type == self.device.type + + def mem_get_info(self) -> tuple[int, int]: + """Free and total device memory in bytes.""" + mem_get_info = getattr(self.module, "mem_get_info", None) + if mem_get_info is not None: + return mem_get_info(self.device.index) + if self._accelerator_serves("2.10"): + return torch.accelerator.get_memory_info(self.device) + raise NotImplementedError( + f"`torch.{self.device.type}` does not implement `mem_get_info()`, and `torch.accelerator.get_memory_info()` " + f"cannot serve `{self.device}` on torch {torch.__version__} (requires torch>=2.10 and the current accelerator)." + ) + + def max_memory_allocated(self) -> int: + """Peak memory allocated on the device in bytes since the last reset; 0 where the backend keeps no statistics.""" + max_memory_allocated = getattr(self.module, "max_memory_allocated", None) + if max_memory_allocated is not None: + return max_memory_allocated(self.device.index) + if self._accelerator_serves("2.9"): + return torch.accelerator.max_memory_allocated(self.device) + logger.warning( + f"`torch.{self.device.type}` keeps no memory statistics on torch {torch.__version__}; " + "`max_memory_allocated()` returns 0." + ) + return 0 + + def reset_peak_memory_stats(self) -> None: + reset_peak_memory_stats = getattr(self.module, "reset_peak_memory_stats", None) + if reset_peak_memory_stats is not None: + reset_peak_memory_stats(self.device.index) + return + if self._accelerator_serves("2.9"): + torch.accelerator.reset_peak_memory_stats(self.device) + return + logger.warning( + f"`torch.{self.device.type}` keeps no memory statistics on torch {torch.__version__}; " + "`reset_peak_memory_stats()` is a no-op." + ) + + +def get_device() -> str: + return TorchDeviceBackend._detect_device_type() + def empty_device_cache(device_type: str | None = None): - if device_type is None: - device_type = get_device() - if device_type in ["cpu"]: - return - device_mod = getattr(torch, device_type, torch.cuda) - device_mod.empty_cache() - - -def device_synchronize(device_type: str | None = None): - if device_type is None: - device_type = get_device() - device_mod = getattr(torch, device_type, torch.cuda) - device_mod.synchronize() + TorchDeviceBackend(device_type).empty_cache() + + +# Function-style spellings of `TorchDeviceBackend` for test code. +def backend_manual_seed(device: str, seed: int): + TorchDeviceBackend(device).manual_seed(seed) + + +def backend_synchronize(device: str): + TorchDeviceBackend(device).synchronize() + + +def backend_empty_cache(device: str): + TorchDeviceBackend(device).empty_cache() + + +def backend_device_count(device: str): + return TorchDeviceBackend(device).device_count() + + +def backend_reset_peak_memory_stats(device: str): + TorchDeviceBackend(device).reset_peak_memory_stats() + + +def backend_reset_max_memory_allocated(device: str): + # `reset_max_memory_allocated` is CUDA's deprecated alias of `reset_peak_memory_stats`. + TorchDeviceBackend(device).reset_peak_memory_stats() + + +def backend_max_memory_allocated(device: str): + return TorchDeviceBackend(device).max_memory_allocated() + + +def backend_supports_training(device: str): + return str(device).split(":")[0] not in ("mps", "neuron") def enable_full_determinism(): diff --git a/tests/hooks/test_group_offloading.py b/tests/hooks/test_group_offloading.py index a903186aa6b4..08d3f4ed10b8 100644 --- a/tests/hooks/test_group_offloading.py +++ b/tests/hooks/test_group_offloading.py @@ -334,14 +334,16 @@ def test_warning_logged_if_group_offloaded_pipe_moved_to_accelerator(self, caplo assert f"The module '{self.model.__class__.__name__}' is group offloaded" in caplog.text def test_error_raised_if_streams_used_and_no_accelerator_device(self): - torch_accelerator_module = getattr(torch, torch_device, torch.cuda) - original_is_available = torch_accelerator_module.is_available - torch_accelerator_module.is_available = lambda: False - with pytest.raises(ValueError): + with pytest.raises(ValueError, match="backend implements streams, got `cpu`"): self.model.enable_group_offload( - onload_device=torch.device(torch_device), offload_type="leaf_level", use_stream=True + onload_device=torch.device("cpu"), offload_type="leaf_level", use_stream=True + ) + + def test_error_raised_if_streams_used_and_backend_has_no_streams(self): + with pytest.raises(ValueError, match="backend implements streams, got `mps`"): + self.model.enable_group_offload( + onload_device=torch.device("mps"), offload_type="leaf_level", use_stream=True ) - torch_accelerator_module.is_available = original_is_available def test_error_raised_if_supports_group_offloading_false(self): self.model._supports_group_offloading = False diff --git a/tests/models/test_parallelism_combined.py b/tests/models/test_parallelism_combined.py new file mode 100644 index 000000000000..76b814c73cca --- /dev/null +++ b/tests/models/test_parallelism_combined.py @@ -0,0 +1,164 @@ +# 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. + +"""Tests for combining context parallelism with tensor parallelism on one device mesh. + +These cover the wiring rather than the numerics: that `ParallelConfig` hands each parallelism its own submesh, and +that the resulting process groups are the intended factorisation of the world. They run on CPU over `gloo`, so they +need no accelerator — `gloo` has no `all_to_all`, so an end-to-end Ulysses forward pass cannot run here; that is +covered by `ContextAndTensorParallelTesterMixin` in `testing_utils/parallelism.py`, which needs four accelerators. +""" + +import os +import socket + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +from diffusers.models._modeling_parallel import ContextParallelConfig, ParallelConfig, TensorParallelConfig + + +def _find_free_port(): + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("", 0)) + s.listen(1) + return s.getsockname()[1] + + +def _mesh_factorization_worker(rank, world_size, master_port, ulysses_degree, tp_degree, return_dict): + """Build a combined mesh, run `ParallelConfig.setup`, and report the process groups each parallelism landed on.""" + try: + os.environ["MASTER_ADDR"] = "localhost" + os.environ["MASTER_PORT"] = str(master_port) + os.environ["RANK"] = str(rank) + os.environ["WORLD_SIZE"] = str(world_size) + + dist.init_process_group(backend="gloo", rank=rank, world_size=world_size) + + mesh = torch.distributed.device_mesh.init_device_mesh( + "cpu", mesh_shape=(1, ulysses_degree, tp_degree), mesh_dim_names=("ring", "ulysses", "tp") + ) + config = ParallelConfig( + context_parallel_config=ContextParallelConfig(ulysses_degree=ulysses_degree), + tensor_parallel_config=TensorParallelConfig(tp_degree=tp_degree), + ) + config.setup(rank, world_size, torch.device("cpu"), mesh=mesh) + + cp_config = config.context_parallel_config + tp_config = config.tensor_parallel_config + return_dict[rank] = { + "status": "success", + "cp_ranks": dist.get_process_group_ranks(cp_config._flattened_mesh.get_group()), + "ulysses_local_rank": cp_config._ulysses_local_rank, + "tp_ranks": dist.get_process_group_ranks(tp_config._mesh.get_group()), + "tp_local_rank": tp_config._mesh.get_local_rank(), + "tp_degree": tp_config._tp_degree, + } + except Exception as e: # noqa: BLE001 — surfaced as a test failure by the caller + return_dict[rank] = {"status": "error", "error": f"{type(e).__name__}: {e}"} + finally: + if dist.is_initialized(): + dist.destroy_process_group() + + +def _missing_tp_dim_worker(rank, world_size, master_port, return_dict): + """A combined config handed a CP-only mesh must be rejected, naming the missing 'tp' dimension.""" + try: + os.environ["MASTER_ADDR"] = "localhost" + os.environ["MASTER_PORT"] = str(master_port) + os.environ["RANK"] = str(rank) + os.environ["WORLD_SIZE"] = str(world_size) + + dist.init_process_group(backend="gloo", rank=rank, world_size=world_size) + + mesh = torch.distributed.device_mesh.init_device_mesh( + "cpu", mesh_shape=(1, world_size), mesh_dim_names=("ring", "ulysses") + ) + config = ParallelConfig( + context_parallel_config=ContextParallelConfig(ulysses_degree=world_size), + tensor_parallel_config=TensorParallelConfig(tp_degree=1), + ) + try: + config.setup(rank, world_size, torch.device("cpu"), mesh=mesh) + except ValueError as e: + return_dict[rank] = {"status": "raised", "message": str(e)} + else: + return_dict[rank] = {"status": "no_raise"} + except Exception as e: # noqa: BLE001 + return_dict[rank] = {"status": "error", "error": f"{type(e).__name__}: {e}"} + finally: + if dist.is_initialized(): + dist.destroy_process_group() + + +def _spawn(worker, world_size, *args): + if not dist.is_available(): + pytest.skip("torch.distributed is not available.") + manager = mp.Manager() + return_dict = manager.dict() + mp.spawn(worker, args=(world_size, _find_free_port(), *args, return_dict), nprocs=world_size, join=True) + return return_dict + + +class TestCombinedParallelConfig: + """`ParallelConfig` accepting both parallelisms at once, and splitting the mesh between them.""" + + def test_both_configs_are_accepted(self): + # Combining the two used to raise outright; the mesh dimension per parallelism is what makes it work. + config = ParallelConfig( + context_parallel_config=ContextParallelConfig(ulysses_degree=2), + tensor_parallel_config=TensorParallelConfig(tp_degree=2), + ) + assert config._is_combined + assert config.context_parallel_config.ulysses_degree == 2 + assert config.tensor_parallel_config.tp_degree == 2 + + def test_single_parallelism_is_not_combined(self): + assert not ParallelConfig(context_parallel_config=ContextParallelConfig(ulysses_degree=2))._is_combined + assert not ParallelConfig(tensor_parallel_config=TensorParallelConfig(tp_degree=2))._is_combined + + def test_mesh_is_factored_between_the_two(self): + """Each parallelism must end up on its own process group: TP within a CP chunk, CP across the TP groups. + + With `mesh_shape=(ring=1, ulysses=2, tp=2)` over four ranks, "tp" is the fastest-varying dimension, so ranks + {0,1} and {2,3} are the TP groups and {0,2} / {1,3} are the CP groups. A TP all-reduce must not reach a rank + holding a different sequence chunk, and a CP collective must not reach a rank holding a different weight + shard — this asserts exactly that split. + """ + world_size = 4 + results = _spawn(_mesh_factorization_worker, world_size, 2, 2) + + for rank in range(world_size): + assert results[rank]["status"] == "success", results[rank].get("error") + + assert [results[r]["tp_ranks"] for r in range(4)] == [[0, 1], [0, 1], [2, 3], [2, 3]] + assert [results[r]["cp_ranks"] for r in range(4)] == [[0, 2], [1, 3], [0, 2], [1, 3]] + assert [results[r]["ulysses_local_rank"] for r in range(4)] == [0, 0, 1, 1] + # The TP shard index is the coordinate inside the TP group, which stops tracking the global rank as soon as + # the mesh has more than one dimension. Weight sharding keys off this, so it is the value that matters. + assert [results[r]["tp_local_rank"] for r in range(4)] == [0, 1, 0, 1] + assert all(results[r]["tp_degree"] == 2 for r in range(4)) + + def test_mesh_without_tp_dimension_is_rejected(self): + world_size = 2 + results = _spawn(_missing_tp_dim_worker, world_size) + + for rank in range(world_size): + assert results[rank]["status"] == "raised", ( + f"rank {rank}: expected a ValueError for a mesh with no 'tp' dimension, got {results[rank]}" + ) + assert "tp" in results[rank]["message"] diff --git a/tests/models/testing_utils/__init__.py b/tests/models/testing_utils/__init__.py index 2d7d5ae23257..965a591f95a5 100644 --- a/tests/models/testing_utils/__init__.py +++ b/tests/models/testing_utils/__init__.py @@ -20,8 +20,10 @@ from .lora import LoraHotSwappingForModelTesterMixin, LoraTesterMixin from .memory import CPUOffloadTesterMixin, GroupOffloadTesterMixin, LayerwiseCastingTesterMixin, MemoryTesterMixin from .parallelism import ( + ContextAndTensorParallelTesterMixin, ContextParallelAttentionBackendsTesterMixin, ContextParallelTesterMixin, + HybridParallelTesterMixin, TensorParallelTesterMixin, ) from .quantization import ( @@ -64,8 +66,10 @@ "BitsAndBytesConfigMixin", "BitsAndBytesTesterMixin", "CacheTesterMixin", + "ContextAndTensorParallelTesterMixin", "ContextParallelTesterMixin", "ContextParallelAttentionBackendsTesterMixin", + "HybridParallelTesterMixin", "TensorParallelTesterMixin", "CPUOffloadTesterMixin", "FasterCacheConfigMixin", diff --git a/tests/models/testing_utils/parallelism.py b/tests/models/testing_utils/parallelism.py index 63575abf6b7b..4c47ae790968 100644 --- a/tests/models/testing_utils/parallelism.py +++ b/tests/models/testing_utils/parallelism.py @@ -21,8 +21,9 @@ import torch.distributed as dist import torch.multiprocessing as mp -from diffusers.models._modeling_parallel import ContextParallelConfig, TensorParallelConfig +from diffusers.models._modeling_parallel import ContextParallelConfig, ParallelConfig, TensorParallelConfig from diffusers.models.attention_dispatch import AttentionBackendName, _AttentionBackendRegistry +from diffusers.utils import SAFE_WEIGHTS_INDEX_NAME from ...testing_utils import ( is_attention, @@ -32,6 +33,7 @@ require_torch_multi_accelerator, torch_device, ) +from .common import calculate_expected_num_shards, compute_module_persistent_sizes from .utils import _maybe_cast_to_bf16 @@ -295,6 +297,56 @@ def _tensor_parallel_worker( dist.destroy_process_group() +def _tensor_parallel_from_pretrained_worker( + rank, world_size, master_port, model_class, checkpoint_dir, inputs_dict, return_dict +): + """Worker for `from_pretrained(..., parallel_config=...)`, i.e. sharding while reading the checkpoint. + + Each rank loads only its own slice of every `_tp_plan` parameter straight into a `DTensor` and runs a forward + pass. Rank 0 reports its output and the local/global shapes of one sharded weight so the caller can check both the + numerics and that sharding actually happened. + """ + try: + os.environ["MASTER_ADDR"] = "localhost" + os.environ["MASTER_PORT"] = str(master_port) + os.environ["RANK"] = str(rank) + os.environ["WORLD_SIZE"] = str(world_size) + + device_config = DEVICE_CONFIG.get(torch_device, DEVICE_CONFIG["cuda"]) + dist.init_process_group(backend=device_config["backend"], rank=rank, world_size=world_size) + device_config["module"].set_device(rank) + + from torch.distributed.tensor import DTensor + + model = model_class.from_pretrained( + checkpoint_dir, parallel_config=TensorParallelConfig(tp_degree=world_size) + ).eval() + + device = torch.device(f"{torch_device}:{rank}") + inputs_on_device = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in inputs_dict.items()} + with torch.no_grad(): + output = model(**inputs_on_device, return_dict=False)[0] + if isinstance(output, DTensor): + output = output.full_tensor() + + if rank == 0: + sharded = {k: v for k, v in model.state_dict().items() if isinstance(v, DTensor)} + assert sharded, "No parameter was sharded into a DTensor by the streaming load." + name, param = next(iter(sharded.items())) + return_dict["status"] = "success" + return_dict["num_sharded"] = len(sharded) + return_dict["shard_example"] = (name, list(param.to_local().shape), list(param.shape)) + return_dict["output"] = output.float().cpu().tolist() + + except Exception as e: + if rank == 0: + return_dict["status"] = "error" + return_dict["error"] = f"{type(e).__name__}: {e}" + finally: + if dist.is_initialized(): + dist.destroy_process_group() + + @is_tensor_parallel @require_torch_multi_accelerator class TensorParallelTesterMixin: @@ -344,6 +396,218 @@ def test_tensor_parallel_inference(self, batch_size: int = 1): def test_tensor_parallel_batch_inputs(self): self.test_tensor_parallel_inference(batch_size=2) + def _tp_checkpoint_and_reference(self, tmp_path, world_size, sharded): + """Write a checkpoint for the sharded loaders to read, and record its single-device output. + + Returns `(checkpoint_dir, cpu_inputs, reference_output)`, or skips when the model cannot be sharded + across `world_size` ranks. + """ + if not torch.distributed.is_available(): + pytest.skip("torch.distributed is not available.") + if getattr(self.model_class, "_tp_plan", None) is None: + pytest.skip("Model does not define a `_tp_plan` for tensor parallel inference.") + + init_dict = self.get_init_dict() + num_heads = init_dict.get("num_attention_heads") + if num_heads is not None and num_heads % world_size != 0: + pytest.skip(f"`num_attention_heads` ({num_heads}) is not divisible by tp_degree ({world_size}).") + + inputs_dict = self.get_dummy_inputs() + model = self.model_class(**init_dict).eval().to(torch_device) + with torch.no_grad(): + reference = model(**inputs_dict, return_dict=False)[0].float().cpu() + + checkpoint_dir = str(tmp_path / "checkpoint") + max_shard_size = int(compute_module_persistent_sizes(model)[""] * 0.75) if sharded else "5GB" + model.save_pretrained(checkpoint_dir, max_shard_size=max_shard_size) + + if sharded: + index_path = os.path.join(checkpoint_dir, SAFE_WEIGHTS_INDEX_NAME) + assert os.path.exists(index_path) + expected_num_shards = calculate_expected_num_shards(index_path) + actual_num_shards = len([file for file in os.listdir(checkpoint_dir) if file.endswith(".safetensors")]) + assert actual_num_shards == expected_num_shards > 1 + + inputs_dict = {k: v.cpu() if isinstance(v, torch.Tensor) else v for k, v in inputs_dict.items()} + return checkpoint_dir, inputs_dict, reference + + @pytest.mark.parametrize("sharded", [False, True], ids=["unsharded", "sharded"]) + def test_tensor_parallel_from_pretrained(self, tmp_path, sharded): + """`from_pretrained(..., parallel_config=...)` shards while reading and matches the single-device reference.""" + world_size = 2 + checkpoint_dir, inputs_dict, reference = self._tp_checkpoint_and_reference(tmp_path, world_size, sharded) + + manager = mp.Manager() + return_dict = manager.dict() + mp.spawn( + _tensor_parallel_from_pretrained_worker, + args=( + world_size, + _find_free_port(), + self.model_class, + checkpoint_dir, + inputs_dict, + return_dict, + ), + nprocs=world_size, + join=True, + ) + assert return_dict.get("status") == "success", ( + f"Tensor parallel `from_pretrained` failed: {return_dict.get('error', 'Unknown error')}" + ) + + name, local_shape, global_shape = return_dict["shard_example"] + assert local_shape != global_shape, ( + f"'{name}' has local shape {local_shape} equal to its global shape, so it was not sharded." + ) + + # Sharded matmuls + all-reduce reorder the summation, so allow a small tolerance over the reference. + torch.testing.assert_close(reference, torch.tensor(return_dict["output"]), atol=1e-3, rtol=1e-3) + + +def _context_and_tensor_parallel_worker( + rank, world_size, master_port, model_class, init_dict, cp_dict, tp_degree, inputs_dict, return_dict, state_dict +): + """Worker for combined context + tensor parallel inference. + + Both configs go into one `ParallelConfig`, which puts them on one mesh with a dimension each. The result should + still match the single-device reference: TP is mathematically equivalent to the unsharded model, and CP splits the + sequence and gathers it back, so neither changes the function being computed. + """ + try: + os.environ["MASTER_ADDR"] = "localhost" + os.environ["MASTER_PORT"] = str(master_port) + os.environ["RANK"] = str(rank) + os.environ["WORLD_SIZE"] = str(world_size) + + device_config = DEVICE_CONFIG.get(torch_device, DEVICE_CONFIG["cuda"]) + dist.init_process_group(backend=device_config["backend"], rank=rank, world_size=world_size) + + device_config["module"].set_device(rank) + device = torch.device(f"{torch_device}:{rank}") + + model = model_class(**init_dict) + model.load_state_dict(state_dict) + model.to(device) + model.eval() + + inputs_on_device = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in inputs_dict.items()} + + model.enable_parallelism( + config=ParallelConfig( + context_parallel_config=ContextParallelConfig(**cp_dict), + tensor_parallel_config=TensorParallelConfig(tp_degree=tp_degree), + ) + ) + + with torch.no_grad(): + output = model(**inputs_on_device, return_dict=False)[0] + + if rank == 0: + return_dict["status"] = "success" + return_dict["output_shape"] = list(output.shape) + return_dict["output"] = output.float().cpu().tolist() + + except Exception as e: + if rank == 0: + return_dict["status"] = "error" + return_dict["error"] = f"{type(e).__name__}: {e}" + finally: + if dist.is_initialized(): + dist.destroy_process_group() + + +@is_context_parallel +@is_tensor_parallel +@require_torch_multi_accelerator +class ContextAndTensorParallelTesterMixin: + """Context and tensor parallelism enabled together, over one device mesh. + + The two are orthogonal — CP cuts the sequence, TP cuts the weights — so composing them should leave the computed + function unchanged. Both axes are covered: Ulysses, which trades sequence for heads inside attention, and ring, + which does not touch the head dimension at all. + + Needs `cp_degree` x `tp_degree` accelerators (four by default). Head-count requirements differ per CP type, so a + model whose dummy config has too few heads skips rather than failing. + """ + + cp_degree = 2 + tp_degree = 2 + + @pytest.mark.parametrize("cp_type", ["ulysses_degree", "ring_degree"], ids=["ulysses", "ring"]) + def test_context_and_tensor_parallel_inference(self, cp_type, batch_size: int = 1): + if not torch.distributed.is_available(): + pytest.skip("torch.distributed is not available.") + + if getattr(self.model_class, "_tp_plan", None) is None: + pytest.skip("Model does not define a `_tp_plan` for tensor parallel inference.") + if getattr(self.model_class, "_cp_plan", None) is None: + pytest.skip("Model does not define a `_cp_plan` for context parallel inference.") + + if cp_type == "ring_degree": + active_backend, _ = _AttentionBackendRegistry.get_active_backend() + if active_backend == AttentionBackendName.NATIVE: + pytest.skip("Ring attention is not supported with the native attention backend.") + + world_size = self.cp_degree * self.tp_degree + device_module = DEVICE_CONFIG.get(torch_device, DEVICE_CONFIG["cuda"])["module"] + if device_module.device_count() < world_size: + pytest.skip(f"Combined CP x TP needs {world_size} accelerators, found {device_module.device_count()}.") + + init_dict = self.get_init_dict() + num_heads = init_dict.get("num_attention_heads") + # TP shards the heads; Ulysses then splits what TP left on each rank, so the two multiply. Ring leaves the + # head dimension alone, so only the TP degree has to divide the head count. + required_head_multiple = self.tp_degree * (self.cp_degree if cp_type == "ulysses_degree" else 1) + if num_heads is not None and num_heads % required_head_multiple != 0: + pytest.skip( + f"`num_attention_heads` ({num_heads}) is not divisible by {required_head_multiple}, required for " + f"{cp_type.removesuffix('_degree')}={self.cp_degree} x tp_degree={self.tp_degree}." + ) + + inputs_dict = self.get_dummy_inputs(batch_size=batch_size) + + # Single-device reference, captured before anything is sharded. + model = self.model_class(**init_dict).eval().to(torch_device) + state_dict = {k: v.cpu() for k, v in model.state_dict().items()} + with torch.no_grad(): + ref_output = model(**inputs_dict, return_dict=False)[0].float().cpu() + + inputs_dict = {k: v.cpu() if isinstance(v, torch.Tensor) else v for k, v in inputs_dict.items()} + + manager = mp.Manager() + return_dict = manager.dict() + mp.spawn( + _context_and_tensor_parallel_worker, + args=( + world_size, + _find_free_port(), + self.model_class, + init_dict, + {cp_type: self.cp_degree}, + self.tp_degree, + inputs_dict, + return_dict, + state_dict, + ), + nprocs=world_size, + join=True, + ) + + assert return_dict.get("status") == "success", ( + f"Combined context + tensor parallel inference failed: {return_dict.get('error', 'Unknown error')}" + ) + + combined_output = torch.tensor(return_dict["output"]) + assert list(ref_output.shape) == return_dict["output_shape"] + # Two sets of collectives reorder the summation on top of the sharded matmuls, so the tolerance matches the + # TP-only test rather than the tighter CP-only one. + torch.testing.assert_close(ref_output, combined_output, atol=1e-3, rtol=1e-3) + + @pytest.mark.parametrize("cp_type", ["ulysses_degree", "ring_degree"], ids=["ulysses", "ring"]) + def test_context_and_tensor_parallel_batch_inputs(self, cp_type): + self.test_context_and_tensor_parallel_inference(cp_type, batch_size=2) + @is_context_parallel @require_torch_multi_accelerator @@ -610,3 +874,132 @@ def test_context_parallel_attn_backend_inference(self, cp_type, attention_backen cp_output = torch.tensor(return_dict["output"], dtype=ref_output.dtype) torch.testing.assert_close(ref_output, cp_output, atol=1e-2, rtol=1e-2) + + +def _hybrid_parallel_worker( + rank, world_size, master_port, model_class, init_dict, cp_dict, tp_degree, inputs_dict, return_dict, state_dict +): + """Worker function for combined tensor + context parallel inference testing. + + Both parallelisms are requested through a single `ParallelConfig`, which shares one device mesh between them, so + each rank holds `1 / tp_degree` of every sharded weight *and* `1 / (ring_degree * ulysses_degree)` of the + sequence. Rank 0 reports its output so the caller can compare it against a single-device reference: the + composition is mathematically equivalent to the unsharded model up to floating-point reduction order. + """ + try: + os.environ["MASTER_ADDR"] = "localhost" + os.environ["MASTER_PORT"] = str(master_port) + os.environ["RANK"] = str(rank) + os.environ["WORLD_SIZE"] = str(world_size) + + device_config = DEVICE_CONFIG.get(torch_device, DEVICE_CONFIG["cuda"]) + backend = device_config["backend"] + device_module = device_config["module"] + + dist.init_process_group(backend=backend, rank=rank, world_size=world_size) + + device_module.set_device(rank) + device = torch.device(f"{torch_device}:{rank}") + + model = model_class(**init_dict) + model.load_state_dict(state_dict) + model.to(device) + model.eval() + + inputs_on_device = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in inputs_dict.items()} + + model.enable_parallelism( + config=ParallelConfig( + tensor_parallel_config=TensorParallelConfig(tp_degree=tp_degree), + context_parallel_config=ContextParallelConfig(**cp_dict), + ) + ) + + with torch.no_grad(): + output = model(**inputs_on_device, return_dict=False)[0] + + if rank == 0: + return_dict["status"] = "success" + return_dict["output_shape"] = list(output.shape) + # Serialise via nested list so the manager dict can transport it across processes. + return_dict["output"] = output.float().cpu().tolist() + + except Exception as e: + if rank == 0: + return_dict["status"] = "error" + return_dict["error"] = str(e) + finally: + if dist.is_initialized(): + dist.destroy_process_group() + + +@is_context_parallel +@is_tensor_parallel +@require_torch_multi_accelerator +class HybridParallelTesterMixin: + """Tensor parallelism and context parallelism together, from one `ParallelConfig`. + + Needs `tp_degree * ulysses_degree` accelerators (4 at the degrees used here), so it skips on a 2-device runner. + """ + + def test_hybrid_parallel_inference(self, batch_size: int = 1): + if not torch.distributed.is_available(): + pytest.skip("torch.distributed is not available.") + + for plan in ("_tp_plan", "_cp_plan"): + if getattr(self.model_class, plan, None) is None: + pytest.skip(f"Model does not define a `{plan}`, which hybrid parallelism requires.") + + tp_degree, ulysses_degree = 2, 2 + world_size = tp_degree * ulysses_degree + device_count = DEVICE_CONFIG.get(torch_device, DEVICE_CONFIG["cuda"])["module"].device_count() + if device_count < world_size: + pytest.skip( + f"tp_degree={tp_degree} x ulysses_degree={ulysses_degree} needs {world_size} accelerators, " + f"found {device_count}." + ) + + init_dict = self.get_init_dict() + num_heads = init_dict.get("num_attention_heads") + # Each rank keeps `num_heads // tp_degree` heads, which Ulysses splits again. + if num_heads is not None and num_heads % world_size != 0: + pytest.skip(f"`num_attention_heads` ({num_heads}) is not divisible by {world_size}.") + + inputs_dict = self.get_dummy_inputs(batch_size=batch_size) + + # Single-device reference + model = self.model_class(**init_dict).eval().to(torch_device) + state_dict = {k: v.cpu() for k, v in model.state_dict().items()} + with torch.no_grad(): + ref_output = model(**inputs_dict, return_dict=False)[0].float().cpu() + + inputs_dict = {k: v.cpu() if isinstance(v, torch.Tensor) else v for k, v in inputs_dict.items()} + + master_port = _find_free_port() + manager = mp.Manager() + return_dict = manager.dict() + + mp.spawn( + _hybrid_parallel_worker, + args=( + world_size, + master_port, + self.model_class, + init_dict, + {"ulysses_degree": ulysses_degree}, + tp_degree, + inputs_dict, + return_dict, + state_dict, + ), + nprocs=world_size, + join=True, + ) + + assert return_dict.get("status") == "success", ( + f"Hybrid parallel inference failed: {return_dict.get('error', 'Unknown error')}" + ) + + output = torch.tensor(return_dict["output"]) + # Sharded matmuls plus the Ulysses all-to-all reorder the summation, hence the tolerance. + torch.testing.assert_close(ref_output, output, atol=1e-3, rtol=1e-3) diff --git a/tests/models/transformers/_neuron_hybrid_worker.py b/tests/models/transformers/_neuron_hybrid_worker.py new file mode 100644 index 000000000000..1c054af44104 --- /dev/null +++ b/tests/models/transformers/_neuron_hybrid_worker.py @@ -0,0 +1,125 @@ +# 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. + +"""Generic torchrun worker: assert a model's Neuron tensor-parallel x context-parallel output matches its reference. + +The counterpart of `_neuron_tp_worker.py` for the two parallelisms composed in a single `ParallelConfig`. Same +contract: the model under test is supplied as a `module:function` spec reference on the command line, and the +referenced factory returns `(model_class, init_dict, inputs)` with CPU tensors. + + torchrun --nproc_per_node=8 _neuron_hybrid_worker.py \\ + tests.models.transformers.test_models_transformer_flux:make_neuron_hybrid_spec + +`tp_degree` and `ulysses_degree` are read from `TP_DEGREE` / `ULYSSES_DEGREE` (defaults 2 and 4, whose product is +the launched world size). `ulysses_degree` cannot be 2 on Neuron: its all-to-all only accepts group sizes of 4, 8, +16 or multiples of 32. + +No mesh is passed, so this also exercises the default mesh layout that `enable_parallelism` builds for the combined +case -- which matters on Neuron, where the all-to-all Ulysses depends on rejects strided replica groups and so the +context-parallel dimensions have to vary fastest. + +Exit code 0 means the composed path is numerically equivalent to the unsharded model; non-zero means failure. +""" + +import argparse +import importlib +import os +import sys +import traceback + + +# Make the in-repo `diffusers` and `tests` packages importable when run via torchrun from an arbitrary CWD. +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "..", "src")) +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "..")) + +import torch +import torch.distributed as dist +import torch_neuronx # noqa: F401 — registers torch.neuron + +from diffusers import ContextParallelConfig, ParallelConfig, TensorParallelConfig + + +def main(): + parser = argparse.ArgumentParser(description="Neuron tensor-parallel x context-parallel correctness worker.") + parser.add_argument( + "spec", + help="`module:function` reference returning (model_class, init_dict, cpu_inputs) for the model under test.", + ) + args = parser.parse_args() + module_name, _, fn_name = args.spec.partition(":") + model_class, init_dict, inputs = getattr(importlib.import_module(module_name), fn_name)() + + tp_degree = int(os.environ.get("TP_DEGREE", "2")) + ulysses_degree = int(os.environ.get("ULYSSES_DEGREE", "4")) + + dist.init_process_group(backend="neuron") + rank = dist.get_rank() + world_size = dist.get_world_size() + device = torch.neuron.current_device() + + if tp_degree * ulysses_degree != world_size: + raise ValueError( + f"tp_degree ({tp_degree}) x ulysses_degree ({ulysses_degree}) must equal the world size ({world_size})." + ) + + # Identical weights on every rank (same seed), kept on CPU as the Neuron pre-shard backend requires. + torch.manual_seed(0) + model = model_class(**init_dict).eval() + + # Single-device (unsharded) reference on CPU, computed before the shard plan mutates the weights in place. + with torch.no_grad(): + ref_output = model(**inputs, return_dict=False)[0].float().cpu() + + model.enable_parallelism( + config=ParallelConfig( + tensor_parallel_config=TensorParallelConfig(tp_degree=tp_degree), + context_parallel_config=ContextParallelConfig(ulysses_degree=ulysses_degree), + ) + ) + model = model.to(device) + torch.neuron.synchronize() + + inputs_on_device = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in inputs.items()} + with torch.no_grad(): + output = model(**inputs_on_device, return_dict=False)[0] + torch.neuron.synchronize() + output = output.float().cpu() + + if rank == 0: + assert output.shape == ref_output.shape, f"shape mismatch: {output.shape} vs {ref_output.shape}" + assert torch.isfinite(output).all(), "output contains non-finite values" + max_abs = (output - ref_output).abs().max().item() + denom = ref_output.abs().max().item() + 1e-6 + print( + f"[rank0] tp_degree={tp_degree} ulysses_degree={ulysses_degree} " + f"output_shape={tuple(output.shape)} max_abs_diff={max_abs:.4e} max_rel_diff={max_abs / denom:.4e}" + ) + # Neuron runs matmuls in bf16 internally, so compare with a bf16-level tolerance, as `_neuron_tp_worker` + # does. A wrong shard plan or a mis-ordered mesh produces grossly different output and is caught well + # inside this bound. + torch.testing.assert_close(output, ref_output, atol=2e-2, rtol=2e-2) + print("[rank0] PASS: Neuron hybrid-parallel output matches single-device reference.") + + dist.barrier() + dist.destroy_process_group() + + +if __name__ == "__main__": + try: + main() + except Exception: + traceback.print_exc() + # Ensure a non-zero exit so the launching pytest sees the failure. + os._exit(1) 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): diff --git a/tests/models/transformers/test_models_transformer_flux.py b/tests/models/transformers/test_models_transformer_flux.py index be76f892fc4c..5eca619cb999 100644 --- a/tests/models/transformers/test_models_transformer_flux.py +++ b/tests/models/transformers/test_models_transformer_flux.py @@ -27,7 +27,13 @@ from diffusers.models.transformers.transformer_flux import FluxIPAdapterAttnProcessor from diffusers.utils.torch_utils import randn_tensor -from ...testing_utils import enable_full_determinism, is_tensor_parallel, require_torch_neuron, torch_device +from ...testing_utils import ( + enable_full_determinism, + is_context_parallel, + is_tensor_parallel, + require_torch_neuron, + torch_device, +) from ..testing_utils import ( AttentionBackendTesterMixin, AttentionTesterMixin, @@ -40,6 +46,7 @@ FirstBlockCacheTesterMixin, GGUFCompileTesterMixin, GGUFTesterMixin, + HybridParallelTesterMixin, IPAdapterTesterMixin, LoraHotSwappingForModelTesterMixin, LoraTesterMixin, @@ -276,6 +283,23 @@ class TestFluxTransformerTensorParallel(FluxTransformerTesterConfig, TensorParal """Tensor Parallel inference tests for Flux Transformer (CUDA/XPU multi-accelerator).""" +class TestFluxTransformerHybridParallel(FluxTransformerTesterConfig, HybridParallelTesterMixin): + """Tensor Parallel x Context Parallel inference tests for Flux Transformer (needs 4 accelerators).""" + + +def make_neuron_hybrid_spec(): + """Model spec consumed by the generic Neuron hybrid worker (`_neuron_hybrid_worker.py`). + + Same contract as `make_neuron_tp_spec`, but `num_attention_heads` is raised to 8 so the head count survives + being divided twice: `tp_degree=2` leaves 4 heads per rank and `ulysses_degree=4` splits those into 1 each. + (`ulysses_degree` cannot be 2 on Neuron, whose all-to-all only accepts group sizes of 4, 8, 16 or multiples + of 32.) + """ + config = FluxTransformerTesterConfig() + init_dict = config.get_init_dict() | {"num_attention_heads": 8} + return FluxTransformer2DModel, init_dict, config.get_dummy_inputs(device="cpu") + + def make_neuron_tp_spec(): """Model spec consumed by the generic Neuron TP worker (`_neuron_tp_worker.py`). @@ -309,6 +333,28 @@ def test_tensor_parallel_neuron_inference(self): ) +@is_context_parallel +@is_tensor_parallel +@require_torch_neuron +class TestFluxTransformerHybridParallelNeuron: + """Tensor Parallel x Context Parallel inference test for Flux Transformer on AWS Neuron. + + Same launching pattern as `TestFluxTransformerTensorParallelNeuron`: Neuron needs `torchrun` with the + `"neuron"` distributed backend, so it cannot use the `torch.multiprocessing`/NCCL spawn path of + `HybridParallelTesterMixin`. Runs at `tp_degree=2 x ulysses_degree=4`, i.e. 8 ranks. + """ + + def test_hybrid_parallel_neuron_inference(self): + worker = os.path.join(os.path.dirname(__file__), "_neuron_hybrid_worker.py") + spec = "tests.models.transformers.test_models_transformer_flux:make_neuron_hybrid_spec" + cmd = [sys.executable, "-m", "torch.distributed.run", "--nproc_per_node=8", worker, spec] + result = subprocess.run(cmd, capture_output=True, text=True) + assert result.returncode == 0, ( + f"Neuron hybrid-parallel worker failed (exit {result.returncode}).\n" + f"--- stdout ---\n{result.stdout}\n--- stderr ---\n{result.stderr}" + ) + + class TestFluxTransformerIPAdapter(FluxTransformerTesterConfig, IPAdapterTesterMixin): """IP Adapter tests for Flux Transformer.""" diff --git a/tests/models/transformers/test_models_transformer_krea2.py b/tests/models/transformers/test_models_transformer_krea2.py index 261bc13e77b9..3b5fb4f37ac2 100644 --- a/tests/models/transformers/test_models_transformer_krea2.py +++ b/tests/models/transformers/test_models_transformer_krea2.py @@ -25,6 +25,7 @@ LoraTesterMixin, MemoryTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, ) @@ -159,3 +160,21 @@ class TestKrea2TransformerAttention(Krea2TransformerTesterConfig, AttentionTeste class TestKrea2TransformerLoRA(Krea2TransformerTesterConfig, LoraTesterMixin): pass + + +class TestKrea2TransformerSingleFile(Krea2TransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/krea/Krea-2-Raw/blob/main/raw.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "krea/Krea-2-Raw" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} + + @property + def torch_dtype(self): + return torch.bfloat16 diff --git a/tests/models/transformers/test_models_transformer_minimax_h3.py b/tests/models/transformers/test_models_transformer_minimax_h3.py index 00baa37c84a0..7c2eaad617f5 100644 --- a/tests/models/transformers/test_models_transformer_minimax_h3.py +++ b/tests/models/transformers/test_models_transformer_minimax_h3.py @@ -13,13 +13,17 @@ # See the License for the specific language governing permissions and # limitations under the License. +import os +import subprocess +import sys + import torch from diffusers import MiniMaxH3Transformer3DModel from diffusers.models.transformers.transformer_minimax_h3 import MiniMaxH3TransformerOutput from diffusers.utils.torch_utils import randn_tensor -from ...testing_utils import enable_full_determinism, torch_device +from ...testing_utils import enable_full_determinism, is_tensor_parallel, require_torch_neuron, torch_device from ..testing_utils import ( AttentionTesterMixin, BaseModelTesterConfig, @@ -27,6 +31,8 @@ LoraTesterMixin, MemoryTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, + TensorParallelTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, ) @@ -84,7 +90,7 @@ def get_init_dict(self) -> dict: "rope_freq_dim": 2, } - def get_packed_layout(self, num_video_tokens: int = NUM_VIDEO_TOKENS) -> dict: + def get_packed_layout(self, num_video_tokens: int = NUM_VIDEO_TOKENS, device: str | torch.device = None) -> dict: r""" Build the structural arguments of one packed sequence. @@ -92,29 +98,33 @@ def get_packed_layout(self, num_video_tokens: int = NUM_VIDEO_TOKENS) -> dict: modality and its noise level, and hands over the `(t, h, w)` grid plus the three index tensors. The layout here mirrors what the pipelines pack, with two distinct timesteps so the `(timestep, modality)` AdaLN table is addressed on more than one row. + + `device` defaults to the test device; the Neuron TP spec asks for CPU because its worker builds the model on + CPU and moves it only after sharding. """ + device = torch_device if device is None else device sequence_length = NUM_TEXT_TOKENS + NUM_AUDIO_TOKENS + num_video_tokens - text_indices = torch.arange(NUM_TEXT_TOKENS, device=torch_device) - audio_indices = torch.arange(NUM_TEXT_TOKENS, NUM_TEXT_TOKENS + NUM_AUDIO_TOKENS, device=torch_device) - video_indices = torch.arange(NUM_TEXT_TOKENS + NUM_AUDIO_TOKENS, sequence_length, device=torch_device) + text_indices = torch.arange(NUM_TEXT_TOKENS, device=device) + audio_indices = torch.arange(NUM_TEXT_TOKENS, NUM_TEXT_TOKENS + NUM_AUDIO_TOKENS, device=device) + video_indices = torch.arange(NUM_TEXT_TOKENS + NUM_AUDIO_TOKENS, sequence_length, device=device) # 0 = video, 1 = text, 2 = audio. - token_tags = torch.empty(sequence_length, dtype=torch.long, device=torch_device) + token_tags = torch.empty(sequence_length, dtype=torch.long, device=device) token_tags[text_indices] = 1 token_tags[audio_indices] = 2 token_tags[video_indices] = 0 # The conditioning-free rows share the video timestep; the audio rows step down their own schedule. - timestep_indices = torch.zeros(sequence_length, dtype=torch.long, device=torch_device) + timestep_indices = torch.zeros(sequence_length, dtype=torch.long, device=device) timestep_indices[audio_indices] = 1 - position_ids = torch.zeros(sequence_length, 3, dtype=torch.float32, device=torch_device) - position_ids[:, 0] = torch.arange(sequence_length, dtype=torch.float32, device=torch_device) - position_ids[video_indices, 1] = torch.arange(num_video_tokens, dtype=torch.float32, device=torch_device) % 4 - position_ids[video_indices, 2] = torch.arange(num_video_tokens, dtype=torch.float32, device=torch_device) % 2 + position_ids = torch.zeros(sequence_length, 3, dtype=torch.float32, device=device) + position_ids[:, 0] = torch.arange(sequence_length, dtype=torch.float32, device=device) + position_ids[video_indices, 1] = torch.arange(num_video_tokens, dtype=torch.float32, device=device) % 4 + position_ids[video_indices, 2] = torch.arange(num_video_tokens, dtype=torch.float32, device=device) % 2 return { - "timestep": torch.tensor([0.7, 0.3], device=torch_device), + "timestep": torch.tensor([0.7, 0.3], device=device), "timestep_indices": timestep_indices, "token_tags": token_tags, "position_ids": position_ids, @@ -123,7 +133,10 @@ def get_packed_layout(self, num_video_tokens: int = NUM_VIDEO_TOKENS) -> dict: "text_indices": text_indices, } - def get_dummy_inputs(self, num_video_tokens: int = NUM_VIDEO_TOKENS, batch_size: int = 2) -> dict: + def get_dummy_inputs( + self, num_video_tokens: int = NUM_VIDEO_TOKENS, batch_size: int = 2, device: str | torch.device = None + ) -> dict: + device = torch_device if device is None else device generator = self.generator init_dict = self.get_init_dict() patch_size = init_dict["patch_size"] @@ -131,17 +144,17 @@ def get_dummy_inputs(self, num_video_tokens: int = NUM_VIDEO_TOKENS, batch_size: return { "hidden_states": randn_tensor( - (batch_size, num_video_tokens, video_patch_dim), generator=generator, device=torch_device + (batch_size, num_video_tokens, video_patch_dim), generator=generator, device=device ), "audio_hidden_states": randn_tensor( (batch_size, NUM_AUDIO_TOKENS, init_dict["audio_in_channels"]), generator=generator, - device=torch_device, + device=device, ), "encoder_hidden_states": randn_tensor( - (batch_size, NUM_TEXT_TOKENS, init_dict["text_dim"]), generator=generator, device=torch_device + (batch_size, NUM_TEXT_TOKENS, init_dict["text_dim"]), generator=generator, device=device ), - **self.get_packed_layout(num_video_tokens), + **self.get_packed_layout(num_video_tokens, device=device), } @@ -189,3 +202,56 @@ class TestMiniMaxH3TransformerContextParallel(MiniMaxH3TransformerTesterConfig, class TestMiniMaxH3TransformerLoRA(MiniMaxH3TransformerTesterConfig, LoraTesterMixin): """LoRA tests for the MiniMax-H3 transformer.""" + + +class TestMiniMaxH3TransformerSingleFile(MiniMaxH3TransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return ( + "https://huggingface.co/Comfy-Org/MiniMax-H3/blob/main/diffusion_models/minimax_h3_fl2va_bf16.safetensors" + ) + + @property + def pretrained_model_name_or_path(self): + return "MiniMaxAI/MiniMax-H3" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} + + +class TestMiniMaxH3TransformerTensorParallel(MiniMaxH3TransformerTesterConfig, TensorParallelTesterMixin): + """Tensor Parallel inference tests for the MiniMax-H3 transformer (CUDA/XPU multi-accelerator).""" + + +def make_neuron_tp_spec(): + """Model spec consumed by the generic Neuron TP worker (`_neuron_tp_worker.py`). + + Returns `(model_class, init_dict, cpu_inputs)`. Defined here so all MiniMax-H3-specific test data lives in this + file while the worker stays model-agnostic. Reuses the shared tester config so the spec never drifts from the + rest of the MiniMax-H3 tests. + """ + config = MiniMaxH3TransformerTesterConfig() + return MiniMaxH3Transformer3DModel, config.get_init_dict(), config.get_dummy_inputs(device="cpu") + + +@is_tensor_parallel +@require_torch_neuron +class TestMiniMaxH3TransformerTensorParallelNeuron: + """Tensor Parallel inference test for the MiniMax-H3 transformer on AWS Neuron. + + Neuron TP runs through `torchrun` with the `"neuron"` distributed backend, so it cannot use the + `torch.multiprocessing`/NCCL spawn path of `TensorParallelTesterMixin`. This launches the generic worker with + the MiniMax-H3 model spec (`make_neuron_tp_spec`); the worker asserts the sharded output matches a + single-device reference, and the test checks its exit code. + """ + + def test_tensor_parallel_neuron_inference(self): + worker = os.path.join(os.path.dirname(__file__), "_neuron_tp_worker.py") + spec = "tests.models.transformers.test_models_transformer_minimax_h3:make_neuron_tp_spec" + cmd = [sys.executable, "-m", "torch.distributed.run", "--nproc_per_node=2", worker, spec] + result = subprocess.run(cmd, capture_output=True, text=True) + assert result.returncode == 0, ( + f"Neuron tensor-parallel worker failed (exit {result.returncode}).\n" + f"--- stdout ---\n{result.stdout}\n--- stderr ---\n{result.stderr}" + ) diff --git a/tests/models/transformers/test_models_transformer_qwenimage.py b/tests/models/transformers/test_models_transformer_qwenimage.py index 5fcf37f6ff3f..1729ca22c08b 100644 --- a/tests/models/transformers/test_models_transformer_qwenimage.py +++ b/tests/models/transformers/test_models_transformer_qwenimage.py @@ -30,6 +30,7 @@ AttentionTesterMixin, BaseModelTesterConfig, BitsAndBytesTesterMixin, + ContextAndTensorParallelTesterMixin, ContextParallelAttentionBackendsTesterMixin, ContextParallelTesterMixin, LoraHotSwappingForModelTesterMixin, @@ -307,6 +308,25 @@ class TestQwenImageTransformerTensorParallel(QwenImageTransformerTesterConfig, T """Tensor Parallel inference tests for QwenImage Transformer (CUDA/XPU multi-accelerator).""" +class TestQwenImageTransformerContextAndTensorParallel( + QwenImageTransformerTesterConfig, ContextAndTensorParallelTesterMixin +): + """Context Parallel x Tensor Parallel inference tests for QwenImage Transformer (4 accelerators). + + QwenImage is the model this runs on because its dummy config has four attention heads, and Ulysses splits the + heads TP already sharded — so `ulysses_degree` x `tp_degree` has to divide the head count. + """ + + def get_dummy_inputs(self, batch_size: int = 1) -> dict[str, torch.Tensor]: + inputs = super().get_dummy_inputs(batch_size=batch_size) + encoder_hidden_states_mask = inputs["encoder_hidden_states_mask"] + encoder_hidden_states_mask[:, 1] = 0 + encoder_hidden_states_mask[:, 3] = 0 + encoder_hidden_states_mask[:, 5:] = 0 + inputs["encoder_hidden_states_mask"] = encoder_hidden_states_mask + return inputs + + def make_neuron_tp_spec(): """Model spec consumed by the generic Neuron TP worker (``_neuron_tp_worker.py``). diff --git a/tests/modular_pipelines/test_modular_pipeline_loading.py b/tests/modular_pipelines/test_modular_pipeline_loading.py index 4e4797106fe4..4871eab4e67d 100644 --- a/tests/modular_pipelines/test_modular_pipeline_loading.py +++ b/tests/modular_pipelines/test_modular_pipeline_loading.py @@ -15,9 +15,11 @@ import json import os +import shutil import pytest import torch +from huggingface_hub import snapshot_download from diffusers import AutoModel, ControlNetModel, ModularPipeline, UNet2DConditionModel from diffusers.modular_pipelines.modular_pipeline_utils import ComponentSpec @@ -276,3 +278,69 @@ def test_init_raises_without_resolvable_blocks(self): # The base class has no `default_blocks_name`, so with no `blocks` there is nothing to build from. with pytest.raises(ValueError, match="No pipeline blocks could be resolved"): ModularPipeline() + + +class TestLoadFromLocalCopy: + def test_local_copy_loads_present_components_locally(self, tmp_path): + """`hf download --local-dir` keeps the index pointing at the Hub; components whose subfolder is present in + the local copy load from it, the rest keep their recorded spec.""" + local_dir = str(tmp_path / "local-copy") + cache_dir = str(tmp_path / "cache") + snapshot_download("hf-internal-testing/tiny-anima-modular-pipe", local_dir=local_dir) + + pipe = ModularPipeline.from_pretrained(local_dir) + for name in ("vae", "transformer", "text_encoder", "scheduler"): + spec = pipe._component_specs[name] + assert spec.pretrained_model_name_or_path == local_dir, f"{name} should load from the local copy" + assert spec.revision is None + assert ( + pipe._component_specs["t5_tokenizer"].pretrained_model_name_or_path == "hf-internal-testing/tiny-random-t5" + ) + + pipe.load_components(names=["vae"], dtype=torch.float32, local_files_only=True, cache_dir=cache_dir) + assert pipe.vae is not None + cached_weights = [p for p in (tmp_path / "cache").rglob("*") if p.suffix in (".safetensors", ".bin")] + assert cached_weights == [], f"weights should not be in the Hub cache: {cached_weights}" + + def test_local_copy_missing_files_keeps_recorded_spec(self, tmp_path): + """A missing subfolder, or a model subfolder without weight files (e.g. a partial download), keeps the + recorded spec instead of shadowing it with an unloadable folder.""" + local_dir = str(tmp_path / "local-copy") + snapshot_download("hf-internal-testing/tiny-anima-modular-pipe", local_dir=local_dir) + shutil.rmtree(os.path.join(local_dir, "transformer")) + for filename in os.listdir(os.path.join(local_dir, "vae")): + if filename.endswith((".safetensors", ".bin")): + os.remove(os.path.join(local_dir, "vae", filename)) + + pipe = ModularPipeline.from_pretrained(local_dir) + assert ( + pipe._component_specs["transformer"].pretrained_model_name_or_path + == "hf-internal-testing/tiny-anima-modular-pipe" + ) + assert ( + pipe._component_specs["vae"].pretrained_model_name_or_path == "hf-internal-testing/tiny-anima-modular-pipe" + ) + assert pipe._component_specs["text_encoder"].pretrained_model_name_or_path == local_dir + + def test_local_copy_loads_components_at_root(self, tmp_path): + """A component recorded without a subfolder is at the root of its repo; when the local copy has its files + there it is loaded from the copy: weights for a model, the saved config file for anything else.""" + local_dir = str(tmp_path / "local-copy") + snapshot_download("hf-internal-testing/tiny-cosmos3-modular-pipe", local_dir=local_dir) + index_path = os.path.join(local_dir, "modular_model_index.json") + with open(index_path) as f: + index = json.load(f) + root_components = ["transformer", "scheduler", "text_tokenizer"] + for name in root_components: + for filename in os.listdir(os.path.join(local_dir, name)): + shutil.move(os.path.join(local_dir, name, filename), os.path.join(local_dir, filename)) + index[name][2]["subfolder"] = None + with open(index_path, "w") as f: + json.dump(index, f) + + pipe = ModularPipeline.from_pretrained(local_dir) + for name in root_components: + assert pipe._component_specs[name].pretrained_model_name_or_path == local_dir, f"{name} not local" + pipe.load_components(names=root_components, dtype=torch.float32, local_files_only=True) + for name in root_components: + assert getattr(pipe, name) is not None, f"{name} did not load from the local copy" diff --git a/tests/modular_pipelines/testing_utils/utils.py b/tests/modular_pipelines/testing_utils/utils.py index 32a69f3dd0d7..8b2f1e71ac26 100644 --- a/tests/modular_pipelines/testing_utils/utils.py +++ b/tests/modular_pipelines/testing_utils/utils.py @@ -21,7 +21,7 @@ import torch from huggingface_hub import hf_hub_download -from ...testing_utils import torch_device +from diffusers.utils.torch_utils import TorchDeviceBackend def backend_memory_allocated(device: str) -> int: @@ -37,15 +37,13 @@ def backend_memory_allocated(device: str) -> int: def patch_free_memory(free_bytes: int, total_bytes: int = 80 * 1024): """ - Simulate `free_bytes` of free device memory on whichever backend module (cuda/xpu/...) backs `torch_device`. + Simulate `free_bytes` of free device memory on whichever backend backs `torch_device`. - `mem_get_info` returns `(free, total)` and is the single point where `AutoOffloadStrategy` learns how much memory - is available, so patching it makes offloading decisions deterministic instead of dependent on the real free memory - of the test hardware (an 80GB GPU never runs low on a handful of KB-sized models). + `TorchDeviceBackend.mem_get_info` returns `(free, total)` and is the single point where `AutoOffloadStrategy` learns how + much memory is available, so patching it makes offloading decisions deterministic instead of dependent on the real + free memory of the test hardware (an 80GB GPU never runs low on a handful of KB-sized models). """ - device_type = torch.device(torch_device).type - device_module = getattr(torch, device_type, torch.cuda) - return mock.patch.object(device_module, "mem_get_info", return_value=(free_bytes, total_bytes)) + return mock.patch.object(TorchDeviceBackend, "mem_get_info", return_value=(free_bytes, total_bytes)) def get_specified_components(path_or_repo_id, cache_dir=None): diff --git a/tests/others/test_utils.py b/tests/others/test_utils.py index d1a59cec52f1..bb7e5c298926 100755 --- a/tests/others/test_utils.py +++ b/tests/others/test_utils.py @@ -18,11 +18,13 @@ import warnings import pytest +import torch from diffusers import __version__ -from diffusers.utils import deprecate +from diffusers.utils import deprecate, torch_utils +from diffusers.utils.torch_utils import TorchDeviceBackend, empty_device_cache, get_device -from ..testing_utils import Expectations, str_to_bool +from ..testing_utils import CaptureLogger, Expectations, str_to_bool # Used to test the hub @@ -269,6 +271,63 @@ def _capture(target_device): assert "moved to" in cuda_out, f"Non-MPS target should still emit the CPU-fallback info log, got: {cuda_out}" +class TestTorchDeviceBackend: + """Tests for :class:`diffusers.utils.torch_utils.TorchDeviceBackend` on the CPU backend.""" + + def test_module_resolves_from_str_and_torch_device(self): + assert TorchDeviceBackend("cpu").module is torch.cpu, "str device type should resolve to torch.cpu" + assert TorchDeviceBackend(torch.device("cpu")).module is torch.cpu, "torch.device should resolve to torch.cpu" + assert TorchDeviceBackend("cuda:1").module is torch.cuda, "device index should be ignored for the module" + + def test_get_device_matches_torch_accelerator(self): + expected = torch.accelerator.current_accelerator().type if torch.accelerator.is_available() else "cpu" + assert get_device() == expected, "get_device() should report what torch.accelerator reports" + + def test_default_device_is_the_detected_accelerator(self): + backend = TorchDeviceBackend() + assert backend.device == torch.device(get_device()), "no-arg backend should target get_device()" + assert backend.module is getattr(torch, get_device()), "module should be the torch namespace of that device" + + def test_unknown_backend_raises_instead_of_falling_back_to_cuda(self): + with pytest.raises(RuntimeError, match="does not have a corresponding module"): + TorchDeviceBackend("privateuseone") + + def test_cpu_operations_without_a_caching_allocator(self): + empty_device_cache("cpu") + backend = TorchDeviceBackend("cpu") + backend.empty_cache() + backend.manual_seed(1234) + assert torch.initial_seed() == 1234, "cpu manual_seed should fall back to torch.manual_seed" + + def test_function_style_backend_helpers_on_cpu(self): + torch_utils.backend_manual_seed("cpu", 0) + torch_utils.backend_synchronize("cpu") + torch_utils.backend_empty_cache("cpu") + assert isinstance(torch_utils.backend_device_count("cpu"), int) + assert torch_utils.backend_supports_training("cpu") is True + assert torch_utils.backend_supports_training("mps") is False + + def test_unsupported_operations_raise_naming_the_backend(self): + backend = TorchDeviceBackend("cpu") + with pytest.raises(NotImplementedError, match="torch.cpu"): + backend.mem_get_info() + assert backend.Stream is torch.cpu.Stream, "undefined attributes forward to the backend module" + backend.synchronize() + with pytest.raises(AttributeError, match="Stream"): + TorchDeviceBackend("mps").Stream() + with CaptureLogger(torch_utils.logger) as cl: + backend.reset_peak_memory_stats() + assert backend.max_memory_allocated() == 0, "cpu keeps no memory statistics" + assert "no memory statistics" in cl.out, f"no-op memory calls should warn, got: {cl.out}" + + def test_memory_statistics_on_the_host_accelerator(self): + if get_device() == "cpu": + pytest.skip("memory statistics need an accelerator") + backend = TorchDeviceBackend() + backend.reset_peak_memory_stats() + assert backend.max_memory_allocated() >= 0 + + # Copied from https://github.com/huggingface/transformers/blob/main/tests/utils/test_expectations.py class TestExpectations: def test_expectations(self): diff --git a/tests/testing_utils.py b/tests/testing_utils.py index 9da89e626198..cce6f15325b4 100644 --- a/tests/testing_utils.py +++ b/tests/testing_utils.py @@ -52,6 +52,7 @@ is_transformers_available, ) from diffusers.utils.logging import get_logger +from diffusers.utils.torch_utils import TorchDeviceBackend if is_torch_available(): @@ -1505,107 +1506,38 @@ def _is_torch_fp64_available(device): else None ) - # Function definitions - BACKEND_EMPTY_CACHE = { - "cuda": torch.cuda.empty_cache, - "xpu": torch.xpu.empty_cache, - "cpu": None, - "mps": torch.mps.empty_cache, - "default": None, - } - BACKEND_DEVICE_COUNT = { - "cuda": torch.cuda.device_count, - "xpu": torch.xpu.device_count, - "cpu": lambda: 0, - "mps": lambda: 0, - "default": 0, - } - BACKEND_MANUAL_SEED = { - "cuda": torch.cuda.manual_seed, - "xpu": torch.xpu.manual_seed, - "cpu": torch.manual_seed, - "mps": torch.mps.manual_seed, - "default": torch.manual_seed, - } - BACKEND_RESET_PEAK_MEMORY_STATS = { - "cuda": torch.cuda.reset_peak_memory_stats, - "xpu": getattr(torch.xpu, "reset_peak_memory_stats", None), - "cpu": None, - "mps": None, - "default": None, - } - BACKEND_RESET_MAX_MEMORY_ALLOCATED = { - "cuda": torch.cuda.reset_max_memory_allocated, - "xpu": getattr(torch.xpu, "reset_peak_memory_stats", None), - "cpu": None, - "mps": None, - "default": None, - } - BACKEND_MAX_MEMORY_ALLOCATED = { - "cuda": torch.cuda.max_memory_allocated, - "xpu": getattr(torch.xpu, "max_memory_allocated", None), - "cpu": 0, - "mps": 0, - "default": 0, - } - BACKEND_SYNCHRONIZE = { - "cuda": torch.cuda.synchronize, - "xpu": getattr(torch.xpu, "synchronize", None), - "cpu": None, - "mps": None, - "default": None, - } - if _neuron_device is not None: - BACKEND_EMPTY_CACHE[_neuron_device] = None - BACKEND_DEVICE_COUNT[_neuron_device] = torch.neuron.device_count - BACKEND_MANUAL_SEED[_neuron_device] = torch.manual_seed - BACKEND_RESET_PEAK_MEMORY_STATS[_neuron_device] = None - BACKEND_RESET_MAX_MEMORY_ALLOCATED[_neuron_device] = None - BACKEND_MAX_MEMORY_ALLOCATED[_neuron_device] = 0 - BACKEND_SYNCHRONIZE[_neuron_device] = torch.neuron.synchronize BACKEND_SUPPORTS_TRAINING[_neuron_device] = False -# This dispatches a defined function according to the accelerator from the function definitions. -def _device_agnostic_dispatch(device: str, dispatch_table: dict[str, Callable], *args, **kwargs): - fn = dispatch_table[device] if device in dispatch_table else dispatch_table["default"] - - # Some device agnostic functions return values. Need to guard against 'None' instead at - # user level - if not callable(fn): - return fn - - return fn(*args, **kwargs) - - -# These are callables which automatically dispatch the function specific to the accelerator +# Device operations go through `TorchDeviceBackend`. def backend_manual_seed(device: str, seed: int): - return _device_agnostic_dispatch(device, BACKEND_MANUAL_SEED, seed) + TorchDeviceBackend(device).manual_seed(seed) def backend_synchronize(device: str): - return _device_agnostic_dispatch(device, BACKEND_SYNCHRONIZE) + TorchDeviceBackend(device).synchronize() def backend_empty_cache(device: str): - return _device_agnostic_dispatch(device, BACKEND_EMPTY_CACHE) + TorchDeviceBackend(device).empty_cache() def backend_device_count(device: str): - return _device_agnostic_dispatch(device, BACKEND_DEVICE_COUNT) + return TorchDeviceBackend(device).device_count() def backend_reset_peak_memory_stats(device: str): - return _device_agnostic_dispatch(device, BACKEND_RESET_PEAK_MEMORY_STATS) + TorchDeviceBackend(device).reset_peak_memory_stats() def backend_reset_max_memory_allocated(device: str): - return _device_agnostic_dispatch(device, BACKEND_RESET_MAX_MEMORY_ALLOCATED) + # `reset_max_memory_allocated` is CUDA's deprecated alias of `reset_peak_memory_stats`. + TorchDeviceBackend(device).reset_peak_memory_stats() def backend_max_memory_allocated(device: str): - return _device_agnostic_dispatch(device, BACKEND_MAX_MEMORY_ALLOCATED) + return TorchDeviceBackend(device).max_memory_allocated() # These are callables which return boolean behaviour flags and can be used to specify some @@ -1659,14 +1591,9 @@ def update_mapping_from_spec(device_fn_dict: dict[str, Callable], attribute_name torch_device = device_name - # Add one entry here for each `BACKEND_*` dictionary. - update_mapping_from_spec(BACKEND_MANUAL_SEED, "MANUAL_SEED_FN") - update_mapping_from_spec(BACKEND_EMPTY_CACHE, "EMPTY_CACHE_FN") - update_mapping_from_spec(BACKEND_DEVICE_COUNT, "DEVICE_COUNT_FN") + # `SUPPORTS_TRAINING` is the only per-device table left. Device operations come from the backend's own + # `torch.` module through `TorchDeviceBackend`, so a spec file no longer supplies them. update_mapping_from_spec(BACKEND_SUPPORTS_TRAINING, "SUPPORTS_TRAINING") - update_mapping_from_spec(BACKEND_RESET_PEAK_MEMORY_STATS, "RESET_PEAK_MEMORY_STATS_FN") - update_mapping_from_spec(BACKEND_RESET_MAX_MEMORY_ALLOCATED, "RESET_MAX_MEMORY_ALLOCATED_FN") - update_mapping_from_spec(BACKEND_MAX_MEMORY_ALLOCATED, "MAX_MEMORY_ALLOCATED_FN") # Modified from https://github.com/huggingface/transformers/blob/cdfb018d0300fef3b07d9220f3efe9c2a9974662/src/transformers..testing_utils.py#L3090 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())