Buckets:
TorchTPU
TorchTPU provides a PyTorch backend for Google's Tensor Processing Units (TPUs), enabling you to run diffusers pipelines on Google Cloud TPUs (v6e, v5p, …) with minimal code changes.
Four execution modes are available:
| Mode | Constant | How to activate | Notes |
|---|---|---|---|
| Strict Eager (default) | EagerMode.DEFER_NEVER |
just import torch_tpu |
Operations dispatched one at a time, asynchronous |
| Compile | — | pipe.enable_tpu_compile() |
AOT compilation with TpuBackend |
Installation
Follow the TorchTPU installation guide. After installation,
import torch_tpu registers the "tpu" device automatically.
Basic usage (strict eager mode)
import torch
import torch_tpu # noqa: F401 — registers torch.tpu
from diffusers import FluxPipeline
pipe = FluxPipeline.from_pretrained(
"black-forest-labs/FLUX.1-schnell",
torch_dtype=torch.bfloat16,
)
# Move only the denoising components to TPU; text encoders stay on CPU.
pipe.transformer.to("tpu")
pipe.vae.to("tpu")
# _execution_device is now "tpu" automatically.
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")
Compiled mode (recommended for production)
torch.compile with TpuBackend traces the transformer statically. The first call (warmup)
is slow because it triggers compilation; subsequent calls reuse the compiled graph.
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")
# Compile TPU components with TpuBackend.
# Also applies AttnProcessor to replace SDP-based attention (required for XLA).
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")
Eager mode
TorchTPU defaults to Strict Eager (EagerMode.DEFER_NEVER): operations are dispatched one
at a time asynchronously, matching standard PyTorch GPU behaviour.
For the best production throughput, prefer
torch.compileviapipe.enable_tpu_compile().
API reference
enable_tpu_compile[[diffusers.DiffusionPipeline.enable_tpu_compile]]
diffusers.DiffusionPipeline.enable_tpu_compile[[diffusers.DiffusionPipeline.enable_tpu_compile]]
diffusers.DiffusionPipeline.enable_tpu_compile(model_names: typing.Optional[typing.List[str]] = None, **compile_kwargs)
Parameters:
model_names (list[str], optional) : Names of pipeline components to compile. Defaults to all torch.nn.Module components currently resident on a TPU device.
- **compile_kwargs : Extra keyword arguments forwarded to
torch.compile.backenddefaults toTpuBackend()anddynamicdefaults toFalse(required for static tracing).
Compile pipeline components that are on TPU using torch.compile with the TpuBackend.
Before compiling, each component that exposes set_attn_processor has AttnProcessor applied. This
replaces AttnProcessor2_0 (SDP-based) which triggers XLA fusion-emitter crashes in eager/lazy mode.
TpuBackend handles the resulting torch.cat layout internally during static tracing, so no additional
wrapper is needed at compile time.
Example:
import torch
import torch_tpu # noqa: F401
pipe.transformer.to("tpu")
pipe.vae.to("tpu")
pipe.enable_tpu_compile()
tpu_warmup[[diffusers.DiffusionPipeline.tpu_warmup]]
diffusers.DiffusionPipeline.tpu_warmup[[diffusers.DiffusionPipeline.tpu_warmup]]
diffusers.DiffusionPipeline.tpu_warmup(*args, **kwargs)
Parameters:
*args : Positional arguments forwarded to the pipeline
__call__.**kwargs : Keyword arguments forwarded to the pipeline
__call__.
Run a single forward pass to trigger XLA / TpuBackend compilation.
Call this after enable_tpu_compile and before timed inference. The warmup pass compiles the static
computation graphs; subsequent calls reuse the compiled graphs and run at full speed.
Example:
pipe.tpu_warmup(
prompt="warmup",
height=1024,
width=1024,
num_inference_steps=4,
guidance_scale=0.0,
)
Xet Storage Details
- Size:
- 5.34 kB
- Xet hash:
- 2ea5547a112d0105076e655cc6707395e14f271876464e07692e99ffd91ea56d
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.