Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
29 commits
Select commit Hold shift + click to select a range
acdf4bf
[core] Shard tensor-parallel checkpoints on load and save
JingyaHuang Aug 20, 2026
0b9686b
Merge branch 'main' into add-shard-ckpt-loading
JingyaHuang Aug 20, 2026
40ddb53
Raise when tensor parallelism is combined with quantization, offloadi…
JingyaHuang Aug 20, 2026
0760934
Merge branch 'add-shard-ckpt-loading' of github.com:JingyaHuang/diffu…
JingyaHuang Aug 21, 2026
a8956f4
Add tensor-parallel support for MiniMax-H3
JingyaHuang Aug 21, 2026
7e6e38f
Shard MiniMax-H3's adaln_proj to fit tensor parallelism on one device
JingyaHuang Aug 21, 2026
bb236dc
Build MiniMax-H3's row timestep plan on CPU, as its caller expects
JingyaHuang Aug 21, 2026
aa50ad3
Merge branch 'main' into add-h3-tp-support
JingyaHuang Aug 26, 2026
0f9057b
feat:combine tp+cp
JingyaHuang Aug 27, 2026
9c81ac6
Allow tensor parallelism and context parallelism in one ParallelConfig
whn09 Sep 7, 2026
2e6afea
Merge branch 'main' into tp-cp-compose
JingyaHuang Sep 16, 2026
1b6b17c
Merge remote-tracking branch 'refs/remotes/whn09/tp-cp-compose' into …
JingyaHuang Sep 16, 2026
76acd09
feat: validate the support on TPU
JingyaHuang Sep 30, 2026
a677eec
Merge branch 'main' into add-h3-tp-support
JingyaHuang Sep 30, 2026
78405ec
[Cosmos3] Fix Transfer SeaCache artifacts with control CFG (#14897)
yzhautouskay Sep 30, 2026
c60830e
[fix] Add return types (#14874)
stevhliu Sep 30, 2026
863092f
[core] Shard tensor-parallel checkpoints on load and save (#14544)
JingyaHuang Oct 1, 2026
a91ff96
[Modular] Avoid downloading weights when loading from an existing loc…
DN6 Oct 1, 2026
4de185d
Add single file support for Minimax H3 (#14839)
DN6 Oct 1, 2026
acbabca
Consolidate torch device backend dispatch (#14792)
DN6 Oct 1, 2026
7997e4e
chore: add one additional copied from in qwenimage 2.1 vae (#14810)
sayakpaul Oct 1, 2026
578c9b2
Add single file support for Krea 2 (#14914)
DN6 Oct 1, 2026
fe7515a
Add tensor-parallel support for MiniMax-H3
JingyaHuang Aug 21, 2026
b618f76
Shard MiniMax-H3's adaln_proj to fit tensor parallelism on one device
JingyaHuang Aug 21, 2026
5a115d3
Build MiniMax-H3's row timestep plan on CPU, as its caller expects
JingyaHuang Aug 21, 2026
b828d9d
feat:combine tp+cp
JingyaHuang Aug 27, 2026
71a668e
Allow tensor parallelism and context parallelism in one ParallelConfig
whn09 Sep 7, 2026
7be6a98
feat: validate the support on TPU
JingyaHuang Sep 30, 2026
7e4a7ef
Merge branch 'add-h3-tp-support' of github.com:JingyaHuang/diffusers …
JingyaHuang Oct 2, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
The table of contents is too big for display.
Diff view
Diff view
  •  
  •  
  •  
1 change: 1 addition & 0 deletions .github/workflows/pr_modular_tests.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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() }}
Expand Down
1 change: 1 addition & 0 deletions .github/workflows/pr_tests.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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() }}
Expand Down
1 change: 1 addition & 0 deletions .github/workflows/pr_tests_gpu.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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() }}
Expand Down
5 changes: 5 additions & 0 deletions Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand Down
11 changes: 11 additions & 0 deletions docs/source/en/api/pipelines/krea2.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
4 changes: 4 additions & 0 deletions docs/source/en/api/utilities.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion docs/source/en/modular_diffusers/modular_pipeline.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:

Expand Down
15 changes: 11 additions & 4 deletions docs/source/en/optimization/cache.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:

Expand All @@ -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

Expand Down
61 changes: 39 additions & 22 deletions docs/source/en/training/distributed_inference.md
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand All @@ -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.
Expand Down Expand Up @@ -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]
Expand Down
20 changes: 9 additions & 11 deletions src/diffusers/hooks/group_offloading.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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`.")
Expand Down
14 changes: 8 additions & 6 deletions src/diffusers/hooks/sea_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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(
Expand Down
Loading
Loading