Buckets:

HuggingFaceDocBuilder's picture
|
download
raw
6.12 kB
# 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` | `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 [apply_tensor_parallel()](/docs/diffusers/pr_14039/en/api/parallel#diffusers.hooks.apply_tensor_parallel), the
same mechanism `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](/docs/diffusers/pr_14039/en/api/parallel#diffusers.TensorParallelConfig) to the `parallel_config` argument of [from_pretrained()](/docs/diffusers/pr_14039/en/api/models/overview#diffusers.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")
```

Xet Storage Details

Size:
6.12 kB
·
Xet hash:
64cc9e0b433ea0f6c2af111593d1e15c2469a6bd3deb447ce5b90de926b1f943

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.