From acdf4bf9cbdef5345535bf1b53359983c226a9cd Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Thu, 20 Aug 2026 15:10:26 +0000 Subject: [PATCH 01/22] [core] Shard tensor-parallel checkpoints on load and save Stream each rank's slice of a tensor-parallel checkpoint straight off disk instead of materializing the full checkpoint on every rank and resharding it afterwards, and gather the shards back on save. - `from_pretrained(..., parallel_config=TensorParallelConfig(...))` resolves the shard specs on the still-meta model, then slices each safetensors tensor before the dtype cast, so host memory peaks at ~1/tp_degree of the checkpoint. - `save_pretrained` all-gathers the DTensors into an ordinary checkpoint, or writes a distributed checkpoint with `dcp=True` so no full tensor is ever formed. The writing `tp_degree` is recorded, since a packed weight's stored layout is interleaved by it. - Factor the plan interpretation out of the Neuron pre-shard path into shared `TPShardSpec` / `resolve_tp_shard_specs` / `_local_shard` / `_hooks_only_styles` helpers, so both backends and both the load and save paths shard identically. --- .../en/training/distributed_inference.md | 60 ++- src/diffusers/hooks/tensor_parallel.py | 261 +++++++-- src/diffusers/hooks/tensor_parallel_neuron.py | 177 ++---- src/diffusers/models/_modeling_parallel.py | 2 + src/diffusers/models/model_loading_utils.py | 115 +++- src/diffusers/models/modeling_utils.py | 502 +++++++++++++++--- src/diffusers/utils/__init__.py | 1 + src/diffusers/utils/constants.py | 1 + tests/models/testing_utils/parallelism.py | 202 +++++++ 9 files changed, 1033 insertions(+), 288 deletions(-) diff --git a/docs/source/en/training/distributed_inference.md b/docs/source/en/training/distributed_inference.md index 856572c2ff08..25d4dba4e339 100644 --- a/docs/source/en/training/distributed_inference.md +++ b/docs/source/en/training/distributed_inference.md @@ -436,43 +436,42 @@ 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. ```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 +483,25 @@ 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. +A tensor-parallel `parallel_config` cannot be combined with `device_map`, `quantization_config`, `low_cpu_mem_usage=False`, `use_flashpack=True`, DDUF checkpoints, or non-safetensors weights; each raises rather than quietly falling back to loading the full checkpoint. 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 + +[`~ModelMixin.save_pretrained`] gathers the shards back into ordinary full tensors, so the result is a normal checkpoint that loads with or without tensor parallelism. Gathering is a collective, so call it on **every** rank; only rank 0 writes. + +```py +# on all ranks +pipeline.transformer.save_pretrained("flux2-transformer") +``` + +For a model too large to gather onto a single rank, pass `dcp=True` to write a [distributed checkpoint](https://pytorch.org/docs/stable/distributed.checkpoint.html) instead. Every rank writes its own shards, so no full tensor is ever formed. + +```py +pipeline.transformer.save_pretrained("flux2-transformer-dcp", dcp=True) +``` + +`from_pretrained` detects such a directory automatically, and reads it back with the same `parallel_config` you saved it under. Because a packed projection's shards are stored interleaved by the writing degree, the checkpoint only loads at that same `tp_degree`, and only with tensor parallelism — anything else raises rather than silently returning wrong weights. It is also local-only: a distributed checkpoint is recognized by the `.metadata` file in its directory, so it cannot be pushed to or loaded from the Hub. To lift any of these restrictions, re-save with the default (gathered) path, which produces an ordinary checkpoint. + ### 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. diff --git a/src/diffusers/hooks/tensor_parallel.py b/src/diffusers/hooks/tensor_parallel.py index b90a5761d043..d3cbd82d7980 100644 --- a/src/diffusers/hooks/tensor_parallel.py +++ b/src/diffusers/hooks/tensor_parallel.py @@ -12,6 +12,8 @@ # 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 @@ -65,6 +67,150 @@ 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: + # An uneven split is rejected rather than silently handed to `Shard`, which pads the tail + # and would break both the paired colwise/rowwise matmul and `_unshard_gathered`. + if block_size % tp_size != 0: + raise ValueError( + f"Cannot shard a block of size {block_size} across {tp_size} tensor-parallel ranks: " + f"{block_size} is not divisible by {tp_size}." + ) + 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 _unshard_gathered(gathered: torch.Tensor, dim: int, block_sizes: "list[int]", tp_size: int) -> torch.Tensor: + """Undo `_local_shard`'s block interleaving on an all-gathered tensor. + + `DTensor.full_tensor()` concatenates the local shards rank-major, so a packed weight comes back as `[block0_rank0, + block1_rank0, block0_rank1, block1_rank1, ...]` and has to be regrouped by block. A single block is already in the + original order and passes through unchanged. + """ + if len(block_sizes) == 1: + return gathered + + local_sizes = [block_size // tp_size for block_size in block_sizes] + stride = sum(local_sizes) + parts = [] + for i, local_size in enumerate(local_sizes): + offset = sum(local_sizes[:i]) + parts.extend(gathered.narrow(dim, rank * stride + offset, local_size) for rank in range(tp_size)) + return torch.cat(parts, dim=dim) + + +def gather_tp_state_dict(state_dict: dict, specs: "dict[str, TPShardSpec]", config: TensorParallelConfig) -> dict: + """Reassemble a tensor-parallel `state_dict` into ordinary full tensors. + + Every `DTensor` is all-gathered back to its full shape and, for the packed styles, reordered by `_unshard_gathered` + — `full_tensor()` alone would leave the fused blocks interleaved by rank. Replicated and unplanned parameters pass + through untouched. + + `full_tensor()` is a collective, so this must run on **every** rank even though usually only rank 0 goes on to + write the result. + """ + from torch.distributed.tensor import DTensor + + tp_size = config._tp_degree + gathered = {} + for key, value in state_dict.items(): + if not isinstance(value, DTensor): + gathered[key] = value + continue + # `full_tensor()` is the collective; the reorder after it is plain tensor arithmetic, so keep it off + # the accelerator — CPU is where this state dict is headed anyway, since it is about to be written. + full = value.full_tensor().cpu() + spec = specs[key] + if spec.dim is not None: + full = _unshard_gathered(full, spec.dim, spec.block_sizes, tp_size) + gathered[key] = full.contiguous() + return gathered + + +def resolve_tp_shard_specs(model: torch.nn.Module, tp_plan: dict) -> "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. + + 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": + weight_spec = TPShardSpec(1, [submodule.weight.shape[1]]) + bias_spec = TPShardSpec(None, None) + elif isinstance(style, PackedColwiseParallel): + blocks = style.blocks if style.blocks is not None else submodule._tp_packed_col_blocks + weight_spec = TPShardSpec(0, _blocks_to_block_sizes(submodule.weight.shape[0], blocks)) + bias_spec = weight_spec + elif isinstance(style, PackedRowwiseParallel): + blocks = style.blocks if style.blocks is not None else submodule._tp_packed_row_blocks + 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, or PackedRowwiseParallel." + ) + + specs[f"{path}.weight"] = weight_spec + if submodule.bias is not None: + specs[f"{path}.bias"] = bias_spec + + 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. @@ -123,32 +269,24 @@ def _make_packed_col(marker: PackedColwiseParallel) -> ColwiseParallel: 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 = _blocks if _blocks is not None else module._tp_packed_col_blocks # 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(): + # Replicate 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. 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, + local = _local_shard(full, 0, _blocks_to_block_sizes(full.shape[0], blocks), device_mesh) + module.register_parameter( + param_name, + 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) return _PackedColwiseImpl() @@ -157,26 +295,14 @@ def _make_packed_row(marker: PackedRowwiseParallel) -> RowwiseParallel: 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 = _blocks if _blocks is not None else module._tp_packed_row_blocks for param_name, param in module.named_parameters(): if param_name == "weight": + # See `_make_packed_col`: replicate first so one rank's weights win. 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) + local = _local_shard(full, 1, _blocks_to_block_sizes(full.shape[1], blocks), device_mesh) dist_param = nn.Parameter( DTensor.from_local(local, device_mesh, [Shard(1)], run_check=False), requires_grad=param.requires_grad, @@ -193,7 +319,7 @@ def _partition_linear_fn(self, name, module, device_mesh): # `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. + # matching what `_local_shard` already does for the packed styles. def _make_checked_col(path: str) -> ColwiseParallel: class _CheckedColwiseImpl(ColwiseParallel): def _partition_linear_fn(self, name, module, device_mesh): @@ -240,16 +366,68 @@ def _partition_linear_fn(self, name, module, device_mesh): 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.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() + else: + raise ValueError( + f"Unsupported tensor-parallel style '{style}' for '{path}'. " + f"Expected 'colwise', 'rowwise', PackedColwiseParallel, or PackedRowwiseParallel." + ) + return resolved + + 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`. + + Set `weights_already_sharded` when the planned parameters are already `DTensor` shards, as they are after a + streaming `from_pretrained` load; only the forward hooks are then 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. + """ + if tp_plan is None: + raise ValueError( + "`_tp_plan` must be set on the model class to use tensor parallelism. " + f"'{model.__class__.__name__}' does not define one." + ) + 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.") + num_heads = getattr(model.config, "num_attention_heads", None) + if num_heads is not None and num_heads % config._tp_degree != 0: + raise ValueError(f"`tp_degree` ({config._tp_degree}) must divide the number of attention heads ({num_heads}).") + if tp_mesh.device_type not in _SUPPORTED_TP_DEVICES: raise ValueError( f"Tensor parallelism is not supported on device type '{tp_mesh.device_type}'. Supported device types are " @@ -261,13 +439,18 @@ def apply_tensor_parallel( 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 + 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, resolve_tp_shard_specs(model, tp_plan)) return - from torch.distributed.tensor.parallel import parallelize_module - 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..ffcba0973d81 100644 --- a/src/diffusers/hooks/tensor_parallel_neuron.py +++ b/src/diffusers/hooks/tensor_parallel_neuron.py @@ -17,162 +17,59 @@ 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, _hooks_only_styles, _local_shard 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() + from torch.distributed.tensor import DTensor, Replicate, Shard + from torch.distributed.tensor.parallel import parallelize_module + device = torch.neuron.current_device() + + 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: - _pre_shard_and_tp(block, tp_mesh, relative_plan, rank, tp_size) + parallelize_module(block, tp_mesh, _hooks_only_styles(relative_plan)) diff --git a/src/diffusers/models/_modeling_parallel.py b/src/diffusers/models/_modeling_parallel.py index 86627284e078..b54e86d6b4f2 100644 --- a/src/diffusers/models/_modeling_parallel.py +++ b/src/diffusers/models/_modeling_parallel.py @@ -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 diff --git a/src/diffusers/models/model_loading_utils.py b/src/diffusers/models/model_loading_utils.py index abbde8082bb5..d0ba37514b9e 100644 --- a/src/diffusers/models/model_loading_utils.py +++ b/src/diffusers/models/model_loading_utils.py @@ -388,6 +388,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, @@ -452,28 +545,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 5af0ca0e6278..e690fff0b058 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -42,6 +42,7 @@ from ..quantizers.quantization_config import QuantizationMethod from ..utils import ( CONFIG_NAME, + DCP_CONFIG_NAME, FLASHPACK_WEIGHTS_NAME, HF_ENABLE_PARALLEL_LOADING, SAFE_WEIGHTS_INDEX_NAME, @@ -76,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, ) @@ -686,6 +688,7 @@ def save_pretrained( max_shard_size: int | str = "10GB", push_to_hub: bool = False, use_flashpack: bool = False, + dcp: bool = False, **kwargs, ): """ @@ -718,8 +721,18 @@ def save_pretrained( Whether or not to push your model to the Hugging Face Hub after saving it. You can specify the repository you want to push to with `repo_id` (will default to the name of `save_directory` in your namespace). + dcp (`bool`, *optional*, defaults to `False`): + Write a [`torch.distributed.checkpoint`](https://pytorch.org/docs/stable/distributed.checkpoint.html) + directory instead of safetensors files. Only valid for a tensor-parallel model: every rank writes its + own shards, so no full tensor is ever materialized, which matters for models too large to gather onto + one rank. Read it back with `from_pretrained`, which detects the directory automatically and can + reshard it to a different `tp_degree`. kwargs (`dict[str, Any]`, *optional*): Additional keyword arguments passed along to the [`~utils.PushToHubMixin.push_to_hub`] method. + + A tensor-parallel model is gathered back into ordinary full tensors before saving, so the result is a normal + checkpoint that loads without tensor parallelism. Gathering is a collective: call `save_pretrained` on every + rank, not just the main process. Only rank 0 writes. """ if os.path.isfile(save_directory): logger.error(f"Provided path ({save_directory}) should be a directory, not a file") @@ -742,6 +755,76 @@ def save_pretrained( " the logger on the traceback to understand the reason why the quantized model is not serializable." ) + tp_config = None + if self._parallel_config is not None: + tp_config = self._parallel_config.tensor_parallel_config + + if dcp: + if tp_config is None: + raise ValueError( + "`dcp=True` is only meaningful for a tensor-parallel model, whose parameters are sharded " + "across ranks. Save an unsharded model with the default safetensors path." + ) + unsupported = [ + name + for name, value in ( + ("use_flashpack", use_flashpack), + ("variant", variant), + ("safe_serialization=False", not safe_serialization), + ("save_function", save_function), + ) + if value + ] + if unsupported: + raise ValueError( + f"{unsupported} cannot be combined with `dcp=True`: a distributed checkpoint is a directory " + "of `.distcp` shards, not a single named weights file." + ) + if push_to_hub: + # `from_pretrained` only recognizes a distributed checkpoint by looking for `.metadata` in a + # local directory, so one cannot be loaded back from the Hub. + raise ValueError( + "`push_to_hub=True` cannot be combined with `dcp=True`: a distributed checkpoint can only " + "be loaded from a local directory. Save it with the default safetensors path to push it." + ) + + import torch.distributed.checkpoint as dcp_api + + os.makedirs(save_directory, exist_ok=True) + if tp_config._mesh.get_local_rank() == 0: + self.save_config(save_directory) + # A packed weight's local shard is `cat(block_0_shard, block_1_shard, ...)`, which DTensor — + # and therefore DCP — records as plain chunk `rank` of the global tensor. The stored layout is + # thus interleaved by the saving `tp_degree`, so the checkpoint can only be read back at that + # same degree. Record it so a mismatch fails clearly instead of silently loading garbage. + with open(os.path.join(save_directory, DCP_CONFIG_NAME), "w", encoding="utf-8") as f: + json.dump({"tp_degree": tp_config._tp_degree}, f, indent=2) + # Written from the sharded state dict, so no rank ever holds a full tensor. Collective, so every + # rank takes part. + dcp_api.save(self.state_dict(), checkpoint_id=save_directory) + logger.info(f"Distributed checkpoint saved in {save_directory}") + return + + # Under tensor parallelism the parameters are DTensor shards, so they have to be gathered before + # anything can be written. `state_dict()` is read here rather than further down because the gather is + # a collective: every rank must reach it, while only rank 0 may go on to touch the filesystem or the + # Hub. Non-TP saves keep the original ordering. + state_dict = None + if tp_config is not None: + if use_flashpack: + raise ValueError( + "`use_flashpack=True` is not supported for a tensor-parallel model. Save it with " + "`safe_serialization=True`, or use `dcp=True` to write a sharded checkpoint." + ) + from ..hooks.tensor_parallel import gather_tp_state_dict, resolve_tp_shard_specs + + state_dict = gather_tp_state_dict( + self.state_dict(), resolve_tp_shard_specs(self, self._tp_plan), tp_config + ) + if tp_config._mesh.get_local_rank() != 0: + # `is_main_process` defaults to True on every rank, so it cannot be used for this. + return + weights_name = WEIGHTS_NAME if use_flashpack: weights_name = FLASHPACK_WEIGHTS_NAME @@ -772,7 +855,8 @@ def save_pretrained( model_to_save.save_config(save_directory) # Save the model - state_dict = model_to_save.state_dict() + if state_dict is None: + state_dict = model_to_save.state_dict() quantization_metadata = {} if hf_quantizer is not None: state_dict, quantization_metadata = hf_quantizer.get_state_dict_and_metadata( @@ -1037,7 +1121,9 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None quantization_config = kwargs.pop("quantization_config", None) dduf_entries: dict[str, DDUFEntry] | None = kwargs.pop("dduf_entries", None) disable_mmap = kwargs.pop("disable_mmap", False) - parallel_config: ParallelConfig | ContextParallelConfig | None = kwargs.pop("parallel_config", None) + parallel_config: ParallelConfig | ContextParallelConfig | TensorParallelConfig | None = kwargs.pop( + "parallel_config", None + ) use_flashpack = kwargs.pop("use_flashpack", False) flashpack_kwargs = kwargs.pop("flashpack_kwargs", {}) @@ -1150,6 +1236,35 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None # no in-place modification of the original config. config = copy.deepcopy(config) + # A `torch.distributed.checkpoint` directory written by `save_pretrained(..., dcp=True)` holds + # `.distcp` shards rather than safetensors, so it bypasses the checkpoint-file resolution below. + if os.path.isdir(pretrained_model_name_or_path): + dcp_dir = os.path.join(pretrained_model_name_or_path, subfolder or "") + if os.path.isfile(os.path.join(dcp_dir, ".metadata")): + # Checked here rather than in `_load_dcp_checkpoint` because this branch returns before the + # quantizer is built and before `_check_tp_streaming_supported` runs, so nothing else would + # look at these. + unsupported = [ + name + for name, value in ( + ("device_map", device_map), + ("quantization_config", quantization_config), + ("use_flashpack", use_flashpack), + ("variant", variant), + ("dduf_entries", dduf_entries), + ("low_cpu_mem_usage=False", not low_cpu_mem_usage), + ) + if value + ] + if unsupported: + raise ValueError( + f"{unsupported} cannot be combined with the distributed checkpoint at {dcp_dir}: its " + "shards are read in place onto each rank's device." + ) + return cls._load_dcp_checkpoint( + dcp_dir, config, unused_kwargs, torch_dtype=torch_dtype, parallel_config=parallel_config + ) + # determine initial quantization config. ####################################### pre_quantized = "quantization_config" in config and config["quantization_config"] is not None @@ -1204,6 +1319,27 @@ 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: + cls._check_tp_streaming_supported( + device_map=device_map, + low_cpu_mem_usage=low_cpu_mem_usage, + use_flashpack=use_flashpack, + hf_quantizer=hf_quantizer, + dduf_entries=dduf_entries, + ) + is_sharded = False resolved_model_file = None @@ -1320,6 +1456,27 @@ 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 + + non_safetensors = [f for f in resolved_model_file 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}." + ) + + 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) + # 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. + is_parallel_loading_enabled = False + if use_flashpack: if is_flashpack_available(): import flashpack @@ -1362,7 +1519,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, dduf_entries=dduf_entries) # We only fix it for non sharded checkpoints as we don't need it yet for sharded one. @@ -1370,6 +1527,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()) @@ -1418,6 +1582,8 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None dduf_entries=dduf_entries, 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, @@ -1457,7 +1623,13 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None # Set model in evaluation mode to deactivate DropOut modules by default model.eval() - if parallel_config is not None: + 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 ..hooks.tensor_parallel import apply_tensor_parallel + + apply_tensor_parallel(model, tp_config, cls._tp_plan, weights_already_sharded=True) + elif parallel_config is not None: model.enable_parallelism(config=parallel_config) if output_loading_info: @@ -1604,25 +1776,182 @@ def compile_repeated_blocks(self, *args, **kwargs): f"Regional compilation failed because {repeated_blocks} classes are not found in the model. " ) - def enable_parallelism( - self, + @classmethod + def _load_dcp_checkpoint( + cls, + checkpoint_dir: str, + config: dict, + unused_kwargs: dict, *, - config: ParallelConfig | ContextParallelConfig | TensorParallelConfig, - cp_plan: dict[str, ContextParallelModelPlan] | None = None, + torch_dtype: torch.dtype | None, + parallel_config: ParallelConfig | ContextParallelConfig | TensorParallelConfig | None, ): - logger.warning( - "`enable_parallelism` is an experimental feature. The API may change in the future and breaking changes may be introduced at any time without warning." - ) + """Load a `torch.distributed.checkpoint` directory written by `save_pretrained(..., dcp=True)`. + + The shards are those of a tensor-parallel model, so a tensor-parallel `parallel_config` is required, at the + `tp_degree` the checkpoint was written with — see the note where it is written. Use the ordinary safetensors + path to move a model between degrees; it streams each rank's slice, so it costs no more memory than this + does. + + DCP loads **in place**, so every parameter has to be allocated first with its local shape and on the device it + will end up on. + """ + import torch.distributed.checkpoint as dcp + from torch.distributed.tensor import DTensor, Replicate, Shard + + from ..hooks.tensor_parallel import apply_tensor_parallel, resolve_tp_shard_specs + + with open(os.path.join(checkpoint_dir, DCP_CONFIG_NAME), encoding="utf-8") as f: + saved_tp_degree = json.load(f)["tp_degree"] + + with ContextManagers([no_init_weights(), accelerate.init_empty_weights()]): + model = cls.from_config(config, **unused_kwargs) + + 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 None: + raise ValueError( + f"The distributed checkpoint at {checkpoint_dir} holds the shards of a tensor-parallel model, so " + f"it can only be read back with a tensor-parallel `parallel_config` of `tp_degree=" + f"{saved_tp_degree}`. To load it without tensor parallelism, re-save the model with " + f"`save_pretrained(...)`, which gathers the shards into ordinary safetensors." + ) + # An explicit `mesh` overrides `tp_degree` (see `TensorParallelConfig`), and `_tp_degree` is only set + # by `setup()`, which has not run yet — so the effective degree has to be resolved by hand here. + requested_tp_degree = tp_config.mesh.size() if tp_config.mesh is not None else tp_config.tp_degree + if requested_tp_degree != saved_tp_degree: + raise ValueError( + f"The distributed checkpoint at {checkpoint_dir} was written with `tp_degree={saved_tp_degree}` " + f"and can only be loaded with the same degree, but {requested_tp_degree} was requested. Packed " + f"projections are stored interleaved by the writing degree, so reading at another degree would " + f"silently produce wrong weights. To change degree, re-save the model with " + f"`save_pretrained(...)` (which gathers to ordinary safetensors) and load that with " + f"`from_pretrained(..., parallel_config=...)`." + ) + 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_mesh = tp_config._mesh + device = torch.neuron.current_device() if tp_mesh.device_type == "neuron" else tp_config._device + + for name, meta_param in model.state_dict().items(): + dtype = torch_dtype if torch_dtype is not None and meta_param.is_floating_point() else meta_param.dtype + spec = tp_shard_specs.get(name) + if spec is None or spec.dim is None: + local = torch.empty(meta_param.shape, dtype=dtype, device=device) + else: + shape = list(meta_param.shape) + shape[spec.dim] //= tp_config._tp_degree + local = torch.empty(shape, dtype=dtype, device=device) + + module_path, _, param_name = name.rpartition(".") + module = model.get_submodule(module_path) if module_path else model + if spec is None: + value = local + else: + placement = Replicate() if spec.dim is None else Shard(spec.dim) + value = DTensor.from_local(local, tp_mesh, [placement], run_check=False) + if param_name in module._buffers: + module._buffers[param_name] = value + else: + module.register_parameter(param_name, torch.nn.Parameter(value, requires_grad=False)) - if not torch.distributed.is_available() and not torch.distributed.is_initialized(): + state_dict = model.state_dict() + dcp.load(state_dict, checkpoint_id=checkpoint_dir) + + # `dcp.load` silently does nothing for a parameter left on `meta`, so a mistake above would + # otherwise produce a model of uninitialized weights with no diagnostic at all. + still_meta = sorted(name for name, value in state_dict.items() if value.device.type == "meta") + if still_meta: raise RuntimeError( - "torch.distributed must be available and initialized before calling `enable_parallelism`." + f"Loading the distributed checkpoint at {checkpoint_dir} left these parameters on the meta " + f"device: {still_meta}." ) - from ..hooks.context_parallel import apply_context_parallel - from .attention import AttentionModuleMixin - from .attention_dispatch import AttentionBackendName, _AttentionBackendRegistry - from .attention_processor import Attention, MochiAttention + # 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. + for name, buffer in model.named_buffers(): + if buffer.device != 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(device) + + model.register_to_config(_name_or_path=checkpoint_dir) + model.eval() + + apply_tensor_parallel(model, tp_config, cls._tp_plan, weights_already_sharded=True) + + return model + + @classmethod + def _check_tp_streaming_supported( + cls, + *, + device_map, + low_cpu_mem_usage: bool, + use_flashpack: bool, + hf_quantizer, + dduf_entries, + ) -> None: + """Reject the `from_pretrained` options that cannot be combined with a tensor-parallel load. + + Sharding on load needs a meta-initialized model and lazily sliceable safetensors files. Rather than silently + falling back to loading the full checkpoint and resharding it — which would quietly give up the memory saving + that is the whole point — each unsupported combination raises. + + Called 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 at the point where the resolved file list is + known. + """ + if cls._tp_plan is None: + raise ValueError( + f"`_tp_plan` must be set on the model class to use tensor parallelism. " + f"'{cls.__name__}' does not define one." + ) + 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`. Load the " + "model unquantized, or shard it after loading with `enable_parallelism`." + ) + 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 dduf_entries: + raise ValueError( + "DDUF checkpoints cannot be combined with a tensor-parallel `parallel_config`; their tensors " + "cannot be sliced per rank." + ) + + 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) @@ -1635,6 +1964,51 @@ def enable_parallelism( device_module = torch.get_device_module(device_type) device = torch.device(device_type, rank % device_module.device_count()) + 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 + return config + + def enable_parallelism( + self, + *, + config: ParallelConfig | ContextParallelConfig | TensorParallelConfig, + cp_plan: dict[str, ContextParallelModelPlan] | None = None, + ): + logger.warning( + "`enable_parallelism` is an experimental feature. The API may change in the future and breaking changes may be introduced at any time without warning." + ) + + 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 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." + ) + + config = self._resolve_parallel_config(config) + attention_classes = (Attention, MochiAttention, AttentionModuleMixin) if config.context_parallel_config is not None: @@ -1665,26 +2039,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. @@ -1705,16 +2059,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) @@ -1739,6 +2083,8 @@ def _load_pretrained_model( dduf_entries: dict[str, DDUFEntry] | 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()) @@ -1755,6 +2101,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: @@ -1790,25 +2147,38 @@ 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, - dduf_entries=dduf_entries, - loaded_keys=loaded_keys, - unexpected_keys=unexpected_keys, - offload_index=offload_index, - offload_folder=offload_folder, - state_dict_index=state_dict_index, - state_dict_folder=state_dict_folder, - ignore_mismatched_sizes=ignore_mismatched_sizes, - low_cpu_mem_usage=low_cpu_mem_usage, - disable_mmap=disable_mmap, - ) + if 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, + dduf_entries=dduf_entries, + loaded_keys=loaded_keys, + unexpected_keys=unexpected_keys, + offload_index=offload_index, + offload_folder=offload_folder, + state_dict_index=state_dict_index, + state_dict_folder=state_dict_folder, + ignore_mismatched_sizes=ignore_mismatched_sizes, + low_cpu_mem_usage=low_cpu_mem_usage, + disable_mmap=disable_mmap, + ) if is_parallel_loading_enabled: offload_index, state_dict_index, _mismatched_keys, _error_msgs = load_fn(resolved_model_file) diff --git a/src/diffusers/utils/__init__.py b/src/diffusers/utils/__init__.py index 0554d341022a..97ddb8a2589c 100644 --- a/src/diffusers/utils/__init__.py +++ b/src/diffusers/utils/__init__.py @@ -20,6 +20,7 @@ from .. import __version__ from .constants import ( CONFIG_NAME, + DCP_CONFIG_NAME, DEFAULT_HF_PARALLEL_LOADING_WORKERS, DEPRECATED_REVISION_ARGS, DIFFUSERS_DYNAMIC_MODULE_NAME, diff --git a/src/diffusers/utils/constants.py b/src/diffusers/utils/constants.py index fcf0e4518800..597ccd3eebd1 100644 --- a/src/diffusers/utils/constants.py +++ b/src/diffusers/utils/constants.py @@ -35,6 +35,7 @@ SAFETENSORS_FILE_EXTENSION = "safetensors" FLASHPACK_WEIGHTS_NAME = "model.flashpack" FLASHPACK_FILE_EXTENSION = "flashpack" +DCP_CONFIG_NAME = "dcp_config.json" GGUF_FILE_EXTENSION = "gguf" ONNX_EXTERNAL_WEIGHTS_NAME = "weights.pb" HUGGINGFACE_CO_RESOLVE_ENDPOINT = os.environ.get("HF_ENDPOINT", "https://huggingface.co") diff --git a/tests/models/testing_utils/parallelism.py b/tests/models/testing_utils/parallelism.py index 63575abf6b7b..c5174c396561 100644 --- a/tests/models/testing_utils/parallelism.py +++ b/tests/models/testing_utils/parallelism.py @@ -20,9 +20,11 @@ import torch import torch.distributed as dist import torch.multiprocessing as mp +from safetensors.torch import load_file from diffusers.models._modeling_parallel import ContextParallelConfig, TensorParallelConfig from diffusers.models.attention_dispatch import AttentionBackendName, _AttentionBackendRegistry +from diffusers.utils.constants import SAFETENSORS_WEIGHTS_NAME from ...testing_utils import ( is_attention, @@ -295,6 +297,108 @@ def _tensor_parallel_worker( dist.destroy_process_group() +def _tensor_parallel_from_pretrained_worker( + rank, world_size, master_port, model_class, checkpoint_dir, resave_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`, runs a forward + pass, and (if `resave_dir` is given) saves the model back out, which has to gather the shards first. 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() + + # Gathering is a collective, so every rank has to reach this even though only rank 0 writes. + if resave_dir is not None: + model.save_pretrained(resave_dir) + + 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() + + +def _tensor_parallel_dcp_worker( + rank, world_size, master_port, model_class, checkpoint_dir, dcp_dir, inputs_dict, return_dict +): + """Worker for the `save_pretrained(..., dcp=True)` round trip. + + Streams the checkpoint into shards, writes them as a distributed checkpoint (no rank ever holding a full + tensor), then loads that back and runs a forward pass. Rank 0 reports the output. + """ + 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 + + tp_config = TensorParallelConfig(tp_degree=world_size) + model_class.from_pretrained(checkpoint_dir, parallel_config=tp_config).save_pretrained(dcp_dir, dcp=True) + + reloaded = model_class.from_pretrained( + dcp_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 = reloaded(**inputs_on_device, return_dict=False)[0] + if isinstance(output, DTensor): + output = output.full_tensor() + + if rank == 0: + return_dict["status"] = "success" + 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 +448,104 @@ 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): + """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") + model.save_pretrained(checkpoint_dir) + + 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 + + def test_tensor_parallel_from_pretrained(self, tmp_path): + """`from_pretrained(..., parallel_config=...)` shards while reading, and `save_pretrained` gathers back. + + Covers both directions in one spawn: the streaming load must match the single-device reference, and the + checkpoint it writes back out must be byte-identical to the one it read. The round trip is what catches + a wrong packed-projection reorder — a plain colwise/rowwise mistake would pass the forward check alone. + """ + world_size = 2 + checkpoint_dir, inputs_dict, reference = self._tp_checkpoint_and_reference(tmp_path, world_size) + resave_dir = str(tmp_path / "resaved") + + 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, + resave_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) + + original = load_file(os.path.join(checkpoint_dir, SAFETENSORS_WEIGHTS_NAME)) + resaved = load_file(os.path.join(resave_dir, SAFETENSORS_WEIGHTS_NAME)) + assert original.keys() == resaved.keys() + for key, value in original.items(): + torch.testing.assert_close(resaved[key], value, atol=0, rtol=0, msg=lambda m, key=key: f"{key}: {m}") + + def test_tensor_parallel_dcp_roundtrip(self, tmp_path): + """`save_pretrained(..., dcp=True)` writes sharded and `from_pretrained` reads it back at the same degree.""" + world_size = 2 + checkpoint_dir, inputs_dict, reference = self._tp_checkpoint_and_reference(tmp_path, world_size) + + manager = mp.Manager() + return_dict = manager.dict() + mp.spawn( + _tensor_parallel_dcp_worker, + args=( + world_size, + _find_free_port(), + self.model_class, + checkpoint_dir, + str(tmp_path / "dcp"), + inputs_dict, + return_dict, + ), + nprocs=world_size, + join=True, + ) + assert return_dict.get("status") == "success", ( + f"Tensor parallel DCP round trip failed: {return_dict.get('error', 'Unknown error')}" + ) + torch.testing.assert_close(reference, torch.tensor(return_dict["output"]), atol=1e-3, rtol=1e-3) + @is_context_parallel @require_torch_multi_accelerator From 40ddb53dc0a7faa6bd43ba3bcfee39c76facec5d Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Thu, 20 Aug 2026 22:51:19 +0000 Subject: [PATCH 02/22] Raise when tensor parallelism is combined with quantization, offloading or LoRA MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Addresses the remaining two items of the review on #13718: tensor parallelism was rejected alongside quantization and `device_map` only on the `from_pretrained` streaming path, while `enable_parallelism` — which the quantization error message itself recommended — accepted a quantized, offloaded or adapter-injected model and sharded it anyway. - Add `_check_tp_model_state`, called from `apply_tensor_parallel`, the one chokepoint every TP entry point funnels through. It rejects a model that is quantized, group-offloaded, placed by accelerate (`device_map` or CPU offload), or has PEFT layers injected. Placed before the device-type check so the reported reason is the useful one. - Guard the reverse order too: `enable_group_offload`, the two pipeline CPU-offload methods, and `load_lora_adapter` now refuse a tensor-parallel model. - `save_pretrained` refuses a quantized tensor-parallel model. Previously the `dcp=True` branch returned before the quantizer's serialization step, writing shards with no quantization metadata and no error. - The DCP load guard checked the `quantization_config` kwarg only, so a pre-quantized checkpoint directory loaded silently; check the config's own entry too, and add the missing `_tp_plan` check that otherwise surfaced as a raw `AttributeError`. - Correct the `from_pretrained` message and the doc sentence that pointed at `enable_parallelism` as a way to shard a quantized model. The new tests are the first tensor-parallel tests that need neither an accelerator nor more than one rank: every case asserts a raise before any collective, so they run single-process on gloo. --- .../en/training/distributed_inference.md | 2 +- src/diffusers/hooks/tensor_parallel.py | 55 +++++- src/diffusers/loaders/peft.py | 10 ++ src/diffusers/models/modeling_utils.py | 43 ++++- src/diffusers/pipelines/pipeline_utils.py | 25 +++ tests/models/test_parallelism_guards.py | 170 ++++++++++++++++++ 6 files changed, 295 insertions(+), 10 deletions(-) create mode 100644 tests/models/test_parallelism_guards.py diff --git a/docs/source/en/training/distributed_inference.md b/docs/source/en/training/distributed_inference.md index 25d4dba4e339..20618280c841 100644 --- a/docs/source/en/training/distributed_inference.md +++ b/docs/source/en/training/distributed_inference.md @@ -483,7 +483,7 @@ 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. -A tensor-parallel `parallel_config` cannot be combined with `device_map`, `quantization_config`, `low_cpu_mem_usage=False`, `use_flashpack=True`, DDUF checkpoints, or non-safetensors weights; each raises rather than quietly falling back to loading the full checkpoint. 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. +A tensor-parallel `parallel_config` cannot be combined with `device_map`, `quantization_config`, `low_cpu_mem_usage=False`, `use_flashpack=True`, DDUF checkpoints, or non-safetensors weights; each raises rather than quietly falling back to loading the full checkpoint. Tensor parallelism also cannot be combined with quantization, offloading, or LoRA adapters at all — the parameters it shards have to be plain parameters owned by the model — 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 diff --git a/src/diffusers/hooks/tensor_parallel.py b/src/diffusers/hooks/tensor_parallel.py index d3cbd82d7980..f3856b438100 100644 --- a/src/diffusers/hooks/tensor_parallel.py +++ b/src/diffusers/hooks/tensor_parallel.py @@ -17,7 +17,7 @@ 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 @@ -401,6 +401,55 @@ def _partition_linear_fn(self, name, module, device_mesh): return resolved +def _check_tp_model_state(model: torch.nn.Module) -> None: + """Reject a model whose parameters tensor parallelism cannot take over. + + Tensor parallelism replaces every planned `weight` and `bias` with a `DTensor` shard. That only works on plain + parameters owned by the model itself, so a model whose parameters are quantized, held elsewhere by an offloading + hook, or wrapped by an adapter is rejected up front rather than failing deep inside `parallelize_module` — or, + worse, sharding successfully and producing wrong numbers. + + `from_pretrained` rejects the same combinations earlier and with a message naming the offending argument; this is + the only guard on the `enable_parallelism` path, where the model already exists and only its state can be read. + """ + 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, @@ -428,6 +477,10 @@ def apply_tensor_parallel( if num_heads is not None and num_heads % config._tp_degree != 0: raise ValueError(f"`tp_degree` ({config._tp_degree}) must divide the number of attention heads ({num_heads}).") + # Before the device-type check below, so that a quantized or offloaded model reports what is actually wrong with + # it rather than being turned away for its device type. + _check_tp_model_state(model) + if tp_mesh.device_type not in _SUPPORTED_TP_DEVICES: raise ValueError( f"Tensor parallelism is not supported on device type '{tp_mesh.device_type}'. Supported device types are " diff --git a/src/diffusers/loaders/peft.py b/src/diffusers/loaders/peft.py index b0494207f48e..0f933f5ba096 100644 --- a/src/diffusers/loaders/peft.py +++ b/src/diffusers/loaders/peft.py @@ -154,6 +154,16 @@ def load_lora_adapter( from ..hooks.group_offloading import _maybe_remove_and_reapply_group_offloading + parallel_config = getattr(self, "_parallel_config", None) + if parallel_config is not None and parallel_config.tensor_parallel_config is not None: + # `_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 ValueError( + f"Cannot load a LoRA adapter into '{self.__class__.__name__}': it is sharded with tensor " + f"parallelism, and the adapter layers are not covered by the model's `_tp_plan`. Load the adapter " + f"before sharding the model." + ) + cache_dir = kwargs.pop("cache_dir", None) force_download = kwargs.pop("force_download", False) proxies = kwargs.pop("proxies", None) diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index e690fff0b058..bd4ec03727dd 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -575,6 +575,12 @@ def enable_group_offload( "2. Or, run a forward pass with tiling disabled (can still use small dummy inputs)." ) logger.warning(msg) + if self._parallel_config is not None and self._parallel_config.tensor_parallel_config is not None: + raise ValueError( + f"'{self.__class__.__name__}' is sharded with tensor parallelism, which cannot be combined with group " + "offloading: 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." + ) 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 " @@ -739,6 +745,22 @@ def save_pretrained( return hf_quantizer = getattr(self, "hf_quantizer", None) + + tp_config = None + if self._parallel_config is not None: + tp_config = self._parallel_config.tensor_parallel_config + + if hf_quantizer is not None and tp_config is not None: + # Checked before the serializability check below, so that the reason reported is this one rather than a + # generic "not serializable". Neither save path can honour both: the `dcp=True` branch returns before + # `hf_quantizer.get_state_dict_and_metadata` runs, which would leave the shards without their + # quantization metadata, and the gathered path would hand the quantizer tensors that have been through a + # DTensor round trip. Tensor parallelism and quantization cannot be combined in the first place. + raise ValueError( + "A quantized tensor-parallel model cannot be saved: tensor parallelism and quantization cannot be " + "combined in the first place." + ) + if hf_quantizer is not None: quantization_serializable = ( hf_quantizer is not None @@ -755,10 +777,6 @@ def save_pretrained( " the logger on the traceback to understand the reason why the quantized model is not serializable." ) - tp_config = None - if self._parallel_config is not None: - tp_config = self._parallel_config.tensor_parallel_config - if dcp: if tp_config is None: raise ValueError( @@ -1249,6 +1267,9 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None for name, value in ( ("device_map", device_map), ("quantization_config", quantization_config), + # The config's own entry, not just the kwarg: this branch returns before `pre_quantized` is + # computed, so a pre-quantized checkpoint directory would otherwise load silently. + ("a quantized checkpoint", config.get("quantization_config") is not None), ("use_flashpack", use_flashpack), ("variant", variant), ("dduf_entries", dduf_entries), @@ -1261,6 +1282,12 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None f"{unsupported} cannot be combined with the distributed checkpoint at {dcp_dir}: its " "shards are read in place onto each rank's device." ) + if cls._tp_plan is None: + raise ValueError( + f"`_tp_plan` must be set on the model class to read the distributed checkpoint at " + f"{dcp_dir}, whose shards are those of a tensor-parallel model. '{cls.__name__}' does not " + f"define one." + ) return cls._load_dcp_checkpoint( dcp_dir, config, unused_kwargs, torch_dtype=torch_dtype, parallel_config=parallel_config ) @@ -1790,8 +1817,7 @@ def _load_dcp_checkpoint( The shards are those of a tensor-parallel model, so a tensor-parallel `parallel_config` is required, at the `tp_degree` the checkpoint was written with — see the note where it is written. Use the ordinary safetensors - path to move a model between degrees; it streams each rank's slice, so it costs no more memory than this - does. + path to move a model between degrees; it streams each rank's slice, so it costs no more memory than this does. DCP loads **in place**, so every parameter has to be allocated first with its local shape and on the device it will end up on. @@ -1920,8 +1946,9 @@ def _check_tp_streaming_supported( ) if hf_quantizer is not None: raise ValueError( - "`quantization_config` cannot be combined with a tensor-parallel `parallel_config`. Load the " - "model unquantized, or shard it after loading with `enable_parallelism`." + "`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( diff --git a/src/diffusers/pipelines/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index 24fe0eabfa6f..37b563f3f79e 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -1208,6 +1208,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(raise_error=True) is_pipeline_device_mapped = self._is_pipeline_device_mapped() if is_pipeline_device_mapped: @@ -1326,6 +1327,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(raise_error=True) if is_accelerate_available() and is_accelerate_version(">=", "0.14.0"): from accelerate import cpu_offload @@ -2272,6 +2274,29 @@ def _maybe_raise_error_if_group_offload_active( return True return False + def _maybe_raise_error_if_tensor_parallel_active( + self, raise_error: bool = False, module: torch.nn.Module | None = None + ) -> bool: + """Whether 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. + """ + components = self.components.values() if module is None else [module] + components = [component for component in components if isinstance(component, torch.nn.Module)] + for component in components: + parallel_config = getattr(component, "_parallel_config", None) + if parallel_config is not None and parallel_config.tensor_parallel_config is not None: + if raise_error: + raise ValueError( + f"You are trying to apply model/sequential CPU offloading to a pipeline whose " + f"'{component.__class__.__name__}' is sharded with tensor parallelism. This is not supported: " + f"tensor parallelism already keeps only one shard of each weight per rank, so offloading is " + f"not needed on top of it." + ) + return True + return False + 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/tests/models/test_parallelism_guards.py b/tests/models/test_parallelism_guards.py new file mode 100644 index 000000000000..0521925b050d --- /dev/null +++ b/tests/models/test_parallelism_guards.py @@ -0,0 +1,170 @@ +# 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. +"""Guards rejecting tensor parallelism combined with quantization, offloading, or LoRA adapters. + +Unlike the rest of the tensor-parallel suite in `testing_utils/parallelism.py`, these tests need neither an +accelerator nor more than one rank: every case asserts that a call raises before any collective is issued. They run +single-process on gloo, so they run in ordinary CI. +""" + +import pytest +import torch +import torch.distributed as dist +import torch.nn as nn + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.models._modeling_parallel import ParallelConfig, TensorParallelConfig +from diffusers.models.modeling_utils import ModelMixin + + +class TinyTPModel(ModelMixin, ConfigMixin): + """Smallest model carrying a `_tp_plan`: one colwise Linear feeding one rowwise Linear.""" + + config_name = "config.json" + _tp_plan = {"linear_1": "colwise", "linear_2": "rowwise"} + _supports_group_offloading = True + + @register_to_config + def __init__(self, hidden_size: int = 8, num_attention_heads: int = 2): + super().__init__() + self.linear_1 = nn.Linear(hidden_size, hidden_size) + self.linear_2 = nn.Linear(hidden_size, hidden_size) + + def forward(self, hidden_states): + return self.linear_2(self.linear_1(hidden_states)) + + +class _SerializableQuantizer: + """Stand-in that passes `save_pretrained`'s serializability check, so the TP guard is what raises.""" + + is_serializable = True + supports_safetensors_serialization = True + + +@pytest.fixture(scope="module") +def gloo_process_group(): + """A single-rank CPU process group, enough for `_resolve_parallel_config` to build a mesh.""" + if not dist.is_available(): + pytest.skip("torch.distributed is not available.") + already_initialized = dist.is_initialized() + if not already_initialized: + dist.init_process_group(backend="gloo", init_method="tcp://127.0.0.1:29591", world_size=1, rank=0) + yield + if not already_initialized: + dist.destroy_process_group() + + +def _shard(model): + model.enable_parallelism(config=TensorParallelConfig(tp_degree=1)) + + +def _mark_as_tensor_parallel(model): + """Put the model in the state it would be in after sharding, without needing a real mesh.""" + model._parallel_config = ParallelConfig(tensor_parallel_config=TensorParallelConfig(tp_degree=2)) + return model + + +class TestTensorParallelModelStateGuards: + """`_check_tp_model_state` — a model whose parameters TP cannot take over.""" + + def test_clean_model_reaches_the_device_check(self, gloo_process_group): + """Ordering guard: with none of the bad states, the device-type check is what rejects CPU. + + This is what keeps the tests below meaningful. If `_check_tp_model_state` ran after the + `_SUPPORTED_TP_DEVICES` check, every case would raise the device error instead of its own. + """ + with pytest.raises(ValueError, match="not supported on device type"): + _shard(TinyTPModel()) + + def test_quantized_via_hf_quantizer(self, gloo_process_group): + model = TinyTPModel() + model.hf_quantizer = object() + with pytest.raises(ValueError, match="is quantized"): + _shard(model) + + def test_quantized_via_is_quantized(self, gloo_process_group): + model = TinyTPModel() + model.is_quantized = True + with pytest.raises(ValueError, match="is quantized"): + _shard(model) + + def test_device_map_dispatched(self, gloo_process_group): + model = TinyTPModel() + model.hf_device_map = {"": 0} + with pytest.raises(ValueError, match="placed by accelerate"): + _shard(model) + + def test_accelerate_hook_on_submodule(self, gloo_process_group): + model = TinyTPModel() + model.linear_1._hf_hook = object() + with pytest.raises(ValueError, match="placed by accelerate"): + _shard(model) + + def test_group_offloaded(self, gloo_process_group, monkeypatch): + import diffusers.hooks.group_offloading as group_offloading + + monkeypatch.setattr(group_offloading, "_is_group_offload_enabled", lambda module: True) + with pytest.raises(ValueError, match="group offloading enabled"): + _shard(TinyTPModel()) + + def test_peft_adapter_injected(self, gloo_process_group): + peft = pytest.importorskip("peft") + + model = TinyTPModel() + peft.inject_adapter_in_model(peft.LoraConfig(r=2, target_modules=["linear_1"]), model) + with pytest.raises(ValueError, match=r"adapter \(LoRA\) layers injected"): + _shard(model) + + +class TestTensorParallelReverseDirectionGuards: + """The other order: a model already sharded, then asked to offload or take an adapter.""" + + def test_enable_group_offload_on_tp_model(self): + model = _mark_as_tensor_parallel(TinyTPModel()) + with pytest.raises(ValueError, match="sharded with tensor parallelism"): + model.enable_group_offload(onload_device=torch.device("cpu")) + + def test_pipeline_offload_helper_detects_tp_component(self): + from diffusers.pipelines.pipeline_utils import DiffusionPipeline + + model = _mark_as_tensor_parallel(TinyTPModel()) + # The helper takes an explicit module, so it runs without building a whole pipeline. + with pytest.raises(ValueError, match="sharded with tensor parallelism"): + DiffusionPipeline._maybe_raise_error_if_tensor_parallel_active( + DiffusionPipeline, raise_error=True, module=model + ) + + def test_pipeline_offload_helper_passes_for_plain_model(self): + from diffusers.pipelines.pipeline_utils import DiffusionPipeline + + assert not DiffusionPipeline._maybe_raise_error_if_tensor_parallel_active( + DiffusionPipeline, raise_error=True, module=TinyTPModel() + ) + + +class TestTensorParallelSaveGuards: + """`save_pretrained` must not write a checkpoint that silently drops quantization.""" + + def test_dcp_save_rejects_quantized_model(self, tmp_path): + model = _mark_as_tensor_parallel(TinyTPModel()) + model.hf_quantizer = _SerializableQuantizer() + with pytest.raises(ValueError, match="quantized tensor-parallel model cannot be saved"): + model.save_pretrained(str(tmp_path / "dcp"), dcp=True) + + def test_tp_save_rejects_quantized_model(self, tmp_path): + model = _mark_as_tensor_parallel(TinyTPModel()) + model.hf_quantizer = _SerializableQuantizer() + with pytest.raises(ValueError, match="quantized tensor-parallel model cannot be saved"): + model.save_pretrained(str(tmp_path / "full")) From a8956f4d31583a41eb28074b792a1da2cecc4077 Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Fri, 21 Aug 2026 11:08:58 +0000 Subject: [PATCH 03/22] Add tensor-parallel support for MiniMax-H3 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Shards `MiniMaxH3Transformer3DModel` across devices, following the plan already established for Flux1/Flux2/Qwen-Image. Validated on Trainium at TP=2 and TP=8. - `_tp_plan` with twelve entries: the same six shapes for the 50 denoiser blocks and for the two token-refiner blocks, which are the same attention + SwiGLU FFN minus AdaLN and rotary. Q/K/V and the attention output are unfused, so they are plain colwise/rowwise; the SwiGLU input `ff.net.0.proj` is one Linear producing `[value; gate]` in equal halves and takes PackedColwiseParallel([1, 1]). - The attention processor reshaped by the config head count, `unflatten(-1, (attn.heads, -1))`, which mis-splits under sharding: each rank holds `inner_dim / tp_degree` columns, so this yields `head_dim / tp_degree` per head instead of `heads / tp_degree` heads. Reshape by the fixed `attn.head_dim` instead and let `-1` absorb the head count, as Flux does. Numerically identical unsharded, since `inner_dim == heads * head_dim`. - Norms, QK-norms (head_dim-shaped, applied after the head split), AdaLN modulation and the patch/text embedders and output heads stay replicated. `attn.to_qkv` is deliberately not in the plan: it exists only after `fuse_projections()`, and the plan is resolved by attribute lookup. No RoPE change was needed — unlike Qwen-Image, H3's rotary is already real sin/cos and already broadcasts over the head axis. Tests mirror the Flux2/Qwen-Image layout: the CUDA/XPU `TensorParallelTesterMixin` class, a `make_neuron_tp_spec()` factory, and a Neuron launcher that shells out to the model-agnostic `_neuron_tp_worker.py`. `get_dummy_inputs` and `get_packed_layout` take an optional `device` so the Neuron spec can ask for CPU tensors, since its worker shards on CPU and moves to device after. --- .../transformers/transformer_minimax_h3.py | 37 ++++++++- .../test_models_transformer_minimax_h3.py | 83 +++++++++++++++---- 2 files changed, 100 insertions(+), 20 deletions(-) diff --git a/src/diffusers/models/transformers/transformer_minimax_h3.py b/src/diffusers/models/transformers/transformer_minimax_h3.py index f49cdaca2eb6..41269215cbf3 100644 --- a/src/diffusers/models/transformers/transformer_minimax_h3.py +++ b/src/diffusers/models/transformers/transformer_minimax_h3.py @@ -19,6 +19,7 @@ import torch.nn as nn from ...configuration_utils import ConfigMixin, register_to_config +from ...hooks.tensor_parallel import PackedColwiseParallel from ...loaders import PeftAdapterMixin from ...utils import BaseOutput, apply_lora_scale, logging from .._modeling_parallel import ContextParallelInput, ContextParallelOutput @@ -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) @@ -449,6 +452,34 @@ 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. + # + # 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; the AdaLN modulation (`adaln_proj.linear`, `norm_out.linear`), which indexes the full hidden + # dim; 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", + # 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/tests/models/transformers/test_models_transformer_minimax_h3.py b/tests/models/transformers/test_models_transformer_minimax_h3.py index 00baa37c84a0..e27cbd68cf21 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,7 @@ LoraTesterMixin, MemoryTesterMixin, ModelTesterMixin, + TensorParallelTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, ) @@ -84,7 +89,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 +97,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 +132,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 +143,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 +201,40 @@ class TestMiniMaxH3TransformerContextParallel(MiniMaxH3TransformerTesterConfig, class TestMiniMaxH3TransformerLoRA(MiniMaxH3TransformerTesterConfig, LoraTesterMixin): """LoRA tests for the MiniMax-H3 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}" + ) From 7e6e38facffe29551ca04f003e034b61e732718c Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Fri, 21 Aug 2026 13:45:24 +0000 Subject: [PATCH 04/22] Shard MiniMax-H3's adaln_proj to fit tensor parallelism on one device MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `transformer_blocks.*.adaln_proj.linear` was left replicated on every rank, and at `[96768, 2688]` bf16 per block it is 24.23 GiB of the denoiser's 61.73 GiB — about 40%. That made the per-rank floor 24.40 GiB of weights (plus 5.13 GiB for the two VAEs) regardless of TP degree, so MiniMax-H3 could not fit a 24 GiB NeuronCore at *any* valid TP: 34.20 GiB/rank at TP=8, and still 30.20 GiB/rank at TP=56. Raising TP only divided the 60% that already sharded. (TP=16 is not an option either — 56 attention heads.) Shard it rowwise, over the `time_embed_dim` input, rather than colwise: the six modulation parameters scale and shift the *full* hidden dim of a sequence that is already all-reduced by the time they are applied, so a colwise split would need an all-gather to rebuild that width. Rowwise keeps the output full-width, leaving the module's `view`/`chunk` untouched, and all-reduces a few hundred KB per block per step. Plain `"rowwise"` could not be reused. It is normally the second half of a colwise/rowwise pair, so it defaults to `input_layouts=Shard(-1)` and would read the full-width `temb` as if it were one rank's shard. Hence `ReplicatedInputRowwiseParallel`: input narrowed locally on the way in (no collective), partial output all-reduced on the way out, bias replicated and added after the reduce. It is wired into `_styles`, `_hooks_only_styles` — the path the Neuron backend takes, since `_apply_tp_neuron` pre-shards on CPU and then registers hooks only — and `resolve_tp_shard_specs`. Replicated weights drop from 24.40 GiB to 0.15 GiB, putting TP=8 at 7.84 GiB of transformer plus 5.13 GiB of VAEs, i.e. 12.97 GiB/rank against a 24 GiB budget. Verified on CPU/gloo that both the generic and the pre-sharded hooks-only path shard the weight on its input dim and match a replicated reference to 3.6e-7. Co-Authored-By: Claude Opus 5 (1M context) --- src/diffusers/hooks/tensor_parallel.py | 45 +++++++++++++++---- .../transformers/transformer_minimax_h3.py | 17 +++++-- 2 files changed, 50 insertions(+), 12 deletions(-) diff --git a/src/diffusers/hooks/tensor_parallel.py b/src/diffusers/hooks/tensor_parallel.py index f3856b438100..a13fa229788c 100644 --- a/src/diffusers/hooks/tensor_parallel.py +++ b/src/diffusers/hooks/tensor_parallel.py @@ -50,6 +50,20 @@ 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 _blocks_to_block_sizes(total_size: int, blocks: "list[int]") -> "list[int]": """Convert proportional block counts to absolute sizes. @@ -187,7 +201,8 @@ def resolve_tp_shard_specs(model: torch.nn.Module, tp_plan: dict) -> "dict[str, if style == "colwise": weight_spec = TPShardSpec(0, [submodule.weight.shape[0]]) bias_spec = weight_spec - elif style == "rowwise": + 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): @@ -201,7 +216,8 @@ def resolve_tp_shard_specs(model: torch.nn.Module, tp_plan: dict) -> "dict[str, 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." ) specs[f"{path}.weight"] = weight_spec @@ -256,9 +272,10 @@ def _resolve_tp_plan(model: torch.nn.Module, tp_plan: dict) -> list: 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() | }`, each subclassed to reject a sharded dim that is not divisible by the TP + degree. """ import torch.nn as nn from torch.distributed.tensor import DTensor, Replicate, Shard, distribute_tensor @@ -334,7 +351,7 @@ def _partition_linear_fn(self, name, module, device_mesh): return _CheckedColwiseImpl() - def _make_checked_row(path: str) -> RowwiseParallel: + def _make_checked_row(path: str, replicated_input: bool = False) -> RowwiseParallel: class _CheckedRowwiseImpl(RowwiseParallel): def _partition_linear_fn(self, name, module, device_mesh): tp_size = device_mesh.size() @@ -346,7 +363,10 @@ def _partition_linear_fn(self, name, module, device_mesh): ) super()._partition_linear_fn(name, module, device_mesh) - return _CheckedRowwiseImpl() + # `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. + return _CheckedRowwiseImpl(input_layouts=Replicate()) if replicated_input else _CheckedRowwiseImpl() resolved = {} for path, style in relative_plan.items(): @@ -354,6 +374,8 @@ def _partition_linear_fn(self, name, module, device_mesh): resolved[path] = _make_checked_col(path) elif style == "rowwise": resolved[path] = _make_checked_row(path) + elif isinstance(style, ReplicatedInputRowwiseParallel): + resolved[path] = _make_checked_row(path, replicated_input=True) elif isinstance(style, PackedColwiseParallel): resolved[path] = _make_packed_col(style) elif isinstance(style, PackedRowwiseParallel): @@ -361,7 +383,8 @@ 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 @@ -377,6 +400,7 @@ def _hooks_only_styles(relative_plan: dict) -> dict: 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): @@ -393,10 +417,13 @@ def _partition_linear_fn(self, name, module, device_mesh): 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, or PackedRowwiseParallel." + f"Expected 'colwise', 'rowwise', PackedColwiseParallel, PackedRowwiseParallel, or " + f"ReplicatedInputRowwiseParallel." ) return resolved diff --git a/src/diffusers/models/transformers/transformer_minimax_h3.py b/src/diffusers/models/transformers/transformer_minimax_h3.py index 41269215cbf3..f6038c817848 100644 --- a/src/diffusers/models/transformers/transformer_minimax_h3.py +++ b/src/diffusers/models/transformers/transformer_minimax_h3.py @@ -19,7 +19,7 @@ import torch.nn as nn from ...configuration_utils import ConfigMixin, register_to_config -from ...hooks.tensor_parallel import PackedColwiseParallel +from ...hooks.tensor_parallel import PackedColwiseParallel, ReplicatedInputRowwiseParallel from ...loaders import PeftAdapterMixin from ...utils import BaseOutput, apply_lora_scale, logging from .._modeling_parallel import ContextParallelInput, ContextParallelOutput @@ -457,10 +457,20 @@ class MiniMaxH3Transformer3DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftA # 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; the AdaLN modulation (`adaln_proj.linear`, `norm_out.linear`), which indexes the full hidden - # dim; and the patch/text embedders and the two output heads. + # 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. @@ -472,6 +482,7 @@ class MiniMaxH3Transformer3DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftA "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", From bb236dccfef5529910711c4340e9c9660cc3bffb Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Fri, 21 Aug 2026 14:32:06 +0000 Subject: [PATCH 05/22] Build MiniMax-H3's row timestep plan on CPU, as its caller expects MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `build_row_timesteps` allocates `row_timesteps` with `torch.full` and no device, then scatters into it at `video_indices` / `audio_indices`. The layout step hands those index tensors over already on the execution device, and indexing a CPU tensor with an accelerator one is an error — on Neuron it surfaces as "Non-scalar tensor arg0 is on cpu device, expected neuron", and on CUDA it would raise "indices should be either on cpu or on the same device". CPU is the right place for this to run, not the accelerator: `torch.unique` has a data-dependent output shape, which is precisely what a tracing backend cannot handle, and the caller already moves the finished `(timestep, timestep_indices)` pair to the device itself. So bring the two index tensors back to CPU for the scatter rather than allocating `row_timesteps` on their device. Only reachable once the denoiser is actually on an accelerator while the pipeline's execution device resolves there too, which is why it went unnoticed. Co-Authored-By: Claude Opus 5 (1M context) --- .../modular_pipelines/minimax_h3/before_denoise.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py b/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py index c670467a9307..c119559ad408 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py +++ b/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py @@ -1208,6 +1208,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 From 0f9057bc0a37d49256d537247a8c7da252b2047a Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Thu, 27 Aug 2026 13:20:19 +0000 Subject: [PATCH 06/22] feat:combine tp+cp (cherry picked from commit a489f03e4e9888c0abb996f45d38885f03772ae4) --- src/diffusers/models/_modeling_parallel.py | 41 ++++- src/diffusers/models/modeling_utils.py | 49 +++++- tests/models/test_parallelism_combined.py | 164 ++++++++++++++++++ tests/models/testing_utils/__init__.py | 2 + tests/models/testing_utils/parallelism.py | 146 +++++++++++++++- .../test_models_transformer_qwenimage.py | 20 +++ 6 files changed, 409 insertions(+), 13 deletions(-) create mode 100644 tests/models/test_parallelism_combined.py diff --git a/src/diffusers/models/_modeling_parallel.py b/src/diffusers/models/_modeling_parallel.py index b54e86d6b4f2..7f12ec71df28 100644 --- a/src/diffusers/models/_modeling_parallel.py +++ b/src/diffusers/models/_modeling_parallel.py @@ -200,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. @@ -215,12 +220,10 @@ class ParallelConfig: _device: torch.device = None _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: - 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`." - ) + @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, @@ -234,10 +237,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/modeling_utils.py b/src/diffusers/models/modeling_utils.py index bd4ec03727dd..70197983ebe2 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -1992,15 +1992,30 @@ def _resolve_parallel_config( device = torch.device(device_type, rank % device_module.device_count()) mesh = None - if config.context_parallel_config is not None: - cp_config = config.context_parallel_config + 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 config.tensor_parallel_config is not None: - tp_config = config.tensor_parallel_config + 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,), @@ -2009,6 +2024,32 @@ def _resolve_parallel_config( # `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 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 0e9b7ebb4aa0..766c41d2e672 100644 --- a/tests/models/testing_utils/__init__.py +++ b/tests/models/testing_utils/__init__.py @@ -18,6 +18,7 @@ from .lora import LoraHotSwappingForModelTesterMixin, LoraTesterMixin from .memory import CPUOffloadTesterMixin, GroupOffloadTesterMixin, LayerwiseCastingTesterMixin, MemoryTesterMixin from .parallelism import ( + ContextAndTensorParallelTesterMixin, ContextParallelAttentionBackendsTesterMixin, ContextParallelTesterMixin, TensorParallelTesterMixin, @@ -65,6 +66,7 @@ "BitsAndBytesConfigMixin", "BitsAndBytesTesterMixin", "CacheTesterMixin", + "ContextAndTensorParallelTesterMixin", "ContextParallelTesterMixin", "ContextParallelAttentionBackendsTesterMixin", "TensorParallelTesterMixin", diff --git a/tests/models/testing_utils/parallelism.py b/tests/models/testing_utils/parallelism.py index c5174c396561..d1c06e015891 100644 --- a/tests/models/testing_utils/parallelism.py +++ b/tests/models/testing_utils/parallelism.py @@ -22,7 +22,7 @@ import torch.multiprocessing as mp from safetensors.torch import load_file -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.constants import SAFETENSORS_WEIGHTS_NAME @@ -547,6 +547,150 @@ def test_tensor_parallel_dcp_roundtrip(self, tmp_path): 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 class ContextParallelTesterMixin: diff --git a/tests/models/transformers/test_models_transformer_qwenimage.py b/tests/models/transformers/test_models_transformer_qwenimage.py index 7a03a8fe2353..dd257d676770 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``). From 9c81ac694d2d60914ce562e0a214717cb6d9425e Mon Sep 17 00:00:00 2001 From: whn09 Date: Mon, 7 Sep 2026 03:08:20 +0000 Subject: [PATCH 07/22] Allow tensor parallelism and context parallelism in one ParallelConfig `ParallelConfig.__post_init__` rejects a config carrying both a `context_parallel_config` and a `tensor_parallel_config`, so a model can only use one form of parallelism at a time. That caps some models well below the hardware they run on: MiniMax-H3 has 56 attention heads, so `tp_degree` cannot exceed 8, and on a 64-accelerator host tensor parallelism alone leaves 56 of them idle even though the model already ships a `_cp_plan`. Both configs already accept a caller-supplied `mesh`, documented as being for "combining CP with other parallelism strategies that share the same mesh", so the intent was there; what was missing was building that shared mesh and splitting it between the two. * `__post_init__` now only rejects a config that carries neither. * `ParallelConfig.setup` hands each config its own dimensions: CP keeps the whole mesh, since its own `setup` selects "ring" and "ulysses" out of it, while TP gets the "tp" dimension alone because it reads its degree from `mesh.size()`. * `enable_parallelism` builds one ("tp", "ring", "ulysses") mesh when both are requested, with the context-parallel dimensions varying fastest. That order is a default rather than a universal optimum, and the comment says so: pass `mesh=` on either config to choose the layout yourself. * The Neuron tensor-parallel backend used the global rank as its shard index. That equals the rank's coordinate on the TP mesh only while that mesh spans the whole world, so sharing a mesh made it index past the end of the weight; the generic backend beside it already reads the coordinate. Measured on a trn2.48xlarge (MiniMax-H3, 1344x768x124f, 30 steps): TP=8 alone is 9.285 s/step on 8 cores, TP=4 x ulysses=4 is 5.66 s/step on 16, and TP=8 x ulysses=4 is 3.941 s/step on 32 -- 2.36x faster than the widest configuration reachable before this change, with output that is visually equivalent to the tensor-parallel-only run. Tests: a `HybridParallelTesterMixin` next to the existing TP and CP mixins, wired into the Flux transformer tests, plus a Neuron `torchrun` worker following the `_neuron_tp_worker.py` convention already in the tree. The Neuron test runs at tp_degree=2 x ulysses_degree=4 and passes on a trn2.48xlarge (max_abs_diff 4.2e-05 against the single-device reference). Co-Authored-By: Claude Opus 5 --- src/diffusers/hooks/tensor_parallel_neuron.py | 7 +- src/diffusers/models/_modeling_parallel.py | 13 +- src/diffusers/models/modeling_utils.py | 26 +++- tests/models/testing_utils/__init__.py | 2 + tests/models/testing_utils/parallelism.py | 131 +++++++++++++++++- .../transformers/_neuron_hybrid_worker.py | 125 +++++++++++++++++ .../test_models_transformer_flux.py | 48 ++++++- 7 files changed, 340 insertions(+), 12 deletions(-) create mode 100644 tests/models/transformers/_neuron_hybrid_worker.py diff --git a/src/diffusers/hooks/tensor_parallel_neuron.py b/src/diffusers/hooks/tensor_parallel_neuron.py index 6b8f219a17ff..9c53afc0f623 100644 --- a/src/diffusers/hooks/tensor_parallel_neuron.py +++ b/src/diffusers/hooks/tensor_parallel_neuron.py @@ -22,7 +22,6 @@ """ import torch -import torch.distributed as dist import torch.nn as nn @@ -171,7 +170,11 @@ def _apply_tp_neuron( Model weights must be on CPU when this is called. """ - rank = dist.get_rank() + # The shard index is this rank's coordinate on the tensor-parallel mesh, which is the global rank only + # when that mesh spans the whole world. Sharing a mesh with another parallelism makes `tp_mesh` a sub-mesh, + # and a global rank then indexes past the end of the weight. The generic backend already reads the + # coordinate (`tensor_parallel.py`), as does `ContextParallelConfig.setup`. + rank = tp_mesh.get_local_rank() tp_size = tp_mesh.size() for block, relative_plan in groups: diff --git a/src/diffusers/models/_modeling_parallel.py b/src/diffusers/models/_modeling_parallel.py index 86627284e078..a98574cc49e8 100644 --- a/src/diffusers/models/_modeling_parallel.py +++ b/src/diffusers/models/_modeling_parallel.py @@ -214,10 +214,10 @@ 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`." ) def setup( @@ -233,9 +233,14 @@ def setup( self._device = device self._mesh = mesh if self.context_parallel_config is not None: + # `ContextParallelConfig.setup` selects its own "ring" and "ulysses" dimensions out of `mesh`. self.context_parallel_config.setup(rank, world_size, device, mesh) if self.tensor_parallel_config is not None: - self.tensor_parallel_config.setup(rank, world_size, device, mesh) + # `TensorParallelConfig.setup` reads `tp_degree` off `mesh.size()`, so when the mesh is shared with + # context parallelism it must be handed the "tp" dimension alone rather than the whole mesh. + dim_names = (mesh.mesh_dim_names or ()) if mesh is not None else () + tp_mesh = mesh["tp"] if "tp" in dim_names and len(dim_names) > 1 else mesh + self.tensor_parallel_config.setup(rank, world_size, device, tp_mesh) @dataclass(frozen=True) diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index 425f2f29235e..6926b6fa89d3 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -1669,15 +1669,33 @@ def enable_parallelism( break mesh = None - if config.context_parallel_config is not None: - cp_config = config.context_parallel_config + cp_config, tp_config = config.context_parallel_config, config.tensor_parallel_config + if cp_config is not None and tp_config is not None: + # A single mesh spanning both parallelisms, with the context parallel dimensions varying fastest: a + # CP group is then a contiguous run of ranks and a TP group takes one rank out of each run. + # + # That order is a default, not a universal optimum. Some accelerator runtimes only accept contiguous + # replica groups for all-to-all -- which Ulysses needs -- while tolerating strided groups for + # all-reduce, which is all TP needs; since only one axis of a 2-D mesh can be contiguous, Ulysses + # gets it. On multi-node CUDA the opposite order is usually preferable, because TP is the most + # bandwidth-hungry collective and wants to stay within one NVLink domain. Pass `mesh=` on either + # config to choose the layout yourself. + mesh = ( + cp_config.mesh + or tp_config.mesh + or torch.distributed.device_mesh.init_device_mesh( + device_type=device_type, + mesh_shape=(tp_config.tp_degree, *cp_config.mesh_shape), + mesh_dim_names=("tp", *cp_config.mesh_dim_names), + ) + ) + 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 config.tensor_parallel_config is not None: - tp_config = config.tensor_parallel_config + 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,), diff --git a/tests/models/testing_utils/__init__.py b/tests/models/testing_utils/__init__.py index 760a3fac04e0..4d70f3a1ae29 100644 --- a/tests/models/testing_utils/__init__.py +++ b/tests/models/testing_utils/__init__.py @@ -20,6 +20,7 @@ from .parallelism import ( ContextParallelAttentionBackendsTesterMixin, ContextParallelTesterMixin, + HybridParallelTesterMixin, TensorParallelTesterMixin, ) from .quantization import ( @@ -64,6 +65,7 @@ "CacheTesterMixin", "ContextParallelTesterMixin", "ContextParallelAttentionBackendsTesterMixin", + "HybridParallelTesterMixin", "TensorParallelTesterMixin", "CPUOffloadTesterMixin", "FasterCacheConfigMixin", diff --git a/tests/models/testing_utils/parallelism.py b/tests/models/testing_utils/parallelism.py index 63575abf6b7b..49c7acda6531 100644 --- a/tests/models/testing_utils/parallelism.py +++ b/tests/models/testing_utils/parallelism.py @@ -21,7 +21,7 @@ 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 ...testing_utils import ( @@ -295,6 +295,63 @@ def _tensor_parallel_worker( dist.destroy_process_group() +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_tensor_parallel @require_torch_multi_accelerator class TensorParallelTesterMixin: @@ -345,6 +402,78 @@ def test_tensor_parallel_batch_inputs(self): self.test_tensor_parallel_inference(batch_size=2) +@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) + + @is_context_parallel @require_torch_multi_accelerator class ContextParallelTesterMixin: 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_flux.py b/tests/models/transformers/test_models_transformer_flux.py index 53af9eedc50c..e4d7ce8b808e 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, @@ -268,6 +275,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`). @@ -301,6 +325,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.""" From 76acd09bc32bb8141695a96cd008ab9b1455b0ca Mon Sep 17 00:00:00 2001 From: Jingya HUANG Date: Wed, 30 Sep 2026 14:01:57 +0000 Subject: [PATCH 08/22] feat: validate the support on TPU --- src/diffusers/hooks/tensor_parallel.py | 51 ++++++++++++++++++- src/diffusers/hooks/tensor_parallel_neuron.py | 31 +---------- src/diffusers/models/_modeling_parallel.py | 2 +- src/diffusers/models/modeling_utils.py | 11 ++++ 4 files changed, 63 insertions(+), 32 deletions(-) diff --git a/src/diffusers/hooks/tensor_parallel.py b/src/diffusers/hooks/tensor_parallel.py index a13fa229788c..4cd8bbf26849 100644 --- a/src/diffusers/hooks/tensor_parallel.py +++ b/src/diffusers/hooks/tensor_parallel.py @@ -22,7 +22,7 @@ logger = get_logger(__name__) # pylint: disable=invalid-name -_SUPPORTED_TP_DEVICES = ("cuda", "neuron") +_SUPPORTED_TP_DEVICES = ("cuda", "neuron", "tpu") class PackedColwiseParallel: @@ -428,6 +428,47 @@ def _partition_linear_fn(self, name, module, device_mesh): 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 _check_tp_model_state(model: torch.nn.Module) -> None: """Reject a model whose parameters tensor parallelism cannot take over. @@ -515,7 +556,7 @@ 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}.") @@ -532,5 +573,11 @@ def apply_tensor_parallel( _apply_tp_neuron(model, tp_mesh, groups, resolve_tp_shard_specs(model, tp_plan)) return + 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, resolve_tp_shard_specs(model, tp_plan), 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 ffcba0973d81..a87f28b1edb1 100644 --- a/src/diffusers/hooks/tensor_parallel_neuron.py +++ b/src/diffusers/hooks/tensor_parallel_neuron.py @@ -28,7 +28,7 @@ import torch import torch.nn as nn -from .tensor_parallel import TPShardSpec, _hooks_only_styles, _local_shard +from .tensor_parallel import TPShardSpec, _pre_shard_and_parallelize def _apply_tp_neuron( @@ -45,31 +45,4 @@ def _apply_tp_neuron( Model weights must be on CPU when this is called. """ - from torch.distributed.tensor import DTensor, Replicate, Shard - from torch.distributed.tensor.parallel import parallelize_module - - device = torch.neuron.current_device() - - 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)) + _pre_shard_and_parallelize(model, tp_mesh, groups, specs, torch.neuron.current_device()) diff --git a/src/diffusers/models/_modeling_parallel.py b/src/diffusers/models/_modeling_parallel.py index 0ce2a005aedd..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`): diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index 5993b93d90a4..791174a92cad 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -1654,8 +1654,19 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None 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) elif parallel_config is not None: model.enable_parallelism(config=parallel_config) From 78405ec9a1c1637c4027136658c96b7e85f915e5 Mon Sep 17 00:00:00 2001 From: yzhautouskay Date: Wed, 30 Sep 2026 17:53:20 +0200 Subject: [PATCH 09/22] [Cosmos3] Fix Transfer SeaCache artifacts with control CFG (#14897) * Fix Cosmos 3 Transfer SeaCache indicators * Clarify SeaCache indicator guidance; tests refactor --------- Co-authored-by: Sayak Paul --- docs/source/en/optimization/cache.md | 15 ++- src/diffusers/hooks/sea_cache.py | 14 ++- .../test_models_transformer_cosmos3.py | 107 +++++++++++++++++- 3 files changed, 121 insertions(+), 15 deletions(-) diff --git a/docs/source/en/optimization/cache.md b/docs/source/en/optimization/cache.md index 9f775ec3b88c..3a771336e4b7 100644 --- a/docs/source/en/optimization/cache.md +++ b/docs/source/en/optimization/cache.md @@ -72,8 +72,14 @@ pipeline.transformer.enable_cache(config) [SeaCache](https://huggingface.co/papers/2602.18993) compares Spectral-Evolution-Aware (SEA) indicators between successive denoising steps. When the accumulated indicator change remains below a threshold, it skips the expensive -transformer block stack and predicts its output from cached residuals. The indicator is computed from the raw vision -latents, including clean conditioning frames for image-to-video generation. +transformer block stack and predicts its output from cached residuals. Build the indicator from the visual latents that +form the generated output. Include clean conditioning frames when they are part of that output trajectory, as in +image-to-video and video-to-video generation. Exclude separate visual hints that condition the generation but are not +part of the output. Text conditioning is excluded because it is not a visual latent. + +Cosmos 3 Transfer packs control hints as separate visual sequences, so its adapter excludes them from the indicator. +Control-CFG branches compare the same output trajectory while retaining their own cached residuals. Control hints still +condition the transformer. The implementation provides built-in adapters for the following models: @@ -87,8 +93,9 @@ Other video transformers can integrate with the generic path when they use `Cach block list, and register the block input/output layout in `TransformerBlockRegistry`. The pipeline must enter a `cache_context` for every transformer call, attach `step_index`, `sigma`, and `num_inference_steps`, and use separate context names for independent trajectories such as conditional and unconditional guidance. Pass a `raw_vision_callback` -that returns the noisy vision latents when no built-in adapter is available. Validate output quality and tune the cache -parameters for each model and scheduler; support and benchmark results do not transfer automatically from Cosmos 3. +that returns the visual latents forming the generated output when no built-in adapter is available. Validate output +quality and tune the cache parameters for each model and scheduler; support and benchmark results do not transfer +automatically from Cosmos 3. ### Cosmos 3 diff --git a/src/diffusers/hooks/sea_cache.py b/src/diffusers/hooks/sea_cache.py index 5cb78db1f6b4..d228b0772e48 100644 --- a/src/diffusers/hooks/sea_cache.py +++ b/src/diffusers/hooks/sea_cache.py @@ -64,8 +64,9 @@ class SeaCacheConfig: power_exp (`float`, defaults to `3.0`): Exponent of the SEA clean-signal power prior. SeaCache uses `3.0` for video features. raw_vision_callback (`Callable`, *optional*): - Advanced model adapter returning raw vision latents with shape `(C, T, H, W)`. When omitted, a built-in - adapter is used if one is available. + Advanced model adapter returning the visual latents forming the generated output, each with shape `(C, T, + H, W)`. Include clean conditioning frames within the output trajectory, but exclude separate visual hints + that are not part of the output. When omitted, a built-in adapter is used if one is available. Example: ```python @@ -326,7 +327,6 @@ def _prepare_cosmos3_raw_vision_metadata( return None raw_vision = [] - has_noisy_vision = False for latent, noisy_frame_indexes in zip(vision_tokens, vision_noisy_frame_indexes): if not isinstance(latent, torch.Tensor) or not isinstance(noisy_frame_indexes, torch.Tensor): return None @@ -340,10 +340,12 @@ def _prepare_cosmos3_raw_vision_metadata( noisy_frame_indexes = noisy_frame_indexes.flatten().to(device=latent.device, dtype=torch.long) if torch.any(noisy_frame_indexes < 0) or torch.any(noisy_frame_indexes >= latent.shape[1]): return None - has_noisy_vision = has_noisy_vision or noisy_frame_indexes.numel() > 0 - raw_vision.append(latent) + # A sequence with noisy frames belongs to the generated output. Keep that sequence whole so clean conditioning + # frames remain in the indicator, but exclude separate clean hints that are not part of the output. + if noisy_frame_indexes.numel() > 0: + raw_vision.append(latent) - return raw_vision if raw_vision and has_noisy_vision else None + return raw_vision or None def _prepare_wan_t2v_raw_vision_metadata( diff --git a/tests/models/transformers/test_models_transformer_cosmos3.py b/tests/models/transformers/test_models_transformer_cosmos3.py index 6b04dd77f8c6..c897d5bfdc69 100644 --- a/tests/models/transformers/test_models_transformer_cosmos3.py +++ b/tests/models/transformers/test_models_transformer_cosmos3.py @@ -108,7 +108,106 @@ def output_shape(self) -> tuple[int, ...]: return (1, 2, 1, 1, 1) -class TestCosmos3OmniTransformerModel(Cosmos3OmniTransformerTesterConfig, ModelTesterMixin): +class TestCosmos3OmniTransformerSeaCache(Cosmos3OmniTransformerTesterConfig, SeaCacheTesterMixin): + cache_input_key = "vision_tokens" + + def test_sea_cache_tracks_output_visual_trajectory(self): + model = self.model_class(**self.get_init_dict()).to(torch_device).eval() + model.enable_cache(SeaCacheConfig(threshold=2.0, cache_end_steps=0)) + target = torch.randn(1, 2, 2, 1, 1, device=torch_device) + control = torch.randn_like(target) + target_with_changed_clean_frame = target.clone() + target_with_changed_clean_frame[:, :, 0] += 10 + inputs = self.get_dummy_inputs() + inputs.update( + sequence_length=6, + position_ids=torch.zeros(3, 6, dtype=torch.long, device=torch_device), + vision_tokens=[control, target], + vision_token_shapes=[(2, 1, 1)] * 2, + vision_sequence_indexes=torch.arange(2, 6, device=torch_device), + vision_mse_loss_indexes=torch.tensor([5], device=torch_device), + vision_noisy_frame_indexes=[ + torch.tensor([], dtype=torch.long, device=torch_device), + torch.tensor([1], device=torch_device), + ], + ) + layer_calls = 0 + + def count_layer_calls(_module, _args, _output): + nonlocal layer_calls + layer_calls += 1 + + model.layers[0].register_forward_hook(count_layer_calls) + decisions = [] + for step, (current_control, current_target) in enumerate( + ( + (control, target), + (control + 100, target), + (control + 100, target_with_changed_clean_frame), + ) + ): + inputs["vision_tokens"] = [current_control, current_target] + with ( + torch.no_grad(), + model.cache_context("cond", step_index=step, sigma=0.9 - step * 0.3, num_inference_steps=3), + ): + model(**inputs) + state = model._diffusers_hook.get_hook(_SEA_CACHE_ROOT_HOOK).state_manager._state_cache["cond"] + decisions.append(state.gate_should_compute) + + assert decisions == [True, False, True] + assert layer_calls == 2 + + @pytest.mark.parametrize("residual_order", [0, 1]) + def test_sea_cache_transfer_branches_share_indicator_with_separate_histories(self, residual_order): + model = self.model_class(**self.get_init_dict()).to(torch_device).eval() + model.enable_cache(SeaCacheConfig(threshold=100.0, residual_order=residual_order, cache_end_steps=0)) + root_hook = model._diffusers_hook.get_hook(_SEA_CACHE_ROOT_HOOK) + target = torch.randn(1, 2, 2, 1, 1, device=torch_device) + control = torch.randn_like(target) + decisions = [] + + for step in range(6): + states = [] + for context, with_control in (("cond", True), ("cond_no_control", False), ("uncond", True)): + inputs = self.get_dummy_inputs() + sequence_length = 6 if with_control else 4 + inputs.update( + sequence_length=sequence_length, + position_ids=torch.zeros(3, sequence_length, dtype=torch.long, device=torch_device), + vision_tokens=[control, target] if with_control else [target], + vision_token_shapes=[(2, 1, 1)] * (2 if with_control else 1), + vision_sequence_indexes=torch.arange(2, sequence_length, device=torch_device), + vision_mse_loss_indexes=torch.tensor([sequence_length - 1], device=torch_device), + vision_noisy_frame_indexes=( + [ + torch.tensor([], dtype=torch.long, device=torch_device), + torch.tensor([1], device=torch_device), + ] + if with_control + else [torch.tensor([1], device=torch_device)] + ), + ) + with ( + torch.no_grad(), + model.cache_context(context, step_index=step, sigma=0.9 - step * 0.1, num_inference_steps=6), + ): + output = model(**inputs) + assert torch.isfinite(output.sample[-1]).all() + states.append(root_hook.state_manager._state_cache[context]) + + assert all(len(state.previous_indicator) == 1 for state in states) + for state in states[1:]: + torch.testing.assert_close(state.previous_indicator[0], states[0].previous_indicator[0]) + assert state.gate_should_compute == states[0].gate_should_compute + assert state.history is not states[0].history + assert states[0].history[-1][2].shape != states[1].history[-1][2].shape + decisions.append(states[0].gate_should_compute) + + assert any(decisions) and not all(decisions) + model._reset_stateful_cache() + assert all(not state.history and state.previous_indicator is None for state in states) + def test_cosmos3_supports_sea_cache_without_changing_state_dict_keys(self): model = self.model_class(**self.get_init_dict()).to(torch_device).eval() state_dict_keys = set(model.state_dict()) @@ -423,6 +522,8 @@ def test_cosmos3_sea_cache_regional_compile_fullgraph_without_recompile(self): assert refreshed.sample[0].shape == self.output_shape assert root_hook.state_manager._state_cache["cond"].history[-1][0] == 2 + +class TestCosmos3OmniTransformerModel(Cosmos3OmniTransformerTesterConfig, ModelTesterMixin): def test_cosmos3_decoder_layer_cache_metadata_tracks_generation_stream(self): metadata = TransformerBlockRegistry.get(Cosmos3VLTextMoTDecoderLayer) @@ -571,10 +672,6 @@ def test_cosmos3_nemotron_rms_norm_multiplies_in_float32(self): torch.testing.assert_close(norm(hidden_states), expected, rtol=0, atol=0) -class TestCosmos3OmniTransformerSeaCache(Cosmos3OmniTransformerTesterConfig, SeaCacheTesterMixin): - cache_input_key = "vision_tokens" - - class TestCosmos3OmniTransformerMemory(Cosmos3OmniTransformerTesterConfig, MemoryTesterMixin): @pytest.mark.skip("The transformer returns one tensor list per generated modality.") def test_layerwise_casting_training(self): From c60830ee365d520ab52b110dda562dd26f7b4d7f Mon Sep 17 00:00:00 2001 From: Steven Liu <59462357+stevhliu@users.noreply.github.com> Date: Wed, 30 Sep 2026 08:59:23 -0700 Subject: [PATCH 10/22] [fix] Add return types (#14874) * add return types * add return blocks * style * modular pipelines * add a test to validate * feedback - skip deprecated stuff --- .github/workflows/pr_modular_tests.yml | 1 + .github/workflows/pr_tests.yml | 1 + .github/workflows/pr_tests_gpu.yml | 1 + Makefile | 5 + src/diffusers/models/activations.py | 8 +- src/diffusers/models/attention.py | 4 +- src/diffusers/models/attention_dispatch.py | 10 +- src/diffusers/models/attention_processor.py | 2 +- .../autoencoders/autoencoder_cosmos3_audio.py | 8 +- .../autoencoder_kl_hunyuanimage.py | 4 +- .../autoencoder_kl_hunyuanimage_refiner.py | 6 +- .../autoencoder_kl_hunyuanvideo15.py | 6 +- .../autoencoders/autoencoder_kl_qwenimage.py | 20 +- .../autoencoder_kl_qwenimage21.py | 24 +-- .../models/autoencoders/autoencoder_kl_wan.py | 24 +-- .../autoencoders/autoencoder_oobleck.py | 12 +- .../models/controlnets/controlnet.py | 2 +- .../models/controlnets/controlnet_hunyuan.py | 13 +- .../models/controlnets/controlnet_union.py | 4 +- .../models/controlnets/controlnet_xs.py | 2 +- .../models/controlnets/controlnet_z_image.py | 14 +- src/diffusers/models/embeddings.py | 58 +++--- src/diffusers/models/lora.py | 2 +- src/diffusers/models/normalization.py | 8 +- src/diffusers/models/resnet.py | 2 +- .../models/transformers/dit_transformer_2d.py | 2 +- .../transformers/dual_transformer_2d.py | 3 +- .../transformers/hunyuan_transformer_2d.py | 6 +- .../transformers/latte_transformer_3d.py | 2 +- .../transformers/pixart_transformer_2d.py | 2 +- .../models/transformers/prior_transformer.py | 2 +- .../models/transformers/sana_transformer.py | 4 +- .../transformers/stable_audio_transformer.py | 2 +- .../transformers/t5_film_transformer.py | 5 +- .../models/transformers/transformer_2d.py | 2 +- .../transformers/transformer_2d_dreamlite.py | 2 +- .../transformers/transformer_allegro.py | 2 +- .../transformers/transformer_anyflow.py | 4 +- .../transformers/transformer_anyflow_far.py | 4 +- .../models/transformers/transformer_bria.py | 4 +- .../transformers/transformer_bria_fibo.py | 6 +- .../models/transformers/transformer_chroma.py | 2 +- .../transformers/transformer_chronoedit.py | 2 +- .../transformers/transformer_cosmos3.py | 4 +- .../transformers/transformer_ernie_image.py | 9 +- .../models/transformers/transformer_helios.py | 6 +- .../transformers/transformer_hidream_image.py | 6 +- .../transformer_hunyuan_video_framepack.py | 6 +- .../transformers/transformer_joyimage.py | 8 +- .../transformer_joyimage_edit_plus.py | 2 +- .../transformers/transformer_kandinsky.py | 20 +- .../transformers/transformer_longcat_image.py | 2 +- .../transformers/transformer_lumina2.py | 4 +- .../models/transformers/transformer_mochi.py | 2 +- .../transformer_nucleusmoe_image.py | 2 +- .../transformers/transformer_omnigen.py | 2 +- .../transformers/transformer_qwenimage.py | 2 +- .../transformers/transformer_sana_video.py | 4 +- .../models/transformers/transformer_sd3.py | 2 +- .../transformers/transformer_skyreels_v2.py | 2 +- .../transformers/transformer_temporal.py | 2 +- .../models/transformers/transformer_wan.py | 2 +- .../transformers/transformer_wan_animate.py | 2 +- .../transformers/transformer_wan_animate_2.py | 4 +- .../transformers/transformer_z_image.py | 14 +- src/diffusers/models/unets/unet_3d_blocks.py | 2 +- src/diffusers/models/unets/unet_kandinsky3.py | 24 ++- .../models/unets/unet_motion_model.py | 4 +- .../models/unets/unet_stable_cascade.py | 16 +- src/diffusers/models/unets/uvit_2d.py | 16 +- .../modular_pipelines/anima/before_denoise.py | 28 ++- .../modular_pipelines/anima/decoders.py | 8 +- .../modular_pipelines/anima/denoise.py | 14 +- .../modular_pipelines/anima/encoders.py | 8 +- .../modular_pipelines/cosmos/after_decode.py | 8 +- .../cosmos/before_denoise.py | 60 ++++-- .../cosmos/before_encoder.py | 4 +- .../modular_pipelines/cosmos/decoders.py | 16 +- .../modular_pipelines/cosmos/denoise.py | 48 +++-- .../modular_pipelines/cosmos/encoders.py | 32 +++- .../cosmos/modular_blocks_cosmos3.py | 4 +- .../ernie_image/before_denoise.py | 12 +- .../modular_pipelines/ernie_image/decoders.py | 4 +- .../modular_pipelines/ernie_image/denoise.py | 16 +- .../modular_pipelines/ernie_image/encoders.py | 8 +- .../modular_pipelines/flux/before_denoise.py | 24 ++- .../modular_pipelines/flux/decoders.py | 3 +- .../modular_pipelines/flux/denoise.py | 12 +- .../modular_pipelines/flux/encoders.py | 16 +- .../modular_pipelines/flux/inputs.py | 16 +- .../modular_pipelines/flux2/before_denoise.py | 24 ++- .../modular_pipelines/flux2/decoders.py | 5 +- .../modular_pipelines/flux2/denoise.py | 14 +- .../modular_pipelines/flux2/encoders.py | 20 +- .../modular_pipelines/flux2/inputs.py | 12 +- .../helios/before_denoise.py | 32 +++- .../modular_pipelines/helios/decoders.py | 3 +- .../modular_pipelines/helios/denoise.py | 40 +++- .../modular_pipelines/helios/encoders.py | 12 +- .../hunyuan_video1_5/before_denoise.py | 16 +- .../hunyuan_video1_5/decoders.py | 3 +- .../hunyuan_video1_5/denoise.py | 16 +- .../hunyuan_video1_5/encoders.py | 12 +- .../ideogram4/before_denoise.py | 16 +- .../modular_pipelines/ideogram4/decoders.py | 4 +- .../modular_pipelines/ideogram4/denoise.py | 20 +- .../modular_pipelines/ideogram4/encoders.py | 8 +- .../modular_pipelines/krea2/before_denoise.py | 24 ++- .../modular_pipelines/krea2/decoders.py | 4 +- .../modular_pipelines/krea2/denoise.py | 20 +- .../modular_pipelines/krea2/encoders.py | 8 +- .../modular_pipelines/ltx/before_denoise.py | 16 +- .../modular_pipelines/ltx/decoders.py | 4 +- .../modular_pipelines/ltx/denoise.py | 24 ++- .../modular_pipelines/ltx/encoders.py | 8 +- .../modular_pipelines/ltx2/before_denoise.py | 25 +-- .../modular_pipelines/ltx2/decoders.py | 9 +- .../modular_pipelines/ltx2/denoise.py | 31 ++- .../modular_pipelines/ltx2/encoders.py | 19 +- .../minimax_h3/before_denoise.py | 32 +++- .../minimax_h3/before_encoder.py | 8 +- .../modular_pipelines/minimax_h3/decoders.py | 12 +- .../modular_pipelines/minimax_h3/denoise.py | 12 +- .../modular_pipelines/minimax_h3/encoders.py | 20 +- .../minimax_music3/before_denoise.py | 4 +- .../minimax_music3/decoders.py | 4 +- .../minimax_music3/denoise.py | 24 ++- .../minimax_music3/encoders.py | 8 +- .../modular_pipelines/modular_pipeline.py | 6 +- .../qwenimage/before_denoise.py | 44 +++-- .../modular_pipelines/qwenimage/decoders.py | 20 +- .../modular_pipelines/qwenimage/denoise.py | 32 +++- .../modular_pipelines/qwenimage/encoders.py | 58 ++++-- .../modular_pipelines/qwenimage/inputs.py | 20 +- .../stable_diffusion_3/before_denoise.py | 16 +- .../stable_diffusion_3/decoders.py | 3 +- .../stable_diffusion_3/denoise.py | 8 +- .../stable_diffusion_3/encoders.py | 12 +- .../stable_diffusion_3/inputs.py | 8 +- .../stable_diffusion_xl/before_denoise.py | 40 +++- .../stable_diffusion_xl/decoders.py | 5 +- .../stable_diffusion_xl/denoise.py | 26 ++- .../stable_diffusion_xl/encoders.py | 16 +- .../modular_pipelines/wan/before_denoise.py | 20 +- .../modular_pipelines/wan/decoders.py | 5 +- .../modular_pipelines/wan/denoise.py | 20 +- .../modular_pipelines/wan/encoders.py | 40 +++- .../wan_animate_2/before_denoise.py | 3 +- .../wan_animate_2/decoders.py | 3 +- .../wan_animate_2/denoise.py | 17 +- .../wan_animate_2/encoders.py | 13 +- .../z_image/before_denoise.py | 24 ++- .../modular_pipelines/z_image/decoders.py | 3 +- .../modular_pipelines/z_image/denoise.py | 14 +- .../modular_pipelines/z_image/encoders.py | 8 +- .../pipelines/ace_step/pipeline_ace_step.py | 2 +- .../animatediff/pipeline_animatediff.py | 2 +- .../pipeline_animatediff_controlnet.py | 2 +- .../animatediff/pipeline_animatediff_sdxl.py | 2 +- .../pipeline_animatediff_sparsectrl.py | 2 +- .../pipeline_animatediff_video2video.py | 2 +- ...line_animatediff_video2video_controlnet.py | 2 +- .../pipelines/anyflow/pipeline_anyflow.py | 2 +- .../pipelines/anyflow/pipeline_anyflow_far.py | 2 +- .../pipelines/audioldm2/pipeline_audioldm2.py | 2 +- src/diffusers/pipelines/bria/pipeline_bria.py | 2 +- .../pipelines/bria_fibo/pipeline_bria_fibo.py | 2 +- .../bria_fibo/pipeline_bria_fibo_edit.py | 2 +- .../pipelines/chroma/pipeline_chroma.py | 2 +- .../chroma/pipeline_chroma_img2img.py | 2 +- .../chroma/pipeline_chroma_inpainting.py | 2 +- .../chronoedit/pipeline_chronoedit.py | 2 +- .../pipeline_consistency_models.py | 2 +- .../controlnet/pipeline_controlnet.py | 2 +- .../pipeline_controlnet_blip_diffusion.py | 2 +- .../controlnet/pipeline_controlnet_img2img.py | 2 +- .../controlnet/pipeline_controlnet_inpaint.py | 2 +- .../pipeline_controlnet_inpaint_sd_xl.py | 2 +- .../controlnet/pipeline_controlnet_sd_xl.py | 2 +- .../pipeline_controlnet_sd_xl_img2img.py | 2 +- ...pipeline_controlnet_union_inpaint_sd_xl.py | 2 +- .../pipeline_controlnet_union_sd_xl.py | 2 +- ...pipeline_controlnet_union_sd_xl_img2img.py | 2 +- .../pipeline_hunyuandit_controlnet.py | 2 +- .../pipeline_stable_diffusion_3_controlnet.py | 2 +- ...table_diffusion_3_controlnet_inpainting.py | 2 +- .../cosmos/pipeline_cosmos2_5_predict.py | 2 +- .../cosmos/pipeline_cosmos2_5_transfer.py | 2 +- .../cosmos/pipeline_cosmos2_text2image.py | 2 +- .../cosmos/pipeline_cosmos2_video2world.py | 2 +- .../pipelines/cosmos/pipeline_cosmos3_omni.py | 2 +- .../cosmos/pipeline_cosmos_text2world.py | 2 +- .../cosmos/pipeline_cosmos_video2world.py | 2 +- .../pipelines/deepfloyd_if/pipeline_if.py | 2 +- .../deepfloyd_if/pipeline_if_img2img.py | 2 +- .../pipeline_if_img2img_superresolution.py | 2 +- .../deepfloyd_if/pipeline_if_inpainting.py | 2 +- .../pipeline_if_inpainting_superresolution.py | 2 +- .../pipeline_if_superresolution.py | 2 +- .../pipeline_diffusion_gemma.py | 2 +- .../pipelines/dreamlite/pipeline_dreamlite.py | 2 +- .../dreamlite/pipeline_dreamlite_mobile.py | 2 +- .../easyanimate/pipeline_easyanimate.py | 2 +- .../pipeline_easyanimate_control.py | 2 +- .../pipeline_easyanimate_inpaint.py | 2 +- .../ernie_image/pipeline_ernie_image.py | 2 +- src/diffusers/pipelines/flux/pipeline_flux.py | 2 +- .../pipelines/flux/pipeline_flux_control.py | 2 +- .../flux/pipeline_flux_control_img2img.py | 2 +- .../flux/pipeline_flux_control_inpaint.py | 2 +- .../flux/pipeline_flux_controlnet.py | 2 +- ...pipeline_flux_controlnet_image_to_image.py | 2 +- .../pipeline_flux_controlnet_inpainting.py | 2 +- .../pipelines/flux/pipeline_flux_fill.py | 2 +- .../pipelines/flux/pipeline_flux_img2img.py | 2 +- .../pipelines/flux/pipeline_flux_inpaint.py | 2 +- .../pipelines/flux/pipeline_flux_kontext.py | 2 +- .../flux/pipeline_flux_kontext_inpaint.py | 2 +- .../flux/pipeline_flux_prior_redux.py | 2 +- .../pipelines/flux2/pipeline_flux2.py | 2 +- .../pipelines/flux2/pipeline_flux2_klein.py | 2 +- .../flux2/pipeline_flux2_klein_inpaint.py | 2 +- .../flux2/pipeline_flux2_klein_kv.py | 2 +- .../pipelines/helios/pipeline_helios.py | 2 +- .../helios/pipeline_helios_pyramid.py | 2 +- .../hidream_image/pipeline_hidream_image.py | 3 +- .../hunyuan_image/pipeline_hunyuanimage.py | 2 +- .../pipeline_hunyuanimage_refiner.py | 2 +- .../pipeline_hunyuan_skyreels_image2video.py | 2 +- .../hunyuan_video/pipeline_hunyuan_video.py | 2 +- .../pipeline_hunyuan_video_framepack.py | 2 +- .../pipeline_hunyuan_video_image2video.py | 2 +- .../pipeline_hunyuan_video1_5.py | 2 +- .../pipeline_hunyuan_video1_5_image2video.py | 2 +- .../hunyuandit/pipeline_hunyuandit.py | 2 +- .../pipelines/ideogram4/pipeline_ideogram4.py | 2 +- .../joyimage/pipeline_joyimage_edit.py | 2 +- .../joyimage/pipeline_joyimage_edit_plus.py | 2 +- .../pipelines/kandinsky/pipeline_kandinsky.py | 2 +- .../kandinsky/pipeline_kandinsky_combined.py | 8 +- .../kandinsky/pipeline_kandinsky_img2img.py | 2 +- .../kandinsky/pipeline_kandinsky_inpaint.py | 2 +- .../kandinsky/pipeline_kandinsky_prior.py | 2 +- .../kandinsky2_2/pipeline_kandinsky2_2.py | 2 +- .../pipeline_kandinsky2_2_combined.py | 8 +- .../pipeline_kandinsky2_2_controlnet.py | 2 +- ...ipeline_kandinsky2_2_controlnet_img2img.py | 2 +- .../pipeline_kandinsky2_2_img2img.py | 2 +- .../pipeline_kandinsky2_2_inpainting.py | 2 +- .../pipeline_kandinsky2_2_prior.py | 2 +- .../pipeline_kandinsky2_2_prior_emb2emb.py | 2 +- .../kandinsky3/pipeline_kandinsky3.py | 2 +- .../kandinsky3/pipeline_kandinsky3_img2img.py | 2 +- .../kandinsky5/pipeline_kandinsky.py | 2 +- .../kandinsky5/pipeline_kandinsky_i2i.py | 2 +- .../kandinsky5/pipeline_kandinsky_i2v.py | 2 +- .../kandinsky5/pipeline_kandinsky_t2i.py | 2 +- .../pipelines/kolors/pipeline_kolors.py | 2 +- .../kolors/pipeline_kolors_img2img.py | 2 +- .../pipelines/krea2/pipeline_krea2.py | 2 +- .../pipeline_latent_consistency_img2img.py | 2 +- .../pipeline_latent_consistency_text2img.py | 2 +- .../pipeline_latent_diffusion.py | 2 +- ...peline_latent_diffusion_superresolution.py | 2 +- .../pipeline_leditspp_stable_diffusion.py | 2 +- .../pipeline_leditspp_stable_diffusion_xl.py | 2 +- .../pipelines/llada2/pipeline_llada2.py | 2 +- .../pipeline_longcat_audio_dit.py | 6 +- .../longcat_image/pipeline_longcat_image.py | 2 +- .../pipeline_longcat_image_edit.py | 2 +- src/diffusers/pipelines/ltx/pipeline_ltx.py | 2 +- .../pipelines/ltx/pipeline_ltx_condition.py | 2 +- .../ltx/pipeline_ltx_i2v_long_multi_prompt.py | 2 +- .../pipelines/ltx/pipeline_ltx_image2video.py | 2 +- .../ltx/pipeline_ltx_latent_upsample.py | 7 +- src/diffusers/pipelines/ltx2/pipeline_ltx2.py | 2 +- .../pipelines/ltx2/pipeline_ltx2_condition.py | 2 +- .../pipelines/ltx2/pipeline_ltx2_dfr.py | 2 +- .../ltx2/pipeline_ltx2_dfr_temporal_refine.py | 2 +- .../ltx2/pipeline_ltx2_diffusion_decode.py | 2 +- .../pipelines/ltx2/pipeline_ltx2_hdr_lora.py | 2 +- .../pipelines/ltx2/pipeline_ltx2_ic_lora.py | 2 +- .../ltx2/pipeline_ltx2_image2video.py | 2 +- .../ltx2/pipeline_ltx2_latent_upsample.py | 2 +- .../pipelines/lucy/pipeline_lucy_edit.py | 2 +- .../marigold/pipeline_marigold_depth.py | 2 +- .../marigold/pipeline_marigold_intrinsics.py | 2 +- .../marigold/pipeline_marigold_normals.py | 2 +- .../pipelines/mochi/pipeline_mochi.py | 2 +- .../motif_video/pipeline_motif_video.py | 2 +- .../pipeline_motif_video_image2video.py | 2 +- .../pipeline_nucleusmoe_image.py | 2 +- .../pipelines/omnigen/pipeline_omnigen.py | 6 +- .../ovis_image/pipeline_ovis_image.py | 2 +- .../pag/pipeline_pag_controlnet_sd.py | 2 +- .../pag/pipeline_pag_controlnet_sd_inpaint.py | 2 +- .../pag/pipeline_pag_controlnet_sd_xl.py | 2 +- .../pipeline_pag_controlnet_sd_xl_img2img.py | 2 +- .../pipelines/pag/pipeline_pag_hunyuandit.py | 2 +- .../pipelines/pag/pipeline_pag_kolors.py | 2 +- .../pipelines/pag/pipeline_pag_sd.py | 2 +- .../pipelines/pag/pipeline_pag_sd_3.py | 2 +- .../pag/pipeline_pag_sd_3_img2img.py | 2 +- .../pag/pipeline_pag_sd_animatediff.py | 2 +- .../pipelines/pag/pipeline_pag_sd_img2img.py | 2 +- .../pipelines/pag/pipeline_pag_sd_inpaint.py | 2 +- .../pipelines/pag/pipeline_pag_sd_xl.py | 2 +- .../pag/pipeline_pag_sd_xl_img2img.py | 2 +- .../pag/pipeline_pag_sd_xl_inpaint.py | 2 +- src/diffusers/pipelines/prx/pipeline_prx.py | 2 +- .../pipelines/prx/pipeline_prx_pixel.py | 2 +- .../pipelines/qwenimage/pipeline_qwenimage.py | 2 +- .../pipeline_qwenimage_controlnet.py | 2 +- .../pipeline_qwenimage_controlnet_inpaint.py | 2 +- .../qwenimage/pipeline_qwenimage_edit.py | 2 +- .../pipeline_qwenimage_edit_inpaint.py | 2 +- .../qwenimage/pipeline_qwenimage_edit_plus.py | 2 +- .../qwenimage/pipeline_qwenimage_img2img.py | 2 +- .../qwenimage/pipeline_qwenimage_inpaint.py | 2 +- .../qwenimage/pipeline_qwenimage_layered.py | 2 +- .../qwenimage21/pipeline_qwenimage21.py | 2 +- .../pipelines/shap_e/pipeline_shap_e.py | 2 +- .../shap_e/pipeline_shap_e_img2img.py | 2 +- .../skyreels_v2/pipeline_skyreels_v2.py | 2 +- .../pipeline_skyreels_v2_diffusion_forcing.py | 2 +- ...eline_skyreels_v2_diffusion_forcing_i2v.py | 2 +- ...eline_skyreels_v2_diffusion_forcing_v2v.py | 2 +- .../skyreels_v2/pipeline_skyreels_v2_i2v.py | 2 +- .../stable_audio/pipeline_stable_audio.py | 2 +- .../stable_audio_3/pipeline_stable_audio_3.py | 2 +- .../pipeline_stable_audio_3_audio2audio.py | 2 +- .../pipeline_stable_audio_3_inpaint.py | 2 +- .../stable_cascade/pipeline_stable_cascade.py | 4 +- .../pipeline_stable_cascade_combined.py | 5 +- .../pipeline_stable_cascade_prior.py | 2 +- .../pipeline_onnx_stable_diffusion.py | 2 +- .../pipeline_onnx_stable_diffusion_img2img.py | 2 +- .../pipeline_onnx_stable_diffusion_inpaint.py | 2 +- .../pipeline_onnx_stable_diffusion_upscale.py | 2 +- .../pipeline_stable_diffusion.py | 2 +- .../pipeline_stable_diffusion_depth2img.py | 2 +- ...peline_stable_diffusion_image_variation.py | 2 +- .../pipeline_stable_diffusion_img2img.py | 2 +- .../pipeline_stable_diffusion_inpaint.py | 2 +- ...eline_stable_diffusion_instruct_pix2pix.py | 2 +- ...ipeline_stable_diffusion_latent_upscale.py | 2 +- .../pipeline_stable_diffusion_upscale.py | 2 +- .../pipeline_stable_unclip.py | 2 +- .../pipeline_stable_unclip_img2img.py | 2 +- .../pipeline_stable_diffusion_3.py | 2 +- .../pipeline_stable_diffusion_3_img2img.py | 2 +- .../pipeline_stable_diffusion_3_inpaint.py | 2 +- .../pipeline_stable_diffusion_xl.py | 2 +- .../pipeline_stable_diffusion_xl_img2img.py | 2 +- .../pipeline_stable_diffusion_xl_inpaint.py | 2 +- ...ne_stable_diffusion_xl_instruct_pix2pix.py | 2 +- .../pipeline_stable_video_diffusion.py | 2 +- .../pipeline_stable_diffusion_adapter.py | 2 +- .../pipeline_stable_diffusion_xl_adapter.py | 2 +- .../pipeline_visualcloze_combined.py | 2 +- .../pipeline_visualcloze_generation.py | 2 +- src/diffusers/pipelines/wan/pipeline_wan.py | 2 +- .../pipelines/wan/pipeline_wan_animate.py | 2 +- .../pipelines/wan/pipeline_wan_i2v.py | 2 +- .../pipelines/wan/pipeline_wan_vace.py | 2 +- .../pipelines/wan/pipeline_wan_video2video.py | 2 +- .../pipelines/z_image/pipeline_z_image.py | 2 +- .../z_image/pipeline_z_image_controlnet.py | 2 +- .../pipeline_z_image_controlnet_inpaint.py | 2 +- .../z_image/pipeline_z_image_img2img.py | 2 +- .../z_image/pipeline_z_image_inpaint.py | 2 +- .../z_image/pipeline_z_image_omni.py | 2 +- utils/check_return_annotations.py | 180 ++++++++++++++++++ 373 files changed, 1700 insertions(+), 809 deletions(-) create mode 100644 utils/check_return_annotations.py diff --git a/.github/workflows/pr_modular_tests.yml b/.github/workflows/pr_modular_tests.yml index f4bf9585bb7d..e6f156235a53 100644 --- a/.github/workflows/pr_modular_tests.yml +++ b/.github/workflows/pr_modular_tests.yml @@ -83,6 +83,7 @@ jobs: python utils/check_dummies.py python utils/check_support_list.py python utils/check_forward_call_docstrings.py + python utils/check_return_annotations.py make deps_table_check_updated - name: Check if failure if: ${{ failure() }} diff --git a/.github/workflows/pr_tests.yml b/.github/workflows/pr_tests.yml index 2034673af942..b5179e82a028 100644 --- a/.github/workflows/pr_tests.yml +++ b/.github/workflows/pr_tests.yml @@ -78,6 +78,7 @@ jobs: python utils/check_dummies.py python utils/check_support_list.py python utils/check_forward_call_docstrings.py + python utils/check_return_annotations.py make deps_table_check_updated - name: Check if failure if: ${{ failure() }} diff --git a/.github/workflows/pr_tests_gpu.yml b/.github/workflows/pr_tests_gpu.yml index 4be060555b8d..2cd321e71160 100644 --- a/.github/workflows/pr_tests_gpu.yml +++ b/.github/workflows/pr_tests_gpu.yml @@ -79,6 +79,7 @@ jobs: python utils/check_dummies.py python utils/check_support_list.py python utils/check_forward_call_docstrings.py + python utils/check_return_annotations.py make deps_table_check_updated - name: Check if failure if: ${{ failure() }} diff --git a/Makefile b/Makefile index 4af58740480e..01dcc65448bc 100644 --- a/Makefile +++ b/Makefile @@ -37,6 +37,7 @@ repo-consistency: python utils/check_repo.py python utils/check_inits.py python utils/check_forward_call_docstrings.py + python utils/check_return_annotations.py # this target runs checks on all files @@ -80,6 +81,10 @@ modular-autodoctrings: check-forward-call-docstrings: python utils/check_forward_call_docstrings.py +# Verify forward() / __call__() have return type annotations +check-return-annotations: + python utils/check_return_annotations.py + # Run tests for the library test: diff --git a/src/diffusers/models/activations.py b/src/diffusers/models/activations.py index 2d1fdb5f7d83..caff731dd599 100644 --- a/src/diffusers/models/activations.py +++ b/src/diffusers/models/activations.py @@ -84,7 +84,7 @@ def gelu(self, gate: torch.Tensor) -> torch.Tensor: return F.gelu(gate.to(dtype=torch.float32), approximate=self.approximate).to(dtype=gate.dtype) return F.gelu(gate, approximate=self.approximate) - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: hidden_states = self.proj(hidden_states) hidden_states = self.gelu(hidden_states) return hidden_states @@ -110,7 +110,7 @@ def gelu(self, gate: torch.Tensor) -> torch.Tensor: return F.gelu(gate.to(dtype=torch.float32)).to(dtype=gate.dtype) return F.gelu(gate) - def forward(self, hidden_states, *args, **kwargs): + def forward(self, hidden_states, *args, **kwargs) -> torch.Tensor: if len(args) > 0 or kwargs.get("scale", None) is not None: deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." deprecate("scale", "1.0.0", deprecation_message) @@ -140,7 +140,7 @@ def __init__(self, dim_in: int, dim_out: int, bias: bool = True): self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias) self.activation = nn.SiLU() - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: hidden_states = self.proj(hidden_states) hidden_states, gate = hidden_states.chunk(2, dim=-1) return hidden_states * self.activation(gate) @@ -173,6 +173,6 @@ def __init__(self, dim_in: int, dim_out: int, bias: bool = True, activation: str self.proj = nn.Linear(dim_in, dim_out, bias=bias) self.activation = get_activation(activation) - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: hidden_states = self.proj(hidden_states) return self.activation(hidden_states) diff --git a/src/diffusers/models/attention.py b/src/diffusers/models/attention.py index 65289e4b5f16..504f4afe3af3 100644 --- a/src/diffusers/models/attention.py +++ b/src/diffusers/models/attention.py @@ -1125,7 +1125,7 @@ def __init__( ) self.silu = FP32SiLU() - def forward(self, x): + def forward(self, x) -> torch.Tensor: return self.linear_2(self.silu(self.linear_1(x)) * self.linear_3(x)) @@ -1302,7 +1302,7 @@ def __init__( out_bias=attention_out_bias, ) - def forward(self, hidden_states, encoder_hidden_states, cross_attention_kwargs): + def forward(self, hidden_states, encoder_hidden_states, cross_attention_kwargs) -> torch.Tensor: cross_attention_kwargs = cross_attention_kwargs.copy() if cross_attention_kwargs is not None else {} if self.kv_mapper is not None: diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index 364b7b057e78..237115bac5a7 100644 --- a/src/diffusers/models/attention_dispatch.py +++ b/src/diffusers/models/attention_dispatch.py @@ -2227,7 +2227,7 @@ class SeqAllToAllDim(torch.autograd.Function): """ @staticmethod - def forward(ctx, group, input, scatter_id=2, gather_id=1): + def forward(ctx, group, input, scatter_id=2, gather_id=1) -> torch.Tensor: ctx.group = group ctx.scatter_id = scatter_id ctx.gather_id = gather_id @@ -2408,7 +2408,7 @@ def forward( forward_op, backward_op, _parallel_config: "ParallelConfig" | None = None, - ): + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: ring_mesh = _parallel_config.context_parallel_config._ring_mesh rank = _parallel_config.context_parallel_config._ring_local_rank world_size = _parallel_config.context_parallel_config.ring_degree @@ -2561,7 +2561,7 @@ def forward( forward_op, backward_op, _parallel_config: "ParallelConfig" | None = None, - ): + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: ulysses_mesh = _parallel_config.context_parallel_config._ulysses_mesh world_size = _parallel_config.context_parallel_config.ulysses_degree group = ulysses_mesh.get_group() @@ -2662,7 +2662,7 @@ def forward( forward_op, backward_op, _parallel_config: "ParallelConfig" | None = None, - ): + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: # Ring attention for arbitrary sequence lengths. if attn_mask is not None: raise ValueError( @@ -2779,7 +2779,7 @@ def forward( backward_op, _parallel_config: "ParallelConfig" | None = None, **kwargs, - ): + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: ulysses_mesh = _parallel_config.context_parallel_config._ulysses_mesh group = ulysses_mesh.get_group() diff --git a/src/diffusers/models/attention_processor.py b/src/diffusers/models/attention_processor.py index 1b923e749663..62484a113f37 100755 --- a/src/diffusers/models/attention_processor.py +++ b/src/diffusers/models/attention_processor.py @@ -985,7 +985,7 @@ def forward( encoder_hidden_states: torch.Tensor | None = None, attention_mask: torch.Tensor | None = None, **kwargs, - ): + ) -> tuple[torch.Tensor, torch.Tensor]: return self.processor( self, hidden_states, diff --git a/src/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py b/src/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py index e5549a47e9f1..a12bcc5113bc 100644 --- a/src/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py +++ b/src/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py @@ -53,7 +53,7 @@ def __init__(self, hidden_dim, logscale=True): self.beta.requires_grad = True self.logscale = logscale - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: shape = hidden_states.shape alpha = self.alpha if not self.logscale else torch.exp(self.alpha) @@ -250,7 +250,7 @@ def __init__(self, dimension: int = 16, dilation: int = 1): self.snake2 = Snake1d(dimension) self.conv2 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=1)) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: """ Forward pass through the residual unit. @@ -300,7 +300,7 @@ def __init__(self, input_dim, output_dim, stride: int = 1, output_padding: int = self.res_unit2 = Cosmos3AudioResidualUnit(output_dim, dilation=3) self.res_unit3 = Cosmos3AudioResidualUnit(output_dim, dilation=9) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.snake1(hidden_state) hidden_state = self.conv_t1(hidden_state) hidden_state = self.res_unit1(hidden_state) @@ -345,7 +345,7 @@ def __init__(self, channels, input_channels, audio_channels, upsampling_ratios, self.snake1 = Snake1d(output_dim) self.conv2 = weight_norm(nn.Conv1d(channels, audio_channels, kernel_size=7, padding=3, bias=False)) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.conv1(hidden_state) for layer in self.block: diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage.py b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage.py index c1d975ae6bb7..4d46b7c32ec4 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage.py @@ -58,7 +58,7 @@ def __init__(self, in_channels: int, out_channels: int, non_linearity: str = "si else: self.conv_shortcut = None - def forward(self, x): + def forward(self, x) -> torch.Tensor: # Apply shortcut connection residual = x @@ -95,7 +95,7 @@ def __init__(self, in_channels: int): self.to_v = nn.Conv2d(in_channels, in_channels, 1) self.proj = nn.Conv2d(in_channels, in_channels, 1) - def forward(self, x): + def forward(self, x) -> torch.Tensor: identity = x x = self.norm(x) diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage_refiner.py b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage_refiner.py index 9737a822236e..091304660ba8 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage_refiner.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanimage_refiner.py @@ -86,7 +86,7 @@ def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bi self.gamma = nn.Parameter(torch.ones(shape)) self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - def forward(self, x): + def forward(self, x) -> torch.Tensor: needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( t in str(x.dtype) for t in ("float4_", "float8_") ) @@ -162,7 +162,7 @@ def _dcae_upsample_rearrange(tensor, r1=1, r2=2, r3=2): tensor = tensor.permute(0, 4, 5, 1, 6, 2, 7, 3) return tensor.reshape(b, c, f * r1, h * r2, w * r3) - def forward(self, x: torch.Tensor): + def forward(self, x: torch.Tensor) -> torch.Tensor: r1 = 2 if self.add_temporal_upsample else 1 h = self.conv(x) if self.add_temporal_upsample: @@ -205,7 +205,7 @@ def _dcae_downsample_rearrange(tensor, r1=1, r2=2, r3=2): tensor = tensor.permute(0, 3, 5, 7, 1, 2, 4, 6) return tensor.reshape(b, r1 * r2 * r3 * c, f, h, w) - def forward(self, x: torch.Tensor): + def forward(self, x: torch.Tensor) -> torch.Tensor: r1 = 2 if self.add_temporal_downsample else 1 h = self.conv(x) if self.add_temporal_downsample: diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanvideo15.py b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanvideo15.py index 9260e1fcbb1d..00e28ee620b6 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanvideo15.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_hunyuanvideo15.py @@ -86,7 +86,7 @@ def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bi self.gamma = nn.Parameter(torch.ones(shape)) self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - def forward(self, x): + def forward(self, x) -> torch.Tensor: needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( t in str(x.dtype) for t in ("float4_", "float8_") ) @@ -189,7 +189,7 @@ def _dcae_upsample_rearrange(tensor, r1=1, r2=2, r3=2): tensor = tensor.permute(0, 4, 5, 1, 6, 2, 7, 3) return tensor.reshape(b, c, f * r1, h * r2, w * r3) - def forward(self, x: torch.Tensor): + def forward(self, x: torch.Tensor) -> torch.Tensor: r1 = 2 if self.add_temporal_upsample else 1 h = self.conv(x) if self.add_temporal_upsample: @@ -240,7 +240,7 @@ def _dcae_downsample_rearrange(tensor, r1=1, r2=2, r3=2): tensor = tensor.permute(0, 3, 5, 7, 1, 2, 4, 6) return tensor.reshape(b, r1 * r2 * r3 * c, f, h, w) - def forward(self, x: torch.Tensor): + def forward(self, x: torch.Tensor) -> torch.Tensor: r1 = 2 if self.add_temporal_downsample else 1 h = self.conv(x) if self.add_temporal_downsample: diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage.py b/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage.py index 220520a12e68..fa9f31ac8512 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage.py @@ -72,7 +72,7 @@ def __init__( self._padding = (self.padding[2], self.padding[2], self.padding[1], self.padding[1], 2 * self.padding[0], 0) self.padding = (0, 0, 0) - def forward(self, x, cache_x=None): + def forward(self, x, cache_x=None) -> torch.Tensor: padding = list(self._padding) if cache_x is not None and self._padding[4] > 0: cache_x = cache_x.to(x.device) @@ -104,7 +104,7 @@ def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bi self.gamma = nn.Parameter(torch.ones(shape)) self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - def forward(self, x): + def forward(self, x) -> torch.Tensor: needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( t in str(x.dtype) for t in ("float4_", "float8_") ) @@ -126,7 +126,7 @@ class QwenImageUpsample(nn.Upsample): torch.Tensor: Upsampled tensor with the same data type as the input. """ - def forward(self, x): + def forward(self, x) -> torch.Tensor: return super().forward(x.float()).type_as(x) @@ -171,7 +171,7 @@ def __init__(self, dim: int, mode: str) -> None: else: self.resample = nn.Identity() - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: b, c, t, h, w = x.size() if self.mode == "upsample3d": if feat_cache is not None: @@ -248,7 +248,7 @@ def __init__( self.conv2 = QwenImageCausalConv3d(out_dim, out_dim, 3, padding=1) self.conv_shortcut = QwenImageCausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity() - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: # Apply shortcut connection h = self.conv_shortcut(x) @@ -308,7 +308,7 @@ def __init__(self, dim): self.to_qkv = nn.Conv2d(dim, dim * 3, 1) self.proj = nn.Conv2d(dim, dim, 1) - def forward(self, x): + def forward(self, x) -> torch.Tensor: identity = x batch_size, channels, time, height, width = x.size() @@ -361,7 +361,7 @@ def __init__(self, dim: int, dropout: float = 0.0, non_linearity: str = "silu", self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: # First residual block x = self.resnets[0](x, feat_cache, feat_idx) @@ -443,7 +443,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: if feat_cache is not None: idx = feat_idx[0] cache_x = x[:, :, -min(CACHE_T, x.shape[2]) :, :, :].clone() @@ -526,7 +526,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: """ Forward pass through the upsampling block. @@ -632,7 +632,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: ## conv1 if feat_cache is not None: idx = feat_idx[0] diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py b/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py index 15575a1c5907..3d99d7a69bf1 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py @@ -169,7 +169,7 @@ def __init__( self._padding = (self.padding[1], self.padding[1], self.padding[0], self.padding[0]) self.padding = (0, 0) - def forward(self, x, cache_x=None): + def forward(self, x, cache_x=None) -> torch.Tensor: padding = list(self._padding) if cache_x is not None: raise ValueError( @@ -206,7 +206,7 @@ def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bi self.gamma = nn.Parameter(torch.ones(shape)) self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - def forward(self, x): + def forward(self, x) -> torch.Tensor: needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( t in str(x.dtype) for t in ("float4_", "float8_") ) @@ -229,7 +229,7 @@ class QwenImage21Upsample(nn.Upsample): torch.Tensor: Upsampled tensor with the same data type as the input. """ - def forward(self, x): + def forward(self, x) -> torch.Tensor: return super().forward(x.float()).type_as(x) @@ -278,7 +278,7 @@ def __init__(self, dim: int, mode: str, upsample_out_dim: int = None) -> None: else: self.resample = nn.Identity() - def forward(self, x, feat_cache=None, feat_idx=None): + def forward(self, x, feat_cache=None, feat_idx=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] b, c, t, h, w = x.size() @@ -355,7 +355,7 @@ def __init__( self.conv2 = QwenImage21CausalConv3d(out_dim, out_dim, 3, padding=1) self.conv_shortcut = QwenImage21CausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity() - def forward(self, x, feat_cache=None, feat_idx=None): + def forward(self, x, feat_cache=None, feat_idx=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] # Apply shortcut connection @@ -418,7 +418,7 @@ def __init__(self, dim): self.to_qkv = nn.Conv2d(dim, dim * 3, 1) self.proj = nn.Conv2d(dim, dim, 1) - def forward(self, x): + def forward(self, x) -> torch.Tensor: identity = x batch_size, channels, time, height, width = x.size() @@ -470,7 +470,7 @@ def __init__(self, dim: int, dropout: float = 0.0, num_layers: int = 1): self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=None): + def forward(self, x, feat_cache=None, feat_idx=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] # First residual block @@ -512,7 +512,7 @@ def __init__(self, in_dim, out_dim, dropout, num_res_blocks, temperal_downsample else: self.downsampler = None - def forward(self, x, feat_cache=None, feat_idx=None): + def forward(self, x, feat_cache=None, feat_idx=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] x_copy = x.clone() @@ -603,7 +603,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=None): + def forward(self, x, feat_cache=None, feat_idx=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] if feat_cache is not None: @@ -700,7 +700,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=False): + def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=False) -> torch.Tensor: if feat_idx is None: feat_idx = [0] """ @@ -775,7 +775,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=None): + def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=None) -> torch.Tensor: if feat_idx is None: feat_idx = [0] """ @@ -889,7 +889,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=False): + def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=False) -> torch.Tensor: if feat_idx is None: feat_idx = [0] ## conv1 diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_wan.py b/src/diffusers/models/autoencoders/autoencoder_kl_wan.py index de8a56edc20e..29bd5738a62b 100644 --- a/src/diffusers/models/autoencoders/autoencoder_kl_wan.py +++ b/src/diffusers/models/autoencoders/autoencoder_kl_wan.py @@ -163,7 +163,7 @@ def __init__( self._padding = (self.padding[2], self.padding[2], self.padding[1], self.padding[1], 2 * self.padding[0], 0) self.padding = (0, 0, 0) - def forward(self, x, cache_x=None): + def forward(self, x, cache_x=None) -> torch.Tensor: padding = list(self._padding) if cache_x is not None and self._padding[4] > 0: cache_x = cache_x.to(x.device) @@ -195,7 +195,7 @@ def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bi self.gamma = nn.Parameter(torch.ones(shape)) self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - def forward(self, x): + def forward(self, x) -> torch.Tensor: needs_fp32_normalize = x.dtype in (torch.float16, torch.bfloat16) or any( t in str(x.dtype) for t in ("float4_", "float8_") ) @@ -217,7 +217,7 @@ class WanUpsample(nn.Upsample): torch.Tensor: Upsampled tensor with the same data type as the input. """ - def forward(self, x): + def forward(self, x) -> torch.Tensor: return super().forward(x.float()).type_as(x) @@ -266,7 +266,7 @@ def __init__(self, dim: int, mode: str, upsample_out_dim: int = None) -> None: else: self.resample = nn.Identity() - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: b, c, t, h, w = x.size() if self.mode == "upsample3d": if feat_cache is not None: @@ -343,7 +343,7 @@ def __init__( self.conv2 = WanCausalConv3d(out_dim, out_dim, 3, padding=1) self.conv_shortcut = WanCausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity() - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: # Apply shortcut connection h = self.conv_shortcut(x) @@ -403,7 +403,7 @@ def __init__(self, dim): self.to_qkv = nn.Conv2d(dim, dim * 3, 1) self.proj = nn.Conv2d(dim, dim, 1) - def forward(self, x): + def forward(self, x) -> torch.Tensor: identity = x batch_size, channels, time, height, width = x.size() @@ -456,7 +456,7 @@ def __init__(self, dim: int, dropout: float = 0.0, non_linearity: str = "silu", self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: # First residual block x = self.resnets[0](x, feat_cache=feat_cache, feat_idx=feat_idx) @@ -496,7 +496,7 @@ def __init__(self, in_dim, out_dim, dropout, num_res_blocks, temperal_downsample else: self.downsampler = None - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: x_copy = x.clone() for resnet in self.resnets: x = resnet(x, feat_cache=feat_cache, feat_idx=feat_idx) @@ -587,7 +587,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0]): + def forward(self, x, feat_cache=None, feat_idx=[0]) -> torch.Tensor: if feat_cache is not None: idx = feat_idx[0] cache_x = x[:, :, -CACHE_T:, :, :].clone() @@ -684,7 +684,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False): + def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False) -> torch.Tensor: """ Forward pass through the upsampling block. @@ -759,7 +759,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=None): + def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=None) -> torch.Tensor: """ Forward pass through the upsampling block. @@ -876,7 +876,7 @@ def __init__( self.gradient_checkpointing = False - def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False): + def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False) -> torch.Tensor: ## conv1 if feat_cache is not None: idx = feat_idx[0] diff --git a/src/diffusers/models/autoencoders/autoencoder_oobleck.py b/src/diffusers/models/autoencoders/autoencoder_oobleck.py index d4251fd9f1a9..639199a8dd36 100644 --- a/src/diffusers/models/autoencoders/autoencoder_oobleck.py +++ b/src/diffusers/models/autoencoders/autoencoder_oobleck.py @@ -41,7 +41,7 @@ def __init__(self, hidden_dim, logscale=True): self.beta.requires_grad = True self.logscale = logscale - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: shape = hidden_states.shape alpha = self.alpha if not self.logscale else torch.exp(self.alpha) @@ -67,7 +67,7 @@ def __init__(self, dimension: int = 16, dilation: int = 1): self.snake2 = Snake1d(dimension) self.conv2 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=1)) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: """ Forward pass through the residual unit. @@ -104,7 +104,7 @@ def __init__(self, input_dim, output_dim, stride: int = 1): nn.Conv1d(input_dim, output_dim, kernel_size=2 * stride, stride=stride, padding=math.ceil(stride / 2)) ) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.res_unit1(hidden_state) hidden_state = self.res_unit2(hidden_state) hidden_state = self.snake1(self.res_unit3(hidden_state)) @@ -133,7 +133,7 @@ def __init__(self, input_dim, output_dim, stride: int = 1): self.res_unit2 = OobleckResidualUnit(output_dim, dilation=3) self.res_unit3 = OobleckResidualUnit(output_dim, dilation=9) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.snake1(hidden_state) hidden_state = self.conv_t1(hidden_state) hidden_state = self.res_unit1(hidden_state) @@ -239,7 +239,7 @@ def __init__(self, encoder_hidden_size, audio_channels, downsampling_ratios, cha self.snake1 = Snake1d(d_model) self.conv2 = weight_norm(nn.Conv1d(d_model, encoder_hidden_size, kernel_size=3, padding=1)) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.conv1(hidden_state) for module in self.block: @@ -279,7 +279,7 @@ def __init__(self, channels, input_channels, audio_channels, upsampling_ratios, self.snake1 = Snake1d(output_dim) self.conv2 = weight_norm(nn.Conv1d(channels, audio_channels, kernel_size=7, padding=3, bias=False)) - def forward(self, hidden_state): + def forward(self, hidden_state) -> torch.Tensor: hidden_state = self.conv1(hidden_state) for layer in self.block: diff --git a/src/diffusers/models/controlnets/controlnet.py b/src/diffusers/models/controlnets/controlnet.py index acd88655c9fe..7dc3e9b18f8d 100644 --- a/src/diffusers/models/controlnets/controlnet.py +++ b/src/diffusers/models/controlnets/controlnet.py @@ -95,7 +95,7 @@ def __init__( nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1) ) - def forward(self, conditioning): + def forward(self, conditioning) -> torch.Tensor: embedding = self.conv_in(conditioning) embedding = F.silu(embedding) diff --git a/src/diffusers/models/controlnets/controlnet_hunyuan.py b/src/diffusers/models/controlnets/controlnet_hunyuan.py index 6ef92d78dd6e..b7cae237e005 100644 --- a/src/diffusers/models/controlnets/controlnet_hunyuan.py +++ b/src/diffusers/models/controlnets/controlnet_hunyuan.py @@ -226,7 +226,7 @@ def forward( style=None, image_rotary_emb=None, return_dict=True, - ): + ) -> HunyuanControlNetOutput | tuple[list[torch.Tensor]]: """ The [`HunyuanDiT2DControlNetModel`] forward method. @@ -257,6 +257,10 @@ def forward( The image rotary embeddings to apply on query and key tensors during attention calculation. return_dict: bool Whether to return a dictionary. + + Returns: + If `return_dict` is True, a [`~models.controlnets.controlnet_hunyuan.HunyuanControlNetOutput`] is returned, + otherwise a `tuple` where the first element is the list of ControlNet block samples. """ height, width = hidden_states.shape[-2:] @@ -339,7 +343,7 @@ def forward( style=None, image_rotary_emb=None, return_dict=True, - ): + ) -> HunyuanControlNetOutput | tuple[list[torch.Tensor]]: """ The [`HunyuanDiT2DControlNetModel`] forward method. @@ -370,6 +374,11 @@ def forward( The image rotary embeddings to apply on query and key tensors during attention calculation. return_dict: bool Whether to return a dictionary. + + Returns: + If `return_dict` is True and only one ControlNet is used, a + [`~models.controlnets.controlnet_hunyuan.HunyuanControlNetOutput`] is returned. Otherwise a `tuple` where + the first element is the list of ControlNet block samples, summed across all ControlNets. """ for i, (image, scale, controlnet) in enumerate(zip(controlnet_cond, conditioning_scale, self.nets)): block_samples = controlnet( diff --git a/src/diffusers/models/controlnets/controlnet_union.py b/src/diffusers/models/controlnets/controlnet_union.py index 8b3ac1c36d85..e24933cace5a 100644 --- a/src/diffusers/models/controlnets/controlnet_union.py +++ b/src/diffusers/models/controlnets/controlnet_union.py @@ -56,7 +56,7 @@ def __init__(self, d_model: int): self.gelu = QuickGELU() self.c_proj = nn.Linear(d_model * 4, d_model) - def forward(self, x: torch.Tensor): + def forward(self, x: torch.Tensor) -> torch.Tensor: x = self.c_fc(x) x = self.gelu(x) x = self.c_proj(x) @@ -76,7 +76,7 @@ def attention(self, x: torch.Tensor): self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None return self.attn(x, x, x, need_weights=False, attn_mask=self.attn_mask)[0] - def forward(self, x: torch.Tensor): + def forward(self, x: torch.Tensor) -> torch.Tensor: x = x + self.attention(self.ln_1(x)) x = x + self.mlp(self.ln_2(x)) return x diff --git a/src/diffusers/models/controlnets/controlnet_xs.py b/src/diffusers/models/controlnets/controlnet_xs.py index a25d5d71a5b1..c3f58ea359a2 100644 --- a/src/diffusers/models/controlnets/controlnet_xs.py +++ b/src/diffusers/models/controlnets/controlnet_xs.py @@ -502,7 +502,7 @@ def from_unet( return model - def forward(self, *args, **kwargs): + def forward(self, *args, **kwargs) -> None: raise ValueError( "A ControlNetXSAdapter cannot be run by itself. Use it together with a UNet2DConditionModel to instantiate a UNetControlNetXSModel." ) diff --git a/src/diffusers/models/controlnets/controlnet_z_image.py b/src/diffusers/models/controlnets/controlnet_z_image.py index a4800b255ef0..904e4bfdf6c5 100644 --- a/src/diffusers/models/controlnets/controlnet_z_image.py +++ b/src/diffusers/models/controlnets/controlnet_z_image.py @@ -62,7 +62,7 @@ def timestep_embedding(t, dim, max_period=10000): embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) return embedding - def forward(self, t): + def forward(self, t) -> torch.Tensor: t_freq = self.timestep_embedding(t, self.frequency_embedding_size) weight_dtype = self.mlp[0].weight.dtype compute_dtype = getattr(self.mlp[0], "compute_dtype", None) @@ -166,7 +166,7 @@ def __init__(self, dim: int, hidden_dim: int): def _forward_silu_gating(self, x1, x3): return F.silu(x1) * x3 - def forward(self, x): + def forward(self, x) -> torch.Tensor: return self.w2(self._forward_silu_gating(self.w1(x), self.w3(x))) @@ -238,7 +238,7 @@ def forward( noise_mask: torch.Tensor | None = None, adaln_noisy: torch.Tensor | None = None, adaln_clean: torch.Tensor | None = None, - ): + ) -> torch.Tensor: if self.modulation: seq_len = x.shape[1] @@ -390,7 +390,7 @@ def forward( attn_mask: torch.Tensor, freqs_cis: torch.Tensor, adaln_input: torch.Tensor | None = None, - ): + ) -> torch.Tensor: # Control if self.block_id == 0: c = self.before_proj(c) + x @@ -660,7 +660,7 @@ def forward( conditioning_scale: float = 1.0, patch_size=2, f_patch_size=1, - ): + ) -> dict[int, torch.Tensor]: r""" Args: x (`list` of `torch.Tensor`): @@ -677,6 +677,10 @@ def forward( Spatial patch size used to tokenize the latent. f_patch_size (`int`, *optional*, defaults to `1`): Temporal (frame) patch size used to tokenize the latent. + + Returns: + `dict[int, torch.Tensor]`: The ControlNet block samples, scaled by `conditioning_scale` and keyed by the + index of the transformer layer each one is added to. """ if ( self.t_scale is None diff --git a/src/diffusers/models/embeddings.py b/src/diffusers/models/embeddings.py index cbebf3de3a50..e1938a99b265 100644 --- a/src/diffusers/models/embeddings.py +++ b/src/diffusers/models/embeddings.py @@ -356,7 +356,7 @@ def cropped_pos_embed(self, height, width): spatial_pos_embed = spatial_pos_embed.reshape(1, -1, spatial_pos_embed.shape[-1]) return spatial_pos_embed - def forward(self, latent): + def forward(self, latent) -> torch.Tensor: if self.pos_embed_max_size is not None: height, width = latent.shape[-2:] else: @@ -408,7 +408,7 @@ def __init__(self, patch_size=2, in_channels=4, embed_dim=768, bias=True): bias=bias, ) - def forward(self, x, freqs_cis): + def forward(self, x, freqs_cis) -> tuple[torch.Tensor, torch.Tensor, list[tuple[int, int]], torch.Tensor]: """ Patchifies and embeds the input tensor(s). @@ -517,7 +517,7 @@ def _get_positional_embeddings( return joint_pos_embedding - def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor): + def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor) -> torch.Tensor: r""" Args: text_embeds (`torch.Tensor`): @@ -1057,7 +1057,7 @@ def __init__( else: self.post_act = get_activation(post_act_fn) - def forward(self, sample, condition=None): + def forward(self, sample, condition=None) -> torch.Tensor: if condition is not None: sample = sample + self.cond_proj(condition) sample = self.linear_1(sample) @@ -1109,7 +1109,7 @@ def __init__( self.weight = self.W del self.W - def forward(self, x): + def forward(self, x) -> torch.Tensor: if self.log: x = torch.log(x) @@ -1143,7 +1143,7 @@ def __init__(self, embed_dim: int, max_seq_length: int = 32): pe[0, :, 1::2] = torch.cos(position * div_term) self.register_buffer("pe", pe) - def forward(self, x): + def forward(self, x) -> torch.Tensor: _, seq_length, _ = x.shape x = x + self.pe[:, :seq_length] return x @@ -1191,7 +1191,7 @@ def __init__( self.height_emb = nn.Embedding(self.height, embed_dim) self.width_emb = nn.Embedding(self.width, embed_dim) - def forward(self, index): + def forward(self, index) -> torch.Tensor: emb = self.emb(index) height_emb = self.height_emb(torch.arange(self.height, device=index.device).view(1, self.height)) @@ -1242,7 +1242,7 @@ def token_drop(self, labels, force_drop_ids=None): labels = torch.where(drop_ids, self.num_classes, labels) return labels - def forward(self, labels: torch.LongTensor, force_drop_ids=None): + def forward(self, labels: torch.LongTensor, force_drop_ids=None) -> torch.Tensor: use_dropout = self.dropout_prob > 0 if (self.training and use_dropout) or (force_drop_ids is not None): labels = self.token_drop(labels, force_drop_ids) @@ -1264,7 +1264,7 @@ def __init__( self.image_embeds = nn.Linear(image_embed_dim, self.num_image_text_embeds * cross_attention_dim) self.text_proj = nn.Linear(text_embed_dim, cross_attention_dim) - def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor): + def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor) -> torch.Tensor: batch_size = text_embeds.shape[0] # image @@ -1290,7 +1290,7 @@ def __init__( self.image_embeds = nn.Linear(image_embed_dim, self.num_image_text_embeds * cross_attention_dim) self.norm = nn.LayerNorm(cross_attention_dim) - def forward(self, image_embeds: torch.Tensor): + def forward(self, image_embeds: torch.Tensor) -> torch.Tensor: batch_size = image_embeds.shape[0] # image @@ -1308,7 +1308,7 @@ def __init__(self, image_embed_dim=1024, cross_attention_dim=1024): self.ff = FeedForward(image_embed_dim, cross_attention_dim, mult=1, activation_fn="gelu") self.norm = nn.LayerNorm(cross_attention_dim) - def forward(self, image_embeds: torch.Tensor): + def forward(self, image_embeds: torch.Tensor) -> torch.Tensor: return self.norm(self.ff(image_embeds)) @@ -1322,7 +1322,7 @@ def __init__(self, image_embed_dim=1024, cross_attention_dim=1024, mult=1, num_t self.ff = FeedForward(image_embed_dim, cross_attention_dim * num_tokens, mult=mult, activation_fn="gelu") self.norm = nn.LayerNorm(cross_attention_dim) - def forward(self, image_embeds: torch.Tensor): + def forward(self, image_embeds: torch.Tensor) -> torch.Tensor: x = self.ff(image_embeds) x = x.reshape(-1, self.num_tokens, self.cross_attention_dim) return self.norm(x) @@ -1336,7 +1336,7 @@ def __init__(self, num_classes, embedding_dim, class_dropout_prob=0.1): self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) self.class_embedder = LabelEmbedding(num_classes, embedding_dim, class_dropout_prob) - def forward(self, timestep, class_labels, hidden_dtype=None): + def forward(self, timestep, class_labels, hidden_dtype=None) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) @@ -1355,7 +1355,7 @@ def __init__(self, embedding_dim, pooled_projection_dim): self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) self.text_embedder = PixArtAlphaTextProjection(pooled_projection_dim, embedding_dim, act_fn="silu") - def forward(self, timestep, pooled_projection): + def forward(self, timestep, pooled_projection) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D) @@ -1375,7 +1375,7 @@ def __init__(self, embedding_dim, pooled_projection_dim): self.guidance_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) self.text_embedder = PixArtAlphaTextProjection(pooled_projection_dim, embedding_dim, act_fn="silu") - def forward(self, timestep, guidance, pooled_projection): + def forward(self, timestep, guidance, pooled_projection) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D) @@ -1435,7 +1435,7 @@ def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim) self.num_heads = num_heads - def forward(self, x): + def forward(self, x) -> torch.Tensor: x = x.permute(1, 0, 2) # NLC -> LNC x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (L+1)NC x = x + self.positional_embedding[:, None, :].to(x.dtype) # (L+1)NC @@ -1498,7 +1498,7 @@ def __init__( act_fn="silu_fp32", ) - def forward(self, timestep, encoder_hidden_states, image_meta_size, style, hidden_dtype=None): + def forward(self, timestep, encoder_hidden_states, image_meta_size, style, hidden_dtype=None) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, 256) @@ -1542,7 +1542,7 @@ def __init__(self, hidden_size=4096, cross_attention_dim=2048, frequency_embeddi ), ) - def forward(self, timestep, caption_feat, caption_mask): + def forward(self, timestep, caption_feat, caption_mask) -> torch.Tensor: # timestep embedding: time_freq = self.time_proj(timestep) time_embed = self.timestep_embedder(time_freq.to(dtype=caption_feat.dtype)) @@ -1582,7 +1582,7 @@ def forward( encoder_hidden_states: torch.Tensor, encoder_attention_mask: torch.Tensor, hidden_dtype: torch.dtype | None = None, - ): + ) -> tuple[torch.Tensor, torch.Tensor]: time_proj = self.time_proj(timestep) time_emb = self.timestep_embedder(time_proj.to(dtype=hidden_dtype)) @@ -1601,7 +1601,7 @@ def __init__(self, encoder_dim: int, time_embed_dim: int, num_heads: int = 64): self.proj = nn.Linear(encoder_dim, time_embed_dim) self.norm2 = nn.LayerNorm(time_embed_dim) - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: hidden_states = self.norm1(hidden_states) hidden_states = self.pool(hidden_states) hidden_states = self.proj(hidden_states) @@ -1616,7 +1616,7 @@ def __init__(self, text_embed_dim: int = 768, image_embed_dim: int = 768, time_e self.text_norm = nn.LayerNorm(time_embed_dim) self.image_proj = nn.Linear(image_embed_dim, time_embed_dim) - def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor): + def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor) -> torch.Tensor: # text time_text_embeds = self.text_proj(text_embeds) time_text_embeds = self.text_norm(time_text_embeds) @@ -1633,7 +1633,7 @@ def __init__(self, image_embed_dim: int = 768, time_embed_dim: int = 1536): self.image_proj = nn.Linear(image_embed_dim, time_embed_dim) self.image_norm = nn.LayerNorm(time_embed_dim) - def forward(self, image_embeds: torch.Tensor): + def forward(self, image_embeds: torch.Tensor) -> torch.Tensor: # image time_image_embeds = self.image_proj(image_embeds) time_image_embeds = self.image_norm(time_image_embeds) @@ -1663,7 +1663,7 @@ def __init__(self, image_embed_dim: int = 768, time_embed_dim: int = 1536): nn.Conv2d(256, 4, 3, padding=1), ) - def forward(self, image_embeds: torch.Tensor, hint: torch.Tensor): + def forward(self, image_embeds: torch.Tensor, hint: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: # image time_image_embeds = self.image_proj(image_embeds) time_image_embeds = self.image_norm(time_image_embeds) @@ -1684,7 +1684,7 @@ def __init__(self, num_heads, embed_dim, dtype=None): self.num_heads = num_heads self.dim_per_head = embed_dim // self.num_heads - def forward(self, x): + def forward(self, x) -> torch.Tensor: bs, length, width = x.size() def shape(x): @@ -1875,7 +1875,7 @@ def forward( image_masks=None, phrases_embeddings=None, image_embeddings=None, - ): + ) -> torch.Tensor: masks = masks.unsqueeze(-1) # embedding position (it may includes padding as placeholder) @@ -1938,7 +1938,7 @@ def __init__(self, embedding_dim, size_emb_dim, use_additional_conditions: bool self.resolution_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=size_emb_dim) self.aspect_ratio_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=size_emb_dim) - def forward(self, timestep, resolution, aspect_ratio, batch_size, hidden_dtype): + def forward(self, timestep, resolution, aspect_ratio, batch_size, hidden_dtype) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) @@ -1976,7 +1976,7 @@ def __init__(self, in_features, hidden_size, out_features=None, act_fn="gelu_tan raise ValueError(f"Unknown activation function: {act_fn}") self.linear_2 = nn.Linear(in_features=hidden_size, out_features=out_features, bias=True) - def forward(self, caption): + def forward(self, caption) -> torch.Tensor: hidden_states = self.linear_1(caption) hidden_states = self.act_1(hidden_states) hidden_states = self.linear_2(hidden_states) @@ -2007,7 +2007,7 @@ def __init__( FeedForward(embed_dims, embed_dims, activation_fn="gelu", mult=ffn_ratio, bias=False), ) - def forward(self, x, latents, residual): + def forward(self, x, latents, residual) -> torch.Tensor: encoder_hidden_states = self.ln0(x) latents = self.ln1(latents) encoder_hidden_states = torch.cat([encoder_hidden_states, latents], dim=-2) @@ -2346,7 +2346,7 @@ def num_ip_adapters(self) -> int: """Number of IP-Adapters loaded.""" return len(self.image_projection_layers) - def forward(self, image_embeds: list[torch.Tensor]): + def forward(self, image_embeds: list[torch.Tensor]) -> list[torch.Tensor]: projected_image_embeds = [] # currently, we accept `image_embeds` as diff --git a/src/diffusers/models/lora.py b/src/diffusers/models/lora.py index 72e285832737..94ed621940dd 100644 --- a/src/diffusers/models/lora.py +++ b/src/diffusers/models/lora.py @@ -162,7 +162,7 @@ def _unfuse_lora(self): self.w_up = None self.w_down = None - def forward(self, input): + def forward(self, input) -> torch.Tensor: if self.lora_scale is None: self.lora_scale = 1.0 if self.lora_linear_layer is None: diff --git a/src/diffusers/models/normalization.py b/src/diffusers/models/normalization.py index dc872417914e..4b4ba78dc679 100644 --- a/src/diffusers/models/normalization.py +++ b/src/diffusers/models/normalization.py @@ -503,7 +503,7 @@ def __init__(self, dim, eps: float = 1e-5, elementwise_affine: bool = True, bias self.weight = None self.bias = None - def forward(self, input): + def forward(self, input) -> torch.Tensor: return F.layer_norm(input, self.dim, self.weight, self.bias, self.eps) @@ -538,7 +538,7 @@ def __init__(self, dim, eps: float, elementwise_affine: bool = True, bias: bool if bias: self.bias = nn.Parameter(torch.zeros(dim)) - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: # `npu_rms_norm` requires a gamma tensor. When `elementwise_affine=False`, # `self.weight` is `None`, so fall back to the pure PyTorch path. if is_torch_npu_available() and self.weight is not None: @@ -586,7 +586,7 @@ def __init__(self, dim, eps: float, elementwise_affine: bool = True): else: self.weight = None - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: input_dtype = hidden_states.dtype variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True) hidden_states = hidden_states * torch.rsqrt(variance + self.eps) @@ -612,7 +612,7 @@ def __init__(self, dim): self.gamma = nn.Parameter(torch.zeros(1, 1, 1, dim)) self.beta = nn.Parameter(torch.zeros(1, 1, 1, dim)) - def forward(self, x): + def forward(self, x) -> torch.Tensor: gx = torch.norm(x, p=2, dim=(1, 2), keepdim=True) nx = gx / (gx.mean(dim=-1, keepdim=True) + 1e-6) return self.gamma * (x * nx) + self.beta + x diff --git a/src/diffusers/models/resnet.py b/src/diffusers/models/resnet.py index d63e4fd0017b..fc5c89efd200 100644 --- a/src/diffusers/models/resnet.py +++ b/src/diffusers/models/resnet.py @@ -692,7 +692,7 @@ def forward( hidden_states: torch.Tensor, temb: torch.Tensor | None = None, image_only_indicator: torch.Tensor | None = None, - ): + ) -> torch.Tensor: num_frames = image_only_indicator.shape[-1] hidden_states = self.spatial_res_block(hidden_states, temb) diff --git a/src/diffusers/models/transformers/dit_transformer_2d.py b/src/diffusers/models/transformers/dit_transformer_2d.py index 0457acf77108..fc37781ddbd5 100644 --- a/src/diffusers/models/transformers/dit_transformer_2d.py +++ b/src/diffusers/models/transformers/dit_transformer_2d.py @@ -152,7 +152,7 @@ def forward( class_labels: torch.LongTensor | None = None, cross_attention_kwargs: dict[str, Any] = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`DiTTransformer2DModel`] forward method. diff --git a/src/diffusers/models/transformers/dual_transformer_2d.py b/src/diffusers/models/transformers/dual_transformer_2d.py index 778d5128ee23..4d400c18e200 100644 --- a/src/diffusers/models/transformers/dual_transformer_2d.py +++ b/src/diffusers/models/transformers/dual_transformer_2d.py @@ -11,6 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. +import torch from torch import nn from ..modeling_outputs import Transformer2DModelOutput @@ -101,7 +102,7 @@ def forward( attention_mask=None, cross_attention_kwargs=None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ Args: hidden_states ( When discrete, `torch.LongTensor` of shape `(batch size, num latent pixels)`. diff --git a/src/diffusers/models/transformers/hunyuan_transformer_2d.py b/src/diffusers/models/transformers/hunyuan_transformer_2d.py index 83b3797c4fc3..ba0745ce4876 100644 --- a/src/diffusers/models/transformers/hunyuan_transformer_2d.py +++ b/src/diffusers/models/transformers/hunyuan_transformer_2d.py @@ -367,7 +367,7 @@ def forward( image_rotary_emb=None, controlnet_block_samples=None, return_dict=True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`HunyuanDiT2DModel`] forward method. @@ -396,6 +396,10 @@ def forward( A list of tensors that if specified are added to the residuals of transformer blocks. return_dict: bool Whether to return a dictionary. + + Returns: + If `return_dict` is True, a [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a + `tuple` where the first element is the sample tensor. """ height, width = hidden_states.shape[-2:] diff --git a/src/diffusers/models/transformers/latte_transformer_3d.py b/src/diffusers/models/transformers/latte_transformer_3d.py index 01a1e608a927..c953cdabc936 100644 --- a/src/diffusers/models/transformers/latte_transformer_3d.py +++ b/src/diffusers/models/transformers/latte_transformer_3d.py @@ -171,7 +171,7 @@ def forward( encoder_attention_mask: torch.Tensor | None = None, enable_temporal_attentions: bool = True, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`LatteTransformer3DModel`] forward method. diff --git a/src/diffusers/models/transformers/pixart_transformer_2d.py b/src/diffusers/models/transformers/pixart_transformer_2d.py index e5e6178eaf4a..7e08abff85f3 100644 --- a/src/diffusers/models/transformers/pixart_transformer_2d.py +++ b/src/diffusers/models/transformers/pixart_transformer_2d.py @@ -234,7 +234,7 @@ def forward( attention_mask: torch.Tensor | None = None, encoder_attention_mask: torch.Tensor | None = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`PixArtTransformer2DModel`] forward method. diff --git a/src/diffusers/models/transformers/prior_transformer.py b/src/diffusers/models/transformers/prior_transformer.py index f3890446e28e..8847fff3a4ba 100644 --- a/src/diffusers/models/transformers/prior_transformer.py +++ b/src/diffusers/models/transformers/prior_transformer.py @@ -188,7 +188,7 @@ def forward( encoder_hidden_states: torch.Tensor | None = None, attention_mask: torch.BoolTensor | None = None, return_dict: bool = True, - ): + ) -> PriorTransformerOutput | tuple[torch.Tensor]: """ The [`PriorTransformer`] forward method. diff --git a/src/diffusers/models/transformers/sana_transformer.py b/src/diffusers/models/transformers/sana_transformer.py index 1451750d50ef..d05042c08300 100644 --- a/src/diffusers/models/transformers/sana_transformer.py +++ b/src/diffusers/models/transformers/sana_transformer.py @@ -108,7 +108,9 @@ def __init__(self, embedding_dim): self.silu = nn.SiLU() self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=True) - def forward(self, timestep: torch.Tensor, guidance: torch.Tensor = None, hidden_dtype: torch.dtype = None): + def forward( + self, timestep: torch.Tensor, guidance: torch.Tensor = None, hidden_dtype: torch.dtype = None + ) -> tuple[torch.Tensor, torch.Tensor]: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) diff --git a/src/diffusers/models/transformers/stable_audio_transformer.py b/src/diffusers/models/transformers/stable_audio_transformer.py index f4974926ec72..72c42048db71 100644 --- a/src/diffusers/models/transformers/stable_audio_transformer.py +++ b/src/diffusers/models/transformers/stable_audio_transformer.py @@ -48,7 +48,7 @@ def __init__( self.weight = self.W del self.W - def forward(self, x): + def forward(self, x) -> torch.Tensor: if self.log: x = torch.log(x) diff --git a/src/diffusers/models/transformers/t5_film_transformer.py b/src/diffusers/models/transformers/t5_film_transformer.py index 547e72089990..8fb368545d16 100644 --- a/src/diffusers/models/transformers/t5_film_transformer.py +++ b/src/diffusers/models/transformers/t5_film_transformer.py @@ -89,7 +89,7 @@ def encoder_decoder_mask(self, query_input: torch.Tensor, key_input: torch.Tenso mask = torch.mul(query_input.unsqueeze(-1), key_input.unsqueeze(-2)) return mask.unsqueeze(-3) - def forward(self, encodings_and_masks, decoder_input_tokens, decoder_noise_time): + def forward(self, encodings_and_masks, decoder_input_tokens, decoder_noise_time) -> torch.Tensor: """ The [`T5FilmDecoder`] forward method. @@ -101,6 +101,9 @@ def forward(self, encodings_and_masks, decoder_input_tokens, decoder_noise_time) Input tokens for the decoder. decoder_noise_time (`torch.Tensor` of shape `(batch_size,)`): Diffusion timesteps in `[0, 1)` used to condition the decoder. + + Returns: + `torch.Tensor`: The decoded spectrogram of shape `(batch_size, seq_length, input_dims)`. """ batch, _, _ = decoder_input_tokens.shape assert decoder_noise_time.shape == (batch,) diff --git a/src/diffusers/models/transformers/transformer_2d.py b/src/diffusers/models/transformers/transformer_2d.py index 6714383b77ab..50f5a082efec 100644 --- a/src/diffusers/models/transformers/transformer_2d.py +++ b/src/diffusers/models/transformers/transformer_2d.py @@ -332,7 +332,7 @@ def forward( attention_mask: torch.Tensor | None = None, encoder_attention_mask: torch.Tensor | None = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`Transformer2DModel`] forward method. diff --git a/src/diffusers/models/transformers/transformer_2d_dreamlite.py b/src/diffusers/models/transformers/transformer_2d_dreamlite.py index 9d66eeafbd00..370c16235170 100644 --- a/src/diffusers/models/transformers/transformer_2d_dreamlite.py +++ b/src/diffusers/models/transformers/transformer_2d_dreamlite.py @@ -522,7 +522,7 @@ def forward( attention_mask: torch.Tensor | None = None, encoder_attention_mask: torch.Tensor | None = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """Forward pass of :class:`DreamLiteTransformer2DModel`. Args: diff --git a/src/diffusers/models/transformers/transformer_allegro.py b/src/diffusers/models/transformers/transformer_allegro.py index abe82ab578de..57ce1da7fd68 100644 --- a/src/diffusers/models/transformers/transformer_allegro.py +++ b/src/diffusers/models/transformers/transformer_allegro.py @@ -311,7 +311,7 @@ def forward( encoder_attention_mask: torch.Tensor | None = None, image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`AllegroTransformer3DModel`] forward method. diff --git a/src/diffusers/models/transformers/transformer_anyflow.py b/src/diffusers/models/transformers/transformer_anyflow.py index 6b0872ffdb01..1388c63096b0 100644 --- a/src/diffusers/models/transformers/transformer_anyflow.py +++ b/src/diffusers/models/transformers/transformer_anyflow.py @@ -287,7 +287,7 @@ def forward( encoder_hidden_states: torch.Tensor, encoder_hidden_states_image: Optional[torch.Tensor] = None, layout_cfg=None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: if self.deltatime_type == "r": delta_timestep = r_timestep elif self.deltatime_type == "t-r": @@ -384,7 +384,7 @@ def _forward_full_frame(self, num_frames, height, width, device) -> torch.Tensor freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1) return freqs - def forward(self, layout_cfg, device): + def forward(self, layout_cfg, device) -> dict[str, torch.Tensor]: freqs = self._forward_full_frame( num_frames=layout_cfg["total_frames"], height=layout_cfg["full_frame_shape"][0], diff --git a/src/diffusers/models/transformers/transformer_anyflow_far.py b/src/diffusers/models/transformers/transformer_anyflow_far.py index 9ecc16bd04e0..a5fdf84ac829 100644 --- a/src/diffusers/models/transformers/transformer_anyflow_far.py +++ b/src/diffusers/models/transformers/transformer_anyflow_far.py @@ -468,7 +468,7 @@ def forward( encoder_hidden_states_image: Optional[torch.Tensor] = None, far_cfg=None, clean_timestep=None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: if self.deltatime_type == "r": delta_timestep = r_timestep elif self.deltatime_type == "t-r": @@ -749,7 +749,7 @@ def _forward_full_frame(self, num_frames, height, width, device) -> torch.Tensor freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1) return freqs - def forward(self, far_cfg, device, clean_hidden_states=None): + def forward(self, far_cfg, device, clean_hidden_states=None) -> dict[str, torch.Tensor]: full_frame_freqs = self._forward_full_frame( num_frames=far_cfg["total_frames"], height=far_cfg["full_frame_shape"][0], diff --git a/src/diffusers/models/transformers/transformer_bria.py b/src/diffusers/models/transformers/transformer_bria.py index ff4261343ab2..a97b5b772552 100644 --- a/src/diffusers/models/transformers/transformer_bria.py +++ b/src/diffusers/models/transformers/transformer_bria.py @@ -304,7 +304,7 @@ def __init__( self.scale = scale self.time_theta = time_theta - def forward(self, timesteps): + def forward(self, timesteps) -> torch.Tensor: t_emb = get_timestep_embedding( timesteps, self.num_channels, @@ -325,7 +325,7 @@ def __init__(self, embedding_dim, time_theta): ) self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - def forward(self, timestep, dtype): + def forward(self, timestep, dtype) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=dtype)) # (N, D) return timesteps_emb diff --git a/src/diffusers/models/transformers/transformer_bria_fibo.py b/src/diffusers/models/transformers/transformer_bria_fibo.py index 9ec0ea1647a6..3dff8bcf55c7 100644 --- a/src/diffusers/models/transformers/transformer_bria_fibo.py +++ b/src/diffusers/models/transformers/transformer_bria_fibo.py @@ -296,7 +296,7 @@ def __init__(self, in_features, hidden_size): super().__init__() self.linear = nn.Linear(in_features=in_features, out_features=hidden_size, bias=False) - def forward(self, caption): + def forward(self, caption) -> torch.Tensor: hidden_states = self.linear(caption) return hidden_states @@ -398,7 +398,7 @@ def __init__( self.scale = scale self.time_theta = time_theta - def forward(self, timesteps): + def forward(self, timesteps) -> torch.Tensor: t_emb = get_timestep_embedding( timesteps, self.num_channels, @@ -419,7 +419,7 @@ def __init__(self, embedding_dim, time_theta): ) self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - def forward(self, timestep, dtype): + def forward(self, timestep, dtype) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=dtype)) # (N, D) return timesteps_emb diff --git a/src/diffusers/models/transformers/transformer_chroma.py b/src/diffusers/models/transformers/transformer_chroma.py index 92190bb0120d..81a3dadc0998 100644 --- a/src/diffusers/models/transformers/transformer_chroma.py +++ b/src/diffusers/models/transformers/transformer_chroma.py @@ -190,7 +190,7 @@ def __init__(self, in_dim: int, out_dim: int, hidden_dim: int, n_layers: int = 5 self.norms = nn.ModuleList([nn.RMSNorm(hidden_dim) for _ in range(n_layers)]) self.out_proj = nn.Linear(hidden_dim, out_dim) - def forward(self, x): + def forward(self, x) -> torch.Tensor: x = self.in_proj(x) for layer, norms in zip(self.layers, self.norms): diff --git a/src/diffusers/models/transformers/transformer_chronoedit.py b/src/diffusers/models/transformers/transformer_chronoedit.py index b39a18a98afb..c676a437cd16 100644 --- a/src/diffusers/models/transformers/transformer_chronoedit.py +++ b/src/diffusers/models/transformers/transformer_chronoedit.py @@ -340,7 +340,7 @@ def forward( encoder_hidden_states: torch.Tensor, encoder_hidden_states_image: torch.Tensor | None = None, timestep_seq_len: int | None = None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: timestep = self.timesteps_proj(timestep) if timestep_seq_len is not None: timestep = timestep.unflatten(0, (-1, timestep_seq_len)) diff --git a/src/diffusers/models/transformers/transformer_cosmos3.py b/src/diffusers/models/transformers/transformer_cosmos3.py index d2cafeca9d27..0295e1c0567b 100644 --- a/src/diffusers/models/transformers/transformer_cosmos3.py +++ b/src/diffusers/models/transformers/transformer_cosmos3.py @@ -144,7 +144,7 @@ def apply_interleaved_mrope(self, freqs, rope_axes_dim): freqs_t[..., idx] = freqs[dim, ..., idx] return freqs_t - def forward(self, position_ids, device, dtype): + def forward(self, position_ids, device, dtype) -> tuple[torch.Tensor, torch.Tensor]: if position_ids.ndim == 2: position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1) # [3,B,N] inv_freq_expanded = ( @@ -188,7 +188,7 @@ def __init__(self, hidden_size: int, intermediate_size: int, hidden_act: str = " self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False) self.act_fn = nn.SiLU() if hidden_act == "silu" else None - def forward(self, x): + def forward(self, x) -> torch.Tensor: if self.hidden_act == "relu2": return self.down_proj(torch.relu(self.up_proj(x)).square()) return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) diff --git a/src/diffusers/models/transformers/transformer_ernie_image.py b/src/diffusers/models/transformers/transformer_ernie_image.py index 0abc5d254bb2..791d4e0bce42 100644 --- a/src/diffusers/models/transformers/transformer_ernie_image.py +++ b/src/diffusers/models/transformers/transformer_ernie_image.py @@ -264,7 +264,7 @@ def forward( rotary_pos_emb, temb: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor], attention_mask: torch.Tensor | None = None, - ): + ) -> torch.Tensor: shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = temb residual = x x = self.adaLN_sa_ln(x) @@ -353,7 +353,7 @@ def forward( text_bth: torch.Tensor, text_lens: torch.Tensor, return_dict: bool = True, - ): + ) -> ErnieImageTransformer2DModelOutput | tuple[torch.Tensor]: """ The [`ErnieImageTransformer2DModel`] forward method. @@ -370,6 +370,11 @@ def forward( return_dict (`bool`, *optional*, defaults to `True`): Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain tuple. + + Returns: + If `return_dict` is True, a + [`~models.transformers.transformer_ernie_image.ErnieImageTransformer2DModelOutput`] is returned, otherwise + a `tuple` where the first element is the sample tensor. """ device, dtype = hidden_states.device, hidden_states.dtype B, C, H, W = hidden_states.shape diff --git a/src/diffusers/models/transformers/transformer_helios.py b/src/diffusers/models/transformers/transformer_helios.py index b99ab1e3f34f..6733004bf1e5 100644 --- a/src/diffusers/models/transformers/transformer_helios.py +++ b/src/diffusers/models/transformers/transformer_helios.py @@ -87,7 +87,7 @@ def __init__(self, dim: int, eps: float = 1e-6, elementwise_affine: bool = False self.scale_shift_table = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5) self.norm = FP32LayerNorm(dim, eps, elementwise_affine=False) - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor, original_context_length: int): + def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor, original_context_length: int) -> torch.Tensor: temb = temb[:, -original_context_length:, :] shift, scale = (self.scale_shift_table.unsqueeze(0).to(temb.device) + temb.unsqueeze(2)).chunk(2, dim=2) shift, scale = shift.squeeze(2).to(hidden_states.device), scale.squeeze(2).to(hidden_states.device) @@ -308,7 +308,7 @@ def forward( timestep: torch.Tensor, encoder_hidden_states: torch.Tensor | None = None, is_return_encoder_hidden_states: bool = True, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]: timestep = self.timesteps_proj(timestep) time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype @@ -354,7 +354,7 @@ def _get_spatial_meshgrid(self, height, width, device_str): return grid_y, grid_x @torch.no_grad() - def forward(self, frame_indices, height, width, device): + def forward(self, frame_indices, height, width, device) -> torch.Tensor: batch_size = frame_indices.shape[0] num_frames = frame_indices.shape[1] diff --git a/src/diffusers/models/transformers/transformer_hidream_image.py b/src/diffusers/models/transformers/transformer_hidream_image.py index 703230562415..afb1f8b8edb9 100644 --- a/src/diffusers/models/transformers/transformer_hidream_image.py +++ b/src/diffusers/models/transformers/transformer_hidream_image.py @@ -295,7 +295,7 @@ def __init__( self._force_inference_output = _force_inference_output - def forward(self, hidden_states): + def forward(self, hidden_states) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]: bsz, seq_len, h = hidden_states.shape ### compute gating score hidden_states = hidden_states.view(-1, h) @@ -362,7 +362,7 @@ def __init__( ) self.num_activated_experts = num_activated_experts - def forward(self, x): + def forward(self, x) -> torch.Tensor: wtype = x.dtype identity = x orig_shape = x.shape @@ -409,7 +409,7 @@ def __init__(self, in_features, hidden_size): super().__init__() self.linear = nn.Linear(in_features=in_features, out_features=hidden_size, bias=False) - def forward(self, caption): + def forward(self, caption) -> torch.Tensor: hidden_states = self.linear(caption) return hidden_states diff --git a/src/diffusers/models/transformers/transformer_hunyuan_video_framepack.py b/src/diffusers/models/transformers/transformer_hunyuan_video_framepack.py index 9a3dbc00f4ec..4f62a5e85841 100644 --- a/src/diffusers/models/transformers/transformer_hunyuan_video_framepack.py +++ b/src/diffusers/models/transformers/transformer_hunyuan_video_framepack.py @@ -47,7 +47,9 @@ def __init__(self, patch_size: int, patch_size_t: int, rope_dim: list[int], thet self.rope_dim = rope_dim self.theta = theta - def forward(self, frame_indices: torch.Tensor, height: int, width: int, device: torch.device): + def forward( + self, frame_indices: torch.Tensor, height: int, width: int, device: torch.device + ) -> tuple[torch.Tensor, torch.Tensor]: height = height // self.patch_size width = width // self.patch_size grid = torch.meshgrid( @@ -94,7 +96,7 @@ def forward( latents_clean: torch.Tensor | None = None, latents_clean_2x: torch.Tensor | None = None, latents_clean_4x: torch.Tensor | None = None, - ): + ) -> tuple[torch.Tensor | None, torch.Tensor | None, torch.Tensor | None]: if latents_clean is not None: latents_clean = self.proj(latents_clean) latents_clean = latents_clean.flatten(2).transpose(1, 2) diff --git a/src/diffusers/models/transformers/transformer_joyimage.py b/src/diffusers/models/transformers/transformer_joyimage.py index b17ddb05f799..d2fec631f1ac 100644 --- a/src/diffusers/models/transformers/transformer_joyimage.py +++ b/src/diffusers/models/transformers/transformer_joyimage.py @@ -350,7 +350,7 @@ def forward( self, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: timestep = self.timesteps_proj(timestep) time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype @@ -525,7 +525,7 @@ def forward( timestep: torch.Tensor, encoder_hidden_states: torch.Tensor = None, return_dict: bool = True, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`JoyImageEditTransformer3DModel`] forward method. @@ -539,6 +539,10 @@ def forward( return_dict (`bool`, *optional*, defaults to `True`): Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain tuple. + + Returns: + If `return_dict` is True, a [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a + `tuple` where the first element is the sample tensor. """ # handle multi-item input (b, n, c, t, h, w) is_multi_item = hidden_states.ndim == 6 diff --git a/src/diffusers/models/transformers/transformer_joyimage_edit_plus.py b/src/diffusers/models/transformers/transformer_joyimage_edit_plus.py index 4a13845faad3..a81027b5f40b 100644 --- a/src/diffusers/models/transformers/transformer_joyimage_edit_plus.py +++ b/src/diffusers/models/transformers/transformer_joyimage_edit_plus.py @@ -300,7 +300,7 @@ def forward( self, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: timestep = self.timesteps_proj(timestep) time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype diff --git a/src/diffusers/models/transformers/transformer_kandinsky.py b/src/diffusers/models/transformers/transformer_kandinsky.py index 88ef70d546c8..a908677457c9 100644 --- a/src/diffusers/models/transformers/transformer_kandinsky.py +++ b/src/diffusers/models/transformers/transformer_kandinsky.py @@ -165,7 +165,7 @@ def __init__(self, model_dim, time_dim, max_period=10000.0): self.activation = nn.SiLU() self.out_layer = nn.Linear(time_dim, time_dim, bias=True) - def forward(self, time): + def forward(self, time) -> torch.Tensor: args = torch.outer(time.to(torch.float32), self.freqs.to(device=time.device)) time_embed = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) time_embed = self.out_layer(self.activation(self.in_layer(time_embed))) @@ -178,7 +178,7 @@ def __init__(self, text_dim, model_dim): self.in_layer = nn.Linear(text_dim, model_dim, bias=True) self.norm = nn.LayerNorm(model_dim, elementwise_affine=True) - def forward(self, text_embed): + def forward(self, text_embed) -> torch.Tensor: text_embed = self.in_layer(text_embed) return self.norm(text_embed).type_as(text_embed) @@ -189,7 +189,7 @@ def __init__(self, visual_dim, model_dim, patch_size): self.patch_size = patch_size self.in_layer = nn.Linear(math.prod(patch_size) * visual_dim, model_dim) - def forward(self, x): + def forward(self, x) -> torch.Tensor: batch_size, duration, height, width, dim = x.shape x = ( x.view( @@ -218,7 +218,7 @@ def __init__(self, dim, max_pos=1024, max_period=10000.0): pos = torch.arange(max_pos, dtype=freq.dtype) self.register_buffer("args", torch.outer(pos, freq), persistent=False) - def forward(self, pos): + def forward(self, pos) -> torch.Tensor: args = self.args[pos] cosine = torch.cos(args) sine = torch.sin(args) @@ -239,7 +239,7 @@ def __init__(self, axes_dims, max_pos=(128, 128, 128), max_period=10000.0): pos = torch.arange(ax_max_pos, dtype=freq.dtype) self.register_buffer(f"args_{i}", torch.outer(pos, freq), persistent=False) - def forward(self, shape, pos, scale_factor=(1.0, 1.0, 1.0)): + def forward(self, shape, pos, scale_factor=(1.0, 1.0, 1.0)) -> torch.Tensor: batch_size, duration, height, width = shape args_t = self.args_0[pos[0]] / scale_factor[0] args_h = self.args_1[pos[1]] / scale_factor[1] @@ -268,7 +268,7 @@ def __init__(self, time_dim, model_dim, num_params): self.out_layer.weight.data.zero_() self.out_layer.bias.data.zero_() - def forward(self, x): + def forward(self, x) -> torch.Tensor: return self.out_layer(self.activation(x)) @@ -397,7 +397,7 @@ def __init__(self, dim, ff_dim): self.activation = nn.GELU() self.out_layer = nn.Linear(ff_dim, dim, bias=False) - def forward(self, x): + def forward(self, x) -> torch.Tensor: return self.out_layer(self.activation(self.in_layer(x))) @@ -409,7 +409,7 @@ def __init__(self, model_dim, time_dim, visual_dim, patch_size): self.norm = nn.LayerNorm(model_dim, elementwise_affine=False) self.out_layer = nn.Linear(model_dim, math.prod(patch_size) * visual_dim, bias=True) - def forward(self, visual_embed, text_embed, time_embed): + def forward(self, visual_embed, text_embed, time_embed) -> torch.Tensor: shift, scale = torch.chunk(self.modulation(time_embed).unsqueeze(dim=1), 2, dim=-1) visual_embed = ( @@ -449,7 +449,7 @@ def __init__(self, model_dim, time_dim, ff_dim, head_dim): self.feed_forward_norm = nn.LayerNorm(model_dim, elementwise_affine=False) self.feed_forward = Kandinsky5FeedForward(model_dim, ff_dim) - def forward(self, x, time_embed, rope): + def forward(self, x, time_embed, rope) -> torch.Tensor: self_attn_params, ff_params = torch.chunk(self.text_modulation(time_embed).unsqueeze(dim=1), 2, dim=-1) shift, scale, gate = torch.chunk(self_attn_params, 3, dim=-1) out = (self.self_attention_norm(x.float()) * (scale.float() + 1.0) + shift.float()).type_as(x) @@ -478,7 +478,7 @@ def __init__(self, model_dim, time_dim, ff_dim, head_dim): self.feed_forward_norm = nn.LayerNorm(model_dim, elementwise_affine=False) self.feed_forward = Kandinsky5FeedForward(model_dim, ff_dim) - def forward(self, visual_embed, text_embed, time_embed, rope, sparse_params): + def forward(self, visual_embed, text_embed, time_embed, rope, sparse_params) -> torch.Tensor: self_attn_params, cross_attn_params, ff_params = torch.chunk( self.visual_modulation(time_embed).unsqueeze(dim=1), 3, dim=-1 ) diff --git a/src/diffusers/models/transformers/transformer_longcat_image.py b/src/diffusers/models/transformers/transformer_longcat_image.py index 7b842c42132d..55d4771090e9 100644 --- a/src/diffusers/models/transformers/transformer_longcat_image.py +++ b/src/diffusers/models/transformers/transformer_longcat_image.py @@ -385,7 +385,7 @@ def __init__(self, embedding_dim): self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) - def forward(self, timestep, hidden_dtype): + def forward(self, timestep, hidden_dtype) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) diff --git a/src/diffusers/models/transformers/transformer_lumina2.py b/src/diffusers/models/transformers/transformer_lumina2.py index ba822730cb32..6d51b2925ef0 100644 --- a/src/diffusers/models/transformers/transformer_lumina2.py +++ b/src/diffusers/models/transformers/transformer_lumina2.py @@ -260,7 +260,9 @@ def _get_freqs_cis(self, ids: torch.Tensor) -> torch.Tensor: result.append(torch.gather(freqs.unsqueeze(0).repeat(index.shape[0], 1, 1), dim=1, index=index)) return torch.cat(result, dim=-1).to(device) - def forward(self, hidden_states: torch.Tensor, attention_mask: torch.Tensor): + def forward( + self, hidden_states: torch.Tensor, attention_mask: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, list[int], list[int]]: batch_size, channels, height, width = hidden_states.shape p = self.patch_size post_patch_height, post_patch_width = height // p, width // p diff --git a/src/diffusers/models/transformers/transformer_mochi.py b/src/diffusers/models/transformers/transformer_mochi.py index a1a1f5e9c900..1e544890a4b8 100644 --- a/src/diffusers/models/transformers/transformer_mochi.py +++ b/src/diffusers/models/transformers/transformer_mochi.py @@ -42,7 +42,7 @@ def __init__(self, eps: float): self.eps = eps self.norm = RMSNorm(0, eps, False) - def forward(self, hidden_states, scale=None): + def forward(self, hidden_states, scale=None) -> torch.Tensor: hidden_states_dtype = hidden_states.dtype hidden_states = hidden_states.to(torch.float32) diff --git a/src/diffusers/models/transformers/transformer_nucleusmoe_image.py b/src/diffusers/models/transformers/transformer_nucleusmoe_image.py index f1c0eee949f7..5488c02c579f 100644 --- a/src/diffusers/models/transformers/transformer_nucleusmoe_image.py +++ b/src/diffusers/models/transformers/transformer_nucleusmoe_image.py @@ -127,7 +127,7 @@ def __init__(self, embedding_dim, use_additional_t_cond=False): if use_additional_t_cond: self.addition_t_embedding = nn.Embedding(2, embedding_dim) - def forward(self, timestep, hidden_states, addition_t_cond=None): + def forward(self, timestep, hidden_states, addition_t_cond=None) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_states.dtype)) diff --git a/src/diffusers/models/transformers/transformer_omnigen.py b/src/diffusers/models/transformers/transformer_omnigen.py index f860f5d5ab3e..9b3ee661bb10 100644 --- a/src/diffusers/models/transformers/transformer_omnigen.py +++ b/src/diffusers/models/transformers/transformer_omnigen.py @@ -150,7 +150,7 @@ def __init__( self.long_factor = rope_scaling["long_factor"] self.original_max_position_embeddings = original_max_position_embeddings - def forward(self, hidden_states, position_ids): + def forward(self, hidden_states, position_ids) -> tuple[torch.Tensor, torch.Tensor]: seq_len = torch.max(position_ids) + 1 if seq_len > self.original_max_position_embeddings: ext_factors = torch.tensor(self.long_factor, dtype=torch.float32, device=hidden_states.device) diff --git a/src/diffusers/models/transformers/transformer_qwenimage.py b/src/diffusers/models/transformers/transformer_qwenimage.py index 5a242c8bd5c0..45722ca0d7b0 100644 --- a/src/diffusers/models/transformers/transformer_qwenimage.py +++ b/src/diffusers/models/transformers/transformer_qwenimage.py @@ -214,7 +214,7 @@ def __init__(self, embedding_dim, use_additional_t_cond=False): if use_additional_t_cond: self.addition_t_embedding = nn.Embedding(2, embedding_dim) - def forward(self, timestep, hidden_states, addition_t_cond=None): + def forward(self, timestep, hidden_states, addition_t_cond=None) -> torch.Tensor: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_states.dtype)) # (N, D) diff --git a/src/diffusers/models/transformers/transformer_sana_video.py b/src/diffusers/models/transformers/transformer_sana_video.py index db1f08a73a81..84c1fa18a9bb 100644 --- a/src/diffusers/models/transformers/transformer_sana_video.py +++ b/src/diffusers/models/transformers/transformer_sana_video.py @@ -263,7 +263,9 @@ def __init__(self, embedding_dim): self.silu = nn.SiLU() self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=True) - def forward(self, timestep: torch.Tensor, guidance: torch.Tensor = None, hidden_dtype: torch.dtype = None): + def forward( + self, timestep: torch.Tensor, guidance: torch.Tensor = None, hidden_dtype: torch.dtype = None + ) -> tuple[torch.Tensor, torch.Tensor]: timesteps_proj = self.time_proj(timestep) timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D) diff --git a/src/diffusers/models/transformers/transformer_sd3.py b/src/diffusers/models/transformers/transformer_sd3.py index 9a56ca4e226d..eaafe057fd82 100644 --- a/src/diffusers/models/transformers/transformer_sd3.py +++ b/src/diffusers/models/transformers/transformer_sd3.py @@ -59,7 +59,7 @@ def __init__( self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") - def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor): + def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor) -> torch.Tensor: # 1. Attention norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) attn_output = self.attn(hidden_states=norm_hidden_states, encoder_hidden_states=None) diff --git a/src/diffusers/models/transformers/transformer_skyreels_v2.py b/src/diffusers/models/transformers/transformer_skyreels_v2.py index a4a3aa3f6ffa..5687cfcf2964 100644 --- a/src/diffusers/models/transformers/transformer_skyreels_v2.py +++ b/src/diffusers/models/transformers/transformer_skyreels_v2.py @@ -355,7 +355,7 @@ def forward( timestep: torch.Tensor, encoder_hidden_states: torch.Tensor, encoder_hidden_states_image: torch.Tensor | None = None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: timestep = self.timesteps_proj(timestep) time_embedder_dtype = get_parameter_dtype(self.time_embedder) diff --git a/src/diffusers/models/transformers/transformer_temporal.py b/src/diffusers/models/transformers/transformer_temporal.py index 10bad499caf3..5141f07b8235 100644 --- a/src/diffusers/models/transformers/transformer_temporal.py +++ b/src/diffusers/models/transformers/transformer_temporal.py @@ -283,7 +283,7 @@ def forward( encoder_hidden_states: torch.Tensor | None = None, image_only_indicator: torch.Tensor | None = None, return_dict: bool = True, - ): + ) -> TransformerTemporalModelOutput | tuple[torch.Tensor]: """ Args: hidden_states (`torch.Tensor` of shape `(batch size, channel, height, width)`): diff --git a/src/diffusers/models/transformers/transformer_wan.py b/src/diffusers/models/transformers/transformer_wan.py index cf1b4ecc5d78..b9e0b75f46e8 100644 --- a/src/diffusers/models/transformers/transformer_wan.py +++ b/src/diffusers/models/transformers/transformer_wan.py @@ -333,7 +333,7 @@ def forward( encoder_hidden_states: torch.Tensor, encoder_hidden_states_image: torch.Tensor | None = None, timestep_seq_len: int | None = None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: timestep = self.timesteps_proj(timestep) if timestep_seq_len is not None: timestep = timestep.unflatten(0, (-1, timestep_seq_len)) diff --git a/src/diffusers/models/transformers/transformer_wan_animate.py b/src/diffusers/models/transformers/transformer_wan_animate.py index 084c3a2aed7d..3e188ce1ec51 100644 --- a/src/diffusers/models/transformers/transformer_wan_animate.py +++ b/src/diffusers/models/transformers/transformer_wan_animate.py @@ -809,7 +809,7 @@ def forward( encoder_hidden_states: torch.Tensor, encoder_hidden_states_image: torch.Tensor | None = None, timestep_seq_len: int | None = None, - ): + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: timestep = self.timesteps_proj(timestep) if timestep_seq_len is not None: timestep = timestep.unflatten(0, (-1, timestep_seq_len)) diff --git a/src/diffusers/models/transformers/transformer_wan_animate_2.py b/src/diffusers/models/transformers/transformer_wan_animate_2.py index c19655e6952e..55cfb8548b61 100644 --- a/src/diffusers/models/transformers/transformer_wan_animate_2.py +++ b/src/diffusers/models/transformers/transformer_wan_animate_2.py @@ -544,7 +544,7 @@ def __init__(self, dim, out_dim, patch_size, eps=1e-6): # modulation self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5) - def forward(self, x, e): + def forward(self, x, e) -> torch.Tensor: shift, scale = (self.modulation + e.float().unsqueeze(1)).chunk(2, dim=1) x = self.head((self.norm(x.float()) * (1 + scale) + shift).type_as(x)) return x @@ -562,7 +562,7 @@ def __init__(self, in_dim, out_dim): torch.nn.LayerNorm(out_dim), ) - def forward(self, image_embeds): + def forward(self, image_embeds) -> torch.Tensor: clip_extra_context_tokens = self.proj(image_embeds) return clip_extra_context_tokens diff --git a/src/diffusers/models/transformers/transformer_z_image.py b/src/diffusers/models/transformers/transformer_z_image.py index 4cea745e5ed5..913075ff9ea5 100644 --- a/src/diffusers/models/transformers/transformer_z_image.py +++ b/src/diffusers/models/transformers/transformer_z_image.py @@ -60,7 +60,7 @@ def timestep_embedding(t, dim, max_period=10000): embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) return embedding - def forward(self, t): + def forward(self, t) -> torch.Tensor: t_freq = self.timestep_embedding(t, self.frequency_embedding_size) weight_dtype = self.mlp[0].weight.dtype compute_dtype = getattr(self.mlp[0], "compute_dtype", None) @@ -176,7 +176,7 @@ def __init__(self, dim: int, hidden_dim: int): def _forward_silu_gating(self, x1, x3): return F.silu(x1) * x3 - def forward(self, x): + def forward(self, x) -> torch.Tensor: return self.w2(self._forward_silu_gating(self.w1(x), self.w3(x))) @@ -232,7 +232,7 @@ def forward( noise_mask: torch.Tensor | None = None, adaln_noisy: torch.Tensor | None = None, adaln_clean: torch.Tensor | None = None, - ): + ) -> torch.Tensor: if self.modulation: seq_len = x.shape[1] @@ -291,7 +291,7 @@ def __init__(self, hidden_size, out_channels): nn.Linear(min(hidden_size, ADALN_EMBED_DIM), hidden_size, bias=True), ) - def forward(self, x, c=None, noise_mask=None, c_noisy=None, c_clean=None): + def forward(self, x, c=None, noise_mask=None, c_noisy=None, c_clean=None) -> torch.Tensor: seq_len = x.shape[1] if noise_mask is not None: @@ -902,7 +902,7 @@ def forward( image_noise_mask: list[list[int]] | None = None, patch_size: int = 2, f_patch_size: int = 1, - ): + ) -> Transformer2DModelOutput | tuple[torch.Tensor]: """ The [`ZImageTransformer2DModel`] forward method. @@ -930,6 +930,10 @@ def forward( Spatial patch size used to patchify the input latents. f_patch_size (`int`, *optional*, defaults to 1): Temporal patch size used to patchify the input latents. + + Returns: + If `return_dict` is True, a [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a + `tuple` where the first element is the sample tensor. """ assert patch_size in self.all_patch_size and f_patch_size in self.all_f_patch_size omni_mode = isinstance(x[0], list) diff --git a/src/diffusers/models/unets/unet_3d_blocks.py b/src/diffusers/models/unets/unet_3d_blocks.py index e0d7f03bea3a..49ea270e67f4 100644 --- a/src/diffusers/models/unets/unet_3d_blocks.py +++ b/src/diffusers/models/unets/unet_3d_blocks.py @@ -936,7 +936,7 @@ def forward( self, hidden_states: torch.Tensor, image_only_indicator: torch.Tensor, - ): + ) -> torch.Tensor: hidden_states = self.resnets[0]( hidden_states, image_only_indicator=image_only_indicator, diff --git a/src/diffusers/models/unets/unet_kandinsky3.py b/src/diffusers/models/unets/unet_kandinsky3.py index 790d255101a4..571cf0c7c4d1 100644 --- a/src/diffusers/models/unets/unet_kandinsky3.py +++ b/src/diffusers/models/unets/unet_kandinsky3.py @@ -39,7 +39,7 @@ def __init__(self, encoder_hid_dim, cross_attention_dim): self.projection_linear = nn.Linear(encoder_hid_dim, cross_attention_dim, bias=False) self.projection_norm = nn.LayerNorm(cross_attention_dim) - def forward(self, x): + def forward(self, x) -> torch.Tensor: x = self.projection_linear(x) x = self.projection_norm(x) return x @@ -146,7 +146,9 @@ def set_default_attn_processor(self): """ self.set_attn_processor(AttnProcessor()) - def forward(self, sample, timestep, encoder_hidden_states=None, encoder_attention_mask=None, return_dict=True): + def forward( + self, sample, timestep, encoder_hidden_states=None, encoder_attention_mask=None, return_dict=True + ) -> Kandinsky3UNetOutput | tuple[torch.Tensor]: r""" Args: sample (`torch.Tensor`): Input sample. @@ -159,6 +161,10 @@ def forward(self, sample, timestep, encoder_hidden_states=None, encoder_attentio return_dict (`bool`, *optional*, defaults to `True`): Whether or not to return a [`~models.unets.unet_2d_condition.UNet2DConditionOutput`] instead of a plain tuple. + + Returns: + If `return_dict` is True, a [`~models.unets.unet_kandinsky3.Kandinsky3UNetOutput`] is returned, otherwise a + `tuple` where the first element is the sample tensor. """ if encoder_attention_mask is not None: encoder_attention_mask = (1 - encoder_attention_mask.to(sample.dtype)) * -10000.0 @@ -260,7 +266,7 @@ def __init__( self.resnets_in = nn.ModuleList(resnets_in) self.resnets_out = nn.ModuleList(resnets_out) - def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None): + def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None) -> torch.Tensor: for attention, resnet_in, resnet_out in zip(self.attentions[1:], self.resnets_in, self.resnets_out): x = resnet_in(x, time_embed) if self.context_dim is not None: @@ -328,7 +334,7 @@ def __init__( self.resnets_in = nn.ModuleList(resnets_in) self.resnets_out = nn.ModuleList(resnets_out) - def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None): + def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None) -> torch.Tensor: if self.self_attention: x = self.attentions[0](x, time_embed, image_mask=image_mask) @@ -348,7 +354,7 @@ def __init__(self, groups, normalized_shape, context_dim): self.context_mlp[1].weight.data.zero_() self.context_mlp[1].bias.data.zero_() - def forward(self, x, context): + def forward(self, x, context) -> torch.Tensor: context = self.context_mlp(context) for _ in range(len(x.shape[2:])): @@ -377,7 +383,7 @@ def __init__(self, in_channels, out_channels, time_embed_dim, kernel_size=3, nor else: self.down_sample = nn.Identity() - def forward(self, x, time_embed): + def forward(self, x, time_embed) -> torch.Tensor: x = self.group_norm(x, time_embed) x = self.activation(x) x = self.up_sample(x) @@ -418,7 +424,7 @@ def __init__( else nn.Identity() ) - def forward(self, x, time_embed): + def forward(self, x, time_embed) -> torch.Tensor: out = x for resnet_block in self.resnet_blocks: out = resnet_block(out, time_embed) @@ -441,7 +447,7 @@ def __init__(self, num_channels, context_dim, head_dim=64): out_bias=False, ) - def forward(self, x, context, context_mask=None): + def forward(self, x, context, context_mask=None) -> torch.Tensor: context_mask = context_mask.to(dtype=context.dtype) context = self.attention(context.mean(dim=1, keepdim=True), context, context_mask) return x + context.squeeze(1) @@ -467,7 +473,7 @@ def __init__(self, num_channels, time_embed_dim, context_dim=None, norm_groups=3 nn.Conv2d(hidden_channels, num_channels, kernel_size=1, bias=False), ) - def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None): + def forward(self, x, time_embed, context=None, context_mask=None, image_mask=None) -> torch.Tensor: height, width = x.shape[-2:] out = self.in_norm(x, time_embed) out = out.reshape(x.shape[0], -1, height * width).permute(0, 2, 1) diff --git a/src/diffusers/models/unets/unet_motion_model.py b/src/diffusers/models/unets/unet_motion_model.py index faa181d9bfd5..23063bc498f1 100644 --- a/src/diffusers/models/unets/unet_motion_model.py +++ b/src/diffusers/models/unets/unet_motion_model.py @@ -484,7 +484,7 @@ def forward( encoder_attention_mask: torch.Tensor | None = None, cross_attention_kwargs: dict[str, Any] | None = None, additional_residuals: torch.Tensor | None = None, - ): + ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]]: if cross_attention_kwargs is not None: if cross_attention_kwargs.get("scale", None) is not None: logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") @@ -1190,7 +1190,7 @@ def __init__( self.down_blocks = nn.ModuleList(down_blocks) self.up_blocks = nn.ModuleList(up_blocks) - def forward(self, sample): + def forward(self, sample) -> None: r""" Args: sample (`torch.Tensor`): Input sample. diff --git a/src/diffusers/models/unets/unet_stable_cascade.py b/src/diffusers/models/unets/unet_stable_cascade.py index e000fdc51e06..f101de40656e 100644 --- a/src/diffusers/models/unets/unet_stable_cascade.py +++ b/src/diffusers/models/unets/unet_stable_cascade.py @@ -46,7 +46,7 @@ def __init__(self, c, c_timestep, conds=[]): for cname in conds: setattr(self, f"mapper_{cname}", nn.Linear(c_timestep, c * 2)) - def forward(self, x, t): + def forward(self, x, t) -> torch.Tensor: t = t.chunk(len(self.conds) + 1, dim=1) a, b = self.mapper(t[0])[:, :, None, None].chunk(2, dim=1) for i, c in enumerate(self.conds): @@ -68,7 +68,7 @@ def __init__(self, c, c_skip=0, kernel_size=3, dropout=0.0): nn.Linear(c * 4, c), ) - def forward(self, x, x_skip=None): + def forward(self, x, x_skip=None) -> torch.Tensor: x_res = x x = self.norm(self.depthwise(x)) if x_skip is not None: @@ -84,7 +84,7 @@ def __init__(self, dim): self.gamma = nn.Parameter(torch.zeros(1, 1, 1, dim)) self.beta = nn.Parameter(torch.zeros(1, 1, 1, dim)) - def forward(self, x): + def forward(self, x) -> torch.Tensor: agg_norm = torch.norm(x, p=2, dim=(1, 2), keepdim=True) stand_div_norm = agg_norm / (agg_norm.mean(dim=-1, keepdim=True) + 1e-6) return self.gamma * (x * stand_div_norm) + self.beta + x @@ -99,7 +99,7 @@ def __init__(self, c, c_cond, nhead, self_attn=True, dropout=0.0): self.attention = Attention(query_dim=c, heads=nhead, dim_head=c // nhead, dropout=dropout, bias=True) self.kv_mapper = nn.Sequential(nn.SiLU(), nn.Linear(c_cond, c)) - def forward(self, x, kv): + def forward(self, x, kv) -> torch.Tensor: kv = self.kv_mapper(kv) norm_x = self.norm(x) if self.self_attn: @@ -122,7 +122,7 @@ def __init__(self, in_channels, out_channels, mode, enabled=True): mapping = nn.Conv2d(in_channels, out_channels, kernel_size=1) self.blocks = nn.ModuleList([interpolation, mapping] if mode == "up" else [mapping, interpolation]) - def forward(self, x): + def forward(self, x) -> torch.Tensor: for block in self.blocks: x = block(x) return x @@ -547,7 +547,7 @@ def forward( sca=None, crp=None, return_dict=True, - ): + ) -> StableCascadeUNetOutput | tuple[torch.Tensor]: r""" Args: sample (`torch.Tensor`): The noisy input sample. @@ -569,6 +569,10 @@ def forward( Optional `crp` conditioning value used to build the timestep embedding. return_dict (`bool`, *optional*, defaults to `True`): Whether or not to return a [`StableCascadeUNetOutput`] instead of a plain tuple. + + Returns: + If `return_dict` is True, a [`~models.unets.unet_stable_cascade.StableCascadeUNetOutput`] is returned, + otherwise a `tuple` where the first element is the sample tensor. """ if pixels is None: pixels = sample.new_zeros(sample.size(0), 3, 8, 8) diff --git a/src/diffusers/models/unets/uvit_2d.py b/src/diffusers/models/unets/uvit_2d.py index 317abe80b1eb..3c4e24642e42 100644 --- a/src/diffusers/models/unets/uvit_2d.py +++ b/src/diffusers/models/unets/uvit_2d.py @@ -148,7 +148,9 @@ def __init__( self.gradient_checkpointing = False @apply_lora_scale("cross_attention_kwargs") - def forward(self, input_ids, encoder_hidden_states, pooled_text_emb, micro_conds, cross_attention_kwargs=None): + def forward( + self, input_ids, encoder_hidden_states, pooled_text_emb, micro_conds, cross_attention_kwargs=None + ) -> torch.Tensor: r""" Args: input_ids (`torch.LongTensor`): @@ -161,6 +163,10 @@ def forward(self, input_ids, encoder_hidden_states, pooled_text_emb, micro_conds Micro-conditioning values that are embedded and combined with `pooled_text_emb`. cross_attention_kwargs (`dict`, *optional*): A kwargs dictionary that if specified is passed along to the `AttentionProcessor`. + + Returns: + `torch.Tensor`: The logits over the codebook for each image token, of shape `(batch_size, codebook_size, + height, width)`. """ encoder_hidden_states = self.encoder_proj(encoder_hidden_states) encoder_hidden_states = self.encoder_proj_layer_norm(encoder_hidden_states) @@ -246,7 +252,7 @@ def __init__(self, in_channels, block_out_channels, vocab_size, elementwise_affi self.layer_norm = RMSNorm(in_channels, eps, elementwise_affine) self.conv = nn.Conv2d(in_channels, block_out_channels, kernel_size=1, bias=bias) - def forward(self, input_ids): + def forward(self, input_ids) -> torch.Tensor: embeddings = self.embeddings(input_ids) embeddings = self.layer_norm(embeddings) embeddings = embeddings.permute(0, 3, 1, 2) @@ -333,7 +339,7 @@ def __init__( else: self.upsample = None - def forward(self, x, pooled_text_emb, encoder_hidden_states, cross_attention_kwargs): + def forward(self, x, pooled_text_emb, encoder_hidden_states, cross_attention_kwargs) -> torch.Tensor: if self.downsample is not None: x = self.downsample(x) @@ -374,7 +380,7 @@ def __init__( self.channelwise_dropout = nn.Dropout(hidden_dropout) self.cond_embeds_mapper = nn.Linear(hidden_size, channels * 2, use_bias) - def forward(self, x, cond_embeds): + def forward(self, x, cond_embeds) -> torch.Tensor: x_res = x x = self.depthwise(x) @@ -413,7 +419,7 @@ def __init__( self.layer_norm = RMSNorm(in_channels, layer_norm_eps, ln_elementwise_affine) self.conv2 = nn.Conv2d(in_channels, codebook_size, kernel_size=1, bias=use_bias) - def forward(self, hidden_states): + def forward(self, hidden_states) -> torch.Tensor: hidden_states = self.conv1(hidden_states) hidden_states = self.layer_norm(hidden_states.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) logits = self.conv2(hidden_states) diff --git a/src/diffusers/modular_pipelines/anima/before_denoise.py b/src/diffusers/modular_pipelines/anima/before_denoise.py index ede832c9d80b..17a94811df6f 100644 --- a/src/diffusers/modular_pipelines/anima/before_denoise.py +++ b/src/diffusers/modular_pipelines/anima/before_denoise.py @@ -214,7 +214,9 @@ def _condition_prompt_embeds( return prompt_embeds.to(dtype=output_dtype, device=device) @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device conditioning_dtype = components.text_conditioner.dtype @@ -300,7 +302,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.batch_size = block_state.prompt_embeds.shape[0] @@ -363,7 +367,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) latent_height, latent_width = block_state.image_latents.shape[-2:] @@ -454,7 +460,9 @@ def prepare_latents( return randn_tensor(shape, generator=generator, device=device, dtype=dtype) @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.height = block_state.height or components.default_height @@ -520,7 +528,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -598,7 +608,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -685,7 +697,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/anima/decoders.py b/src/diffusers/modular_pipelines/anima/decoders.py index f1f4b475a4b8..3be9fc7c2ecb 100644 --- a/src/diffusers/modular_pipelines/anima/decoders.py +++ b/src/diffusers/modular_pipelines/anima/decoders.py @@ -46,7 +46,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("images", note="tensor output of the VAE decoder")] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) latents = block_state.latents.to(components.vae.dtype) @@ -107,7 +109,9 @@ def check_inputs(output_type): raise ValueError(f"Invalid output_type: {output_type}") @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.output_type) diff --git a/src/diffusers/modular_pipelines/anima/denoise.py b/src/diffusers/modular_pipelines/anima/denoise.py index d8146beefe72..86a3569648bc 100644 --- a/src/diffusers/modular_pipelines/anima/denoise.py +++ b/src/diffusers/modular_pipelines/anima/denoise.py @@ -40,7 +40,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[AnimaModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) timestep = t.expand(block_state.latents.shape[0]).to(block_state.dtype) @@ -117,7 +119,7 @@ def inputs(self) -> list[InputParam]: @torch.no_grad() def __call__( self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[AnimaModularPipeline, BlockState]: components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) guider_state = components.guider.prepare_inputs_from_block_state(block_state, self._guider_input_fields) @@ -151,7 +153,9 @@ def description(self) -> str: return "Step within the denoising loop that updates Anima latents." @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: AnimaModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[AnimaModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, t, block_state.latents, return_dict=False @@ -181,7 +185,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) num_warmup_steps = len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order diff --git a/src/diffusers/modular_pipelines/anima/encoders.py b/src/diffusers/modular_pipelines/anima/encoders.py index 68950f97be83..726a68280d82 100644 --- a/src/diffusers/modular_pipelines/anima/encoders.py +++ b/src/diffusers/modular_pipelines/anima/encoders.py @@ -235,7 +235,9 @@ def encode_prompt( } @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -379,7 +381,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: AnimaModularPipeline, state: PipelineState + ) -> tuple[AnimaModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/cosmos/after_decode.py b/src/diffusers/modular_pipelines/cosmos/after_decode.py index 7f8dd903d615..37c213fcc5ce 100644 --- a/src/diffusers/modular_pipelines/cosmos/after_decode.py +++ b/src/diffusers/modular_pipelines/cosmos/after_decode.py @@ -38,7 +38,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("action", type_hint=list[torch.Tensor], description="Generated action vectors.")] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) action_output = None if block_state.action_mode in {"inverse_dynamics", "policy"} and block_state.action_latents is not None: @@ -92,7 +94,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("output_path", type_hint=str, description="Path of the exported video file.")] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) output_path = str(block_state.output_path) fps = int(round(block_state.fps)) diff --git a/src/diffusers/modular_pipelines/cosmos/before_denoise.py b/src/diffusers/modular_pipelines/cosmos/before_denoise.py index 7e9d83fa6316..e6f9dbf05abb 100644 --- a/src/diffusers/modular_pipelines/cosmos/before_denoise.py +++ b/src/diffusers/modular_pipelines/cosmos/before_denoise.py @@ -48,7 +48,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device block_state.cond_text_segment = components._prepare_text_segment(block_state.cond_input_ids, device=device) @@ -124,7 +126,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device sampling_dtype = torch.float32 @@ -225,7 +229,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device sampling_dtype = torch.float32 @@ -324,7 +330,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device sampling_dtype = torch.float32 @@ -464,7 +472,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device has_image_condition = bool(block_state.vision_condition_indexes_for_pack) @@ -556,7 +566,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -652,7 +664,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -737,7 +751,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.cond_position_ids = torch.cat( [ @@ -838,7 +854,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.cond_position_ids = torch.cat( [block_state.cond_position_ids, block_state.cond_sound_segment["sound_mrope_ids"]], dim=1 @@ -927,7 +945,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.cond_position_ids = torch.cat( [block_state.cond_position_ids, block_state.cond_action_segment["action_mrope_ids"]], dim=1 @@ -970,7 +990,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device if components.config.use_native_flow_schedule: @@ -1045,7 +1067,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device sampling_dtype = torch.float32 @@ -1142,7 +1166,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device num_hints = len(block_state.control_latents) @@ -1244,7 +1270,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device components.scheduler.set_timesteps(block_state.num_inference_steps, device=device) @@ -1307,7 +1335,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/cosmos/before_encoder.py b/src/diffusers/modular_pipelines/cosmos/before_encoder.py index 2cdf68712cdf..e186559bbcd3 100644 --- a/src/diffusers/modular_pipelines/cosmos/before_encoder.py +++ b/src/diffusers/modular_pipelines/cosmos/before_encoder.py @@ -82,7 +82,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype diff --git a/src/diffusers/modular_pipelines/cosmos/decoders.py b/src/diffusers/modular_pipelines/cosmos/decoders.py index a76e48501d85..688c74629d77 100644 --- a/src/diffusers/modular_pipelines/cosmos/decoders.py +++ b/src/diffusers/modular_pipelines/cosmos/decoders.py @@ -44,7 +44,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("videos")] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -106,7 +108,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) if components.sound_tokenizer is None: raise ValueError("Sound decoding requires a sound-capable checkpoint with a sound_tokenizer.") @@ -171,7 +175,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) latents = block_state.latents vae_dtype = components.vae.dtype @@ -234,7 +240,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/cosmos/denoise.py b/src/diffusers/modular_pipelines/cosmos/denoise.py index 6a369357e96f..294213f48f4b 100644 --- a/src/diffusers/modular_pipelines/cosmos/denoise.py +++ b/src/diffusers/modular_pipelines/cosmos/denoise.py @@ -51,7 +51,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: device = components._execution_device block_state.vision_tokens = [block_state.latents.to(device=device, dtype=components.transformer.dtype)] block_state.vision_timesteps = torch.full( @@ -95,7 +97,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: device = components._execution_device block_state.sound_tokens = [block_state.sound_latents.to(device=device, dtype=components.transformer.dtype)] block_state.sound_timesteps = torch.full( @@ -141,7 +145,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: device = components._execution_device block_state.action_tokens = [block_state.action_latents.to(device=device, dtype=components.transformer.dtype)] block_state.action_timesteps = torch.full( @@ -189,7 +195,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: denoiser_input_fields = block_state.denoiser_input_fields loop_input_fields = block_state.as_dict() has_sound = "sound_tokens" in loop_input_fields @@ -297,7 +305,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("latents")] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: block_state.latents = components.scheduler.step( block_state.velocity_vision.unsqueeze(0), t, block_state.latents.unsqueeze(0), return_dict=False )[0].squeeze(0) @@ -348,7 +358,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("latents")] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: velocity_vision = block_state.velocity_vision.float() latents = block_state.latents.float() @@ -404,7 +416,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("sound_latents", type_hint=torch.Tensor, description="Updated sound latents.")] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: block_state.sound_latents = block_state.sound_scheduler.step( block_state.velocity_sound.unsqueeze(0), t, block_state.sound_latents.unsqueeze(0), return_dict=False )[0].squeeze(0) @@ -455,7 +469,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("action_latents", type_hint=torch.Tensor, description="Updated action latents.")] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: has_noisy_action = block_state.action_condition_mask.sum() < block_state.action_condition_mask.numel() if has_noisy_action: block_state.action_latents = block_state.action_scheduler.step( @@ -515,7 +531,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) mixed_precision = Cosmos3MixedPrecisionConfig.resolve( components.transformer, @@ -686,7 +704,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: device = components._execution_device dtype = components.transformer.dtype block_state.vision_tokens_full = [c.to(device=device, dtype=dtype) for c in block_state.control_latents] + [ @@ -798,7 +818,9 @@ def _forward(components, static, vision_tokens, vision_timesteps, context_name, return preds_vision[-1] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: # active-at: a None interval is always active; otherwise the timestep must fall within [lo, hi]. guidance_interval = block_state.guidance_interval guidance_active = guidance_interval is None or ( @@ -915,7 +937,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("latents", type_hint=torch.Tensor, description="Updated target latents for this chunk.")] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Cosmos3OmniModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Cosmos3OmniModularPipeline, BlockState]: block_state.latents = components.scheduler.step( block_state.velocity.unsqueeze(0), t, block_state.latents.unsqueeze(0), return_dict=False )[0].squeeze(0) diff --git a/src/diffusers/modular_pipelines/cosmos/encoders.py b/src/diffusers/modular_pipelines/cosmos/encoders.py index c29b6174edda..735bf7d61f90 100644 --- a/src/diffusers/modular_pipelines/cosmos/encoders.py +++ b/src/diffusers/modular_pipelines/cosmos/encoders.py @@ -162,7 +162,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.num_frames is None: block_state.num_frames = 189 @@ -273,7 +275,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) self._check_inputs(block_state) @@ -412,7 +416,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.num_frames is None: block_state.num_frames = 189 @@ -574,7 +580,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) self._check_inputs(block_state) if block_state.use_system_prompt is None: @@ -668,7 +676,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -772,7 +782,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -941,7 +953,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.vae.dtype @@ -1056,7 +1070,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py b/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py index e0bd89e77ded..ab6ebac2c8cd 100644 --- a/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py +++ b/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py @@ -882,7 +882,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Cosmos3OmniModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Cosmos3OmniModularPipeline, state: PipelineState + ) -> tuple[Cosmos3OmniModularPipeline, PipelineState]: num_chunks = state.get("num_chunks") state.set("output_chunks", []) state.set("previous_output", None) diff --git a/src/diffusers/modular_pipelines/ernie_image/before_denoise.py b/src/diffusers/modular_pipelines/ernie_image/before_denoise.py index 034230632396..f1301ec3fcd3 100644 --- a/src/diffusers/modular_pipelines/ernie_image/before_denoise.py +++ b/src/diffusers/modular_pipelines/ernie_image/before_denoise.py @@ -118,7 +118,9 @@ def _expand(hiddens: list[torch.Tensor], num_images_per_prompt: int) -> list[tor return [h for h in hiddens for _ in range(num_images_per_prompt)] @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype @@ -177,7 +179,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device num_inference_steps = block_state.num_inference_steps @@ -243,7 +247,9 @@ def _check_inputs(components: ErnieImageModularPipeline, height: int, width: int ) @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype diff --git a/src/diffusers/modular_pipelines/ernie_image/decoders.py b/src/diffusers/modular_pipelines/ernie_image/decoders.py index d7d056b82584..b31c08987eb1 100644 --- a/src/diffusers/modular_pipelines/ernie_image/decoders.py +++ b/src/diffusers/modular_pipelines/ernie_image/decoders.py @@ -73,7 +73,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("images", type_hint=list, description="The generated images.")] @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae device = block_state.latents.device diff --git a/src/diffusers/modular_pipelines/ernie_image/denoise.py b/src/diffusers/modular_pipelines/ernie_image/denoise.py index 3a2a2e312486..150947944dd5 100644 --- a/src/diffusers/modular_pipelines/ernie_image/denoise.py +++ b/src/diffusers/modular_pipelines/ernie_image/denoise.py @@ -59,7 +59,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[ErnieImageModularPipeline, BlockState]: latents = block_state.latents block_state.latent_model_input = latents.to(components.transformer.dtype) block_state.timestep = t.expand(latents.shape[0]).to(components.transformer.dtype) @@ -122,7 +124,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[ErnieImageModularPipeline, BlockState]: guider_inputs = { "text_bth": (block_state.text_bth, block_state.negative_text_bth), "text_lens": (block_state.text_lens, block_state.negative_text_lens), @@ -159,7 +163,9 @@ def description(self) -> str: return "Step within the denoising loop that updates the latents using the scheduler step." @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: ErnieImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[ErnieImageModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, t, block_state.latents, return_dict=False @@ -208,7 +214,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("latents", type_hint=torch.Tensor, description="The denoised latents.")] @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: for i, t in enumerate(block_state.timesteps): diff --git a/src/diffusers/modular_pipelines/ernie_image/encoders.py b/src/diffusers/modular_pipelines/ernie_image/encoders.py index 161646d181be..ec016bf520b7 100644 --- a/src/diffusers/modular_pipelines/ernie_image/encoders.py +++ b/src/diffusers/modular_pipelines/ernie_image/encoders.py @@ -121,7 +121,9 @@ def _enhance_prompt( return pe_tokenizer.decode(generated_ids, skip_special_tokens=True).strip() @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -223,7 +225,9 @@ def _encode( return text_hiddens @torch.no_grad() - def __call__(self, components: ErnieImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ErnieImageModularPipeline, state: PipelineState + ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/flux/before_denoise.py b/src/diffusers/modular_pipelines/flux/before_denoise.py index 243f9e927d74..8e5a285dbe11 100644 --- a/src/diffusers/modular_pipelines/flux/before_denoise.py +++ b/src/diffusers/modular_pipelines/flux/before_denoise.py @@ -194,7 +194,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -283,7 +285,9 @@ def get_timesteps(scheduler, num_inference_steps, strength, device): return timesteps, num_inference_steps - t_start @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -395,7 +399,9 @@ def prepare_latents( return latents @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.height = block_state.height or components.default_height block_state.width = block_state.width or components.default_width @@ -477,7 +483,9 @@ def check_inputs(image_latents, latents): raise ValueError(f"`image_latents` must have 3 dimensions (patchified), but got {image_latents.ndim}") @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(image_latents=block_state.image_latents, latents=block_state.latents) @@ -530,7 +538,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_embeds = block_state.prompt_embeds @@ -582,7 +592,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_embeds = block_state.prompt_embeds diff --git a/src/diffusers/modular_pipelines/flux/decoders.py b/src/diffusers/modular_pipelines/flux/decoders.py index 5fcde5008680..796bff258a6b 100644 --- a/src/diffusers/modular_pipelines/flux/decoders.py +++ b/src/diffusers/modular_pipelines/flux/decoders.py @@ -24,6 +24,7 @@ from ...video_processor import VaeImageProcessor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import FluxModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -89,7 +90,7 @@ def intermediate_outputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/flux/denoise.py b/src/diffusers/modular_pipelines/flux/denoise.py index 490ef6d88f57..7f2e20ffcec9 100644 --- a/src/diffusers/modular_pipelines/flux/denoise.py +++ b/src/diffusers/modular_pipelines/flux/denoise.py @@ -92,7 +92,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: FluxModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[FluxModularPipeline, BlockState]: noise_pred = components.transformer( hidden_states=block_state.latents, timestep=t.flatten() / 1000, @@ -174,7 +174,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: FluxModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[FluxModularPipeline, BlockState]: latents = block_state.latents latent_model_input = latents image_latents = block_state.image_latents @@ -219,7 +219,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("latents", type_hint=torch.Tensor, description="The denoised latents")] @torch.no_grad() - def __call__(self, components: FluxModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: FluxModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[FluxModularPipeline, BlockState]: # Perform scheduler step using the predicted output latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( @@ -270,7 +272,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/flux/encoders.py b/src/diffusers/modular_pipelines/flux/encoders.py index 5f7e61a535b7..fb7b40024b0a 100644 --- a/src/diffusers/modular_pipelines/flux/encoders.py +++ b/src/diffusers/modular_pipelines/flux/encoders.py @@ -116,7 +116,9 @@ def check_inputs(height, width, vae_scale_factor): raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState): + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.resized_image is None and block_state.image is None: @@ -169,7 +171,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam(name="processed_image")] @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState): + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: from ...pipelines.flux.pipeline_flux_kontext import PREFERRED_KONTEXT_RESOLUTIONS block_state = self.get_block_state(state) @@ -260,7 +264,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) image = getattr(block_state, self._image_input_name) @@ -451,7 +457,9 @@ def encode_prompt( return prompt_embeds, pooled_prompt_embeds @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: # Get inputs and intermediates block_state = self.get_block_state(state) self.check_inputs(block_state) diff --git a/src/diffusers/modular_pipelines/flux/inputs.py b/src/diffusers/modular_pipelines/flux/inputs.py index c513d237bee2..24cf107d0128 100644 --- a/src/diffusers/modular_pipelines/flux/inputs.py +++ b/src/diffusers/modular_pipelines/flux/inputs.py @@ -98,7 +98,9 @@ def check_inputs(self, components, block_state): ) @torch.no_grad() - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: # TODO: consider adding negative embeddings? block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -187,7 +189,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam(name="image_width", type_hint=int, description="The width of the image latents"), ] - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs (height/width calculation, patchify, and batch expansion) @@ -246,7 +250,9 @@ def __call__(self, components: FluxModularPipeline, state: PipelineState) -> Pip class FluxKontextAdditionalInputsStep(FluxAdditionalInputsStep): model_name = "flux-kontext" - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs (height/width calculation, patchify, and batch expansion) @@ -334,7 +340,9 @@ def check_inputs(height, width, vae_scale_factor): if width is not None and width % (vae_scale_factor * 2) != 0: raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") - def __call__(self, components: FluxModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: FluxModularPipeline, state: PipelineState + ) -> tuple[FluxModularPipeline, PipelineState]: block_state = self.get_block_state(state) height = block_state.height or components.default_height diff --git a/src/diffusers/modular_pipelines/flux2/before_denoise.py b/src/diffusers/modular_pipelines/flux2/before_denoise.py index 87a6b568a258..5ae79ffaedd9 100644 --- a/src/diffusers/modular_pipelines/flux2/before_denoise.py +++ b/src/diffusers/modular_pipelines/flux2/before_denoise.py @@ -145,7 +145,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -293,7 +295,9 @@ def prepare_latents( return latents @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.height = block_state.height or components.default_height block_state.width = block_state.width or components.default_width @@ -368,7 +372,9 @@ def _prepare_text_ids(x: torch.Tensor, t_coord: torch.Tensor | None = None): return torch.stack(out_ids) - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_embeds = block_state.prompt_embeds @@ -429,7 +435,9 @@ def _prepare_text_ids(x: torch.Tensor, t_coord: torch.Tensor | None = None): return torch.stack(out_ids) - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_embeds = block_state.prompt_embeds @@ -516,7 +524,9 @@ def _pack_latents(latents): return latents @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) image_latents = block_state.image_latents @@ -579,7 +589,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device batch_size = block_state.batch_size * block_state.num_images_per_prompt diff --git a/src/diffusers/modular_pipelines/flux2/decoders.py b/src/diffusers/modular_pipelines/flux2/decoders.py index 81f5ca00dc33..4d8490c45e32 100644 --- a/src/diffusers/modular_pipelines/flux2/decoders.py +++ b/src/diffusers/modular_pipelines/flux2/decoders.py @@ -26,6 +26,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import Flux2ModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -97,7 +98,7 @@ def _unpack_latents_with_ids(x: torch.Tensor, x_ids: torch.Tensor) -> torch.Tens return torch.stack(x_list, dim=0) @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) latents = block_state.latents @@ -162,7 +163,7 @@ def _unpatchify_latents(latents): return latents @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/flux2/denoise.py b/src/diffusers/modular_pipelines/flux2/denoise.py index 675f14b03c63..fa6877180057 100644 --- a/src/diffusers/modular_pipelines/flux2/denoise.py +++ b/src/diffusers/modular_pipelines/flux2/denoise.py @@ -106,7 +106,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: Flux2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[Flux2ModularPipeline, BlockState]: latents = block_state.latents latent_model_input = latents.to(components.transformer.dtype) img_ids = block_state.latent_ids @@ -195,7 +195,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: Flux2KleinModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[Flux2KleinModularPipeline, BlockState]: latents = block_state.latents latent_model_input = latents.to(components.transformer.dtype) img_ids = block_state.latent_ids @@ -310,7 +310,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: Flux2KleinModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[Flux2KleinModularPipeline, BlockState]: latents = block_state.latents latent_model_input = latents.to(components.transformer.dtype) img_ids = block_state.latent_ids @@ -379,7 +379,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("latents", type_hint=torch.Tensor, description="The denoised latents")] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Flux2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Flux2ModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, @@ -430,7 +432,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/flux2/encoders.py b/src/diffusers/modular_pipelines/flux2/encoders.py index 215f33b60ea8..df5e7c218ba7 100644 --- a/src/diffusers/modular_pipelines/flux2/encoders.py +++ b/src/diffusers/modular_pipelines/flux2/encoders.py @@ -152,7 +152,9 @@ def _get_mistral_3_prompt_embeds( return prompt_embeds @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -214,7 +216,9 @@ def check_inputs(block_state): raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(block_state.prompt)}") @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: import io import requests @@ -353,7 +357,9 @@ def _get_qwen3_prompt_embeds( return prompt_embeds @torch.no_grad() - def __call__(self, components: Flux2KleinModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2KleinModularPipeline, state: PipelineState + ) -> tuple[Flux2KleinModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -495,7 +501,9 @@ def _get_qwen3_prompt_embeds( return prompt_embeds @torch.no_grad() - def __call__(self, components: Flux2KleinModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2KleinModularPipeline, state: PipelineState + ) -> tuple[Flux2KleinModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -585,7 +593,9 @@ def _encode_vae_image(self, vae: AutoencoderKLFlux2, image: torch.Tensor, genera return image_latents @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) condition_images = block_state.condition_images diff --git a/src/diffusers/modular_pipelines/flux2/inputs.py b/src/diffusers/modular_pipelines/flux2/inputs.py index 6bfe6aec97fd..ea868ca795cb 100644 --- a/src/diffusers/modular_pipelines/flux2/inputs.py +++ b/src/diffusers/modular_pipelines/flux2/inputs.py @@ -71,7 +71,9 @@ def intermediate_outputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.batch_size = block_state.prompt_embeds.shape[0] @@ -146,7 +148,9 @@ def intermediate_outputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.batch_size = block_state.prompt_embeds.shape[0] @@ -202,7 +206,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam(name="condition_images", type_hint=list[torch.Tensor])] @torch.no_grad() - def __call__(self, components: Flux2ModularPipeline, state: PipelineState): + def __call__( + self, components: Flux2ModularPipeline, state: PipelineState + ) -> tuple[Flux2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) images = block_state.image diff --git a/src/diffusers/modular_pipelines/helios/before_denoise.py b/src/diffusers/modular_pipelines/helios/before_denoise.py index 593843d48272..bb60ce7f11f3 100644 --- a/src/diffusers/modular_pipelines/helios/before_denoise.py +++ b/src/diffusers/modular_pipelines/helios/before_denoise.py @@ -93,7 +93,9 @@ def check_inputs(self, components, block_state): ) @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -296,7 +298,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return outputs - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) for input_param in self._image_latent_inputs: @@ -400,7 +404,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -504,7 +510,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -614,7 +622,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) batch_size = block_state.batch_size @@ -719,7 +729,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.history_latents = torch.cat([block_state.history_latents, block_state.fake_image_latents], dim=2) @@ -757,7 +769,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) history_latents = block_state.history_latents @@ -809,7 +823,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) patch_size = components.transformer.config.patch_size diff --git a/src/diffusers/modular_pipelines/helios/decoders.py b/src/diffusers/modular_pipelines/helios/decoders.py index c448d36136e6..0ab55da4ac63 100644 --- a/src/diffusers/modular_pipelines/helios/decoders.py +++ b/src/diffusers/modular_pipelines/helios/decoders.py @@ -22,6 +22,7 @@ from ...video_processor import VideoProcessor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import HeliosModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -72,7 +73,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/helios/denoise.py b/src/diffusers/modular_pipelines/helios/denoise.py index 5fcf01a73ffc..32a6f0b11b79 100644 --- a/src/diffusers/modular_pipelines/helios/denoise.py +++ b/src/diffusers/modular_pipelines/helios/denoise.py @@ -132,7 +132,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: keep_first_frame = block_state.keep_first_frame history_sizes = block_state.history_sizes image_latents = block_state.image_latents @@ -218,7 +220,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: keep_first_frame = block_state.keep_first_frame history_sizes = block_state.history_sizes image_latents = block_state.image_latents @@ -257,7 +261,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: device = components._execution_device block_state.latents = randn_tensor( block_state.latent_shape, generator=block_state.generator, device=device, dtype=torch.float32 @@ -291,7 +297,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: device = components._execution_device batch_size, num_channels_latents, num_latent_frames, h_latent, w_latent = block_state.latent_shape @@ -337,7 +345,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: device = components._execution_device components.scheduler.set_timesteps( block_state.num_inference_steps, device=device, sigmas=block_state.sigmas, mu=block_state.mu @@ -392,7 +402,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: latents = block_state.latents timesteps = block_state.timesteps num_inference_steps = block_state.num_inference_steps @@ -511,7 +523,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: device = components._execution_device transformer_dtype = components.transformer.dtype latents = block_state.latents @@ -685,7 +699,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: # e. Collect denoised latents for this chunk block_state.latent_chunks.append(block_state.latents) @@ -733,7 +749,9 @@ def loop_intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.latent_chunks = [] @@ -848,7 +866,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: HeliosModularPipeline, block_state: BlockState, k: int + ) -> tuple[HeliosModularPipeline, BlockState]: device = components._execution_device transformer_dtype = components.transformer.dtype latents = block_state.latents diff --git a/src/diffusers/modular_pipelines/helios/encoders.py b/src/diffusers/modular_pipelines/helios/encoders.py index ce11f1b58762..15c55579218f 100644 --- a/src/diffusers/modular_pipelines/helios/encoders.py +++ b/src/diffusers/modular_pipelines/helios/encoders.py @@ -160,7 +160,9 @@ def check_inputs(prompt, negative_prompt): ) @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt = block_state.prompt @@ -248,7 +250,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae @@ -336,7 +340,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HeliosModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HeliosModularPipeline, state: PipelineState + ) -> tuple[HeliosModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/hunyuan_video1_5/before_denoise.py b/src/diffusers/modular_pipelines/hunyuan_video1_5/before_denoise.py index 4c02eb9dd084..f478c25b82e3 100644 --- a/src/diffusers/modular_pipelines/hunyuan_video1_5/before_denoise.py +++ b/src/diffusers/modular_pipelines/hunyuan_video1_5/before_denoise.py @@ -112,7 +112,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.batch_size = getattr(block_state, "batch_size", None) or block_state.prompt_embeds.shape[0] self.set_block_state(state, block_state) @@ -145,7 +147,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -202,7 +206,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype @@ -297,7 +303,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype diff --git a/src/diffusers/modular_pipelines/hunyuan_video1_5/decoders.py b/src/diffusers/modular_pipelines/hunyuan_video1_5/decoders.py index 630af85c1b10..ce5ecd4c9529 100644 --- a/src/diffusers/modular_pipelines/hunyuan_video1_5/decoders.py +++ b/src/diffusers/modular_pipelines/hunyuan_video1_5/decoders.py @@ -21,6 +21,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import HunyuanVideo15ModularPipeline logger = logging.get_logger(__name__) @@ -59,7 +60,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) latents = block_state.latents.to(components.vae.dtype) / components.vae.config.scaling_factor diff --git a/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py b/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py index 293fad57c93f..223d5e2eee9b 100644 --- a/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py +++ b/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py @@ -49,7 +49,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[HunyuanVideo15ModularPipeline, BlockState]: block_state.latent_model_input = torch.cat( [block_state.latents, block_state.cond_latents_concat, block_state.mask_concat], dim=1 ) @@ -131,7 +133,7 @@ def inputs(self) -> list[InputParam]: @torch.no_grad() def __call__( self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[HunyuanVideo15ModularPipeline, BlockState]: timestep = t.expand(block_state.latent_model_input.shape[0]).to(block_state.latent_model_input.dtype) # Step 1: Collect model inputs @@ -185,7 +187,9 @@ def description(self) -> str: return "Step within the denoising loop that updates the latents" @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[HunyuanVideo15ModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, t, block_state.latents, return_dict=False @@ -220,7 +224,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( @@ -335,7 +341,7 @@ def inputs(self) -> list[InputParam]: @torch.no_grad() def __call__( self, components: HunyuanVideo15ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[HunyuanVideo15ModularPipeline, BlockState]: timestep = t.expand(block_state.latent_model_input.shape[0]).to(block_state.latent_model_input.dtype) # MeanFlow timestep_r (lines 855-862) diff --git a/src/diffusers/modular_pipelines/hunyuan_video1_5/encoders.py b/src/diffusers/modular_pipelines/hunyuan_video1_5/encoders.py index 9d340cc88194..11f511d7e4e6 100644 --- a/src/diffusers/modular_pipelines/hunyuan_video1_5/encoders.py +++ b/src/diffusers/modular_pipelines/hunyuan_video1_5/encoders.py @@ -259,7 +259,9 @@ def encode_prompt( return prompt_embeds, prompt_embeds_mask, prompt_embeds_2, prompt_embeds_mask_2 @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device dtype = components.transformer.dtype @@ -363,7 +365,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -424,7 +428,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: HunyuanVideo15ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: HunyuanVideo15ModularPipeline, state: PipelineState + ) -> tuple[HunyuanVideo15ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ideogram4/before_denoise.py b/src/diffusers/modular_pipelines/ideogram4/before_denoise.py index 98be3b141aec..c29ee38e085b 100644 --- a/src/diffusers/modular_pipelines/ideogram4/before_denoise.py +++ b/src/diffusers/modular_pipelines/ideogram4/before_denoise.py @@ -178,7 +178,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_batch = block_state.text_features.shape[0] @@ -256,7 +258,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -351,7 +355,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -520,7 +526,9 @@ def _prepare_ids( return position_ids.to(device), segment_ids.to(device), indicator.to(device) @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ideogram4/decoders.py b/src/diffusers/modular_pipelines/ideogram4/decoders.py index bf5d69270b7c..710734ee7fb2 100644 --- a/src/diffusers/modular_pipelines/ideogram4/decoders.py +++ b/src/diffusers/modular_pipelines/ideogram4/decoders.py @@ -85,7 +85,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("images")] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) z = block_state.latents diff --git a/src/diffusers/modular_pipelines/ideogram4/denoise.py b/src/diffusers/modular_pipelines/ideogram4/denoise.py index 871db69d344c..db1a708c4315 100644 --- a/src/diffusers/modular_pipelines/ideogram4/denoise.py +++ b/src/diffusers/modular_pipelines/ideogram4/denoise.py @@ -56,7 +56,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Ideogram4ModularPipeline, BlockState]: # Conditional packed sequence is [text-padding][image latents]; text region length = total - image tokens. max_text_tokens = block_state.position_ids.shape[1] - block_state.latents.shape[1] text_z_padding = torch.zeros( @@ -150,7 +152,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Ideogram4ModularPipeline, BlockState]: transformer = components.transformer unconditional_transformer = components.unconditional_transformer @@ -200,7 +204,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam(name="latents", type_hint=torch.Tensor, description="The denoised latents.")] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Ideogram4ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Ideogram4ModularPipeline, BlockState]: block_state.latents = components.scheduler.step( block_state.noise_pred, t, block_state.latents, return_dict=False )[0] @@ -280,7 +286,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: @@ -344,7 +352,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) z = block_state.latents diff --git a/src/diffusers/modular_pipelines/ideogram4/encoders.py b/src/diffusers/modular_pipelines/ideogram4/encoders.py index 6e149fa8392e..fa7fc765ea9b 100644 --- a/src/diffusers/modular_pipelines/ideogram4/encoders.py +++ b/src/diffusers/modular_pipelines/ideogram4/encoders.py @@ -133,7 +133,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.prompt_upsampling: @@ -280,7 +282,9 @@ def _get_text_encoder_hidden_states( return [captured[i] for i in QWEN3_VL_ACTIVATION_LAYERS] @torch.no_grad() - def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Ideogram4ModularPipeline, state: PipelineState + ) -> tuple[Ideogram4ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/krea2/before_denoise.py b/src/diffusers/modular_pipelines/krea2/before_denoise.py index 63810d30a903..17ed6a3cd376 100644 --- a/src/diffusers/modular_pipelines/krea2/before_denoise.py +++ b/src/diffusers/modular_pipelines/krea2/before_denoise.py @@ -138,7 +138,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_batch, seq_len, num_layers, dim = block_state.prompt_embeds.shape @@ -233,7 +235,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) prompt_batch, seq_len, num_layers, dim = block_state.prompt_embeds.shape @@ -317,7 +321,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -415,7 +421,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam(name="timesteps", type_hint=torch.Tensor, description="The denoising timesteps.")] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -491,7 +499,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam(name="timesteps", type_hint=torch.Tensor, description="The denoising timesteps.")] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -575,7 +585,9 @@ def prepare_position_ids(text_seq_len: int, grid_height: int, grid_width: int, d return torch.cat([text_ids, image_ids], dim=0) @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/krea2/decoders.py b/src/diffusers/modular_pipelines/krea2/decoders.py index fd308b5ef648..2c3fe4c0bf20 100644 --- a/src/diffusers/modular_pipelines/krea2/decoders.py +++ b/src/diffusers/modular_pipelines/krea2/decoders.py @@ -92,7 +92,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("images")] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/krea2/denoise.py b/src/diffusers/modular_pipelines/krea2/denoise.py index 88c6cdca7aba..ccb972ca74c1 100644 --- a/src/diffusers/modular_pipelines/krea2/denoise.py +++ b/src/diffusers/modular_pipelines/krea2/denoise.py @@ -55,7 +55,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Krea2ModularPipeline, BlockState]: num_train_timesteps = components.scheduler.config.num_train_timesteps block_state.timestep = (t / num_train_timesteps).expand(block_state.batch_size) return components, block_state @@ -113,7 +115,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Krea2ModularPipeline, BlockState]: transformer = components.transformer latents = block_state.latents.to(transformer.dtype) @@ -190,7 +194,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Krea2ModularPipeline, BlockState]: transformer = components.transformer latents = block_state.latents.to(transformer.dtype) @@ -224,7 +230,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam(name="latents", type_hint=torch.Tensor, description="The denoised latents.")] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[Krea2ModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, t, block_state.latents, return_dict=False @@ -261,7 +269,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: diff --git a/src/diffusers/modular_pipelines/krea2/encoders.py b/src/diffusers/modular_pipelines/krea2/encoders.py index 7640222e9ad2..a07f75305557 100644 --- a/src/diffusers/modular_pipelines/krea2/encoders.py +++ b/src/diffusers/modular_pipelines/krea2/encoders.py @@ -170,7 +170,9 @@ def _encode_prompt(self, components, prompt, max_sequence_length, device): return hidden_states, attention_mask @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -262,7 +264,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: Krea2ModularPipeline, state: PipelineState + ) -> tuple[Krea2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ltx/before_denoise.py b/src/diffusers/modular_pipelines/ltx/before_denoise.py index cd8b3ea82b82..5543742a4d46 100644 --- a/src/diffusers/modular_pipelines/ltx/before_denoise.py +++ b/src/diffusers/modular_pipelines/ltx/before_denoise.py @@ -131,7 +131,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.batch_size = block_state.prompt_embeds.shape[0] @@ -196,7 +198,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -288,7 +292,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -353,7 +359,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ltx/decoders.py b/src/diffusers/modular_pipelines/ltx/decoders.py index 8664dee25bfe..9d5238069bcf 100644 --- a/src/diffusers/modular_pipelines/ltx/decoders.py +++ b/src/diffusers/modular_pipelines/ltx/decoders.py @@ -23,7 +23,7 @@ from ...video_processor import VideoProcessor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam -from .modular_pipeline import LTXVideoPachifier +from .modular_pipeline import LTXModularPipeline, LTXVideoPachifier logger = logging.get_logger(__name__) @@ -84,7 +84,7 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("videos")] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/ltx/denoise.py b/src/diffusers/modular_pipelines/ltx/denoise.py index b3ed86b51679..44a3c91d471b 100644 --- a/src/diffusers/modular_pipelines/ltx/denoise.py +++ b/src/diffusers/modular_pipelines/ltx/denoise.py @@ -49,7 +49,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTXModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) return components, block_state @@ -115,7 +117,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[LTXModularPipeline, BlockState]: components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) latent_num_frames = (block_state.num_frames - 1) // components.vae_temporal_compression_ratio + 1 @@ -171,7 +173,9 @@ def description(self) -> str: ) @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTXModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, @@ -211,7 +215,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( @@ -275,7 +281,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTXModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) block_state.timestep_adjusted = t.expand(block_state.latent_model_input.shape[0]).unsqueeze(-1) * ( 1 - block_state.conditioning_mask @@ -342,7 +350,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[LTXModularPipeline, BlockState]: components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) latent_num_frames = (block_state.num_frames - 1) // components.vae_temporal_compression_ratio + 1 @@ -411,7 +419,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: LTXModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTXModularPipeline, BlockState]: latent_num_frames = (block_state.num_frames - 1) // components.vae_temporal_compression_ratio + 1 latent_height = block_state.height // components.vae_spatial_compression_ratio latent_width = block_state.width // components.vae_spatial_compression_ratio diff --git a/src/diffusers/modular_pipelines/ltx/encoders.py b/src/diffusers/modular_pipelines/ltx/encoders.py index 55405ad0aefe..f8e1851eb5bc 100644 --- a/src/diffusers/modular_pipelines/ltx/encoders.py +++ b/src/diffusers/modular_pipelines/ltx/encoders.py @@ -148,7 +148,9 @@ def encode_prompt( return prompt_embeds, prompt_attention_mask, negative_prompt_embeds, negative_prompt_attention_mask @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -235,7 +237,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: LTXModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: LTXModularPipeline, state: PipelineState + ) -> tuple[LTXModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ltx2/before_denoise.py b/src/diffusers/modular_pipelines/ltx2/before_denoise.py index 81ffc28188ea..1878c01305fb 100644 --- a/src/diffusers/modular_pipelines/ltx2/before_denoise.py +++ b/src/diffusers/modular_pipelines/ltx2/before_denoise.py @@ -30,6 +30,7 @@ from ...utils.torch_utils import randn_tensor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import LTX2ModularPipeline logger = logging.get_logger(__name__) @@ -318,7 +319,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) # `repeat_interleave` keeps each prompt's copies contiguous, matching how the latents are laid out @@ -376,7 +377,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -484,7 +485,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -580,7 +581,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -682,7 +683,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -779,7 +780,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -924,7 +925,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1225,7 +1226,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1452,7 +1453,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1523,7 +1524,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1632,7 +1633,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1739,7 +1740,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/ltx2/decoders.py b/src/diffusers/modular_pipelines/ltx2/decoders.py index fc957a3f9925..cbeb1a66826d 100644 --- a/src/diffusers/modular_pipelines/ltx2/decoders.py +++ b/src/diffusers/modular_pipelines/ltx2/decoders.py @@ -30,6 +30,7 @@ from ...video_processor import VideoProcessor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import LTX2ModularPipeline logger = logging.get_logger(__name__) @@ -118,7 +119,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.latents = block_state.latents[:, : block_state.base_token_count] self.set_block_state(state, block_state) @@ -168,7 +169,7 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("videos")] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) decoder = components.diffusion_decoder @@ -258,7 +259,7 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("videos")] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae @@ -362,7 +363,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) audio_vae = components.audio_vae diff --git a/src/diffusers/modular_pipelines/ltx2/denoise.py b/src/diffusers/modular_pipelines/ltx2/denoise.py index b1c4657d4d04..2fb4951377a1 100644 --- a/src/diffusers/modular_pipelines/ltx2/denoise.py +++ b/src/diffusers/modular_pipelines/ltx2/denoise.py @@ -29,6 +29,7 @@ PipelineState, ) from ..modular_pipeline_utils import ComponentSpec, InputParam +from .modular_pipeline import LTX2ModularPipeline # Velocity-space helpers, mirrored from `diffusers.pipelines.ltx2.pipeline_ltx2.LTX2Pipeline` and redefined here @@ -94,7 +95,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) block_state.audio_latent_model_input = block_state.audio_latents.to(block_state.dtype) timestep = t.expand(block_state.latents.shape[0]) @@ -130,7 +133,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) block_state.audio_latent_model_input = block_state.audio_latents.to(block_state.dtype) timestep = t.expand(block_state.latents.shape[0]) @@ -172,7 +177,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) block_state.audio_latent_model_input = block_state.audio_latents.to(block_state.dtype) timestep = t.expand(block_state.latents.shape[0]) @@ -337,7 +344,9 @@ def inputs(self) -> list[InputParam]: return inputs @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: latent_num_frames = (block_state.num_frames - 1) // components.vae_temporal_compression_ratio + 1 latent_height = block_state.height // components.vae_spatial_compression_ratio latent_width = block_state.width // components.vae_spatial_compression_ratio @@ -461,7 +470,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: noise_pred_video = convert_x0_to_velocity( block_state.latents, block_state.noise_pred_video, i, components.scheduler ) @@ -514,7 +525,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: spatial_patch = components.transformer_spatial_patch_size temporal_patch = components.transformer_temporal_patch_size latent_num_frames = (block_state.num_frames - 1) // components.vae_temporal_compression_ratio + 1 @@ -587,7 +600,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[LTX2ModularPipeline, BlockState]: # Conditioning strengths run from 0 (always use the denoised sample) to 1 (always use the condition), with # intermediate values specifying how strongly to follow the condition. Applied in x0 space, not velocity # space (which is what the transformer outputs). @@ -632,7 +647,7 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/ltx2/encoders.py b/src/diffusers/modular_pipelines/ltx2/encoders.py index b261597c0f68..ae3d29b5cfe3 100644 --- a/src/diffusers/modular_pipelines/ltx2/encoders.py +++ b/src/diffusers/modular_pipelines/ltx2/encoders.py @@ -50,6 +50,7 @@ from ...video_processor import VideoProcessor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import LTX2ModularPipeline logger = logging.get_logger(__name__) @@ -221,7 +222,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not block_state.enable_prompt_enhancement: @@ -318,7 +319,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not block_state.enable_prompt_enhancement: @@ -427,7 +428,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not block_state.enable_prompt_enhancement: @@ -543,7 +544,7 @@ def check_inputs(block_state): raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(block_state.prompt)}") @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -631,7 +632,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) padding_side = components.tokenizer.padding_side @@ -721,7 +722,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if getattr(components, "duration_head", None) is None: @@ -887,7 +888,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1006,7 +1007,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1228,7 +1229,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py b/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py index 247b9e88d761..94395a7c3ebe 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py +++ b/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py @@ -155,7 +155,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.keyframe_anchors = () @@ -371,7 +373,9 @@ def build_packed_sequence( return position_ids, token_tags, video_indices, audio_indices, text_indices, num_condition_rows, 0 @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -725,7 +729,9 @@ def build_ref2va_packed_sequence( ) @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -845,7 +851,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device patch_size = components.patch_size @@ -944,7 +952,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device patch_size = components.patch_size @@ -1012,7 +1022,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.latents = torch.cat([block_state.condition_rows, block_state.latents]) @@ -1088,7 +1100,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1222,7 +1236,9 @@ def build_row_timesteps( return torch.unique(row_timesteps, sorted=True, return_inverse=True) @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/minimax_h3/before_encoder.py b/src/diffusers/modular_pipelines/minimax_h3/before_encoder.py index ebae839fda65..dea178af4cb7 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/before_encoder.py +++ b/src/diffusers/modular_pipelines/minimax_h3/before_encoder.py @@ -112,7 +112,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) keyframes = [keyframe for keyframe in (block_state.image, block_state.last_image) if keyframe is not None] @@ -383,7 +385,9 @@ def _normalize_audio_condition( return torchaudio.transforms.Resample(sample_rate, target_sample_rate)(waveform) @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) # 1. Validate the request. diff --git a/src/diffusers/modular_pipelines/minimax_h3/decoders.py b/src/diffusers/modular_pipelines/minimax_h3/decoders.py index e5b624cdb6c4..bdc7adf7c5e7 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/decoders.py +++ b/src/diffusers/modular_pipelines/minimax_h3/decoders.py @@ -96,7 +96,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) patch_t, patch_h, patch_w = components.patch_size channels = components.vae_latent_channels @@ -169,7 +171,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("videos", description="The generated video.")] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -234,7 +238,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/minimax_h3/denoise.py b/src/diffusers/modular_pipelines/minimax_h3/denoise.py index 2f2ce59bfda6..efcf36c1130b 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/denoise.py +++ b/src/diffusers/modular_pipelines/minimax_h3/denoise.py @@ -108,7 +108,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: MiniMaxH3ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[MiniMaxH3ModularPipeline, BlockState]: transformer = getattr(components, self.transformer_name) unique_timesteps, timestep_indices = block_state.row_timestep_plan[i] # The layout tags its outputs `denoiser_input_fields`, and their names are the transformer's own parameter @@ -218,7 +220,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: MiniMaxH3ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[MiniMaxH3ModularPipeline, BlockState]: num_condition_video_rows = block_state.num_condition_video_rows num_condition_audio_rows = block_state.num_condition_audio_rows @@ -258,7 +262,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) with self.progress_bar(total=len(block_state.timesteps)) as progress_bar: for i, t in enumerate(block_state.timesteps): diff --git a/src/diffusers/modular_pipelines/minimax_h3/encoders.py b/src/diffusers/modular_pipelines/minimax_h3/encoders.py index 8ed4d548e9d1..80e84895f6e0 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/encoders.py +++ b/src/diffusers/modular_pipelines/minimax_h3/encoders.py @@ -181,7 +181,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not isinstance(block_state.prompt, str): raise ValueError( @@ -258,7 +260,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not isinstance(block_state.prompt, str): raise ValueError( @@ -354,7 +358,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -606,7 +612,9 @@ def emit(segment: tuple[list[int], list[int]]) -> None: return token_ids, token_tags @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not isinstance(block_state.prompt, str): raise ValueError( @@ -705,7 +713,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxH3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxH3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/minimax_music3/before_denoise.py b/src/diffusers/modular_pipelines/minimax_music3/before_denoise.py index be58527d681f..31f8411b5357 100644 --- a/src/diffusers/modular_pipelines/minimax_music3/before_denoise.py +++ b/src/diffusers/modular_pipelines/minimax_music3/before_denoise.py @@ -61,7 +61,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxMusic3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxMusic3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) num_frames = block_state.frame_hiddens.shape[1] diff --git a/src/diffusers/modular_pipelines/minimax_music3/decoders.py b/src/diffusers/modular_pipelines/minimax_music3/decoders.py index 2472542f1563..c74773862e59 100644 --- a/src/diffusers/modular_pipelines/minimax_music3/decoders.py +++ b/src/diffusers/modular_pipelines/minimax_music3/decoders.py @@ -73,7 +73,9 @@ def check_inputs(block_state): raise ValueError(f"Invalid output_type: {block_state.output_type}") @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxMusic3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxMusic3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) diff --git a/src/diffusers/modular_pipelines/minimax_music3/denoise.py b/src/diffusers/modular_pipelines/minimax_music3/denoise.py index d65155af07c4..0a709b8bf85b 100644 --- a/src/diffusers/modular_pipelines/minimax_music3/denoise.py +++ b/src/diffusers/modular_pipelines/minimax_music3/denoise.py @@ -74,7 +74,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int + ) -> tuple[MiniMaxMusic3ModularPipeline, BlockState]: device = components._execution_device chunk_start = block_state.chunk_starts[k] @@ -111,7 +113,9 @@ def inputs(self) -> list[InputParam]: return [InputParam.template("generator")] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int + ) -> tuple[MiniMaxMusic3ModularPipeline, BlockState]: device = components._execution_device latents = randn_tensor( @@ -148,7 +152,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int + ) -> tuple[MiniMaxMusic3ModularPipeline, BlockState]: device = components._execution_device sigmas = np.linspace(1.0, 1.0 / block_state.num_inference_steps, block_state.num_inference_steps) @@ -194,7 +200,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int + ) -> tuple[MiniMaxMusic3ModularPipeline, BlockState]: latents = block_state.latents timesteps = block_state.timesteps overlap = block_state.overlap @@ -246,7 +254,9 @@ def description(self) -> str: ) @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int): + def __call__( + self, components: MiniMaxMusic3ModularPipeline, block_state: BlockState, k: int + ) -> tuple[MiniMaxMusic3ModularPipeline, BlockState]: latents = block_state.latents if block_state.overlap > 0: latents[..., : block_state.overlap] = block_state.previous_latent[..., : block_state.overlap] @@ -292,7 +302,9 @@ def loop_intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxMusic3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxMusic3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.latent_chunks = [] diff --git a/src/diffusers/modular_pipelines/minimax_music3/encoders.py b/src/diffusers/modular_pipelines/minimax_music3/encoders.py index 6fae35be3ffe..fafd8cccc0ea 100644 --- a/src/diffusers/modular_pipelines/minimax_music3/encoders.py +++ b/src/diffusers/modular_pipelines/minimax_music3/encoders.py @@ -206,7 +206,9 @@ def check_inputs(block_state): raise ValueError(f"`lyrics` must be a non-empty string, got {block_state.lyrics!r}") @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxMusic3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxMusic3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -285,7 +287,9 @@ def check_inputs(block_state): raise ValueError(f"`audio_duration` must be positive, got {block_state.audio_duration}") @torch.no_grad() - def __call__(self, components: MiniMaxMusic3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: MiniMaxMusic3ModularPipeline, state: PipelineState + ) -> tuple[MiniMaxMusic3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) diff --git a/src/diffusers/modular_pipelines/modular_pipeline.py b/src/diffusers/modular_pipelines/modular_pipeline.py index e9e5463c1e72..0cca16877623 100644 --- a/src/diffusers/modular_pipelines/modular_pipeline.py +++ b/src/diffusers/modular_pipelines/modular_pipeline.py @@ -777,7 +777,7 @@ def select_block(self, **kwargs) -> str | None: raise NotImplementedError(f"Subclass {self.__class__.__name__} must implement the `select_block` method.") @torch.no_grad() - def __call__(self, pipeline, state: PipelineState) -> PipelineState: + def __call__(self, pipeline, state: PipelineState) -> tuple["ModularPipeline", PipelineState]: trigger_kwargs = {name: state.get(name) for name in self.block_trigger_inputs if name is not None} block_name = self.select_block(**trigger_kwargs) @@ -1149,7 +1149,7 @@ def outputs(self) -> list[str]: return self.intermediate_outputs @torch.no_grad() - def __call__(self, pipeline, state: PipelineState) -> PipelineState: + def __call__(self, pipeline, state: PipelineState) -> tuple["ModularPipeline", PipelineState]: for block_name, block in self.sub_blocks.items(): try: pipeline, state = block(pipeline, state) @@ -1533,7 +1533,7 @@ def loop_step(self, components, state: PipelineState, **kwargs): raise return components, state - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple["ModularPipeline", PipelineState]: raise NotImplementedError("`__call__` method needs to be implemented by the subclass") @property diff --git a/src/diffusers/modular_pipelines/qwenimage/before_denoise.py b/src/diffusers/modular_pipelines/qwenimage/before_denoise.py index b928bf7fce9e..e4720b16244d 100644 --- a/src/diffusers/modular_pipelines/qwenimage/before_denoise.py +++ b/src/diffusers/modular_pipelines/qwenimage/before_denoise.py @@ -196,7 +196,9 @@ def check_inputs(height, width, vae_scale_factor): raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -315,7 +317,9 @@ def check_inputs(height, width, vae_scale_factor): raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -433,7 +437,9 @@ def check_inputs(image_latents, latents): raise ValueError(f"`image_latents` must have 3 dimensions (patchified), but got {image_latents.ndim}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -515,7 +521,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -601,7 +609,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -683,7 +693,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -780,7 +790,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -875,7 +887,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.img_shapes = [ @@ -960,7 +974,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # for edit, image size can be different from the target size (height/width) @@ -1072,7 +1088,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae_scale_factor = components.vae_scale_factor @@ -1182,7 +1200,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1281,7 +1299,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) controlnet = unwrap_module(components.controlnet) diff --git a/src/diffusers/modular_pipelines/qwenimage/decoders.py b/src/diffusers/modular_pipelines/qwenimage/decoders.py index e4ccb6b8e047..9ba2365416b2 100644 --- a/src/diffusers/modular_pipelines/qwenimage/decoders.py +++ b/src/diffusers/modular_pipelines/qwenimage/decoders.py @@ -90,7 +90,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae_scale_factor = components.vae_scale_factor @@ -158,7 +160,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Unpack: (B, seq, C*4) -> (B, C, layers+1, H, W) @@ -225,7 +227,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("images", note="tensor output of the vae decoder.")] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # YiYi Notes: remove support for output_type = "latents', we can just skip decode/encode step in modular @@ -307,7 +311,7 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam.template("images")] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) latents = block_state.latents @@ -409,7 +413,9 @@ def check_inputs(output_type): raise ValueError(f"Invalid output_type: {output_type}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.output_type) @@ -492,7 +498,9 @@ def check_inputs(output_type, mask_overlay_kwargs): raise ValueError("only support output_type 'pil' for mask overlay") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.output_type, block_state.mask_overlay_kwargs) diff --git a/src/diffusers/modular_pipelines/qwenimage/denoise.py b/src/diffusers/modular_pipelines/qwenimage/denoise.py index de8ea05c5047..7f271782f82b 100644 --- a/src/diffusers/modular_pipelines/qwenimage/denoise.py +++ b/src/diffusers/modular_pipelines/qwenimage/denoise.py @@ -57,7 +57,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: # one timestep block_state.timestep = t.expand(block_state.latents.shape[0]).to(block_state.latents.dtype) block_state.latent_model_input = block_state.latents @@ -88,7 +90,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: # one timestep block_state.latent_model_input = torch.cat([block_state.latents, block_state.image_latents], dim=1) @@ -138,7 +142,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[QwenImageModularPipeline, BlockState]: # cond_scale for the timestep (controlnet input) if isinstance(block_state.controlnet_keep[i], list): block_state.cond_scale = [ @@ -205,7 +211,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: guider_inputs = { "encoder_hidden_states": ( getattr(block_state, "prompt_embeds", None), @@ -290,7 +298,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: guider_inputs = { "encoder_hidden_states": ( getattr(block_state, "prompt_embeds", None), @@ -366,7 +376,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, @@ -419,7 +431,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: QwenImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[QwenImageModularPipeline, BlockState]: block_state.init_latents_proper = block_state.image_latents if i < len(block_state.timesteps) - 1: block_state.noise_timestep = block_state.timesteps[i + 1] @@ -466,7 +480,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/qwenimage/encoders.py b/src/diffusers/modular_pipelines/qwenimage/encoders.py index 5dade5716a49..1c414bb1f3c5 100644 --- a/src/diffusers/modular_pipelines/qwenimage/encoders.py +++ b/src/diffusers/modular_pipelines/qwenimage/encoders.py @@ -325,7 +325,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) images = block_state.image @@ -414,7 +416,9 @@ def check_inputs(resolution: int): raise ValueError(f"Resolution must be 1024 or 640 but is {resolution}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(resolution=block_state.resolution) @@ -505,7 +509,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) images = block_state.image @@ -619,7 +625,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -743,7 +751,9 @@ def check_inputs(prompt, negative_prompt, max_sequence_length): raise ValueError(f"`max_sequence_length` cannot be greater than 1024 but is {max_sequence_length}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -874,7 +884,9 @@ def check_inputs(prompt, negative_prompt): raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.prompt, block_state.negative_prompt) @@ -1002,7 +1014,9 @@ def check_inputs(prompt, negative_prompt): raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.prompt, block_state.negative_prompt) @@ -1132,7 +1146,9 @@ def check_inputs(height, width, vae_scale_factor): raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -1228,7 +1244,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) width, height = block_state.resized_image[0].size @@ -1312,7 +1330,9 @@ def check_inputs(height, width, vae_scale_factor): raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -1387,7 +1407,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) width, height = block_state.resized_image[0].size @@ -1459,7 +1481,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState): + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) image = block_state.resized_image @@ -1563,7 +1587,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return [self._output] # default is "image_latents" @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -1670,7 +1696,9 @@ def check_inputs(height, width, vae_scale_factor): raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}") @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state.height, block_state.width, components.vae_scale_factor) @@ -1769,7 +1797,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Permute: (B, C, 1, H, W) -> (B, 1, C, H, W) diff --git a/src/diffusers/modular_pipelines/qwenimage/inputs.py b/src/diffusers/modular_pipelines/qwenimage/inputs.py index 38a49e07345f..2d35d833c60f 100644 --- a/src/diffusers/modular_pipelines/qwenimage/inputs.py +++ b/src/diffusers/modular_pipelines/qwenimage/inputs.py @@ -205,7 +205,9 @@ def check_inputs( ): raise ValueError("`negative_prompt_embeds_mask` must have the same batch size as `prompt_embeds`") - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs( @@ -411,7 +413,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return outputs - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs @@ -626,7 +630,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return outputs - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs @@ -852,7 +858,9 @@ def intermediate_outputs(self) -> list[OutputParam]: return outputs - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs @@ -969,7 +977,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: QwenImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: QwenImageModularPipeline, state: PipelineState + ) -> tuple[QwenImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) if isinstance(components.controlnet, QwenImageMultiControlNetModel): diff --git a/src/diffusers/modular_pipelines/stable_diffusion_3/before_denoise.py b/src/diffusers/modular_pipelines/stable_diffusion_3/before_denoise.py index 5007faa12f67..09f5a5f4f54c 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_3/before_denoise.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_3/before_denoise.py @@ -190,7 +190,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -285,7 +287,9 @@ def get_timesteps(scheduler, num_inference_steps, strength): return timesteps, num_inference_steps - t_start @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -372,7 +376,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device batch_size = block_state.batch_size * block_state.num_images_per_prompt @@ -446,7 +452,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) latent_timestep = block_state.timesteps[:1].repeat(block_state.latents.shape[0]) block_state.initial_noise = block_state.latents diff --git a/src/diffusers/modular_pipelines/stable_diffusion_3/decoders.py b/src/diffusers/modular_pipelines/stable_diffusion_3/decoders.py index b1a8df1c7fa7..f2665ed06e47 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_3/decoders.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_3/decoders.py @@ -21,6 +21,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import StableDiffusion3ModularPipeline logger = logging.get_logger(__name__) @@ -62,7 +63,7 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("images", type_hint=list[PIL.Image.Image] | torch.Tensor)] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae = components.vae diff --git a/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py b/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py index 33bd98095d8a..cde6ace66245 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py @@ -102,7 +102,7 @@ def __call__( block_state: BlockState, i: int, t: torch.Tensor, - ) -> PipelineState: + ) -> tuple[StableDiffusion3ModularPipeline, BlockState]: do_cfg = block_state.negative_prompt_embeds is not None guider_inputs = { @@ -174,7 +174,7 @@ def __call__( block_state: BlockState, i: int, t: torch.Tensor, - ): + ) -> tuple[StableDiffusion3ModularPipeline, BlockState]: latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, @@ -207,7 +207,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order, diff --git a/src/diffusers/modular_pipelines/stable_diffusion_3/encoders.py b/src/diffusers/modular_pipelines/stable_diffusion_3/encoders.py index bef2a0f812ec..ea163b3b9cf9 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_3/encoders.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_3/encoders.py @@ -364,7 +364,9 @@ def check_inputs(height, width, vae_scale_factor, patch_size): raise ValueError(f"Width must be divisible by {vae_scale_factor * patch_size} but is {width}") @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState): + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.image is None: @@ -432,7 +434,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) image = getattr(block_state, self._image_input_name) @@ -526,7 +530,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device diff --git a/src/diffusers/modular_pipelines/stable_diffusion_3/inputs.py b/src/diffusers/modular_pipelines/stable_diffusion_3/inputs.py index d7e88b571612..8dd997b07b57 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_3/inputs.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_3/inputs.py @@ -187,7 +187,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.batch_size = block_state.prompt_embeds.shape[0] @@ -282,7 +284,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ), ] - def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusion3ModularPipeline, state: PipelineState + ) -> tuple[StableDiffusion3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) for input_name in self._image_latent_inputs: diff --git a/src/diffusers/modular_pipelines/stable_diffusion_xl/before_denoise.py b/src/diffusers/modular_pipelines/stable_diffusion_xl/before_denoise.py index 92c74219bd06..f88256fbd9f4 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_xl/before_denoise.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_xl/before_denoise.py @@ -334,7 +334,9 @@ def check_inputs(self, components, block_state): ) @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -482,7 +484,9 @@ def get_timesteps(components, num_inference_steps, strength, device, denoising_s return timesteps, num_inference_steps @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -567,7 +571,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -823,7 +829,9 @@ def prepare_mask_latents( return mask, masked_image_latents @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.dtype = block_state.dtype if block_state.dtype is not None else components.vae.dtype @@ -926,7 +934,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.dtype = block_state.dtype if block_state.dtype is not None else components.vae.dtype @@ -1027,7 +1037,9 @@ def prepare_latents(comp, batch_size, num_channels_latents, height, width, dtype return latents @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.dtype is None: @@ -1217,7 +1229,9 @@ def get_guidance_scale_embedding( return emb @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -1395,7 +1409,9 @@ def get_guidance_scale_embedding( return emb @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.device = components._execution_device @@ -1560,7 +1576,9 @@ def prepare_control_image( return image @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) # (1) prepare controlnet inputs @@ -1789,7 +1807,9 @@ def prepare_control_image( return image @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) controlnet = unwrap_module(components.controlnet) diff --git a/src/diffusers/modular_pipelines/stable_diffusion_xl/decoders.py b/src/diffusers/modular_pipelines/stable_diffusion_xl/decoders.py index b4f15df8b411..ea6fdfa1df98 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_xl/decoders.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_xl/decoders.py @@ -27,6 +27,7 @@ PipelineState, ) from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import StableDiffusionXLModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -84,7 +85,7 @@ def upcast_vae(components): components.vae.to(dtype=torch.float32) @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) if not block_state.output_type == "latent": @@ -182,7 +183,7 @@ def inputs(self) -> list[tuple[str, Any]]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) if block_state.padding_mask_crop is not None and block_state.crops_coords is not None: diff --git a/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py b/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py index ec344fe0ad37..16a8b236ce2e 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py @@ -66,7 +66,9 @@ def inputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: block_state.scaled_latents = components.scheduler.scale_model_input(block_state.latents, t) return components, block_state @@ -131,7 +133,9 @@ def check_inputs(components, block_state): ) @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: self.check_inputs(components, block_state) block_state.scaled_latents = components.scheduler.scale_model_input(block_state.latents, t) @@ -198,7 +202,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int - ) -> PipelineState: + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: # Map the keys we'll see on each `guider_state_batch` (e.g. guider_state_batch.prompt_embeds) # to the corresponding (cond, uncond) fields on block_state. (e.g. block_state.prompt_embeds, block_state.negative_prompt_embeds) guider_inputs = { @@ -351,7 +355,9 @@ def prepare_extra_kwargs(func, exclude_kwargs=[], **kwargs): return extra_kwargs @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: extra_controlnet_kwargs = self.prepare_extra_kwargs( components.controlnet.forward, **block_state.controlnet_kwargs ) @@ -508,7 +514,9 @@ def prepare_extra_kwargs(func, exclude_kwargs=[], **kwargs): return extra_kwargs @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: # Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline block_state.extra_step_kwargs = self.prepare_extra_kwargs( components.scheduler.step, generator=block_state.generator, eta=block_state.eta @@ -603,7 +611,9 @@ def check_inputs(self, components, block_state): raise ValueError(f"noise is required for this step {self.__class__.__name__}") @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int): + def __call__( + self, components: StableDiffusionXLModularPipeline, block_state: BlockState, i: int, t: int + ) -> tuple[StableDiffusionXLModularPipeline, BlockState]: self.check_inputs(components, block_state) # Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline @@ -684,7 +694,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.disable_guidance = True if components.unet.config.time_cond_proj_dim is not None else False diff --git a/src/diffusers/modular_pipelines/stable_diffusion_xl/encoders.py b/src/diffusers/modular_pipelines/stable_diffusion_xl/encoders.py index 26e5524309f1..381badb2204d 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_xl/encoders.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_xl/encoders.py @@ -188,7 +188,9 @@ def prepare_ip_adapter_image_embeds( return ip_adapter_image_embeds @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.prepare_unconditional_embeds = components.guider.num_conditions > 1 @@ -534,7 +536,9 @@ def encode_prompt( return prompt_embeds, negative_prompt_embeds, pooled_prompt_embeds, negative_pooled_prompt_embeds @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: # Get inputs and intermediates block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -655,7 +659,9 @@ def _encode_vae_image(self, components, image: torch.Tensor, generator: torch.Ge return image_latents @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.preprocess_kwargs = block_state.preprocess_kwargs or {} block_state.device = components._execution_device @@ -825,7 +831,9 @@ def prepare_mask_latents( return mask, masked_image_latents @torch.no_grad() - def __call__(self, components: StableDiffusionXLModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: StableDiffusionXLModularPipeline, state: PipelineState + ) -> tuple[StableDiffusionXLModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.dtype = block_state.dtype if block_state.dtype is not None else components.vae.dtype diff --git a/src/diffusers/modular_pipelines/wan/before_denoise.py b/src/diffusers/modular_pipelines/wan/before_denoise.py index 1d90c20d8124..0e7aec5364fb 100644 --- a/src/diffusers/modular_pipelines/wan/before_denoise.py +++ b/src/diffusers/modular_pipelines/wan/before_denoise.py @@ -245,7 +245,9 @@ def check_inputs(self, components, block_state): ) @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -355,7 +357,9 @@ def inputs(self) -> list[InputParam]: return inputs - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs (height/width calculation, patchify, and batch expansion) @@ -436,7 +440,9 @@ def check_inputs(block_state): "Generating multiple videos per prompt is not yet supported. This may be supported in the future." ) - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -469,7 +475,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -568,7 +576,9 @@ def prepare_latents( return latents @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) diff --git a/src/diffusers/modular_pipelines/wan/decoders.py b/src/diffusers/modular_pipelines/wan/decoders.py index 529c9291c250..f820d86339ef 100644 --- a/src/diffusers/modular_pipelines/wan/decoders.py +++ b/src/diffusers/modular_pipelines/wan/decoders.py @@ -24,6 +24,7 @@ from ...video_processor import VideoProcessor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import WanModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -54,7 +55,7 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.latents = block_state.latents[:, :, block_state.num_reference_images :] @@ -107,7 +108,7 @@ def intermediate_outputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae_dtype = components.vae.dtype diff --git a/src/diffusers/modular_pipelines/wan/denoise.py b/src/diffusers/modular_pipelines/wan/denoise.py index 4beb0425a5fc..3036b33868d1 100644 --- a/src/diffusers/modular_pipelines/wan/denoise.py +++ b/src/diffusers/modular_pipelines/wan/denoise.py @@ -64,7 +64,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[WanModularPipeline, BlockState]: block_state.latent_model_input = block_state.latents.to(block_state.dtype) return components, block_state @@ -104,7 +106,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[WanModularPipeline, BlockState]: block_state.latent_model_input = torch.cat( [block_state.latents, block_state.image_condition_latents], dim=1 ).to(block_state.dtype) @@ -181,7 +185,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[WanModularPipeline, BlockState]: components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) # The guider splits model inputs into separate batches for conditional/unconditional predictions. @@ -315,7 +319,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[WanModularPipeline, BlockState]: boundary_timestep = components.config.boundary_ratio * components.num_train_timesteps if t >= boundary_timestep: block_state.current_model = components.transformer @@ -391,7 +395,9 @@ def description(self) -> str: ) @torch.no_grad() - def __call__(self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: WanModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[WanModularPipeline, BlockState]: # Perform scheduler step using the predicted output latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( @@ -441,7 +447,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/wan/encoders.py b/src/diffusers/modular_pipelines/wan/encoders.py index 3bebe341a511..6eaf42ce4f6a 100644 --- a/src/diffusers/modular_pipelines/wan/encoders.py +++ b/src/diffusers/modular_pipelines/wan/encoders.py @@ -273,7 +273,9 @@ def encode_prompt( return prompt_embeds, negative_prompt_embeds @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: # Get inputs and intermediates block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -319,7 +321,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("resized_image", type_hint=PIL.Image.Image), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) max_area = block_state.height * block_state.width @@ -356,7 +360,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("resized_last_image", type_hint=PIL.Image.Image), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) height = block_state.resized_image.height @@ -403,7 +409,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("image_embeds", type_hint=torch.Tensor, description="The image embeddings"), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -448,7 +456,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("image_embeds", type_hint=torch.Tensor, description="The image embeddings"), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -521,7 +531,9 @@ def check_inputs(components, block_state): f"`num_frames` has to be greater than 0, and (num_frames - 1) must be divisible by {components.vae_scale_factor_temporal}, but got {block_state.num_frames}." ) - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -822,7 +834,9 @@ def prepare_masks(components, mask, reference_images): return torch.stack(mask_list) @torch.no_grad() - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -897,7 +911,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("image_condition_latents", type_hint=torch.Tensor | None), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) batch_size, _, _, latent_height, latent_width = block_state.first_frame_latents.shape @@ -976,7 +992,9 @@ def check_inputs(components, block_state): f"`num_frames` has to be greater than 0, and (num_frames - 1) must be divisible by {components.vae_scale_factor_temporal}, but got {block_state.num_frames}." ) - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -1046,7 +1064,9 @@ def intermediate_outputs(self) -> list[OutputParam]: OutputParam("image_condition_latents", type_hint=torch.Tensor | None), ] - def __call__(self, components: WanModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: WanModularPipeline, state: PipelineState + ) -> tuple[WanModularPipeline, PipelineState]: block_state = self.get_block_state(state) batch_size, _, _, latent_height, latent_width = block_state.first_last_frame_latents.shape diff --git a/src/diffusers/modular_pipelines/wan_animate_2/before_denoise.py b/src/diffusers/modular_pipelines/wan_animate_2/before_denoise.py index 0ad038e8ccaa..7cf13afea769 100644 --- a/src/diffusers/modular_pipelines/wan_animate_2/before_denoise.py +++ b/src/diffusers/modular_pipelines/wan_animate_2/before_denoise.py @@ -20,6 +20,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import InputParam, OutputParam +from .modular_pipeline import WanAnimate2ModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -82,7 +83,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) latent_height, latent_width = block_state.reference_image_latents.shape[-2:] diff --git a/src/diffusers/modular_pipelines/wan_animate_2/decoders.py b/src/diffusers/modular_pipelines/wan_animate_2/decoders.py index 5317ee85e67a..cc9e6fb4f6e2 100644 --- a/src/diffusers/modular_pipelines/wan_animate_2/decoders.py +++ b/src/diffusers/modular_pipelines/wan_animate_2/decoders.py @@ -20,6 +20,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import WanAnimate2ModularPipeline from .video_processor import WanAnimate2VideoProcessor @@ -87,7 +88,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) video = torch.cat(block_state.segment_frames, dim=2)[:, :, : block_state.real_frame_len] diff --git a/src/diffusers/modular_pipelines/wan_animate_2/denoise.py b/src/diffusers/modular_pipelines/wan_animate_2/denoise.py index d96b8f814239..1e18b3a09eeb 100644 --- a/src/diffusers/modular_pipelines/wan_animate_2/denoise.py +++ b/src/diffusers/modular_pipelines/wan_animate_2/denoise.py @@ -28,6 +28,7 @@ from ..modular_pipeline import BlockState, LoopSequentialPipelineBlocks, ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam from .encoders import encode_vae, get_i2v_mask +from .modular_pipeline import WanAnimate2ModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -115,7 +116,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: device = components._execution_device latent_height, latent_width = block_state.reference_image_latents.shape[-2:] @@ -190,7 +191,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: # `block_state.out_frames` is seeded by the loop wrapper and written by the decode step of the # previous iteration. device = components._execution_device @@ -270,7 +271,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: device = components._execution_device block_state.latents = randn_tensor( @@ -319,7 +320,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: device = components._execution_device components.scheduler.set_timesteps(block_state.num_inference_steps, device=device) @@ -400,7 +401,7 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: device = components._execution_device transformer_dtype = components.transformer.dtype @@ -527,7 +528,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: transformer_dtype = components.transformer.dtype guider_inputs = { @@ -673,7 +674,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, block_state: BlockState, k: int): + def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnimate2ModularPipeline, BlockState]: latents = block_state.latents.to(torch.float32) # The first latent frame is the reference image's slot, not video content. out_frames = decode_vae(components.vae, latents[:, 1:]) @@ -730,7 +731,7 @@ def loop_intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Seed the loop-carried state: `segment_frames` collects each segment's decoded frames (the decode step diff --git a/src/diffusers/modular_pipelines/wan_animate_2/encoders.py b/src/diffusers/modular_pipelines/wan_animate_2/encoders.py index 21b70f636f7d..ffba4c672691 100644 --- a/src/diffusers/modular_pipelines/wan_animate_2/encoders.py +++ b/src/diffusers/modular_pipelines/wan_animate_2/encoders.py @@ -25,6 +25,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import WanAnimate2ModularPipeline from .video_processor import WanAnimate2VideoProcessor @@ -169,7 +170,7 @@ def check_inputs(block_state): raise ValueError(f"`prompt` has to be of type `str` but is {type(block_state.prompt)}") @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -269,7 +270,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -387,7 +388,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -473,7 +474,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -523,7 +524,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -578,7 +579,7 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2ModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device diff --git a/src/diffusers/modular_pipelines/z_image/before_denoise.py b/src/diffusers/modular_pipelines/z_image/before_denoise.py index 5216529d460f..e96d80b4bfca 100644 --- a/src/diffusers/modular_pipelines/z_image/before_denoise.py +++ b/src/diffusers/modular_pipelines/z_image/before_denoise.py @@ -258,7 +258,9 @@ def check_inputs(self, components, block_state): ) @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -366,7 +368,9 @@ def inputs(self) -> list[InputParam]: return inputs - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) # Process image latent inputs (height/width calculation, patchify, and batch expansion) @@ -467,7 +471,9 @@ def prepare_latents( return latents @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -524,7 +530,9 @@ def intermediate_outputs(self) -> list[OutputParam]: ] @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) device = components._execution_device @@ -580,7 +588,9 @@ def check_inputs(self, components, block_state): raise ValueError(f"Strength must be between 0.0 and 1.0, but got {block_state.strength}") @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) @@ -613,7 +623,9 @@ def inputs(self) -> list[InputParam]: InputParam("timesteps", required=True), ] - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) latent_timestep = block_state.timesteps[:1].repeat(block_state.latents.shape[0]) diff --git a/src/diffusers/modular_pipelines/z_image/decoders.py b/src/diffusers/modular_pipelines/z_image/decoders.py index 353253102376..21e307d53df4 100644 --- a/src/diffusers/modular_pipelines/z_image/decoders.py +++ b/src/diffusers/modular_pipelines/z_image/decoders.py @@ -24,6 +24,7 @@ from ...utils import logging from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import ZImageModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -74,7 +75,7 @@ def intermediate_outputs(self) -> list[str]: ] @torch.no_grad() - def __call__(self, components, state: PipelineState) -> PipelineState: + def __call__(self, components, state: PipelineState) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) vae_dtype = components.vae.dtype diff --git a/src/diffusers/modular_pipelines/z_image/denoise.py b/src/diffusers/modular_pipelines/z_image/denoise.py index 863df312389a..899800a5019a 100644 --- a/src/diffusers/modular_pipelines/z_image/denoise.py +++ b/src/diffusers/modular_pipelines/z_image/denoise.py @@ -63,7 +63,9 @@ def inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: ZImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[ZImageModularPipeline, BlockState]: latents = block_state.latents.unsqueeze(2).to( block_state.dtype ) # [batch_size, num_channels, 1, height, width] @@ -152,7 +154,7 @@ def inputs(self) -> list[tuple[str, Any]]: @torch.no_grad() def __call__( self, components: ZImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor - ) -> PipelineState: + ) -> tuple[ZImageModularPipeline, BlockState]: components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) # The guider splits model inputs into separate batches for conditional/unconditional predictions. @@ -219,7 +221,9 @@ def description(self) -> str: ) @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): + def __call__( + self, components: ZImageModularPipeline, block_state: BlockState, i: int, t: torch.Tensor + ) -> tuple[ZImageModularPipeline, BlockState]: # Perform scheduler step using the predicted output latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( @@ -269,7 +273,9 @@ def loop_inputs(self) -> list[InputParam]: ] @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( diff --git a/src/diffusers/modular_pipelines/z_image/encoders.py b/src/diffusers/modular_pipelines/z_image/encoders.py index 06deb8236893..c3ab717078e1 100644 --- a/src/diffusers/modular_pipelines/z_image/encoders.py +++ b/src/diffusers/modular_pipelines/z_image/encoders.py @@ -244,7 +244,9 @@ def encode_prompt( return prompt_embeds, negative_prompt_embeds @torch.no_grad() - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: # Get inputs and intermediates block_state = self.get_block_state(state) self.check_inputs(block_state) @@ -316,7 +318,9 @@ def check_inputs(components, block_state): f"`height` and `width` have to be divisible by {components.vae_scale_factor_spatial} but are {block_state.height} and {block_state.width}." ) - def __call__(self, components: ZImageModularPipeline, state: PipelineState) -> PipelineState: + def __call__( + self, components: ZImageModularPipeline, state: PipelineState + ) -> tuple[ZImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) self.check_inputs(components, block_state) diff --git a/src/diffusers/pipelines/ace_step/pipeline_ace_step.py b/src/diffusers/pipelines/ace_step/pipeline_ace_step.py index b11e08208eb1..8f8e2e1be95b 100644 --- a/src/diffusers/pipelines/ace_step/pipeline_ace_step.py +++ b/src/diffusers/pipelines/ace_step/pipeline_ace_step.py @@ -831,7 +831,7 @@ def __call__( cfg_interval_end: float = 1.0, timesteps: Optional[List[float]] = None, attention_kwargs: Optional[dict] = None, - ): + ) -> AudioPipelineOutput | tuple: r""" The call function to the pipeline for music generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff.py index b9e5b40b65ff..fe97aedf627a 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff.py @@ -595,7 +595,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], decode_chunk_size: int = 16, **kwargs, - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff_controlnet.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff_controlnet.py index a9630cc3c00f..4cc3136a21b9 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff_controlnet.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff_controlnet.py @@ -748,7 +748,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], decode_chunk_size: int = 16, - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff_sdxl.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff_sdxl.py index 70c6a5dc5cd6..4bb6bee294c1 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff_sdxl.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff_sdxl.py @@ -905,7 +905,7 @@ def __call__( clip_skip: int | None = None, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff_sparsectrl.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff_sparsectrl.py index fcf260b47f3a..6be516c0938f 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff_sparsectrl.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff_sparsectrl.py @@ -738,7 +738,7 @@ def __call__( clip_skip: int | None = None, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video.py index b8aa82ab9d2f..e3f6966a159a 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video.py @@ -770,7 +770,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], decode_chunk_size: int = 16, - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video_controlnet.py b/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video_controlnet.py index 7c649b501f32..35821b09cd52 100644 --- a/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video_controlnet.py +++ b/src/diffusers/pipelines/animatediff/pipeline_animatediff_video2video_controlnet.py @@ -940,7 +940,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], decode_chunk_size: int = 16, - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/anyflow/pipeline_anyflow.py b/src/diffusers/pipelines/anyflow/pipeline_anyflow.py index 61240b40e6b5..152183bbb311 100644 --- a/src/diffusers/pipelines/anyflow/pipeline_anyflow.py +++ b/src/diffusers/pipelines/anyflow/pipeline_anyflow.py @@ -406,7 +406,7 @@ def __call__( callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 512, use_mean_velocity: bool = True, - ): + ) -> AnyFlowPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/anyflow/pipeline_anyflow_far.py b/src/diffusers/pipelines/anyflow/pipeline_anyflow_far.py index 63ac8fa4d6bd..829bc4a236d8 100644 --- a/src/diffusers/pipelines/anyflow/pipeline_anyflow_far.py +++ b/src/diffusers/pipelines/anyflow/pipeline_anyflow_far.py @@ -476,7 +476,7 @@ def __call__( use_mean_velocity: bool = True, use_kv_cache: bool = True, chunk_partition: Optional[List[int]] = None, - ): + ) -> AnyFlowPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/audioldm2/pipeline_audioldm2.py b/src/diffusers/pipelines/audioldm2/pipeline_audioldm2.py index 110ddd1bfef3..b9e64c18c468 100644 --- a/src/diffusers/pipelines/audioldm2/pipeline_audioldm2.py +++ b/src/diffusers/pipelines/audioldm2/pipeline_audioldm2.py @@ -864,7 +864,7 @@ def __call__( callback_steps: int | None = 1, cross_attention_kwargs: dict[str, Any] | None = None, output_type: str | None = "np", - ): + ) -> AudioPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/bria/pipeline_bria.py b/src/diffusers/pipelines/bria/pipeline_bria.py index 9b80278af21e..a3f262ce935e 100644 --- a/src/diffusers/pipelines/bria/pipeline_bria.py +++ b/src/diffusers/pipelines/bria/pipeline_bria.py @@ -469,7 +469,7 @@ def __call__( max_sequence_length: int = 128, clip_value: None | float = None, normalize: bool = False, - ): + ) -> BriaPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo.py b/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo.py index 398758294bcc..96f713851cf4 100644 --- a/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo.py +++ b/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo.py @@ -453,7 +453,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 3000, do_patching=False, - ): + ) -> BriaFiboPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo_edit.py b/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo_edit.py index 28858322fb96..3b7ae9b1c206 100644 --- a/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo_edit.py +++ b/src/diffusers/pipelines/bria_fibo/pipeline_bria_fibo_edit.py @@ -621,7 +621,7 @@ def __call__( max_sequence_length: int = 3000, do_patching=False, _auto_resize: bool = True, - ): + ) -> BriaFiboPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/chroma/pipeline_chroma.py b/src/diffusers/pipelines/chroma/pipeline_chroma.py index 2f60d21091dd..37562f4b6a52 100644 --- a/src/diffusers/pipelines/chroma/pipeline_chroma.py +++ b/src/diffusers/pipelines/chroma/pipeline_chroma.py @@ -612,7 +612,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ChromaPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py b/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py index f1a3cc6c7b24..da6615e71c7f 100644 --- a/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py +++ b/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py @@ -673,7 +673,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ChromaPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py b/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py index 4900ef14e3f1..fe95a103cd4f 100644 --- a/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py +++ b/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py @@ -795,7 +795,7 @@ def __call__( max_sequence_length: int = 256, prompt_attention_mask: torch.Tensor | None = None, negative_prompt_attention_mask: torch.Tensor | None = None, - ): + ) -> ChromaPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py b/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py index 3fefa951c9b8..95a3ad3c20e3 100644 --- a/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py +++ b/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py @@ -493,7 +493,7 @@ def __call__( max_sequence_length: int = 512, enable_temporal_reasoning: bool = False, num_temporal_reasoning_steps: int = 0, - ): + ) -> ChronoEditPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/consistency_models/pipeline_consistency_models.py b/src/diffusers/pipelines/consistency_models/pipeline_consistency_models.py index b2b18b52e824..d8fb22f1c262 100644 --- a/src/diffusers/pipelines/consistency_models/pipeline_consistency_models.py +++ b/src/diffusers/pipelines/consistency_models/pipeline_consistency_models.py @@ -182,7 +182,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, - ): + ) -> ImagePipelineOutput | tuple: r""" Args: batch_size (`int`, *optional*, defaults to 1): diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet.py index 7ef287ae8154..bc95c1efb940 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet.py @@ -936,7 +936,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_blip_diffusion.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_blip_diffusion.py index 3b91229814cc..1e0dfd3fb94f 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_blip_diffusion.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_blip_diffusion.py @@ -255,7 +255,7 @@ def __call__( prompt_reps: int = 20, output_type: str | None = "pil", return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_img2img.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_img2img.py index b40dda940959..15aecfa4d4b4 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_img2img.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_img2img.py @@ -934,7 +934,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint.py index c981dd14abd6..7f21b964417d 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint.py @@ -1025,7 +1025,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint_sd_xl.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint_sd_xl.py index 86905fb39b29..0503974c101e 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint_sd_xl.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint_sd_xl.py @@ -1210,7 +1210,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl.py index f0174cb6ce98..3c75d15ea6e9 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl.py @@ -1039,7 +1039,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl_img2img.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl_img2img.py index 92366e465673..ba24ed70db0e 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl_img2img.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_sd_xl_img2img.py @@ -1119,7 +1119,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_inpaint_sd_xl.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_inpaint_sd_xl.py index 3dd12ec72989..7e7aeddcbbac 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_inpaint_sd_xl.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_inpaint_sd_xl.py @@ -1190,7 +1190,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl.py index 4802279b81f4..2177f30b3e72 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl.py @@ -1016,7 +1016,7 @@ def __call__( clip_skip: int | None = None, callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl_img2img.py b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl_img2img.py index 62c9c06f46a4..0cbc923fce10 100644 --- a/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl_img2img.py +++ b/src/diffusers/pipelines/controlnet/pipeline_controlnet_union_sd_xl_img2img.py @@ -1110,7 +1110,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet_hunyuandit/pipeline_hunyuandit_controlnet.py b/src/diffusers/pipelines/controlnet_hunyuandit/pipeline_hunyuandit_controlnet.py index ba241bf4feb6..bea400c19f86 100644 --- a/src/diffusers/pipelines/controlnet_hunyuandit/pipeline_hunyuandit_controlnet.py +++ b/src/diffusers/pipelines/controlnet_hunyuandit/pipeline_hunyuandit_controlnet.py @@ -662,7 +662,7 @@ def __call__( target_size: tuple[int, int] | None = None, crops_coords_top_left: tuple[int, int] = (0, 0), use_resolution_binning: bool = True, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation with HunyuanDiT. diff --git a/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet.py b/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet.py index 4530a424adb4..5a7ee29a1c23 100644 --- a/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet.py +++ b/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet.py @@ -852,7 +852,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet_inpainting.py b/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet_inpainting.py index d2890d55811c..60b5828cb8ac 100644 --- a/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet_inpainting.py +++ b/src/diffusers/pipelines/controlnet_sd3/pipeline_stable_diffusion_3_controlnet_inpainting.py @@ -1020,7 +1020,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_predict.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_predict.py index c2c5e6d2c824..74c708f25444 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_predict.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_predict.py @@ -566,7 +566,7 @@ def __call__( max_sequence_length: int = 512, conditional_frame_timestep: float = 0.0001, num_latent_conditional_frames: int = 2, - ): + ) -> CosmosPipelineOutput | tuple: r""" The call function to the pipeline for generation. Supports three modes: diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_transfer.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_transfer.py index e38d926bbd28..d8d084946a98 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_transfer.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_5_transfer.py @@ -595,7 +595,7 @@ def __call__( conditional_frame_timestep: float = 0.1, num_ar_conditional_frames: Optional[int] = 1, num_ar_latent_conditional_frames: Optional[int] = None, - ): + ) -> CosmosPipelineOutput | tuple: r""" `controls` drive the conditioning through ControlNet. Controls are assumed to be pre-processed, e.g. edge maps are pre-computed. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_text2image.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_text2image.py index 8c6de18b3a9a..e29fde090c22 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_text2image.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_text2image.py @@ -434,7 +434,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> CosmosImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_video2world.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_video2world.py index 2a708e1118e0..01b8d7b73680 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos2_video2world.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos2_video2world.py @@ -507,7 +507,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, sigma_conditioning: float = 0.0001, - ): + ) -> CosmosPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py index 02dc70b29cfc..edd8f44409c2 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py @@ -1350,7 +1350,7 @@ def __call__( mixed_precision_first_steps: int | None = None, mixed_precision_last_steps: int | None = None, mixed_precision_reasoner_policy: str | None = None, - ) -> Cosmos3OmniPipelineOutput: + ) -> Cosmos3OmniPipelineOutput | tuple: r""" Run the Cosmos 3 omni pipeline end-to-end: encode the (optional) conditioning image/video, denoise vision and (optional) sound latents jointly, and decode them back into a video and audio waveform. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos_text2world.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos_text2world.py index 61d9ec8f0574..55eec93d67a9 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos_text2world.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos_text2world.py @@ -420,7 +420,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> CosmosPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos_video2world.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos_video2world.py index bf7e28584967..7298218a4770 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos_video2world.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos_video2world.py @@ -536,7 +536,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> CosmosPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if.py index b8c70fc6528c..aaf177b31edc 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if.py @@ -566,7 +566,7 @@ def __call__( callback_steps: int = 1, clean_caption: bool = True, cross_attention_kwargs: dict[str, Any] | None = None, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img.py index 3dadc63f4952..6bc077aef03d 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img.py @@ -685,7 +685,7 @@ def __call__( callback_steps: int = 1, clean_caption: bool = True, cross_attention_kwargs: dict[str, Any] | None = None, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img_superresolution.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img_superresolution.py index 4839a0860462..ed28d7b59fa4 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img_superresolution.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_img2img_superresolution.py @@ -770,7 +770,7 @@ def __call__( cross_attention_kwargs: dict[str, Any] | None = None, noise_level: int = 250, clean_caption: bool = True, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting.py index 03a9d6f7c5e8..1ed581f56a0e 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting.py @@ -783,7 +783,7 @@ def __call__( callback_steps: int = 1, clean_caption: bool = True, cross_attention_kwargs: dict[str, Any] | None = None, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting_superresolution.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting_superresolution.py index 841382ad9c63..b7a1736c61a6 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting_superresolution.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_inpainting_superresolution.py @@ -864,7 +864,7 @@ def __call__( cross_attention_kwargs: dict[str, Any] | None = None, noise_level: int = 0, clean_caption: bool = True, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_superresolution.py b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_superresolution.py index 52ebebb6f9b4..45f1ee6b1610 100644 --- a/src/diffusers/pipelines/deepfloyd_if/pipeline_if_superresolution.py +++ b/src/diffusers/pipelines/deepfloyd_if/pipeline_if_superresolution.py @@ -635,7 +635,7 @@ def __call__( cross_attention_kwargs: dict[str, Any] | None = None, noise_level: int = 250, clean_caption: bool = True, - ): + ) -> IFPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py b/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py index 5d608d7c49fb..207142dccb67 100644 --- a/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py +++ b/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py @@ -184,7 +184,7 @@ def __call__( | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] | None = None, - ) -> DiffusionGemmaPipelineOutput | tuple[torch.LongTensor, list[str] | None]: + ) -> DiffusionGemmaPipelineOutput | tuple: """ Generate text with block diffusion. diff --git a/src/diffusers/pipelines/dreamlite/pipeline_dreamlite.py b/src/diffusers/pipelines/dreamlite/pipeline_dreamlite.py index e9a0e3c2a767..85b303b32dec 100644 --- a/src/diffusers/pipelines/dreamlite/pipeline_dreamlite.py +++ b/src/diffusers/pipelines/dreamlite/pipeline_dreamlite.py @@ -403,7 +403,7 @@ def __call__( return_dict: bool = True, max_sequence_length: int = 200, text_pad_embedding: Optional[torch.Tensor] = None, - ): + ) -> DreamLitePipelineOutput | tuple: r"""Run the DreamLite pipeline. Args: diff --git a/src/diffusers/pipelines/dreamlite/pipeline_dreamlite_mobile.py b/src/diffusers/pipelines/dreamlite/pipeline_dreamlite_mobile.py index ca9e6b7b4c40..339b58c9cbf9 100644 --- a/src/diffusers/pipelines/dreamlite/pipeline_dreamlite_mobile.py +++ b/src/diffusers/pipelines/dreamlite/pipeline_dreamlite_mobile.py @@ -398,7 +398,7 @@ def __call__( return_dict: bool = True, max_sequence_length: int = 200, text_pad_embedding: Optional[torch.Tensor] = None, - ): + ) -> DreamLitePipelineOutput | tuple: r"""Run the distilled DreamLite Mobile pipeline. Args: diff --git a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate.py b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate.py index 72e19a8cce1f..428a20039778 100755 --- a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate.py +++ b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate.py @@ -546,7 +546,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], guidance_rescale: float = 0.0, - ): + ) -> EasyAnimatePipelineOutput | tuple: r""" Generates images or video using the EasyAnimate pipeline based on the provided prompts. diff --git a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_control.py b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_control.py index 4ad3a48b70ec..9f773e18b5c3 100755 --- a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_control.py +++ b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_control.py @@ -695,7 +695,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], guidance_rescale: float = 0.0, timesteps: list[int] | None = None, - ): + ) -> EasyAnimatePipelineOutput | tuple: r""" Generates images or video using the EasyAnimate pipeline based on the provided prompts. diff --git a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_inpaint.py b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_inpaint.py index 69bb332944d6..768c4730fcbb 100755 --- a/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_inpaint.py +++ b/src/diffusers/pipelines/easyanimate/pipeline_easyanimate_inpaint.py @@ -815,7 +815,7 @@ def __call__( strength: float = 1.0, noise_aug_strength: float = 0.0563, timesteps: list[int] | None = None, - ): + ) -> EasyAnimatePipelineOutput | tuple: r""" The call function to the pipeline for generation with HunyuanDiT. diff --git a/src/diffusers/pipelines/ernie_image/pipeline_ernie_image.py b/src/diffusers/pipelines/ernie_image/pipeline_ernie_image.py index 11fce6a204bf..a0097cae660f 100644 --- a/src/diffusers/pipelines/ernie_image/pipeline_ernie_image.py +++ b/src/diffusers/pipelines/ernie_image/pipeline_ernie_image.py @@ -222,7 +222,7 @@ def __call__( callback_on_step_end: Optional[Callable[[int, int, dict], None]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], use_pe: bool = True, # 默认使用PE进行改写 - ): + ) -> ErnieImagePipelineOutput | tuple: """ Generate images from text prompts. diff --git a/src/diffusers/pipelines/flux/pipeline_flux.py b/src/diffusers/pipelines/flux/pipeline_flux.py index eb831a7975ba..2e580388874c 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux.py +++ b/src/diffusers/pipelines/flux/pipeline_flux.py @@ -628,7 +628,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_control.py b/src/diffusers/pipelines/flux/pipeline_flux_control.py index d483161b33b2..b50861cb437f 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_control.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_control.py @@ -603,7 +603,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py b/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py index 56bec7a637ce..15d876caf6ca 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py @@ -658,7 +658,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py b/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py index 5a8c7ff2900e..6fa0c1dff9f6 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py @@ -777,7 +777,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py b/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py index 177a1e5e4ef2..ab2cebb1e765 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py @@ -710,7 +710,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_controlnet_image_to_image.py b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_image_to_image.py index 70e4df7acacd..482fca3b59c3 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_controlnet_image_to_image.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_image_to_image.py @@ -663,7 +663,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py index 02a2e93420bf..26a434807694 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py @@ -770,7 +770,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_fill.py b/src/diffusers/pipelines/flux/pipeline_flux_fill.py index 929f7530bb86..55d3464cfb4e 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_fill.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_fill.py @@ -723,7 +723,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_img2img.py b/src/diffusers/pipelines/flux/pipeline_flux_img2img.py index 81bb499ac4ae..3288bba94772 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_img2img.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_img2img.py @@ -708,7 +708,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_inpaint.py b/src/diffusers/pipelines/flux/pipeline_flux_inpaint.py index 466fba8ac7e6..15d9bb9868fd 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_inpaint.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_inpaint.py @@ -810,7 +810,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_kontext.py b/src/diffusers/pipelines/flux/pipeline_flux_kontext.py index e5fc95e5a1c1..3579c3abd8d1 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_kontext.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_kontext.py @@ -726,7 +726,7 @@ def __call__( max_sequence_length: int = 512, max_area: int = 1024**2, _auto_resize: bool = True, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py b/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py index 020f9761e121..05bf8200f588 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py @@ -919,7 +919,7 @@ def __call__( max_sequence_length: int = 512, max_area: int = 1024**2, _auto_resize: bool = True, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux/pipeline_flux_prior_redux.py b/src/diffusers/pipelines/flux/pipeline_flux_prior_redux.py index fb39ad2a1583..e74d6d4e13b6 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_prior_redux.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_prior_redux.py @@ -387,7 +387,7 @@ def __call__( prompt_embeds_scale: float | list[float] | None = 1.0, pooled_prompt_embeds_scale: float | list[float] | None = 1.0, return_dict: bool = True, - ): + ) -> FluxPriorReduxPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux2/pipeline_flux2.py b/src/diffusers/pipelines/flux2/pipeline_flux2.py index aa73c53fb58d..7bb716772e3a 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2.py @@ -766,7 +766,7 @@ def __call__( max_sequence_length: int = 512, text_encoder_out_layers: tuple[int] = (10, 20, 30), caption_upsample_temperature: float = None, - ): + ) -> Flux2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py b/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py index 92005750e551..50fd51e283f6 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py @@ -634,7 +634,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, text_encoder_out_layers: tuple[int] = (9, 18, 27), - ): + ) -> Flux2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py index 0f9051a99b12..3150c1144f67 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py @@ -849,7 +849,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, text_encoder_out_layers: tuple[int, ...] = (9, 18, 27), - ): + ) -> Flux2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for inpainting. diff --git a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py index 82a33a84568d..711246db71d5 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py @@ -628,7 +628,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, text_encoder_out_layers: tuple[int] = (9, 18, 27), - ): + ) -> Flux2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/helios/pipeline_helios.py b/src/diffusers/pipelines/helios/pipeline_helios.py index 90ac654bc77c..c80158ce8625 100644 --- a/src/diffusers/pipelines/helios/pipeline_helios.py +++ b/src/diffusers/pipelines/helios/pipeline_helios.py @@ -483,7 +483,7 @@ def __call__( num_latent_frames_per_chunk: int = 9, keep_first_frame: bool = True, is_skip_first_chunk: bool = False, - ): + ) -> HeliosPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py b/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py index c187e436a857..c32108d613ed 100644 --- a/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py +++ b/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py @@ -552,7 +552,7 @@ def __call__( zero_steps: int | None = 1, # ------------ DMD ------------ is_amplify_first_chunk: bool = False, - ): + ) -> HeliosPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py b/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py index 1bf3ef3699e4..e4cb9f7c12ea 100644 --- a/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py +++ b/src/diffusers/pipelines/hidream_image/pipeline_hidream_image.py @@ -703,7 +703,8 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 128, - ): + **kwargs, + ) -> HiDreamImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py index 50239e9afa22..4faa73d40247 100644 --- a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py +++ b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py @@ -528,7 +528,7 @@ def __call__( attention_kwargs: dict[str, Any] | None = None, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> HunyuanImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py index efdb5505e604..4b926b82c474 100644 --- a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py +++ b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py @@ -457,7 +457,7 @@ def __call__( attention_kwargs: dict[str, Any] | None = None, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> HunyuanImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py index bd54f2563b52..22e8346b8a14 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py @@ -510,7 +510,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], prompt_template: dict[str, Any] = DEFAULT_PROMPT_TEMPLATE, max_sequence_length: int = 256, - ): + ) -> HunyuanVideoPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py index 9e7c198c19cc..1edd4f6db9f6 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py @@ -466,7 +466,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], prompt_template: dict[str, Any] = DEFAULT_PROMPT_TEMPLATE, max_sequence_length: int = 256, - ): + ) -> HunyuanVideoPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py index 349481492ac0..0208559bb56e 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py @@ -619,7 +619,7 @@ def __call__( prompt_template: dict[str, Any] = DEFAULT_PROMPT_TEMPLATE, max_sequence_length: int = 256, sampling_type: FramepackSamplingType = FramepackSamplingType.INVERTED_ANTI_DRIFTING, - ): + ) -> HunyuanVideoFramepackPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py index 13eb35386001..929118cc0af8 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py @@ -651,7 +651,7 @@ def __call__( prompt_template: dict[str, Any] = DEFAULT_PROMPT_TEMPLATE, max_sequence_length: int = 256, image_embed_interleave: int | None = None, - ): + ) -> HunyuanVideoPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py index 7232ebbee5b8..cba5f0aa5c68 100644 --- a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py +++ b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py @@ -564,7 +564,7 @@ def __call__( output_type: str | None = "np", return_dict: bool = True, attention_kwargs: dict[str, Any] | None = None, - ): + ) -> HunyuanVideo15PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py index 71a36a1c51cd..b171f83c3668 100644 --- a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py +++ b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py @@ -671,7 +671,7 @@ def __call__( output_type: str | None = "np", return_dict: bool = True, attention_kwargs: dict[str, Any] | None = None, - ): + ) -> HunyuanVideo15PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/hunyuandit/pipeline_hunyuandit.py b/src/diffusers/pipelines/hunyuandit/pipeline_hunyuandit.py index 5d656a3c370a..1e9a78130ec5 100644 --- a/src/diffusers/pipelines/hunyuandit/pipeline_hunyuandit.py +++ b/src/diffusers/pipelines/hunyuandit/pipeline_hunyuandit.py @@ -596,7 +596,7 @@ def __call__( target_size: tuple[int, int] | None = None, crops_coords_top_left: tuple[int, int] = (0, 0), use_resolution_binning: bool = True, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation with HunyuanDiT. diff --git a/src/diffusers/pipelines/ideogram4/pipeline_ideogram4.py b/src/diffusers/pipelines/ideogram4/pipeline_ideogram4.py index 7577e0463ca7..8aa76660be15 100644 --- a/src/diffusers/pipelines/ideogram4/pipeline_ideogram4.py +++ b/src/diffusers/pipelines/ideogram4/pipeline_ideogram4.py @@ -501,7 +501,7 @@ def __call__( attention_kwargs: dict[str, Any] | None = None, callback_on_step_end: Callable[["Ideogram4Pipeline", int, int, dict[str, Any]], dict[str, Any]] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ) -> Ideogram4PipelineOutput | tuple[Any]: + ) -> Ideogram4PipelineOutput | tuple: r""" Run text-to-image generation. diff --git a/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit.py b/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit.py index 2eff843d5225..29b2269ef018 100644 --- a/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit.py +++ b/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit.py @@ -629,7 +629,7 @@ def __call__( callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 4096, enable_denormalization: bool = True, - ): + ) -> JoyImageEditPipelineOutput | tuple: r""" Generate an edited image conditioned on a reference image and a text prompt. diff --git a/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit_plus.py b/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit_plus.py index ac8e01278e3e..0f4d1c207325 100644 --- a/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit_plus.py +++ b/src/diffusers/pipelines/joyimage/pipeline_joyimage_edit_plus.py @@ -465,7 +465,7 @@ def __call__( | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 4096, - ): + ) -> JoyImageEditPlusPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky.py b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky.py index aa6006ddc082..5274453551ca 100644 --- a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky.py +++ b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky.py @@ -252,7 +252,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_combined.py b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_combined.py index 3a0c3c07e8ca..feb552e26afb 100644 --- a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_combined.py +++ b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_combined.py @@ -28,7 +28,7 @@ from ...utils import ( replace_example_docstring, ) -from ..pipeline_utils import DiffusionPipeline +from ..pipeline_utils import DiffusionPipeline, ImagePipelineOutput from .pipeline_kandinsky import KandinskyPipeline from .pipeline_kandinsky_img2img import KandinskyImg2ImgPipeline from .pipeline_kandinsky_inpaint import KandinskyInpaintPipeline @@ -231,7 +231,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. @@ -452,7 +452,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. @@ -693,7 +693,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_img2img.py b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_img2img.py index f374741bd9fb..740dad385c79 100644 --- a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_img2img.py +++ b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_img2img.py @@ -314,7 +314,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_inpaint.py b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_inpaint.py index 9bc52ce871f5..61d0c371ea89 100644 --- a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_inpaint.py +++ b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_inpaint.py @@ -419,7 +419,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_prior.py b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_prior.py index 5708286b7d5a..c21589d15043 100644 --- a/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_prior.py +++ b/src/diffusers/pipelines/kandinsky/pipeline_kandinsky_prior.py @@ -415,7 +415,7 @@ def __call__( guidance_scale: float = 4.0, output_type: str | None = "pt", return_dict: bool = True, - ): + ) -> KandinskyPriorPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2.py index 668df4ac5980..2ef2f245c841 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2.py @@ -145,7 +145,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_combined.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_combined.py index d5616958a563..3fc6e68b7c66 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_combined.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_combined.py @@ -21,7 +21,7 @@ from ...models import PriorTransformer, UNet2DConditionModel, VQModel from ...schedulers import DDPMScheduler, UnCLIPScheduler from ...utils import deprecate, logging, replace_example_docstring -from ..pipeline_utils import DiffusionPipeline +from ..pipeline_utils import DiffusionPipeline, ImagePipelineOutput from .pipeline_kandinsky2_2 import KandinskyV22Pipeline from .pipeline_kandinsky2_2_img2img import KandinskyV22Img2ImgPipeline from .pipeline_kandinsky2_2_inpainting import KandinskyV22InpaintPipeline @@ -222,7 +222,7 @@ def __call__( prior_callback_on_step_end_tensor_inputs: list[str] = ["latents"], callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. @@ -468,7 +468,7 @@ def __call__( prior_callback_on_step_end_tensor_inputs: list[str] = ["latents"], callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. @@ -723,7 +723,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet.py index a9b7be516114..4a8c19784f2d 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet.py @@ -174,7 +174,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet_img2img.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet_img2img.py index f77a40595194..99a7ee527068 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet_img2img.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_controlnet_img2img.py @@ -215,7 +215,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, return_dict: bool = True, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_img2img.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_img2img.py index dc17f49bbfe0..649c1d7c678b 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_img2img.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_img2img.py @@ -198,7 +198,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_inpainting.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_inpainting.py index f258dfc07094..e4ba5558f63e 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_inpainting.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_inpainting.py @@ -318,7 +318,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior.py index 8095f79280d4..c26afa61ff6f 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior.py @@ -387,7 +387,7 @@ def __call__( return_dict: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> KandinskyPriorPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior_emb2emb.py b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior_emb2emb.py index 72f1d8556ec5..9ed4ef589789 100644 --- a/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior_emb2emb.py +++ b/src/diffusers/pipelines/kandinsky2_2/pipeline_kandinsky2_2_prior_emb2emb.py @@ -410,7 +410,7 @@ def __call__( guidance_scale: float = 4.0, output_type: str | None = "pt", # pt only return_dict: bool = True, - ): + ) -> KandinskyPriorPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3.py b/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3.py index ca8f124c74cf..ba26777f4366 100644 --- a/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3.py +++ b/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3.py @@ -353,7 +353,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3_img2img.py b/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3_img2img.py index beb4caafb6d3..9ce1da92f459 100644 --- a/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3_img2img.py +++ b/src/diffusers/pipelines/kandinsky3/pipeline_kandinsky3_img2img.py @@ -418,7 +418,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py index 7a5dc2cd1ac1..1ad08729e8eb 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py @@ -702,7 +702,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> KandinskyPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py index 1784b4e42972..9fa3378d0143 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py @@ -589,7 +589,7 @@ def __call__( callback_on_step_end: Callable[[int, int, None], PipelineCallback | MultiPipelineCallbacks] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> KandinskyImagePipelineOutput | tuple: r""" The call function to the pipeline for image-to-image generation. diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py index d39478547d8e..34099f191891 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py @@ -771,7 +771,7 @@ def __call__( callback_on_step_end: Callable[[int, int, None], PipelineCallback | MultiPipelineCallbacks] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> KandinskyPipelineOutput | tuple: r""" The call function to the pipeline for image-to-video generation. diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py index d86fff668771..46002e086a28 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py @@ -555,7 +555,7 @@ def __call__( callback_on_step_end: Callable[[int, int, None], PipelineCallback | MultiPipelineCallbacks] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> KandinskyImagePipelineOutput | tuple: r""" The call function to the pipeline for text-to-image generation. diff --git a/src/diffusers/pipelines/kolors/pipeline_kolors.py b/src/diffusers/pipelines/kolors/pipeline_kolors.py index 1e11faf8b9b6..a4d92e278d70 100644 --- a/src/diffusers/pipelines/kolors/pipeline_kolors.py +++ b/src/diffusers/pipelines/kolors/pipeline_kolors.py @@ -679,7 +679,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> KolorsPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/kolors/pipeline_kolors_img2img.py b/src/diffusers/pipelines/kolors/pipeline_kolors_img2img.py index d9b519267216..39ed6e37ffaf 100644 --- a/src/diffusers/pipelines/kolors/pipeline_kolors_img2img.py +++ b/src/diffusers/pipelines/kolors/pipeline_kolors_img2img.py @@ -814,7 +814,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> KolorsPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/krea2/pipeline_krea2.py b/src/diffusers/pipelines/krea2/pipeline_krea2.py index 51d33cb48619..dbc9a19e74be 100644 --- a/src/diffusers/pipelines/krea2/pipeline_krea2.py +++ b/src/diffusers/pipelines/krea2/pipeline_krea2.py @@ -466,7 +466,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], attention_kwargs: dict[str, Any] | None = None, max_sequence_length: int = 512, - ): + ) -> Krea2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_img2img.py b/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_img2img.py index 424a2c46e06b..6a0fdb96147e 100644 --- a/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_img2img.py +++ b/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_img2img.py @@ -730,7 +730,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_text2img.py b/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_text2img.py index 60f59ec7f9d3..947421628577 100644 --- a/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_text2img.py +++ b/src/diffusers/pipelines/latent_consistency_models/pipeline_latent_consistency_text2img.py @@ -661,7 +661,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion.py b/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion.py index 2b0e85d393a7..63d6dc6c7117 100644 --- a/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion.py +++ b/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion.py @@ -85,7 +85,7 @@ def __call__( output_type: str | None = "pil", return_dict: bool = True, **kwargs, - ) -> tuple | ImagePipelineOutput: + ) -> ImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion_superresolution.py b/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion_superresolution.py index c44d49944ea3..13f28e3ee8c7 100644 --- a/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion_superresolution.py +++ b/src/diffusers/pipelines/latent_diffusion/pipeline_latent_diffusion_superresolution.py @@ -78,7 +78,7 @@ def __call__( generator: torch.Generator | list[torch.Generator] | None = None, output_type: str | None = "pil", return_dict: bool = True, - ) -> tuple | ImagePipelineOutput: + ) -> ImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion.py b/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion.py index c6cdb127309b..a4dd6ddd1dd5 100644 --- a/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion.py +++ b/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion.py @@ -745,7 +745,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> LEditsPPDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for editing. The [`~pipelines.ledits_pp.LEditsPPPipelineStableDiffusion.invert`] method has to be called beforehand. Edits will diff --git a/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion_xl.py b/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion_xl.py index 6e97b1ea2ed4..039bc276aa6c 100644 --- a/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion_xl.py +++ b/src/diffusers/pipelines/ledits_pp/pipeline_leditspp_stable_diffusion_xl.py @@ -813,7 +813,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> LEditsPPDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for editing. The [`~pipelines.ledits_pp.LEditsPPPipelineStableDiffusionXL.invert`] method has to be called beforehand. Edits diff --git a/src/diffusers/pipelines/llada2/pipeline_llada2.py b/src/diffusers/pipelines/llada2/pipeline_llada2.py index 06b4875f18a9..a9b8f351144e 100644 --- a/src/diffusers/pipelines/llada2/pipeline_llada2.py +++ b/src/diffusers/pipelines/llada2/pipeline_llada2.py @@ -271,7 +271,7 @@ def __call__( | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] | None = None, - ) -> LLaDA2PipelineOutput | tuple[torch.LongTensor, list[str] | None]: + ) -> LLaDA2PipelineOutput | tuple: """ Generate text with block-wise refinement. diff --git a/src/diffusers/pipelines/longcat_audio_dit/pipeline_longcat_audio_dit.py b/src/diffusers/pipelines/longcat_audio_dit/pipeline_longcat_audio_dit.py index e6478535b373..f48d3376916c 100644 --- a/src/diffusers/pipelines/longcat_audio_dit/pipeline_longcat_audio_dit.py +++ b/src/diffusers/pipelines/longcat_audio_dit/pipeline_longcat_audio_dit.py @@ -231,7 +231,7 @@ def __call__( return_dict: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> AudioPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. @@ -254,6 +254,10 @@ def __call__( Tensor inputs passed to `callback_on_step_end`. Examples: + + Returns: + [`~pipelines.AudioPipelineOutput`] or `tuple`: [`~pipelines.AudioPipelineOutput`] if `return_dict` is True, + otherwise a `tuple`. When returning a tuple, the first element is the generated audio waveform. """ if prompt is None: prompt = [] diff --git a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py index 41ca3eb54f83..eed7b40c48fc 100644 --- a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py +++ b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py @@ -490,7 +490,7 @@ def __call__( enable_cfg_renorm: bool | None = True, cfg_renorm_min: float | None = 0.0, enable_prompt_rewrite: bool | None = True, - ): + ) -> LongCatImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py index 9f35bb685d9f..8f107bd456c4 100644 --- a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py +++ b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py @@ -546,7 +546,7 @@ def __call__( output_type: str | None = "pil", return_dict: bool = True, joint_attention_kwargs: dict[str, Any] | None = None, - ): + ) -> LongCatImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx.py b/src/diffusers/pipelines/ltx/pipeline_ltx.py index ce9177547c52..e66f8dc18cf4 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx.py @@ -561,7 +561,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 128, - ): + ) -> LTXPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py b/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py index 28d296695998..dc2dde5aad93 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py @@ -881,7 +881,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> LTXPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py b/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py index 838d5afc5c5a..f057f8d7908e 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py @@ -972,7 +972,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 128, - ): + ) -> LTXPipelineOutput | tuple: r""" Generate an image-to-video sequence via temporal sliding windows and multi-prompt scheduling. diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py b/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py index 81ecfce50efa..11e65270cfb6 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py @@ -623,7 +623,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 128, - ): + ) -> LTXPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx_latent_upsample.py b/src/diffusers/pipelines/ltx/pipeline_ltx_latent_upsample.py index 6f325237a248..c6eff670908f 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx_latent_upsample.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx_latent_upsample.py @@ -199,7 +199,7 @@ def __call__( generator: torch.Generator | list[torch.Generator] | None = None, output_type: str | None = "pil", return_dict: bool = True, - ): + ) -> LTXPipelineOutput | tuple: r""" Function invoked when calling the pipeline for latent upsampling. @@ -227,6 +227,11 @@ def __call__( The output format of the generated video. Choose between `PIL.Image`, `np.array`, or `latent`. return_dict (`bool`, *optional*, defaults to `True`): Whether or not to return a [`~pipelines.ltx.LTXPipelineOutput`] instead of a plain tuple. + + Returns: + [`~pipelines.ltx.LTXPipelineOutput`] or `tuple`: [`~pipelines.ltx.LTXPipelineOutput`] if `return_dict` is + True, otherwise a `tuple`. When returning a tuple, the first element is the upsampled video (or the latents + if `output_type="latent"`). """ self.check_inputs( video=video, diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2.py index 22948a7ecf3a..4e5ced0b4ec8 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2.py @@ -970,7 +970,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py index bd2ee3ec6708..e1bff845302c 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py @@ -1390,7 +1390,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py index 0cda9c6079b6..936af16d8805 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py @@ -1563,7 +1563,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2DFRPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py index bebf3739956a..bdf9c8e67677 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py @@ -1377,7 +1377,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2DFRPipelineOutput | tuple: r""" Run one temporal refine round. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_diffusion_decode.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_diffusion_decode.py index 2f1137830a57..28d740991a54 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_diffusion_decode.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_diffusion_decode.py @@ -84,7 +84,7 @@ def __call__( output_type: str = "pil", return_dict: bool = True, denormalize: bool = True, - ): + ) -> LTX2VideoDecodeOutput | tuple: r""" Args: latents (`torch.Tensor`): diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py index 91173bc6e161..d9d95cdc0a0b 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py @@ -1080,7 +1080,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2PipelineOutput | tuple: r""" Run HDR IC-LoRA video generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py index dc92b6eb965a..2924a086c721 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py @@ -1822,7 +1822,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py index c7c81d26cb45..36b7effdfd6c 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py @@ -1026,7 +1026,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ): + ) -> LTX2PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_latent_upsample.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_latent_upsample.py index 8aa72a425e1c..09d842927d12 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_latent_upsample.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_latent_upsample.py @@ -280,7 +280,7 @@ def __call__( generator: torch.Generator | list[torch.Generator] | None = None, output_type: str | None = "pil", return_dict: bool = True, - ): + ) -> LTXPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py b/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py index 1bd7ab4ca675..b0d2f91736ea 100644 --- a/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py +++ b/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py @@ -472,7 +472,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> LucyPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/marigold/pipeline_marigold_depth.py b/src/diffusers/pipelines/marigold/pipeline_marigold_depth.py index a81d1c51742c..f8eaf683e0c0 100644 --- a/src/diffusers/pipelines/marigold/pipeline_marigold_depth.py +++ b/src/diffusers/pipelines/marigold/pipeline_marigold_depth.py @@ -363,7 +363,7 @@ def __call__( output_uncertainty: bool = False, output_latent: bool = False, return_dict: bool = True, - ): + ) -> MarigoldDepthOutput | tuple: """ Function invoked when calling the pipeline. diff --git a/src/diffusers/pipelines/marigold/pipeline_marigold_intrinsics.py b/src/diffusers/pipelines/marigold/pipeline_marigold_intrinsics.py index 9488d8f5c9b8..e0fe93d561b2 100644 --- a/src/diffusers/pipelines/marigold/pipeline_marigold_intrinsics.py +++ b/src/diffusers/pipelines/marigold/pipeline_marigold_intrinsics.py @@ -375,7 +375,7 @@ def __call__( output_uncertainty: bool = False, output_latent: bool = False, return_dict: bool = True, - ): + ) -> MarigoldIntrinsicsOutput | tuple: """ Function invoked when calling the pipeline. diff --git a/src/diffusers/pipelines/marigold/pipeline_marigold_normals.py b/src/diffusers/pipelines/marigold/pipeline_marigold_normals.py index 3f94ce441232..14a46be58ea6 100644 --- a/src/diffusers/pipelines/marigold/pipeline_marigold_normals.py +++ b/src/diffusers/pipelines/marigold/pipeline_marigold_normals.py @@ -348,7 +348,7 @@ def __call__( output_uncertainty: bool = False, output_latent: bool = False, return_dict: bool = True, - ): + ) -> MarigoldNormalsOutput | tuple: """ Function invoked when calling the pipeline. diff --git a/src/diffusers/pipelines/mochi/pipeline_mochi.py b/src/diffusers/pipelines/mochi/pipeline_mochi.py index c146d2d1e564..cbb8f3b2c6c7 100644 --- a/src/diffusers/pipelines/mochi/pipeline_mochi.py +++ b/src/diffusers/pipelines/mochi/pipeline_mochi.py @@ -466,7 +466,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> MochiPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/motif_video/pipeline_motif_video.py b/src/diffusers/pipelines/motif_video/pipeline_motif_video.py index 8ad37932e970..8bd2eb5bb5c9 100644 --- a/src/diffusers/pipelines/motif_video/pipeline_motif_video.py +++ b/src/diffusers/pipelines/motif_video/pipeline_motif_video.py @@ -516,7 +516,7 @@ def __call__( callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 512, vae_batch_size: int | None = None, - ): + ) -> MotifVideoPipelineOutput | tuple: r""" The call function to the pipeline for text-to-video generation. diff --git a/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py b/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py index 1b32ba74f24b..d30aebc610b8 100644 --- a/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py +++ b/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py @@ -644,7 +644,7 @@ def __call__( ] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> MotifVideoPipelineOutput | tuple: r""" The call function to the pipeline for image-to-video generation. diff --git a/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py b/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py index f50f11c8c152..71b7a82de4ea 100644 --- a/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py +++ b/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py @@ -401,7 +401,7 @@ def __call__( attention_kwargs: dict[str, Any] | None = None, callback_on_step_end: Callable[[int, int, dict], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> NucleusMoEImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/omnigen/pipeline_omnigen.py b/src/diffusers/pipelines/omnigen/pipeline_omnigen.py index 6564b2a672a0..69d369c0ab34 100644 --- a/src/diffusers/pipelines/omnigen/pipeline_omnigen.py +++ b/src/diffusers/pipelines/omnigen/pipeline_omnigen.py @@ -293,7 +293,7 @@ def __call__( return_dict: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> ImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. @@ -356,6 +356,10 @@ def __call__( Returns: [`~pipelines.ImagePipelineOutput`] or `tuple`: If `return_dict` is `True`, [`~pipelines.ImagePipelineOutput`] is returned, otherwise a `tuple` is returned where the first element is a list with the generated images. + + Returns: + [`~pipelines.ImagePipelineOutput`] or `tuple`: [`~pipelines.ImagePipelineOutput`] if `return_dict` is True, + otherwise a `tuple`. When returning a tuple, the first element is a list with the generated images. """ height = height or self.default_sample_size * self.vae_scale_factor diff --git a/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py b/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py index b22f2f0cec2d..841f4e1471e2 100644 --- a/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py +++ b/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py @@ -476,7 +476,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, - ): + ) -> OvisImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd.py b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd.py index 3a88272f24f4..7b6abb32acc5 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd.py @@ -893,7 +893,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_inpaint.py b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_inpaint.py index 98221e4c30ba..3cdd6bcdc6e9 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_inpaint.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_inpaint.py @@ -1003,7 +1003,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl.py b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl.py index 72060682c196..27c3c2f4ded9 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl.py @@ -1043,7 +1043,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl_img2img.py b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl_img2img.py index 4f2b8b2044fc..133acd82168d 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl_img2img.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_controlnet_sd_xl_img2img.py @@ -1122,7 +1122,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_hunyuandit.py b/src/diffusers/pipelines/pag/pipeline_pag_hunyuandit.py index a443a19bd952..6c96e16f51cb 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_hunyuandit.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_hunyuandit.py @@ -612,7 +612,7 @@ def __call__( use_resolution_binning: bool = True, pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation with HunyuanDiT. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_kolors.py b/src/diffusers/pipelines/pag/pipeline_pag_kolors.py index 4f138d91d9c6..9c588c2c168a 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_kolors.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_kolors.py @@ -699,7 +699,7 @@ def __call__( pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, max_sequence_length: int = 256, - ): + ) -> KolorsPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd.py b/src/diffusers/pipelines/pag/pipeline_pag_sd.py index b12597460f65..5548ec8acdd2 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd.py @@ -771,7 +771,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_3.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_3.py index d86adccc2ccf..aebdd2495c3b 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_3.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_3.py @@ -712,7 +712,7 @@ def __call__( max_sequence_length: int = 256, pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_3_img2img.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_3_img2img.py index 24f3d828bd81..def0cc3c3b03 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_3_img2img.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_3_img2img.py @@ -765,7 +765,7 @@ def __call__( max_sequence_length: int = 256, pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_animatediff.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_animatediff.py index 2baeda5649ad..4303faa08671 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_animatediff.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_animatediff.py @@ -600,7 +600,7 @@ def __call__( decode_chunk_size: int = 16, pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> AnimateDiffPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_img2img.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_img2img.py index de6dfbc585fa..dacd2c066a24 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_img2img.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_img2img.py @@ -806,7 +806,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_inpaint.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_inpaint.py index 426419f12f73..2b10667c5cdb 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_inpaint.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_inpaint.py @@ -941,7 +941,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl.py index ca9c6b5aadd9..dd873edc63a6 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl.py @@ -873,7 +873,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_img2img.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_img2img.py index 31fdf19cbade..4566225a240f 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_img2img.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_img2img.py @@ -1029,7 +1029,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_inpaint.py b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_inpaint.py index 77933867631c..3d7d3500e669 100644 --- a/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_inpaint.py +++ b/src/diffusers/pipelines/pag/pipeline_pag_sd_xl_inpaint.py @@ -1125,7 +1125,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], pag_scale: float = 3.0, pag_adaptive_scale: float = 0.0, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/prx/pipeline_prx.py b/src/diffusers/pipelines/prx/pipeline_prx.py index f4ec214313e3..b07fc13d2f0e 100644 --- a/src/diffusers/pipelines/prx/pipeline_prx.py +++ b/src/diffusers/pipelines/prx/pipeline_prx.py @@ -576,7 +576,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], tokenizer_max_length: int | None = None, skip_text_cleaning: bool = False, - ): + ) -> PRXPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/prx/pipeline_prx_pixel.py b/src/diffusers/pipelines/prx/pipeline_prx_pixel.py index 22a4d8dd4b18..4aa6e41aac5d 100644 --- a/src/diffusers/pipelines/prx/pipeline_prx_pixel.py +++ b/src/diffusers/pipelines/prx/pipeline_prx_pixel.py @@ -426,7 +426,7 @@ def __call__( use_resolution_binning: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> PRXPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py index 1da0518a4f65..34bfb8637f68 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py @@ -430,7 +430,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py index f946fdf27d00..ff619a781585 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py @@ -537,7 +537,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py index 97f510a6dbf4..8c09df5574c1 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py @@ -603,7 +603,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py index 85abb815cf23..1861141810f6 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py @@ -527,7 +527,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py index 57d1fdaaf99f..6413205424eb 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py @@ -664,7 +664,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py index 84d1b60152b1..8b17f17c0e3e 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py @@ -549,7 +549,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py index 9b9af83737e5..9beaa769fae2 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py @@ -506,7 +506,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py index 3d5f0040932a..a7b0e4d9912b 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py @@ -619,7 +619,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py index 7e06a7d36ffd..e7053794e9b6 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py @@ -562,7 +562,7 @@ def __call__( resolution: int = 640, cfg_normalize: bool = False, use_en_prompt: bool = False, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py b/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py index 786b09b4e3cd..6aecf130d64f 100644 --- a/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py +++ b/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py @@ -526,7 +526,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], output_resolution: int = 1024, use_kv_cache: bool = True, - ): + ) -> QwenImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/shap_e/pipeline_shap_e.py b/src/diffusers/pipelines/shap_e/pipeline_shap_e.py index eea83aff9e10..56e8f57dab35 100644 --- a/src/diffusers/pipelines/shap_e/pipeline_shap_e.py +++ b/src/diffusers/pipelines/shap_e/pipeline_shap_e.py @@ -200,7 +200,7 @@ def __call__( frame_size: int = 64, output_type: str | None = "pil", # pil, np, latent, mesh return_dict: bool = True, - ): + ) -> ShapEPipelineOutput | tuple: """ The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/shap_e/pipeline_shap_e_img2img.py b/src/diffusers/pipelines/shap_e/pipeline_shap_e_img2img.py index f59fd298c684..403d5fa8a743 100644 --- a/src/diffusers/pipelines/shap_e/pipeline_shap_e_img2img.py +++ b/src/diffusers/pipelines/shap_e/pipeline_shap_e_img2img.py @@ -182,7 +182,7 @@ def __call__( frame_size: int = 64, output_type: str | None = "pil", # pil, np, latent, mesh return_dict: bool = True, - ): + ) -> ShapEPipelineOutput | tuple: """ The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py index 0c9e6add9937..b4a4242b0036 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py @@ -396,7 +396,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> SkyReelsV2PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py index 31b75bfb336f..73b1541f9969 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py @@ -623,7 +623,7 @@ def __call__( ar_step: int = 0, causal_block_size: int | None = None, fps: int = 24, - ): + ) -> SkyReelsV2PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py index 576681b1b957..c19997155d66 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py @@ -672,7 +672,7 @@ def __call__( ar_step: int = 0, causal_block_size: int | None = None, fps: int = 24, - ): + ) -> SkyReelsV2PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py index df6076263238..c18a80fdae0e 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py @@ -710,7 +710,7 @@ def __call__( ar_step: int = 0, causal_block_size: int | None = None, fps: int = 24, - ): + ) -> SkyReelsV2PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py index b1f70b60b22a..977073a60a1e 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py @@ -501,7 +501,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> SkyReelsV2PipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_audio/pipeline_stable_audio.py b/src/diffusers/pipelines/stable_audio/pipeline_stable_audio.py index 475f4032edab..b5f7cbfdf93f 100644 --- a/src/diffusers/pipelines/stable_audio/pipeline_stable_audio.py +++ b/src/diffusers/pipelines/stable_audio/pipeline_stable_audio.py @@ -483,7 +483,7 @@ def __call__( callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int | None = 1, output_type: str | None = "pt", - ): + ) -> AudioPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3.py b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3.py index a1ef3cf3cca7..11004b31b70b 100644 --- a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3.py +++ b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3.py @@ -418,7 +418,7 @@ def __call__( callback_on_step_end: Optional[Callable[[int, int, dict], dict]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], output_type: str = "pt", - ) -> Union[AudioPipelineOutput, tuple]: + ) -> AudioPipelineOutput | tuple: r""" Generate audio from a text prompt. diff --git a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_audio2audio.py b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_audio2audio.py index 50cd4775d5f1..b84388adb12a 100644 --- a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_audio2audio.py +++ b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_audio2audio.py @@ -442,7 +442,7 @@ def __call__( callback_on_step_end: Optional[Callable[[int, int, dict], dict]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], output_type: str = "pt", - ) -> Union[AudioPipelineOutput, tuple]: + ) -> AudioPipelineOutput | tuple: r""" Generate an audio variation conditioned on a text prompt and a reference waveform. diff --git a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_inpaint.py b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_inpaint.py index 687546947eb5..14bafdadbab3 100644 --- a/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_inpaint.py +++ b/src/diffusers/pipelines/stable_audio_3/pipeline_stable_audio_3_inpaint.py @@ -475,7 +475,7 @@ def __call__( callback_on_step_end: Optional[Callable[[int, int, dict], dict]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], output_type: str = "pt", - ) -> Union[AudioPipelineOutput, tuple]: + ) -> AudioPipelineOutput | tuple: r""" Generate inpainted audio conditioned on a text prompt and reference. diff --git a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade.py b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade.py index 0961d1c46e94..c7af35aa955e 100644 --- a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade.py +++ b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade.py @@ -14,6 +14,8 @@ from typing import Callable +import numpy as np +import PIL.Image import torch from transformers import CLIPTextModelWithProjection, CLIPTokenizer @@ -321,7 +323,7 @@ def __call__( return_dict: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> ImagePipelineOutput | list[PIL.Image.Image] | np.ndarray | torch.Tensor: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_combined.py b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_combined.py index 71bfc3f7bab6..a1d2ca6c56b3 100644 --- a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_combined.py +++ b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_combined.py @@ -13,6 +13,7 @@ # limitations under the License. from typing import Callable +import numpy as np import PIL import torch from transformers import CLIPImageProcessor, CLIPTextModelWithProjection, CLIPTokenizer, CLIPVisionModelWithProjection @@ -21,7 +22,7 @@ from ...schedulers import DDPMWuerstchenScheduler from ...utils import is_torch_version, replace_example_docstring from ..deprecated.wuerstchen.modeling_paella_vq_model import PaellaVQModel -from ..pipeline_utils import DeprecatedPipelineMixin, DiffusionPipeline +from ..pipeline_utils import DeprecatedPipelineMixin, DiffusionPipeline, ImagePipelineOutput from .pipeline_stable_cascade import StableCascadeDecoderPipeline from .pipeline_stable_cascade_prior import StableCascadePriorPipeline @@ -181,7 +182,7 @@ def __call__( prior_callback_on_step_end_tensor_inputs: list[str] = ["latents"], callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> ImagePipelineOutput | list[PIL.Image.Image] | np.ndarray | torch.Tensor: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_prior.py b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_prior.py index fb58094f964b..b2fc2799ff31 100644 --- a/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_prior.py +++ b/src/diffusers/pipelines/stable_cascade/pipeline_stable_cascade_prior.py @@ -396,7 +396,7 @@ def __call__( return_dict: bool = True, callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], - ): + ) -> StableCascadePriorPipelineOutput | tuple: """ Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion.py b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion.py index d7776c7b5196..3c47b5e45b62 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion.py @@ -282,7 +282,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, np.ndarray], None] | None = None, callback_steps: int = 1, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_img2img.py b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_img2img.py index 88b84cff804d..dd27a0de50d2 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_img2img.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_img2img.py @@ -330,7 +330,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, np.ndarray], None] | None = None, callback_steps: int = 1, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_inpaint.py b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_inpaint.py index cf04bdf4da7b..1eb0ea87a8a1 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_inpaint.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_inpaint.py @@ -339,7 +339,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, np.ndarray], None] | None = None, callback_steps: int = 1, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_upscale.py b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_upscale.py index 8494a253f54f..548fdd9e1734 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_upscale.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_onnx_stable_diffusion_upscale.py @@ -366,7 +366,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, np.ndarray], None] | None = None, callback_steps: int | None = 1, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py index d28bb2a9fe59..36b350bd8258 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py @@ -803,7 +803,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_depth2img.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_depth2img.py index 977de5d7fb39..5ffa48daf6d2 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_depth2img.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_depth2img.py @@ -653,7 +653,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> ImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_image_variation.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_image_variation.py index 15b8daf334ed..914af27b6a97 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_image_variation.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_image_variation.py @@ -272,7 +272,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_img2img.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_img2img.py index 719be9258341..54a6005d17fe 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_img2img.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_img2img.py @@ -881,7 +881,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_inpaint.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_inpaint.py index 96794eaa297a..dc1f74c74b3d 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_inpaint.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_inpaint.py @@ -908,7 +908,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_instruct_pix2pix.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_instruct_pix2pix.py index 7a24e6008351..c0df751cf054 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_instruct_pix2pix.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_instruct_pix2pix.py @@ -192,7 +192,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], cross_attention_kwargs: dict[str, Any] | None = None, **kwargs, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_latent_upscale.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_latent_upscale.py index 1920f033c126..c057434b8f8b 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_latent_upscale.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_latent_upscale.py @@ -411,7 +411,7 @@ def __call__( return_dict: bool = True, callback: Callable[[int, int, torch.Tensor], None] | None = None, callback_steps: int = 1, - ): + ) -> ImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_upscale.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_upscale.py index 1a0a7412e5d7..402f3f67844e 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_upscale.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion_upscale.py @@ -553,7 +553,7 @@ def __call__( callback_steps: int = 1, cross_attention_kwargs: dict[str, Any] | None = None, clip_skip: int = None, - ): + ) -> StableDiffusionPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip.py index d2c9bf0c4162..1d6feec05846 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip.py @@ -670,7 +670,7 @@ def __call__( prior_guidance_scale: float = 4.0, prior_latents: torch.Tensor | None = None, clip_skip: int | None = None, - ): + ) -> ImagePipelineOutput | tuple: """ The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip_img2img.py b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip_img2img.py index 059ae1e6fd4d..c560b411620e 100644 --- a/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip_img2img.py +++ b/src/diffusers/pipelines/stable_diffusion/pipeline_stable_unclip_img2img.py @@ -646,7 +646,7 @@ def __call__( noise_level: int = 0, image_embeds: torch.Tensor | None = None, clip_skip: int | None = None, - ): + ) -> ImagePipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3.py b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3.py index 5c05b469660f..9509adde741b 100644 --- a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3.py +++ b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3.py @@ -805,7 +805,7 @@ def __call__( skip_layer_guidance_stop: float = 0.2, skip_layer_guidance_start: float = 0.01, mu: float | None = None, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_img2img.py b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_img2img.py index c0ab805a4ef4..54ae68d19fb4 100644 --- a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_img2img.py +++ b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_img2img.py @@ -860,7 +860,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, mu: float | None = None, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_inpaint.py b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_inpaint.py index 321e9f8dd80e..9f623ec02091 100644 --- a/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_inpaint.py +++ b/src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3_inpaint.py @@ -955,7 +955,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 256, mu: float | None = None, - ): + ) -> StableDiffusion3PipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py index c116a49d81c6..d170e3e805af 100644 --- a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py +++ b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py @@ -859,7 +859,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_img2img.py b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_img2img.py index aedd131aae3c..7ea71e76844b 100644 --- a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_img2img.py +++ b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_img2img.py @@ -1014,7 +1014,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_inpaint.py b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_inpaint.py index 407b1a856216..15b978fab8ab 100644 --- a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_inpaint.py +++ b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_inpaint.py @@ -1124,7 +1124,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], **kwargs, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_instruct_pix2pix.py b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_instruct_pix2pix.py index bcd337414bac..10afb1b89ee0 100644 --- a/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_instruct_pix2pix.py +++ b/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl_instruct_pix2pix.py @@ -624,7 +624,7 @@ def __call__( original_size: tuple[int, int] = None, crops_coords_top_left: tuple[int, int] = (0, 0), target_size: tuple[int, int] = None, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/stable_video_diffusion/pipeline_stable_video_diffusion.py b/src/diffusers/pipelines/stable_video_diffusion/pipeline_stable_video_diffusion.py index 007d2b8da0cb..eb05eea7837f 100644 --- a/src/diffusers/pipelines/stable_video_diffusion/pipeline_stable_video_diffusion.py +++ b/src/diffusers/pipelines/stable_video_diffusion/pipeline_stable_video_diffusion.py @@ -406,7 +406,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], return_dict: bool = True, - ): + ) -> StableVideoDiffusionPipelineOutput | list[list[PIL.Image.Image]] | np.ndarray | torch.Tensor: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_adapter.py b/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_adapter.py index ffb877cfd0f6..3e111fa752ba 100644 --- a/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_adapter.py +++ b/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_adapter.py @@ -712,7 +712,7 @@ def __call__( cross_attention_kwargs: dict[str, Any] | None = None, adapter_conditioning_scale: float | list[float] = 1.0, clip_skip: int | None = None, - ): + ) -> StableDiffusionAdapterPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_xl_adapter.py b/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_xl_adapter.py index 1e7966192650..56c68ae8c9d0 100644 --- a/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_xl_adapter.py +++ b/src/diffusers/pipelines/t2i_adapter/pipeline_stable_diffusion_xl_adapter.py @@ -896,7 +896,7 @@ def __call__( adapter_conditioning_scale: float | list[float] = 1.0, adapter_conditioning_factor: float = 1.0, clip_skip: int | None = None, - ): + ) -> StableDiffusionXLPipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_combined.py b/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_combined.py index 2d881e22a783..6b5317e723c9 100644 --- a/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_combined.py +++ b/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_combined.py @@ -273,7 +273,7 @@ def __call__( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, upsampling_strength: float = 1.0, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the VisualCloze pipeline for generation. diff --git a/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py b/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py index b34f4c3faeab..cd6e61cff1df 100644 --- a/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py +++ b/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py @@ -677,7 +677,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> FluxPipelineOutput | tuple: r""" Function invoked when calling the VisualCloze pipeline for generation. diff --git a/src/diffusers/pipelines/wan/pipeline_wan.py b/src/diffusers/pipelines/wan/pipeline_wan.py index b33a2a7af3db..452911cf899d 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan.py +++ b/src/diffusers/pipelines/wan/pipeline_wan.py @@ -402,7 +402,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> WanPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/wan/pipeline_wan_animate.py b/src/diffusers/pipelines/wan/pipeline_wan_animate.py index a6b340c2d19f..e96729826e70 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_animate.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_animate.py @@ -791,7 +791,7 @@ def __call__( callback_on_step_end: Callable[[int, int, None], PipelineCallback | MultiPipelineCallbacks] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> WanPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/wan/pipeline_wan_i2v.py b/src/diffusers/pipelines/wan/pipeline_wan_i2v.py index a98f0324e0f0..2d1f7b94750e 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_i2v.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_i2v.py @@ -533,7 +533,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> WanPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/wan/pipeline_wan_vace.py b/src/diffusers/pipelines/wan/pipeline_wan_vace.py index 9186304b5953..f7689496c968 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_vace.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_vace.py @@ -714,7 +714,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> WanPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/wan/pipeline_wan_video2video.py b/src/diffusers/pipelines/wan/pipeline_wan_video2video.py index cfb26f5cb3b1..b192147acb64 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_video2video.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_video2video.py @@ -502,7 +502,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> WanPipelineOutput | tuple: r""" The call function to the pipeline for generation. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image.py b/src/diffusers/pipelines/z_image/pipeline_z_image.py index 3e2055c6257f..b97bdda170bf 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image.py @@ -318,7 +318,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet.py b/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet.py index 81373ffb56ff..6ba698a30db3 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet.py @@ -410,7 +410,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet_inpaint.py b/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet_inpaint.py index 178e74dea4fa..cb72c807700c 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet_inpaint.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image_controlnet_inpaint.py @@ -419,7 +419,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image_img2img.py b/src/diffusers/pipelines/z_image/pipeline_z_image_img2img.py index b5c7740bb0c1..efa6f50962db 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image_img2img.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image_img2img.py @@ -392,7 +392,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for image-to-image generation. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image_inpaint.py b/src/diffusers/pipelines/z_image/pipeline_z_image_inpaint.py index 132c22c0cff3..62a53ea261d1 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image_inpaint.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image_inpaint.py @@ -560,7 +560,7 @@ def __call__( callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for inpainting. diff --git a/src/diffusers/pipelines/z_image/pipeline_z_image_omni.py b/src/diffusers/pipelines/z_image/pipeline_z_image_omni.py index 50776ceaf34d..48ecc6e93163 100644 --- a/src/diffusers/pipelines/z_image/pipeline_z_image_omni.py +++ b/src/diffusers/pipelines/z_image/pipeline_z_image_omni.py @@ -383,7 +383,7 @@ def __call__( callback_on_step_end: Callable[[int, int], None] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> ZImagePipelineOutput | tuple: r""" Function invoked when calling the pipeline for generation. diff --git a/utils/check_return_annotations.py b/utils/check_return_annotations.py new file mode 100644 index 000000000000..be8b1dc61ffc --- /dev/null +++ b/utils/check_return_annotations.py @@ -0,0 +1,180 @@ +# coding=utf-8 +# Copyright 2026 The HuggingFace Inc. team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" +Check that these methods have a return type annotation: + +* `forward()` on every class in `src/diffusers/models` +* `__call__()` on every pipeline in `src/diffusers/pipelines` +* `__call__()` on every modular pipeline block in `src/diffusers/modular_pipelines` + +A class counts as a pipeline if it inherits from `DiffusionPipeline`, either directly or through another class. A class +counts as a modular pipeline block if it inherits from `ModularPipelineBlocks` in the same way. + +Deprecated code is skipped: + +* anything in a folder named `deprecated`, such as `pipelines/deprecated` +* pipelines that inherit from `DeprecatedPipelineMixin` +* classes and methods whose `# Copied from` comment points to deprecated code, because they can't change unless the + deprecated code changes too + +A method is only checked on the class where it's written, not on classes that inherit it. Any annotation passes, +including `-> None`. + +Run from the repository root: + + python utils/check_return_annotations.py +""" + +from __future__ import annotations + +import ast +import sys +from collections import defaultdict +from pathlib import Path + + +REPO_ROOT = Path(__file__).resolve().parents[1] +SRC_DIR = REPO_ROOT / "src" / "diffusers" +MODELS_DIR = SRC_DIR / "models" +PIPELINES_DIR = SRC_DIR / "pipelines" +MODULAR_DIR = SRC_DIR / "modular_pipelines" + +PIPELINE_BASE = "DiffusionPipeline" +DEPRECATED_PIPELINE_BASE = "DeprecatedPipelineMixin" +BLOCKS_BASE = "ModularPipelineBlocks" + + +def _base_names(class_def: ast.ClassDef) -> list[str]: + """Return the names of the classes this class inherits from. For a name like `nn.Module`, keep only `Module`.""" + names = [] + for base in class_def.bases: + if isinstance(base, ast.Name): + names.append(base.id) + elif isinstance(base, ast.Attribute): + names.append(base.attr) + return names + + +def _find_method(class_def: ast.ClassDef, method_name: str) -> ast.FunctionDef | ast.AsyncFunctionDef | None: + for node in class_def.body: + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) and node.name == method_name: + return node + return None + + +def _parse_classes(paths: list[Path]) -> list[tuple[Path, ast.ClassDef, list[str]]]: + """Return every class in `paths`, along with its file and the lines of that file.""" + classes = [] + for path in paths: + try: + source = path.read_text(encoding="utf-8") + tree = ast.parse(source) + except (SyntaxError, UnicodeDecodeError): + continue + lines = source.splitlines() + classes.extend((path, node, lines) for node in ast.walk(tree) if isinstance(node, ast.ClassDef)) + return classes + + +def _is_deprecated_path(path: Path) -> bool: + """Return whether the file is inside a folder named `deprecated`.""" + return "deprecated" in path.relative_to(SRC_DIR).parts[:-1] + + +def _copied_from_deprecated(lines: list[str], node: ast.ClassDef | ast.FunctionDef | ast.AsyncFunctionDef) -> bool: + """Return whether the `# Copied from` comment right above a class or method points to deprecated code.""" + first_line = min([decorator.lineno for decorator in node.decorator_list] + [node.lineno]) + if first_line < 2: + return False + comment = lines[first_line - 2].strip() + return comment.startswith("# Copied from") and ".deprecated." in comment + + +def _subclass_checker(classes: list[tuple[Path, ast.ClassDef, list[str]]]): + """ + Return a function `is_subclass(name, base)` that tells whether the class `name` inherits from the class `base`, + either directly or through other classes. + + Classes are matched by name only, so two classes with the same name in different files are treated as one class. + """ + bases_by_name: dict[str, set[str]] = defaultdict(set) + for _, class_def, _ in classes: + bases_by_name[class_def.name].update(_base_names(class_def)) + + cache: dict[tuple[str, str], bool] = {} + + def is_subclass(name: str, base: str, _seen: frozenset[str] = frozenset()) -> bool: + if name == base: + return True + if (name, base) in cache: + return cache[(name, base)] + if name in _seen: # stop if this class was already visited, so a loop in the class names can't run forever + return False + result = any(is_subclass(parent, base, _seen | {name}) for parent in bases_by_name.get(name, ())) + cache[(name, base)] = result + return result + + return is_subclass + + +def _is_under(path: Path, directory: Path) -> bool: + return directory in path.parents + + +def main() -> int: + classes = _parse_classes(sorted(SRC_DIR.rglob("*.py"))) + is_subclass = _subclass_checker(classes) + + errors = [] + for path, class_def, lines in classes: + if _is_deprecated_path(path): + continue + if _is_under(path, MODELS_DIR): + method_name = "forward" + elif _is_under(path, PIPELINES_DIR): + if not is_subclass(class_def.name, PIPELINE_BASE) or is_subclass(class_def.name, DEPRECATED_PIPELINE_BASE): + continue + method_name = "__call__" + elif _is_under(path, MODULAR_DIR): + if not is_subclass(class_def.name, BLOCKS_BASE): + continue + method_name = "__call__" + else: + continue + + method = _find_method(class_def, method_name) + if method is None or method.returns is not None: + continue + if _copied_from_deprecated(lines, class_def) or _copied_from_deprecated(lines, method): + continue + rel = path.relative_to(REPO_ROOT).as_posix() + errors.append(f"{rel}:{method.lineno}: {class_def.name}.{method_name} has no return type annotation") + + if errors: + print("\n".join(errors)) + sys.stdout.flush() # print the list before the summary, even when both end up in the same log + if len(errors) == 1: + summary = "Found 1 method without a return type annotation. Add one to the method above." + else: + summary = f"Found {len(errors)} methods without a return type annotation. Add one to each method above." + print(f"\n{summary}", file=sys.stderr) + return 1 + + print("All forward/__call__ methods have return type annotations.") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) From 863092ffec10c65e3bc27193e069197a395f5eb5 Mon Sep 17 00:00:00 2001 From: Jingya HUANG <44135271+JingyaHuang@users.noreply.github.com> Date: Thu, 1 Oct 2026 06:46:29 +0200 Subject: [PATCH 11/22] [core] Shard tensor-parallel checkpoints on load and save (#14544) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * [core] Shard tensor-parallel checkpoints on load and save Stream each rank's slice of a tensor-parallel checkpoint straight off disk instead of materializing the full checkpoint on every rank and resharding it afterwards, and gather the shards back on save. - `from_pretrained(..., parallel_config=TensorParallelConfig(...))` resolves the shard specs on the still-meta model, then slices each safetensors tensor before the dtype cast, so host memory peaks at ~1/tp_degree of the checkpoint. - `save_pretrained` all-gathers the DTensors into an ordinary checkpoint, or writes a distributed checkpoint with `dcp=True` so no full tensor is ever formed. The writing `tp_degree` is recorded, since a packed weight's stored layout is interleaved by it. - Factor the plan interpretation out of the Neuron pre-shard path into shared `TPShardSpec` / `resolve_tp_shard_specs` / `_local_shard` / `_hooks_only_styles` helpers, so both backends and both the load and save paths shard identically. * Raise when tensor parallelism is combined with quantization, offloading or LoRA Addresses the remaining two items of the review on #13718: tensor parallelism was rejected alongside quantization and `device_map` only on the `from_pretrained` streaming path, while `enable_parallelism` — which the quantization error message itself recommended — accepted a quantized, offloaded or adapter-injected model and sharded it anyway. - Add `_check_tp_model_state`, called from `apply_tensor_parallel`, the one chokepoint every TP entry point funnels through. It rejects a model that is quantized, group-offloaded, placed by accelerate (`device_map` or CPU offload), or has PEFT layers injected. Placed before the device-type check so the reported reason is the useful one. - Guard the reverse order too: `enable_group_offload`, the two pipeline CPU-offload methods, and `load_lora_adapter` now refuse a tensor-parallel model. - `save_pretrained` refuses a quantized tensor-parallel model. Previously the `dcp=True` branch returned before the quantizer's serialization step, writing shards with no quantization metadata and no error. - The DCP load guard checked the `quantization_config` kwarg only, so a pre-quantized checkpoint directory loaded silently; check the config's own entry too, and add the missing `_tp_plan` check that otherwise surfaced as a raw `AttributeError`. - Correct the `from_pretrained` message and the doc sentence that pointed at `enable_parallelism` as a way to shard a quantized model. The new tests are the first tensor-parallel tests that need neither an accelerator nor more than one rank: every case asserts a raise before any collective, so they run single-process on gloo. * review: remove tp save/ dcp related * review: add a new helper . * review: moved the TP checks into enable_parallelism * review: move the warning under caution block * review: add loading time numbers * review: move tp size dividende check * review: delete test * review: consolidate tp checks * review: consolidate the tp checks, all in tensor_parallel.py * review: improve conditionals * review: debug info * review: check weights format before instantiating the model * review: restore duplicated _find_mismatched_keys since irrelevant * revert: delete duplicated helper * review: update doc for loading benchmark of flux2-dev * review: improve the raising check * review: add an args section * Test tensor-parallel loading from sharded checkpoints --------- Co-authored-by: Claude Opus 5 (1M context) Co-authored-by: Sayak Paul --- .../en/training/distributed_inference.md | 61 ++- src/diffusers/hooks/tensor_parallel.py | 437 ++++++++++++++---- src/diffusers/hooks/tensor_parallel_neuron.py | 177 ++----- src/diffusers/loaders/peft.py | 10 + src/diffusers/models/_modeling_parallel.py | 2 + src/diffusers/models/model_loading_utils.py | 115 ++++- src/diffusers/models/modeling_utils.py | 269 ++++++++--- src/diffusers/pipelines/pipeline_utils.py | 19 + tests/models/testing_utils/parallelism.py | 120 +++++ 9 files changed, 867 insertions(+), 343 deletions(-) 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/tensor_parallel.py b/src/diffusers/hooks/tensor_parallel.py index b90a5761d043..0f56c3644ec2 100644 --- a/src/diffusers/hooks/tensor_parallel.py +++ b/src/diffusers/hooks/tensor_parallel.py @@ -12,10 +12,12 @@ # 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 @@ -48,6 +50,18 @@ def __init__(self, blocks: "list[int] | None" = None): self.blocks = blocks +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 +79,108 @@ 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": + 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, or PackedRowwiseParallel." + ) + + 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 +223,68 @@ 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. + 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, PackedColwiseParallel): resolved[path] = _make_packed_col(style) elif isinstance(style, PackedRowwiseParallel): @@ -240,12 +297,196 @@ def _partition_linear_fn(self, name, module, device_mesh): 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.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() + else: + raise ValueError( + f"Unsupported tensor-parallel style '{style}' for '{path}'. " + f"Expected 'colwise', 'rowwise', PackedColwiseParallel, or PackedRowwiseParallel." + ) + return resolved + + +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.") @@ -261,13 +502,21 @@ def apply_tensor_parallel( 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 - 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..ffcba0973d81 100644 --- a/src/diffusers/hooks/tensor_parallel_neuron.py +++ b/src/diffusers/hooks/tensor_parallel_neuron.py @@ -17,162 +17,59 @@ 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, _hooks_only_styles, _local_shard 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() + from torch.distributed.tensor import DTensor, Replicate, Shard + from torch.distributed.tensor.parallel import parallelize_module + device = torch.neuron.current_device() + + 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: - _pre_shard_and_tp(block, tp_mesh, relative_plan, rank, tp_size) + parallelize_module(block, tp_mesh, _hooks_only_styles(relative_plan)) 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/models/_modeling_parallel.py b/src/diffusers/models/_modeling_parallel.py index 86627284e078..b54e86d6b4f2 100644 --- a/src/diffusers/models/_modeling_parallel.py +++ b/src/diffusers/models/_modeling_parallel.py @@ -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 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..a6881e971edf 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,14 @@ 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 ..hooks.tensor_parallel import apply_tensor_parallel + + 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 +1694,52 @@ 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 + 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 + return config + def enable_parallelism( self, *, @@ -1617,26 +1750,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 +1811,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 +1831,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 +1854,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 +1872,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,24 +1918,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) diff --git a/src/diffusers/pipelines/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index 82e6c4c2aff4..d7e9d7729fb5 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -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: @@ -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/tests/models/testing_utils/parallelism.py b/tests/models/testing_utils/parallelism.py index 63575abf6b7b..b525637c953e 100644 --- a/tests/models/testing_utils/parallelism.py +++ b/tests/models/testing_utils/parallelism.py @@ -23,6 +23,7 @@ from diffusers.models._modeling_parallel import ContextParallelConfig, 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,74 @@ 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) + @is_context_parallel @require_torch_multi_accelerator From a91ff967f5b3930652461248946052aec128bc36 Mon Sep 17 00:00:00 2001 From: Dhruv Nair Date: Thu, 1 Oct 2026 11:13:15 +0530 Subject: [PATCH 12/22] [Modular] Avoid downloading weights when loading from an existing local path (#14797) * avoid redownload from local path * update * update * update --------- Co-authored-by: Sayak Paul --- .../en/modular_diffusers/modular_pipeline.md | 2 +- .../modular_pipelines/modular_pipeline.py | 51 +++++++++++++- .../pipelines/pipeline_loading_utils.py | 13 ---- src/diffusers/pipelines/pipeline_utils.py | 2 +- src/diffusers/utils/__init__.py | 1 + src/diffusers/utils/constants.py | 11 +++ .../test_modular_pipeline_loading.py | 68 +++++++++++++++++++ 7 files changed, 132 insertions(+), 16 deletions(-) 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/src/diffusers/modular_pipelines/modular_pipeline.py b/src/diffusers/modular_pipelines/modular_pipeline.py index 0cca16877623..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 @@ -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/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 d7e9d7729fb5..959252ff68a2 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, 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/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" From 4de185d6ec51a54ae12a07bba1fccef598e86147 Mon Sep 17 00:00:00 2001 From: Dhruv Nair Date: Thu, 1 Oct 2026 13:14:41 +0530 Subject: [PATCH 13/22] Add single file support for Minimax H3 (#14839) update Co-authored-by: Sayak Paul --- src/diffusers/loaders/single_file_model.py | 5 ++ src/diffusers/loaders/single_file_utils.py | 90 +++++++++++++++++++ .../transformers/transformer_minimax_h3.py | 6 +- .../test_models_transformer_minimax_h3.py | 17 ++++ 4 files changed, 116 insertions(+), 2 deletions(-) diff --git a/src/diffusers/loaders/single_file_model.py b/src/diffusers/loaders/single_file_model.py index cd49ddef69f2..57ac5f71fb19 100644 --- a/src/diffusers/loaders/single_file_model.py +++ b/src/diffusers/loaders/single_file_model.py @@ -50,6 +50,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, @@ -225,6 +226,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..35d75d678087 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", @@ -244,6 +245,7 @@ "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"}, + "minimax-h3": {"pretrained_model_name_or_path": "MiniMaxAI/MiniMax-H3"}, } # Use to configure model sample size when original config is provided @@ -829,6 +831,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 +4234,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 = {} diff --git a/src/diffusers/models/transformers/transformer_minimax_h3.py b/src/diffusers/models/transformers/transformer_minimax_h3.py index f49cdaca2eb6..95ee1f2ac511 100644 --- a/src/diffusers/models/transformers/transformer_minimax_h3.py +++ b/src/diffusers/models/transformers/transformer_minimax_h3.py @@ -19,7 +19,7 @@ import torch.nn as nn from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin +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 @@ -373,7 +373,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. diff --git a/tests/models/transformers/test_models_transformer_minimax_h3.py b/tests/models/transformers/test_models_transformer_minimax_h3.py index 00baa37c84a0..4165f9b5392d 100644 --- a/tests/models/transformers/test_models_transformer_minimax_h3.py +++ b/tests/models/transformers/test_models_transformer_minimax_h3.py @@ -27,6 +27,7 @@ LoraTesterMixin, MemoryTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, ) @@ -189,3 +190,19 @@ 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"} From acbabca385070cbcb7554a91d1a2722dc0265c70 Mon Sep 17 00:00:00 2001 From: Dhruv Nair Date: Thu, 1 Oct 2026 13:41:30 +0530 Subject: [PATCH 14/22] Consolidate torch device backend dispatch (#14792) * introduce TorchDeviceBackend * update --------- Co-authored-by: Sayak Paul --- docs/source/en/api/utilities.md | 4 + src/diffusers/hooks/group_offloading.py | 20 +- src/diffusers/loaders/lora_pipeline.py | 3 +- .../modular_pipelines/components_manager.py | 23 +- .../pipelines/pag/pipeline_pag_sana.py | 10 +- src/diffusers/pipelines/pipeline_utils.py | 2 +- src/diffusers/pipelines/sana/pipeline_sana.py | 10 +- .../sana/pipeline_sana_controlnet.py | 10 +- .../pipelines/sana/pipeline_sana_sprint.py | 10 +- .../sana/pipeline_sana_sprint_img2img.py | 10 +- .../sana_video/pipeline_sana_video.py | 10 +- .../sana_video/pipeline_sana_video_i2v.py | 10 +- .../quantizers/gguf/gguf_quantizer.py | 8 +- src/diffusers/training_utils.py | 11 +- src/diffusers/utils/torch_utils.py | 277 ++++++++---------- tests/hooks/test_group_offloading.py | 14 +- .../modular_pipelines/testing_utils/utils.py | 14 +- tests/others/test_utils.py | 63 +++- tests/testing_utils.py | 97 +----- 19 files changed, 252 insertions(+), 354 deletions(-) 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/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/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/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/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/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index 959252ff68a2..750fed48cbd0 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -1288,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""" 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/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/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/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 From 7997e4eabb1f880fa5697e2b24de09a892fffad5 Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Thu, 1 Oct 2026 13:51:28 +0530 Subject: [PATCH 15/22] chore: add one additional copied from in qwenimage 2.1 vae (#14810) * add one additional copied from in qwenimage 2.1 vae. * formattinmg --- .../autoencoders/autoencoder_kl_qwenimage21.py | 2 +- .../models/autoencoders/autoencoder_kl_wan.py | 12 ++++++++---- 2 files changed, 9 insertions(+), 5 deletions(-) diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py b/src/diffusers/models/autoencoders/autoencoder_kl_qwenimage21.py index 3d99d7a69bf1..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 diff --git a/src/diffusers/models/autoencoders/autoencoder_kl_wan.py b/src/diffusers/models/autoencoders/autoencoder_kl_wan.py index 29bd5738a62b..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 From 578c9b2c6636ab2424a0e56186268b83623656b2 Mon Sep 17 00:00:00 2001 From: Dhruv Nair Date: Thu, 1 Oct 2026 16:37:20 +0530 Subject: [PATCH 16/22] Add single file support for Krea 2 (#14914) update --- docs/source/en/api/pipelines/krea2.md | 11 ++++ src/diffusers/loaders/single_file_model.py | 5 ++ src/diffusers/loaders/single_file_utils.py | 53 +++++++++++++++++++ .../models/transformers/transformer_krea2.py | 4 +- .../test_models_transformer_krea2.py | 19 +++++++ 5 files changed, 90 insertions(+), 2 deletions(-) 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/src/diffusers/loaders/single_file_model.py b/src/diffusers/loaders/single_file_model.py index 57ac5f71fb19..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, @@ -199,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", diff --git a/src/diffusers/loaders/single_file_utils.py b/src/diffusers/loaders/single_file_utils.py index 35d75d678087..3e160be02b4b 100644 --- a/src/diffusers/loaders/single_file_utils.py +++ b/src/diffusers/loaders/single_file_utils.py @@ -160,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 = { @@ -245,6 +246,7 @@ "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"}, } @@ -791,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" @@ -4335,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/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/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 From fe7515a45ee736b433cee53bbab26ba0b322ff8c Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Fri, 21 Aug 2026 11:08:58 +0000 Subject: [PATCH 17/22] Add tensor-parallel support for MiniMax-H3 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Shards `MiniMaxH3Transformer3DModel` across devices, following the plan already established for Flux1/Flux2/Qwen-Image. Validated on Trainium at TP=2 and TP=8. - `_tp_plan` with twelve entries: the same six shapes for the 50 denoiser blocks and for the two token-refiner blocks, which are the same attention + SwiGLU FFN minus AdaLN and rotary. Q/K/V and the attention output are unfused, so they are plain colwise/rowwise; the SwiGLU input `ff.net.0.proj` is one Linear producing `[value; gate]` in equal halves and takes PackedColwiseParallel([1, 1]). - The attention processor reshaped by the config head count, `unflatten(-1, (attn.heads, -1))`, which mis-splits under sharding: each rank holds `inner_dim / tp_degree` columns, so this yields `head_dim / tp_degree` per head instead of `heads / tp_degree` heads. Reshape by the fixed `attn.head_dim` instead and let `-1` absorb the head count, as Flux does. Numerically identical unsharded, since `inner_dim == heads * head_dim`. - Norms, QK-norms (head_dim-shaped, applied after the head split), AdaLN modulation and the patch/text embedders and output heads stay replicated. `attn.to_qkv` is deliberately not in the plan: it exists only after `fuse_projections()`, and the plan is resolved by attribute lookup. No RoPE change was needed — unlike Qwen-Image, H3's rotary is already real sin/cos and already broadcasts over the head axis. Tests mirror the Flux2/Qwen-Image layout: the CUDA/XPU `TensorParallelTesterMixin` class, a `make_neuron_tp_spec()` factory, and a Neuron launcher that shells out to the model-agnostic `_neuron_tp_worker.py`. `get_dummy_inputs` and `get_packed_layout` take an optional `device` so the Neuron spec can ask for CPU tensors, since its worker shards on CPU and moves to device after. --- .../transformers/transformer_minimax_h3.py | 37 ++++++++- .../test_models_transformer_minimax_h3.py | 83 +++++++++++++++---- 2 files changed, 100 insertions(+), 20 deletions(-) diff --git a/src/diffusers/models/transformers/transformer_minimax_h3.py b/src/diffusers/models/transformers/transformer_minimax_h3.py index 95ee1f2ac511..9d816aed2a15 100644 --- a/src/diffusers/models/transformers/transformer_minimax_h3.py +++ b/src/diffusers/models/transformers/transformer_minimax_h3.py @@ -19,6 +19,7 @@ import torch.nn as nn from ...configuration_utils import ConfigMixin, register_to_config +from ...hooks.tensor_parallel import PackedColwiseParallel from ...loaders import FromOriginalModelMixin, PeftAdapterMixin from ...utils import BaseOutput, apply_lora_scale, logging from .._modeling_parallel import ContextParallelInput, ContextParallelOutput @@ -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) @@ -451,6 +454,34 @@ class MiniMaxH3Transformer3DModel( "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. + # + # 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; the AdaLN modulation (`adaln_proj.linear`, `norm_out.linear`), which indexes the full hidden + # dim; 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", + # 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/tests/models/transformers/test_models_transformer_minimax_h3.py b/tests/models/transformers/test_models_transformer_minimax_h3.py index 4165f9b5392d..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, @@ -28,6 +32,7 @@ MemoryTesterMixin, ModelTesterMixin, SingleFileTesterMixin, + TensorParallelTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, ) @@ -85,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. @@ -93,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, @@ -124,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"] @@ -132,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), } @@ -206,3 +218,40 @@ def pretrained_model_name_or_path(self): @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}" + ) From b618f76bc0dc4f9fa66e6b442850fab474b27cc4 Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Fri, 21 Aug 2026 13:45:24 +0000 Subject: [PATCH 18/22] Shard MiniMax-H3's adaln_proj to fit tensor parallelism on one device MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `transformer_blocks.*.adaln_proj.linear` was left replicated on every rank, and at `[96768, 2688]` bf16 per block it is 24.23 GiB of the denoiser's 61.73 GiB — about 40%. That made the per-rank floor 24.40 GiB of weights (plus 5.13 GiB for the two VAEs) regardless of TP degree, so MiniMax-H3 could not fit a 24 GiB NeuronCore at *any* valid TP: 34.20 GiB/rank at TP=8, and still 30.20 GiB/rank at TP=56. Raising TP only divided the 60% that already sharded. (TP=16 is not an option either — 56 attention heads.) Shard it rowwise, over the `time_embed_dim` input, rather than colwise: the six modulation parameters scale and shift the *full* hidden dim of a sequence that is already all-reduced by the time they are applied, so a colwise split would need an all-gather to rebuild that width. Rowwise keeps the output full-width, leaving the module's `view`/`chunk` untouched, and all-reduces a few hundred KB per block per step. Plain `"rowwise"` could not be reused. It is normally the second half of a colwise/rowwise pair, so it defaults to `input_layouts=Shard(-1)` and would read the full-width `temb` as if it were one rank's shard. Hence `ReplicatedInputRowwiseParallel`: input narrowed locally on the way in (no collective), partial output all-reduced on the way out, bias replicated and added after the reduce. It is wired into `_styles`, `_hooks_only_styles` — the path the Neuron backend takes, since `_apply_tp_neuron` pre-shards on CPU and then registers hooks only — and `resolve_tp_shard_specs`. Replicated weights drop from 24.40 GiB to 0.15 GiB, putting TP=8 at 7.84 GiB of transformer plus 5.13 GiB of VAEs, i.e. 12.97 GiB/rank against a 24 GiB budget. Verified on CPU/gloo that both the generic and the pre-sharded hooks-only path shard the weight on its input dim and match a replicated reference to 3.6e-7. Co-Authored-By: Claude Opus 5 (1M context) --- src/diffusers/hooks/tensor_parallel.py | 40 +++++++++++++++---- .../transformers/transformer_minimax_h3.py | 17 ++++++-- 2 files changed, 47 insertions(+), 10 deletions(-) diff --git a/src/diffusers/hooks/tensor_parallel.py b/src/diffusers/hooks/tensor_parallel.py index 0f56c3644ec2..9f74fa2ac790 100644 --- a/src/diffusers/hooks/tensor_parallel.py +++ b/src/diffusers/hooks/tensor_parallel.py @@ -50,6 +50,20 @@ 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" @@ -149,7 +163,8 @@ def resolve_tp_shard_specs(model: torch.nn.Module, tp_plan: dict, tp_degree: int if style == "colwise": weight_spec = TPShardSpec(0, [submodule.weight.shape[0]]) bias_spec = weight_spec - elif style == "rowwise": + 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): @@ -163,7 +178,8 @@ def resolve_tp_shard_specs(model: torch.nn.Module, tp_plan: dict, tp_degree: int 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." ) specs[f"{path}.weight"] = weight_spec @@ -242,9 +258,9 @@ def _shard_packed_param(param, dim: int, blocks: "list[int]", device_mesh, src_d 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() | }`. Divisibility by the TP - degree is checked earlier, by `resolve_tp_shard_specs`. + 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`. """ from torch.distributed.tensor import Replicate, distribute_tensor from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel @@ -285,6 +301,11 @@ def _partition_linear_fn(self, name, module, device_mesh): resolved[path] = ColwiseParallel() elif style == "rowwise": 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): @@ -292,7 +313,8 @@ 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 @@ -308,6 +330,7 @@ def _hooks_only_styles(relative_plan: dict) -> dict: 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): @@ -324,10 +347,13 @@ def _partition_linear_fn(self, name, module, device_mesh): 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, or PackedRowwiseParallel." + f"Expected 'colwise', 'rowwise', PackedColwiseParallel, PackedRowwiseParallel, or " + f"ReplicatedInputRowwiseParallel." ) return resolved diff --git a/src/diffusers/models/transformers/transformer_minimax_h3.py b/src/diffusers/models/transformers/transformer_minimax_h3.py index 9d816aed2a15..0f9e38bdb965 100644 --- a/src/diffusers/models/transformers/transformer_minimax_h3.py +++ b/src/diffusers/models/transformers/transformer_minimax_h3.py @@ -19,7 +19,7 @@ import torch.nn as nn from ...configuration_utils import ConfigMixin, register_to_config -from ...hooks.tensor_parallel import PackedColwiseParallel +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 @@ -459,10 +459,20 @@ class MiniMaxH3Transformer3DModel( # 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; the AdaLN modulation (`adaln_proj.linear`, `norm_out.linear`), which indexes the full hidden - # dim; and the patch/text embedders and the two output heads. + # 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. @@ -474,6 +484,7 @@ class MiniMaxH3Transformer3DModel( "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", From 5a115d36f8294471019ee195399af3694309c8e8 Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Fri, 21 Aug 2026 14:32:06 +0000 Subject: [PATCH 19/22] Build MiniMax-H3's row timestep plan on CPU, as its caller expects MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `build_row_timesteps` allocates `row_timesteps` with `torch.full` and no device, then scatters into it at `video_indices` / `audio_indices`. The layout step hands those index tensors over already on the execution device, and indexing a CPU tensor with an accelerator one is an error — on Neuron it surfaces as "Non-scalar tensor arg0 is on cpu device, expected neuron", and on CUDA it would raise "indices should be either on cpu or on the same device". CPU is the right place for this to run, not the accelerator: `torch.unique` has a data-dependent output shape, which is precisely what a tracing backend cannot handle, and the caller already moves the finished `(timestep, timestep_indices)` pair to the device itself. So bring the two index tensors back to CPU for the scatter rather than allocating `row_timesteps` on their device. Only reachable once the denoiser is actually on an accelerator while the pipeline's execution device resolves there too, which is why it went unnoticed. Co-Authored-By: Claude Opus 5 (1M context) --- .../modular_pipelines/minimax_h3/before_denoise.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py b/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py index 94395a7c3ebe..bb1c1cd54b4f 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py +++ b/src/diffusers/modular_pipelines/minimax_h3/before_denoise.py @@ -1228,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 From b828d9daaa178e1290725891d7aa3563e891b948 Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Thu, 27 Aug 2026 13:20:19 +0000 Subject: [PATCH 20/22] feat:combine tp+cp (cherry picked from commit a489f03e4e9888c0abb996f45d38885f03772ae4) --- src/diffusers/models/_modeling_parallel.py | 41 ++++- src/diffusers/models/modeling_utils.py | 49 +++++- tests/models/test_parallelism_combined.py | 164 ++++++++++++++++++ tests/models/testing_utils/__init__.py | 2 + tests/models/testing_utils/parallelism.py | 146 +++++++++++++++- .../test_models_transformer_qwenimage.py | 20 +++ 6 files changed, 409 insertions(+), 13 deletions(-) create mode 100644 tests/models/test_parallelism_combined.py diff --git a/src/diffusers/models/_modeling_parallel.py b/src/diffusers/models/_modeling_parallel.py index b54e86d6b4f2..7f12ec71df28 100644 --- a/src/diffusers/models/_modeling_parallel.py +++ b/src/diffusers/models/_modeling_parallel.py @@ -200,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. @@ -215,12 +220,10 @@ class ParallelConfig: _device: torch.device = None _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: - 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`." - ) + @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, @@ -234,10 +237,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/modeling_utils.py b/src/diffusers/models/modeling_utils.py index a6881e971edf..268226d918af 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -1720,15 +1720,30 @@ def _resolve_parallel_config( device = torch.device(device_type, rank % device_module.device_count()) mesh = None - if config.context_parallel_config is not None: - cp_config = config.context_parallel_config + 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 config.tensor_parallel_config is not None: - tp_config = config.tensor_parallel_config + 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,), @@ -1737,6 +1752,32 @@ def _resolve_parallel_config( # `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 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..c1d32fc89844 100644 --- a/tests/models/testing_utils/__init__.py +++ b/tests/models/testing_utils/__init__.py @@ -20,6 +20,7 @@ from .lora import LoraHotSwappingForModelTesterMixin, LoraTesterMixin from .memory import CPUOffloadTesterMixin, GroupOffloadTesterMixin, LayerwiseCastingTesterMixin, MemoryTesterMixin from .parallelism import ( + ContextAndTensorParallelTesterMixin, ContextParallelAttentionBackendsTesterMixin, ContextParallelTesterMixin, TensorParallelTesterMixin, @@ -64,6 +65,7 @@ "BitsAndBytesConfigMixin", "BitsAndBytesTesterMixin", "CacheTesterMixin", + "ContextAndTensorParallelTesterMixin", "ContextParallelTesterMixin", "ContextParallelAttentionBackendsTesterMixin", "TensorParallelTesterMixin", diff --git a/tests/models/testing_utils/parallelism.py b/tests/models/testing_utils/parallelism.py index b525637c953e..5a204526c0c4 100644 --- a/tests/models/testing_utils/parallelism.py +++ b/tests/models/testing_utils/parallelism.py @@ -21,7 +21,7 @@ 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 @@ -465,6 +465,150 @@ def test_tensor_parallel_from_pretrained(self, tmp_path, sharded): 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 class ContextParallelTesterMixin: 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``). From 71a668e396c32832264427a1ba5cfbb338aaac6f Mon Sep 17 00:00:00 2001 From: whn09 Date: Mon, 7 Sep 2026 03:08:20 +0000 Subject: [PATCH 21/22] Allow tensor parallelism and context parallelism in one ParallelConfig `ParallelConfig.__post_init__` rejects a config carrying both a `context_parallel_config` and a `tensor_parallel_config`, so a model can only use one form of parallelism at a time. That caps some models well below the hardware they run on: MiniMax-H3 has 56 attention heads, so `tp_degree` cannot exceed 8, and on a 64-accelerator host tensor parallelism alone leaves 56 of them idle even though the model already ships a `_cp_plan`. Both configs already accept a caller-supplied `mesh`, documented as being for "combining CP with other parallelism strategies that share the same mesh", so the intent was there; what was missing was building that shared mesh and splitting it between the two. * `__post_init__` now only rejects a config that carries neither. * `ParallelConfig.setup` hands each config its own dimensions: CP keeps the whole mesh, since its own `setup` selects "ring" and "ulysses" out of it, while TP gets the "tp" dimension alone because it reads its degree from `mesh.size()`. * `enable_parallelism` builds one ("tp", "ring", "ulysses") mesh when both are requested, with the context-parallel dimensions varying fastest. That order is a default rather than a universal optimum, and the comment says so: pass `mesh=` on either config to choose the layout yourself. * The Neuron tensor-parallel backend used the global rank as its shard index. That equals the rank's coordinate on the TP mesh only while that mesh spans the whole world, so sharing a mesh made it index past the end of the weight; the generic backend beside it already reads the coordinate. Measured on a trn2.48xlarge (MiniMax-H3, 1344x768x124f, 30 steps): TP=8 alone is 9.285 s/step on 8 cores, TP=4 x ulysses=4 is 5.66 s/step on 16, and TP=8 x ulysses=4 is 3.941 s/step on 32 -- 2.36x faster than the widest configuration reachable before this change, with output that is visually equivalent to the tensor-parallel-only run. Tests: a `HybridParallelTesterMixin` next to the existing TP and CP mixins, wired into the Flux transformer tests, plus a Neuron `torchrun` worker following the `_neuron_tp_worker.py` convention already in the tree. The Neuron test runs at tp_degree=2 x ulysses_degree=4 and passes on a trn2.48xlarge (max_abs_diff 4.2e-05 against the single-device reference). Co-Authored-By: Claude Opus 5 --- src/diffusers/models/_modeling_parallel.py | 7 + src/diffusers/models/modeling_utils.py | 1 - tests/models/testing_utils/__init__.py | 2 + tests/models/testing_utils/parallelism.py | 129 ++++++++++++++++++ .../transformers/_neuron_hybrid_worker.py | 125 +++++++++++++++++ .../test_models_transformer_flux.py | 48 ++++++- 6 files changed, 310 insertions(+), 2 deletions(-) create mode 100644 tests/models/transformers/_neuron_hybrid_worker.py diff --git a/src/diffusers/models/_modeling_parallel.py b/src/diffusers/models/_modeling_parallel.py index 7f12ec71df28..0ce2a005aedd 100644 --- a/src/diffusers/models/_modeling_parallel.py +++ b/src/diffusers/models/_modeling_parallel.py @@ -220,6 +220,13 @@ class ParallelConfig: _device: torch.device = None _mesh: torch.distributed.device_mesh.DeviceMesh = None + def __post_init__(self): + if self.context_parallel_config is None and self.tensor_parallel_config is None: + raise ValueError( + "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.""" diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index 268226d918af..a1477589944e 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -1990,7 +1990,6 @@ def _load_pretrained_model( 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/tests/models/testing_utils/__init__.py b/tests/models/testing_utils/__init__.py index c1d32fc89844..965a591f95a5 100644 --- a/tests/models/testing_utils/__init__.py +++ b/tests/models/testing_utils/__init__.py @@ -23,6 +23,7 @@ ContextAndTensorParallelTesterMixin, ContextParallelAttentionBackendsTesterMixin, ContextParallelTesterMixin, + HybridParallelTesterMixin, TensorParallelTesterMixin, ) from .quantization import ( @@ -68,6 +69,7 @@ "ContextAndTensorParallelTesterMixin", "ContextParallelTesterMixin", "ContextParallelAttentionBackendsTesterMixin", + "HybridParallelTesterMixin", "TensorParallelTesterMixin", "CPUOffloadTesterMixin", "FasterCacheConfigMixin", diff --git a/tests/models/testing_utils/parallelism.py b/tests/models/testing_utils/parallelism.py index 5a204526c0c4..4c47ae790968 100644 --- a/tests/models/testing_utils/parallelism.py +++ b/tests/models/testing_utils/parallelism.py @@ -874,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_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.""" From 7be6a981f20948aef6800816dad5970ccbffc6d4 Mon Sep 17 00:00:00 2001 From: Jingya HUANG Date: Wed, 30 Sep 2026 14:01:57 +0000 Subject: [PATCH 22/22] feat: validate the support on TPU --- src/diffusers/hooks/tensor_parallel.py | 51 ++++++++++++++++++- src/diffusers/hooks/tensor_parallel_neuron.py | 31 +---------- src/diffusers/models/_modeling_parallel.py | 2 +- src/diffusers/models/modeling_utils.py | 13 +++++ 4 files changed, 65 insertions(+), 32 deletions(-) diff --git a/src/diffusers/hooks/tensor_parallel.py b/src/diffusers/hooks/tensor_parallel.py index 9f74fa2ac790..ce8abd844ab5 100644 --- a/src/diffusers/hooks/tensor_parallel.py +++ b/src/diffusers/hooks/tensor_parallel.py @@ -22,7 +22,7 @@ logger = get_logger(__name__) # pylint: disable=invalid-name -_SUPPORTED_TP_DEVICES = ("cuda", "neuron") +_SUPPORTED_TP_DEVICES = ("cuda", "neuron", "tpu") class PackedColwiseParallel: @@ -358,6 +358,47 @@ def _partition_linear_fn(self, name, module, device_mesh): 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 @@ -524,7 +565,7 @@ 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}.") @@ -544,5 +585,11 @@ def apply_tensor_parallel( _apply_tp_neuron(model, tp_mesh, groups, specs) return + 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 ffcba0973d81..a87f28b1edb1 100644 --- a/src/diffusers/hooks/tensor_parallel_neuron.py +++ b/src/diffusers/hooks/tensor_parallel_neuron.py @@ -28,7 +28,7 @@ import torch import torch.nn as nn -from .tensor_parallel import TPShardSpec, _hooks_only_styles, _local_shard +from .tensor_parallel import TPShardSpec, _pre_shard_and_parallelize def _apply_tp_neuron( @@ -45,31 +45,4 @@ def _apply_tp_neuron( Model weights must be on CPU when this is called. """ - from torch.distributed.tensor import DTensor, Replicate, Shard - from torch.distributed.tensor.parallel import parallelize_module - - device = torch.neuron.current_device() - - 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)) + _pre_shard_and_parallelize(model, tp_mesh, groups, specs, torch.neuron.current_device()) diff --git a/src/diffusers/models/_modeling_parallel.py b/src/diffusers/models/_modeling_parallel.py index 0ce2a005aedd..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`): diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index a1477589944e..262e5a8ae699 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -1544,8 +1544,21 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None 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)