Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
74 commits
Select commit Hold shift + click to select a range
38007b4
feat: add torchtpu
JingyaHuang May 29, 2026
8ed15f9
feat:draft TorchTPU support
JingyaHuang Jun 2, 2026
e343c0a
fix: wan overflow issue + compile mode error on sdxl
JingyaHuang Jun 9, 2026
84b4049
doc: enhance with TorchTPU doc
JingyaHuang Jun 25, 2026
bb3ec1e
doc: enhance with TorchTPU doc
JingyaHuang Jun 25, 2026
339be41
Merge branch 'main' into add-torchtpu-support
JingyaHuang Jun 25, 2026
92193a7
Merge branch 'main' into add-torchtpu-support
JingyaHuang Jul 17, 2026
1a2369b
Merge branch 'main' into add-torchtpu-support
JingyaHuang Jul 17, 2026
66254b4
style: remove unused imports flagged by ruff
JingyaHuang Jul 17, 2026
7c241be
docs: remove Debug Eager and Fused Eager sections from tpu.md
JingyaHuang Jul 17, 2026
05d56d5
Merge branch 'main' into add-torchtpu-support
JingyaHuang Jul 27, 2026
b7a8c0b
Merge branch 'main' into add-torchtpu-support
JingyaHuang Sep 2, 2026
b0c5595
Merge branch 'huggingface:main' into add-torchtpu-support
JingyaHuang Sep 3, 2026
7c3df6c
Merge branch 'main' into add-torchtpu-support
JingyaHuang Sep 7, 2026
7c1e379
feat: add TP for TPU
JingyaHuang Jun 26, 2026
f944c48
style: fix import sorting in TPU test scripts
JingyaHuang Jul 17, 2026
82d20ab
style: ruff format TPU test scripts
JingyaHuang Jul 17, 2026
f6c17d2
fix: style
Sep 7, 2026
e8fe48c
fix: test for native 4 devices
JingyaHuang Sep 7, 2026
9b62254
fix: propagate TPU device fixes to Flux/Flux2/Wan-family copies; regi…
JingyaHuang Sep 8, 2026
744fa95
tests: cleanup
JingyaHuang Sep 8, 2026
bfc7d17
fix: fix compile mode
JingyaHuang Sep 8, 2026
10d4ad1
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 9, 2026
a999a3d
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 9, 2026
cb2dcf3
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 9, 2026
9f11d1e
Merge branch 'main' into add-torchtpu-support
JingyaHuang Sep 9, 2026
c48d41c
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 9, 2026
aa70a45
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 9, 2026
6cb37d2
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 9, 2026
3e06f71
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 9, 2026
09d5ce2
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 9, 2026
7b46e08
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 9, 2026
e82e795
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 9, 2026
d2bf529
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 9, 2026
22d595a
test: remove flux2 e2e test
JingyaHuang Sep 9, 2026
042a88a
doc: apply suggestions
JingyaHuang Sep 9, 2026
c6a4c80
doc: apply suggestions
JingyaHuang Sep 9, 2026
c82fce7
Merge branch 'add-torchtpu-support' of github.com:JingyaHuang/diffuse…
JingyaHuang Sep 9, 2026
cb314bd
review: remove monkey patch
JingyaHuang Sep 10, 2026
bf8934f
review: revert neuron-specific changes in the tests
JingyaHuang Sep 10, 2026
c35e23a
review: apply suggestions
JingyaHuang Sep 11, 2026
419d657
Merge branch 'main' of https://github.com/huggingface/diffusers into …
JingyaHuang Sep 11, 2026
a2a770a
Update docs/source/en/optimization/tpu.md
JingyaHuang Sep 19, 2026
7f73095
Merge branch 'main' into add-torchtpu-support
JingyaHuang Sep 19, 2026
406ac74
doc: add tpu to tp doc
JingyaHuang Sep 19, 2026
67a84d2
Merge branch 'add-torchtpu-support' of github.com:JingyaHuang/diffuse…
JingyaHuang Sep 19, 2026
a638798
Merge branch 'main' into add-torchtpu-support
JingyaHuang Sep 21, 2026
2bf3619
review: remove unnecessary for tpu
JingyaHuang Sep 21, 2026
5d6ad22
Merge branch 'main' into add-torchtpu-support
JingyaHuang Sep 21, 2026
eb49a61
Merge branch 'main' into add-torchtpu-support
JingyaHuang Oct 1, 2026
3152988
test: assert every _tp_plan parameter is sharded after a TP load
JingyaHuang Oct 1, 2026
7257a79
removal: delete redundant tp shard helpers
JingyaHuang Oct 1, 2026
fd0f9a5
removal: drop TPU workarounds no longer needed after recent main changes
JingyaHuang Oct 1, 2026
6ee32b4
doc: keep the FLUX.2-dev text encoder off a single TPU chip in the TP…
JingyaHuang Oct 2, 2026
2d38ce7
doc: encode the prompt on CPU in the TPU TP example so it runs end to…
JingyaHuang Oct 2, 2026
1eddd6f
test: tighten TPU TP tolerance and simplify the TPU TP worker
JingyaHuang Oct 2, 2026
1351263
doc: use enable_model_cpu_offload() in the TPU eager example
JingyaHuang Oct 5, 2026
a170a50
doc: shard both FLUX.2-dev text encoder and transformer in the TPU TP…
JingyaHuang Oct 5, 2026
5c9e426
doc: run the text encoders on TPU in the compiled example
JingyaHuang Oct 5, 2026
a3f9ab1
test: drop the redundant TPU sync in the TP worker
JingyaHuang Oct 6, 2026
e7e6c49
test: trim comments in the TPU TP tests
JingyaHuang Oct 6, 2026
5c262c7
test: run the TPU TP test with mp.spawn like the CUDA one
JingyaHuang Oct 6, 2026
ea2eaca
test: pin the TPU TP test to 4 chips
JingyaHuang Oct 6, 2026
866e165
Merge branch 'main' into add-torchtpu-support
JingyaHuang Oct 6, 2026
6c5c580
Merge branch 'main' into add-torchtpu-support
sayakpaul Oct 7, 2026
2de295b
Update docs/source/en/optimization/tpu.md
JingyaHuang Oct 7, 2026
69d07a3
doc: use DistributedConfig and dtype, and drop the torch_tpu import i…
JingyaHuang Oct 7, 2026
a9412fe
test: skip the TPU TP test for models without a _tp_plan
JingyaHuang Oct 7, 2026
6c89944
doc: correct the warning of recompilation
JingyaHuang Oct 7, 2026
91bba1a
fix: clone the SD3 timestep so torch.compile doesn't recompile every …
JingyaHuang Oct 7, 2026
935d2d1
doc: update advice for compile mode when one chip doesn't fit
JingyaHuang Oct 7, 2026
fb27377
test: pass the TPU TP tolerances as test arguments
JingyaHuang Oct 7, 2026
04a043a
Merge branch 'add-torchtpu-support' of https://github.com/JingyaHuang…
JingyaHuang Oct 7, 2026
b3b529a
Merge branch 'main' into add-torchtpu-support
JingyaHuang Oct 7, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions docs/source/en/_toctree.yml
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,8 @@
title: Intel Gaudi
- local: optimization/neuron
title: AWS Neuron
- local: optimization/tpu
title: TPU
title: Hardware-specific acceleration
- isExpanded: false
sections:
Expand Down
151 changes: 151 additions & 0 deletions docs/source/en/optimization/tpu.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,151 @@
<!--Copyright 2026 The HuggingFace Team. All rights reserved.

Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except in compliance with
the License. You may obtain a copy of the License at

