Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
56 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
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 @@ -122,6 +122,8 @@
title: Intel Gaudi
- local: optimization/neuron
title: AWS Neuron
- local: optimization/tpu
title: TPU
title: Hardware-specific acceleration
- isExpanded: false
sections:
Expand Down
163 changes: 163 additions & 0 deletions docs/source/en/optimization/tpu.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,163 @@
<!--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` | `import torch_tpu` | Operations dispatched one at a time, asynchronous |
| Compile | — | `torch.compile(module, backend="tpu")` | AOT compilation with `TpuBackend` |

Follow the [TorchTPU installation guide](https://github.com/google-pytorch/torch_tpu/). After installation,
`import torch_tpu` registers the `"tpu"` device automatically.

## Eager mode

```python
import gc
import torch
import torch_tpu # noqa: F401

from diffusers import FluxPipeline

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

# 1. Encode on TPU.
pipe.text_encoder.to("tpu")
pipe.text_encoder_2.to("tpu")
with torch.no_grad():
prompt_embeds, pooled_prompt_embeds, _ = pipe.encode_prompt(
prompt="a golden retriever surfing a wave, photorealistic",
prompt_2="a golden retriever surfing a wave, photorealistic",
device=torch.device("tpu"),
max_sequence_length=512,
)

# 2. Free the text encoders — nothing below needs them.
pipe.text_encoder = None
pipe.text_encoder_2 = None
gc.collect()

# 3. Move the transformer and VAE in, then denoise with the precomputed embeddings.
pipe.transformer.to("tpu")
pipe.vae.to("tpu")
image = pipe(
prompt_embeds=prompt_embeds,
pooled_prompt_embeds=pooled_prompt_embeds,
height=1024,
width=1024,
num_inference_steps=4,
guidance_scale=0.0,
).images[0]

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

If the text encoder alone is too large for a single chip(eg. FLUX.2-dev's Mistral-3-Small is ~45GB),
shard it across multiple chips with [`~diffusers.hooks.tensor_parallel.apply_tensor_parallel`], the
same mechanism [`~ModelMixin.enable_parallelism`] uses for the transformer (see [Tensor
parallelism](../training/distributed_inference#tensor-parallelism)). It only requires `model:
torch.nn.Module`, so it works directly on a `transformers.PreTrainedModel` text encoder too, not
just a diffusers `ModelMixin`. The text encoder doesn't define a `_tp_plan`, so supply one: pair
each attention/MLP projection that expands the hidden dimension (`"colwise"`) with the one that
contracts it back (`"rowwise"`), matching the `transformers` model's actual module names.

## Compiled mode

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

> [!IMPORTANT]
> TorchTPU requires **static shapes** — pass `dynamic=False`. Every time `height`, `width`, or
> `num_inference_steps` changes, the graph is recompiled from scratch. Keep these values constant
> across all calls after warmup, or run another warmup pass before changing them.

```python
import torch
import torch_tpu # noqa: F401 — registers the "tpu" torch.compile backend

from diffusers import FluxPipeline

pipe = FluxPipeline.from_pretrained(
"black-forest-labs/FLUX.1-schnell",
torch_dtype=torch.bfloat16,
)
pipe.transformer.to("tpu")
pipe.vae.to("tpu")

pipe.transformer = torch.compile(pipe.transformer, backend="tpu", fullgraph=True, dynamic=False)
pipe.vae = torch.compile(pipe.vae, backend="tpu", fullgraph=True, dynamic=False)

# Warmup — triggers static graph compilation.
with torch.no_grad():
pipe(
prompt="warmup",
height=1024,
width=1024,
num_inference_steps=4,
guidance_scale=0.0,
)

# Timed inference reuses the compiled graph.
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")
```

## Tensor parallelism

Shard a transformer too large for one chip across several by passing a [`TensorParallelConfig`] 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. 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
import torch_tpu # noqa: F401
from torch.distributed.device_mesh import DeviceMesh

from diffusers import DiffusionPipeline, Flux2Transformer2DModel, TensorParallelConfig

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

transformer = Flux2Transformer2DModel.from_pretrained(
"black-forest-labs/FLUX.2-dev",
subfolder="transformer",
torch_dtype=torch.bfloat16,
parallel_config=TensorParallelConfig(mesh=tp_mesh),
)
pipe = DiffusionPipeline.from_pretrained(
"black-forest-labs/FLUX.2-dev", transformer=transformer, torch_dtype=torch.bfloat16
)
# The transformer is already sharded across the chips; move the remaining components individually. The ~45GB
# text encoder doesn't fit on one chip, so leave it on CPU (or shard it as described in the eager mode section)
# and encode the prompt there.
pipe.vae.to("tpu")
with torch.no_grad():
prompt_embeds, _ = pipe.encode_prompt(
prompt="a golden retriever surfing a wave, photorealistic", device=torch.device("cpu")
)

image = pipe(prompt_embeds=prompt_embeds.to("tpu"), num_inference_steps=28).images[0]
if dist.get_rank() == 0:
image.save("output.png")
```
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
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