Buckets:
| # TorchTPU | |
| [TorchTPU](https://github.com/google-pytorch/torch_tpu/) 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](https://github.com/google-pytorch/torch_tpu/). After installation, | |
| `import torch_tpu` registers the `"tpu"` device automatically. | |
| ## Eager mode | |
| ```python | |
| 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()](/docs/diffusers/pr_14039/en/api/parallel#diffusers.hooks.apply_tensor_parallel), the | |
| same mechanism `enable_parallelism()` uses for the transformer (see [Tensor | |
| parallelism](../training/distributed_inference#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. | |
| > [!IMPORTANT] | |
| > 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. | |
| ```python | |
| 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](/docs/diffusers/pr_14039/en/api/parallel#diffusers.TensorParallelConfig) to the `parallel_config` argument of [from_pretrained()](/docs/diffusers/pr_14039/en/api/models/overview#diffusers.ModelMixin.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](../training/distributed_inference#tensor-parallelism) guide. On TPU, initialize the process group with `backend="tpu_dist"` and build the mesh with `DeviceMesh("tpu", ...)`. | |
| ```python | |
| 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.