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 | — | torch.compile(module, backend="tpu") |
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
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.
TorchTPU requires static shapes — pass
dynamic=False. Every timeheight,width, ornum_inference_stepschanges, the graph is recompiled from scratch. Keep these values constant across all calls after warmup, or run another warmup pass before changing them.
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 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 guide. On TPU, initialize the process group with backend="tpu_dist" and build the mesh with DeviceMesh("tpu", ...).
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.