Buckets:
TorchTPU
TorchTPU 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 | — | pipe.enable_tpu_compile() |
AOT compilation with TpuBackend |
Follow the TorchTPU installation guide. After installation,
import torch_tpu registers the "tpu" device automatically.
Eager mode
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(), the
same mechanism enable_parallelism() uses for the transformer (see 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
enable_tpu_compile runs torch.compile with TpuBackend on each pipeline module that is already on TPU. The first call (warmup) is slow because it compiles. Later calls reuse the compiled graph. Where it's supported, it replaces SDP-based attention with AttnProcessor for XLA tracing.
TorchTPU requires static shapes —
torch.compileis called withdynamic=Falseinternally. Every timeheight,width, ornum_inference_stepschanges, the graph is recompiled from scratch. Keep these values constant across all calls after warmup, or calltpu_warmupagain before changing them.
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,
)
pipe.transformer.to("tpu")
pipe.vae.to("tpu")
pipe.enable_tpu_compile()
# Warmup — triggers static graph compilation.
pipe.tpu_warmup(
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")
Xet Storage Details
- Size:
- 3.93 kB
- Xet hash:
- a203a39b11ed3b3a93299e613d6b53c5dc2caa2551fd6e722d7a8de81957cfd7
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.