http://www.apache.org/licenses/LICENSE-2.0

Unless required by applicable law or agreed to in writing, software distributed under the License is distributed on
an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the License for the
specific language governing permissions and limitations under the License.
-->

# TorchTPU
Comment thread
JingyaHuang marked this conversation as resolved.

[TorchTPU](https://github.com/google-pytorch/torch_tpu/) is a PyTorch backend for Google's Tensor Processing Units (TPUs), which lets you run Diffusers pipelines on Cloud TPUs (v6e, v5p, etc.) with minimal code changes.

Two execution modes are available:

| Mode | Constant | How to activate | Notes |
|---|---|---|---|
| Strict eager (default) | `EagerMode.DEFER_NEVER` | `pipe.to("tpu")` | Operations dispatched one at a time, asynchronous |
| Compile | — | `torch.compile(module, backend="tpu")` | AOT compilation with `TpuBackend` |
Comment thread
JingyaHuang marked this conversation as resolved.

Follow the [TorchTPU installation guide](https://github.com/google-pytorch/torch_tpu/). Once installed, `import torch`
loads it automatically and registers the `"tpu"` device, so `pipe.to("tpu")` is the only change needed. Add
`import torch_tpu` only if you disabled backend autoloading with `TORCH_DEVICE_BACKEND_AUTOLOAD=0`.

## Eager mode
Comment thread
JingyaHuang marked this conversation as resolved.

FLUX.1-schnell doesn't fit on a single v6e chip all at once, so use [`~DiffusionPipeline.enable_model_cpu_offload`] to
move each model to the TPU only while it runs. It detects the `"tpu"` device automatically.

```python
import torch

from diffusers import FluxPipeline

pipe = FluxPipeline.from_pretrained("black-forest-labs/FLUX.1-schnell", dtype=torch.bfloat16)
pipe.enable_model_cpu_offload()

image = pipe(
prompt="a golden retriever surfing a wave, photorealistic",
height=1024,
width=1024,
num_inference_steps=4,
guidance_scale=0.0,
).images[0]

image.save("output.png")
```

If a model is too large for a single chip, or you have several chips and want lower latency, shard the models
across chips instead. See the [Tensor parallelism](#tensor-parallelism) section.

## Compiled mode

TorchTPU registers `"tpu"` as a `torch.compile` backend name (`TpuBackend` under the hood), so
components compile like any other `torch.compile` target. The first
call (warmup) is slow because it compiles; later calls with the same shapes reuse the compiled graph.

> [!IMPORTANT]
> TorchTPU requires **static shapes**, so pass `dynamic=False`. A new `height` or `width` compiles again for that
> shape, once; shapes already seen are reused. Changing `num_inference_steps` doesn't recompile.

When the whole pipeline fits on one chip, move it to the TPU and compile the full transformer. Stable Diffusion 3.5
Medium (~15GB in bf16) fits on a single v6e chip.

```python
import torch

from diffusers import StableDiffusion3Pipeline

pipe = StableDiffusion3Pipeline.from_pretrained("stabilityai/stable-diffusion-3.5-medium", dtype=torch.bfloat16)
pipe.to("tpu")
pipe.transformer.compile(backend="tpu", fullgraph=True, dynamic=False)

# Warmup — triggers static graph compilation.
pipe(prompt="warmup", height=1024, width=1024, num_inference_steps=40, guidance_scale=4.5)

# Later calls with the same shapes reuse the compiled graph.
image = pipe(
prompt="a golden retriever surfing a wave, photorealistic",
height=1024,
width=1024,
num_inference_steps=40,
guidance_scale=4.5,
).images[0]

image.save("output.png")
```

If the pipeline doesn't fit on one chip:

- With several chips, shard it with [tensor parallelism](#tensor-parallelism) instead. Everything stays on the TPU,
and `pipe.transformer.compile(...)` works the same way on the sharded transformer.
- With [`~DiffusionPipeline.enable_model_cpu_offload`], the offload hooks can't be traced by `torch.compile`. Compile
only the transformer's repeated blocks instead, with
`pipe.transformer.compile_repeated_blocks(backend="tpu", fullgraph=True, dynamic=False)`.

## Tensor parallelism

Shard models too large for one chip across several. FLUX.2-dev's text encoder (~48GB) and transformer (~64GB) each
exceed a single chip, so the example below shards both:

- the transformer with [`TensorParallelConfig`], passed to the `parallel_config` argument of [`~ModelMixin.from_pretrained`]. Each rank reads only its own slice of every sharded weight, so the full model is never materialized. For general TP details (`_tp_plan`, colwise/rowwise), see the [Tensor parallelism](../training/distributed_inference#tensor-parallelism) guide.
- the text encoder with Transformers' own [tensor parallelism](https://huggingface.co/docs/transformers/perf_infer_gpu_multi), passing a [`~transformers.DistributedConfig`] and the same mesh.

On TPU, initialize the process group with `backend="tpu_dist"` and build the mesh with `DeviceMesh("tpu", ...)`.

```python
import torch
import torch.distributed as dist
from torch.distributed.device_mesh import DeviceMesh
from transformers import DistributedConfig, Mistral3ForConditionalGeneration

from diffusers import Flux2Pipeline, Flux2Transformer2DModel, TensorParallelConfig

dist.init_process_group(backend="tpu_dist")
mesh = DeviceMesh("tpu", list(range(dist.get_world_size())))

repo_id = "black-forest-labs/FLUX.2-dev"
text_encoder = Mistral3ForConditionalGeneration.from_pretrained(
repo_id,
subfolder="text_encoder",
dtype=torch.bfloat16,
distributed_config=DistributedConfig(tp_plan="auto"),
device_mesh=mesh,
)
transformer = Flux2Transformer2DModel.from_pretrained(
repo_id, subfolder="transformer", dtype=torch.bfloat16, parallel_config=TensorParallelConfig(mesh=mesh)
)
pipe = Flux2Pipeline.from_pretrained(
repo_id, text_encoder=text_encoder, transformer=transformer, dtype=torch.bfloat16
)
pipe.vae.to("tpu")

image = pipe(
prompt="a golden retriever surfing a wave, photorealistic",
num_inference_steps=28,
generator=torch.Generator("cpu").manual_seed(0),
).images[0]
if dist.get_rank() == 0:
image.save("output.png")
```

Launch one process per chip. Set `--nproc_per_node` to use all the number of TPU chips on your host.

```bash
eval $(python -m torch_tpu._internal.distributed.launchers.singlehost_wrapper | sed 's/^/export /')
torchrun --nproc_per_node=8 flux2_tp.py
```
2 changes: 1 addition & 1 deletion src/diffusers/hooks/tensor_parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
2 changes: 1 addition & 1 deletion src/diffusers/models/_modeling_parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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`):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1066,7 +1066,8 @@ def __call__(
# expand the latents if we are doing classifier free guidance
latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0])
# `clone()` so it isn't a view into `timesteps`, which would recompile `torch.compile` every step
timestep = t.expand(latent_model_input.shape[0]).clone()
Comment on lines +1069 to +1070

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Do you have a reproducer?

On a DGX-Spark, I didn't get any recompilation:

from diffusers import DiffusionPipeline 
import torch

pipe = DiffusionPipeline.from_pretrained(
    "DavyMorgan/tiny-sd3-pipe", dtype=torch.bfloat16
).to("cuda")
pipe.transformer.compile(fullgraph=True)

with torch._dynamo.config.patch(error_on_recompile=True):
    pipe(
        prompt="A painting of a squirrel eating a burger", num_inference_steps=4,
    )


noise_pred = self.transformer(
hidden_states=latent_model_input,
Expand All @@ -1088,7 +1089,7 @@ def __call__(
else False
)
if skip_guidance_layers is not None and should_skip_layers:
timestep = t.expand(latents.shape[0])
timestep = t.expand(latents.shape[0]).clone()
latent_model_input = latents
noise_pred_skip_layers = self.transformer(
hidden_states=latent_model_input,
Expand Down
1 change: 1 addition & 0 deletions src/diffusers/utils/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -116,6 +116,7 @@
is_torch_mlu_available,
is_torch_neuronx_available,
is_torch_npu_available,
is_torch_tpu_available,
is_torch_version,
is_torch_xla_available,
is_torch_xla_version,
Expand Down
5 changes: 5 additions & 0 deletions src/diffusers/utils/import_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,7 @@ def _is_package_available(pkg_name: str, get_dist_name: bool = False) -> tuple[b
_torch_xla_available, _torch_xla_version = _is_package_available("torch_xla")
_torch_npu_available, _torch_npu_version = _is_package_available("torch_npu")
_torch_mlu_available, _torch_mlu_version = _is_package_available("torch_mlu")
_torch_tpu_available, _torch_tpu_version = _is_package_available("torch_tpu")
_torch_neuronx_available, _torch_neuronx_version = _is_package_available("torch_neuronx")
_transformers_available, _transformers_version = _is_package_available("transformers")
_hf_hub_available, _hf_hub_version = _is_package_available("huggingface_hub")
Expand Down Expand Up @@ -238,6 +239,10 @@ def is_torch_mlu_available():
return _torch_mlu_available


def is_torch_tpu_available():
return _torch_tpu_available


def is_torch_neuronx_available():
return _torch_neuronx_available

Expand Down
2 changes: 2 additions & 0 deletions tests/models/testing_utils/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
ContextParallelAttentionBackendsTesterMixin,
ContextParallelTesterMixin,
TensorParallelTesterMixin,
TensorParallelTPUTesterMixin,
)
from .quantization import (
AutoRoundCompileTesterMixin,
Expand Down Expand Up @@ -67,6 +68,7 @@
"ContextParallelTesterMixin",
"ContextParallelAttentionBackendsTesterMixin",
"TensorParallelTesterMixin",
"TensorParallelTPUTesterMixin",
"CPUOffloadTesterMixin",
"FasterCacheConfigMixin",
"FasterCacheTesterMixin",
Expand Down
Loading
Loading