Diffusers documentation

TorchTPU

You are viewing main version, which requires installation from source. If you'd like regular pip install, checkout the latest stable version (v0.41.0).
Hugging Face's logo
Join the Hugging Face community

and get access to the augmented documentation experience

to get started

TorchTPU

TorchTPU is a PyTorch backend for Google’s Tensor Processing Units (TPUs), which lets you run Diffusers pipelines on Cloud TPUs (v6e, v7x) with minimal code changes.

Two execution modes are available:

ModeConstantHow to activateNotes
Strict eager (default)EagerMode.DEFER_NEVERpipe.to("tpu")Operations dispatched one at a time, asynchronous
Compile—torch.compile(module, backend="tpu")AOT compilation with TpuBackend

Follow the TorchTPU installation guide. 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

FLUX.1-schnell doesn’t fit on a single v6e chip all at once, so use enable_model_cpu_offload() to move each model to the TPU only while it runs. It detects the "tpu" device automatically.

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 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.

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.

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 instead. Everything stays on the TPU, and pipe.transformer.compile(...) works the same way on the sharded transformer.
  • With 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:

On TPU, initialize the process group with backend="tpu_dist" and build the mesh with DeviceMesh("tpu", ...).

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.

eval $(python -m torch_tpu._internal.distributed.launchers.singlehost_wrapper | sed 's/^/export /')
torchrun --nproc_per_node=8 flux2_tp.py

The launch command above uses a TorchTPU internal module (torch_tpu._internal), not a stable API, and is expected to change in a future TorchTPU release. Check the TorchTPU repository for the current way to launch multi-chip jobs.

Update on GitHub