Repository navigation
[TPU] TorchTPU backend integration - eager / torch.compile / tp #14039
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Open
JingyaHuang
wants to merge
74
commits into
huggingface:main
Choose a base branch
from
JingyaHuang:add-torchtpu-support
base: main
Could not load branches
Branch not found: {{ refName }}
Loading
Could not load tags
Nothing to show
Loading
Are you sure you want to change the base?
Some commits from the old base branch may be removed from the timeline,
and old review comments may become outdated.
+320
−25
Open
Changes from all commits
Commits
Show all changes
74 commits
Select commit
Hold shift + click to select a range
38007b4
feat: add torchtpu
JingyaHuang 8ed15f9
feat:draft TorchTPU support
JingyaHuang e343c0a
fix: wan overflow issue + compile mode error on sdxl
JingyaHuang 84b4049
doc: enhance with TorchTPU doc
JingyaHuang bb3ec1e
doc: enhance with TorchTPU doc
JingyaHuang 339be41
Merge branch 'main' into add-torchtpu-support
JingyaHuang 92193a7
Merge branch 'main' into add-torchtpu-support
JingyaHuang 1a2369b
Merge branch 'main' into add-torchtpu-support
JingyaHuang 66254b4
style: remove unused imports flagged by ruff
JingyaHuang 7c241be
docs: remove Debug Eager and Fused Eager sections from tpu.md
JingyaHuang 05d56d5
Merge branch 'main' into add-torchtpu-support
JingyaHuang b7a8c0b
Merge branch 'main' into add-torchtpu-support
JingyaHuang b0c5595
Merge branch 'huggingface:main' into add-torchtpu-support
JingyaHuang 7c3df6c
Merge branch 'main' into add-torchtpu-support
JingyaHuang 7c1e379
feat: add TP for TPU
JingyaHuang f944c48
style: fix import sorting in TPU test scripts
JingyaHuang 82d20ab
style: ruff format TPU test scripts
JingyaHuang f6c17d2
fix: style
e8fe48c
fix: test for native 4 devices
JingyaHuang 9b62254
fix: propagate TPU device fixes to Flux/Flux2/Wan-family copies; regi…
JingyaHuang 744fa95
tests: cleanup
JingyaHuang bfc7d17
fix: fix compile mode
JingyaHuang 10d4ad1
Update docs/source/en/optimization/tpu.md
JingyaHuang a999a3d
Update docs/source/en/optimization/tpu.md
JingyaHuang cb2dcf3
Update docs/source/en/optimization/tpu.md
JingyaHuang 9f11d1e
Merge branch 'main' into add-torchtpu-support
JingyaHuang c48d41c
Update docs/source/en/optimization/tpu.md
JingyaHuang aa70a45
Update docs/source/en/optimization/tpu.md
JingyaHuang 6cb37d2
Update docs/source/en/optimization/tpu.md
JingyaHuang 3e06f71
Update docs/source/en/optimization/tpu.md
JingyaHuang 09d5ce2
Update docs/source/en/optimization/tpu.md
JingyaHuang 7b46e08
Update docs/source/en/optimization/tpu.md
JingyaHuang e82e795
Update docs/source/en/optimization/tpu.md
JingyaHuang d2bf529
Update docs/source/en/optimization/tpu.md
JingyaHuang 22d595a
test: remove flux2 e2e test
JingyaHuang 042a88a
doc: apply suggestions
JingyaHuang c6a4c80
doc: apply suggestions
JingyaHuang c82fce7
Merge branch 'add-torchtpu-support' of github.com:JingyaHuang/diffuse…
JingyaHuang cb314bd
review: remove monkey patch
JingyaHuang bf8934f
review: revert neuron-specific changes in the tests
JingyaHuang c35e23a
review: apply suggestions
JingyaHuang 419d657
Merge branch 'main' of https://github.com/huggingface/diffusers into …
JingyaHuang a2a770a
Update docs/source/en/optimization/tpu.md
JingyaHuang 7f73095
Merge branch 'main' into add-torchtpu-support
JingyaHuang 406ac74
doc: add tpu to tp doc
JingyaHuang 67a84d2
Merge branch 'add-torchtpu-support' of github.com:JingyaHuang/diffuse…
JingyaHuang a638798
Merge branch 'main' into add-torchtpu-support
JingyaHuang 2bf3619
review: remove unnecessary for tpu
JingyaHuang 5d6ad22
Merge branch 'main' into add-torchtpu-support
JingyaHuang eb49a61
Merge branch 'main' into add-torchtpu-support
JingyaHuang 3152988
test: assert every _tp_plan parameter is sharded after a TP load
JingyaHuang 7257a79
removal: delete redundant tp shard helpers
JingyaHuang fd0f9a5
removal: drop TPU workarounds no longer needed after recent main changes
JingyaHuang 6ee32b4
doc: keep the FLUX.2-dev text encoder off a single TPU chip in the TP…
JingyaHuang 2d38ce7
doc: encode the prompt on CPU in the TPU TP example so it runs end to…
JingyaHuang 1eddd6f
test: tighten TPU TP tolerance and simplify the TPU TP worker
JingyaHuang 1351263
doc: use enable_model_cpu_offload() in the TPU eager example
JingyaHuang a170a50
doc: shard both FLUX.2-dev text encoder and transformer in the TPU TP…
JingyaHuang 5c9e426
doc: run the text encoders on TPU in the compiled example
JingyaHuang a3f9ab1
test: drop the redundant TPU sync in the TP worker
JingyaHuang e7e6c49
test: trim comments in the TPU TP tests
JingyaHuang 5c262c7
test: run the TPU TP test with mp.spawn like the CUDA one
JingyaHuang ea2eaca
test: pin the TPU TP test to 4 chips
JingyaHuang 866e165
Merge branch 'main' into add-torchtpu-support
JingyaHuang 6c5c580
Merge branch 'main' into add-torchtpu-support
sayakpaul 2de295b
Update docs/source/en/optimization/tpu.md
JingyaHuang 69d07a3
doc: use DistributedConfig and dtype, and drop the torch_tpu import i…
JingyaHuang a9412fe
test: skip the TPU TP test for models without a _tp_plan
JingyaHuang 6c89944
doc: correct the warning of recompilation
JingyaHuang 91bba1a
fix: clone the SD3 timestep so torch.compile doesn't recompile every …
JingyaHuang 935d2d1
doc: update advice for compile mode when one chip doesn't fit
JingyaHuang fb27377
test: pass the TPU TP tolerances as test arguments
JingyaHuang 04a043a
Merge branch 'add-torchtpu-support' of https://github.com/JingyaHuang…
JingyaHuang b3b529a
Merge branch 'main' into add-torchtpu-support
JingyaHuang File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| 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 | ||
|
|
||
| [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` | | ||
|
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 | ||
|
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 | ||
| ``` | ||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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, | ||
|
|
@@ -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, | ||
|
|
||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.