Buckets:

HuggingFaceDocBuilder's picture
|
download
raw
6.12 kB

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

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.