diff --git a/docs/source/en/_toctree.yml b/docs/source/en/_toctree.yml index cd79d8d39266..d987a1b85fa4 100644 --- a/docs/source/en/_toctree.yml +++ b/docs/source/en/_toctree.yml @@ -122,6 +122,8 @@ title: Intel Gaudi - local: optimization/neuron title: AWS Neuron + - local: optimization/tpu + title: TPU title: Hardware-specific acceleration - isExpanded: false sections: diff --git a/docs/source/en/optimization/tpu.md b/docs/source/en/optimization/tpu.md new file mode 100644 index 000000000000..b449598b8622 --- /dev/null +++ b/docs/source/en/optimization/tpu.md @@ -0,0 +1,163 @@ + + +# 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 [`~diffusers.hooks.tensor_parallel.apply_tensor_parallel`], the +same mechanism [`~ModelMixin.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`] to the `parallel_config` argument of [`~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") +``` diff --git a/src/diffusers/hooks/tensor_parallel.py b/src/diffusers/hooks/tensor_parallel.py index 0f56c3644ec2..ce0c76f97081 100644 --- a/src/diffusers/hooks/tensor_parallel.py +++ b/src/diffusers/hooks/tensor_parallel.py @@ -22,7 +22,7 @@ logger = get_logger(__name__) # pylint: disable=invalid-name -_SUPPORTED_TP_DEVICES = ("cuda", "neuron") +_SUPPORTED_TP_DEVICES = ("cuda", "neuron", "tpu") class PackedColwiseParallel: diff --git a/src/diffusers/models/_modeling_parallel.py b/src/diffusers/models/_modeling_parallel.py index b54e86d6b4f2..58cbc822f8bf 100644 --- a/src/diffusers/models/_modeling_parallel.py +++ b/src/diffusers/models/_modeling_parallel.py @@ -161,7 +161,7 @@ class TensorParallelConfig: Tensor parallelism shards weight matrices (column-wise and row-wise) across devices. Each device computes a partial result; an AllReduce/AllGather at layer boundaries reconstructs the full output. Uses `torch.distributed.tensor.parallelize_module` with `ColwiseParallel` / `RowwiseParallel` sharding styles. Supported - device types are `"cuda"` and `"neuron"`. + device types are `"cuda"`, `"neuron"` and `"tpu"`. Args: tp_degree (`int`, defaults to `1`): diff --git a/src/diffusers/utils/__init__.py b/src/diffusers/utils/__init__.py index b3051dfcb9d1..5f89eb318d92 100644 --- a/src/diffusers/utils/__init__.py +++ b/src/diffusers/utils/__init__.py @@ -116,6 +116,7 @@ is_torch_mlu_available, is_torch_neuronx_available, is_torch_npu_available, + is_torch_tpu_available, is_torch_version, is_torch_xla_available, is_torch_xla_version, diff --git a/src/diffusers/utils/import_utils.py b/src/diffusers/utils/import_utils.py index d2cf394cd9a7..af765021f409 100644 --- a/src/diffusers/utils/import_utils.py +++ b/src/diffusers/utils/import_utils.py @@ -178,6 +178,7 @@ def _is_package_available(pkg_name: str, get_dist_name: bool = False) -> tuple[b _torch_xla_available, _torch_xla_version = _is_package_available("torch_xla") _torch_npu_available, _torch_npu_version = _is_package_available("torch_npu") _torch_mlu_available, _torch_mlu_version = _is_package_available("torch_mlu") +_torch_tpu_available, _torch_tpu_version = _is_package_available("torch_tpu") _torch_neuronx_available, _torch_neuronx_version = _is_package_available("torch_neuronx") _transformers_available, _transformers_version = _is_package_available("transformers") _hf_hub_available, _hf_hub_version = _is_package_available("huggingface_hub") @@ -238,6 +239,10 @@ def is_torch_mlu_available(): return _torch_mlu_available +def is_torch_tpu_available(): + return _torch_tpu_available + + def is_torch_neuronx_available(): return _torch_neuronx_available diff --git a/tests/models/testing_utils/__init__.py b/tests/models/testing_utils/__init__.py index 2d7d5ae23257..0932e46f9b93 100644 --- a/tests/models/testing_utils/__init__.py +++ b/tests/models/testing_utils/__init__.py @@ -23,6 +23,7 @@ ContextParallelAttentionBackendsTesterMixin, ContextParallelTesterMixin, TensorParallelTesterMixin, + TensorParallelTPUTesterMixin, ) from .quantization import ( AutoRoundCompileTesterMixin, @@ -67,6 +68,7 @@ "ContextParallelTesterMixin", "ContextParallelAttentionBackendsTesterMixin", "TensorParallelTesterMixin", + "TensorParallelTPUTesterMixin", "CPUOffloadTesterMixin", "FasterCacheConfigMixin", "FasterCacheTesterMixin", diff --git a/tests/models/testing_utils/parallelism.py b/tests/models/testing_utils/parallelism.py index b525637c953e..52aa09c30b04 100644 --- a/tests/models/testing_utils/parallelism.py +++ b/tests/models/testing_utils/parallelism.py @@ -15,6 +15,8 @@ import os import socket +import subprocess +import sys import pytest import torch @@ -31,6 +33,7 @@ is_kernels_available, is_tensor_parallel, require_torch_multi_accelerator, + require_torch_tpu, torch_device, ) from .common import calculate_expected_num_shards, compute_module_persistent_sizes @@ -303,8 +306,8 @@ def _tensor_parallel_from_pretrained_worker( """Worker for `from_pretrained(..., parallel_config=...)`, i.e. sharding while reading the checkpoint. Each rank loads only its own slice of every `_tp_plan` parameter straight into a `DTensor` and runs a forward - pass. Rank 0 reports its output and the local/global shapes of one sharded weight so the caller can check both the - numerics and that sharding actually happened. + pass. Rank 0 checks that exactly the parameters `_tp_plan` covers were loaded as `DTensor`s, each with the + placement and local shape its shard spec implies, and reports its output so the caller can check the numerics. """ try: os.environ["MASTER_ADDR"] = "localhost" @@ -316,7 +319,9 @@ def _tensor_parallel_from_pretrained_worker( dist.init_process_group(backend=device_config["backend"], rank=rank, world_size=world_size) device_config["module"].set_device(rank) - from torch.distributed.tensor import DTensor + from torch.distributed.tensor import DTensor, Replicate, Shard + + from diffusers.hooks.tensor_parallel import resolve_tp_shard_specs model = model_class.from_pretrained( checkpoint_dir, parallel_config=TensorParallelConfig(tp_degree=world_size) @@ -330,12 +335,31 @@ def _tensor_parallel_from_pretrained_worker( output = output.full_tensor() if rank == 0: - sharded = {k: v for k, v in model.state_dict().items() if isinstance(v, DTensor)} - assert sharded, "No parameter was sharded into a DTensor by the streaming load." - name, param = next(iter(sharded.items())) + specs = resolve_tp_shard_specs(model, model_class._tp_plan, world_size) + state_dict = model.state_dict() + + # A planned parameter the streaming load left as a full tensor would still compute the right output, so + # the numerics check alone would not catch it. + for name, spec in specs.items(): + param = state_dict[name] + assert isinstance(param, DTensor), ( + f"'{name}' is covered by `_tp_plan` but was not loaded as a DTensor." + ) + placement = Replicate() if spec.dim is None else Shard(spec.dim) + assert param.placements == (placement,), ( + f"'{name}' has placements {param.placements}, not {placement}." + ) + expected_local_shape = list(param.shape) + if spec.dim is not None: + expected_local_shape[spec.dim] //= world_size + assert list(param.to_local().shape) == expected_local_shape, ( + f"'{name}' has local shape {list(param.to_local().shape)}, not {expected_local_shape}." + ) + + unplanned = sorted(k for k, v in state_dict.items() if isinstance(v, DTensor) and k not in specs) + assert not unplanned, f"Parameters not covered by `_tp_plan` were loaded as DTensors: {unplanned}" + return_dict["status"] = "success" - return_dict["num_sharded"] = len(sharded) - return_dict["shard_example"] = (name, list(param.to_local().shape), list(param.shape)) return_dict["output"] = output.float().cpu().tolist() except Exception as e: @@ -456,15 +480,114 @@ def test_tensor_parallel_from_pretrained(self, tmp_path, sharded): f"Tensor parallel `from_pretrained` failed: {return_dict.get('error', 'Unknown error')}" ) - name, local_shape, global_shape = return_dict["shard_example"] - assert local_shape != global_shape, ( - f"'{name}' has local shape {local_shape} equal to its global shape, so it was not sharded." - ) - # Sharded matmuls + all-reduce reorder the summation, so allow a small tolerance over the reference. torch.testing.assert_close(reference, torch.tensor(return_dict["output"]), atol=1e-3, rtol=1e-3) +def _run_tp_worker_subprocess( + worker_filename: str, spec: str, world_size: int, timeout_s: int = 900, extra_args: "list[str] | None" = None +) -> None: + """Launch a `torchrun` TP-correctness worker subprocess and assert it exits cleanly. + + Args: + worker_filename: Name of the worker script, resolved relative to `tests/models/transformers/` (e.g. + `"_tpu_tp_worker.py"`). + spec: `module:function` reference forwarded to the worker, returning `(model_class, init_dict, cpu_inputs)`. + world_size: Number of ranks to launch (`torchrun --nproc_per_node`). + timeout_s: Seconds to wait for the subprocess before failing the test. The worker itself only needs a couple + of minutes even from a cold compile; this generously bounds it so a real hang (e.g. a distributed-runtime + barrier timeout) fails the test loudly instead of stalling the run. + extra_args: Further command-line arguments forwarded to the worker. + """ + worker = os.path.join(os.path.dirname(__file__), "..", "transformers", worker_filename) + cmd = [sys.executable, "-m", "torch.distributed.run", f"--nproc_per_node={world_size}", worker, spec] + cmd += extra_args or [] + try: + result = subprocess.run(cmd, capture_output=True, text=True, timeout=timeout_s) + except subprocess.TimeoutExpired as e: + raise AssertionError( + f"TP worker did not finish within {timeout_s}s (likely stuck on a distributed-runtime barrier).\n" + f"--- stdout ---\n{e.stdout}\n--- stderr ---\n{e.stderr}" + ) from e + assert result.returncode == 0, ( + f"TP worker failed (exit {result.returncode}).\n--- stdout ---\n{result.stdout}\n--- stderr ---\n{result.stderr}" + ) + + +@is_tensor_parallel +@require_torch_tpu +class TensorParallelTPUTesterMixin: + """Mixin for a tensor-parallel correctness test on TPU, run via `_tpu_tp_worker.py`. + + TPU TP runs through `torchrun` with the `"tpu_dist"` distributed backend, so — like `TestFlux2TransformerTensorParallelNeuron` + for Neuron — it cannot use `TensorParallelTesterMixin`'s `torch.multiprocessing.spawn`/NCCL path above and instead + launches a subprocess worker script and checks its exit code. + + Subclasses set `TP_SPEC` to a `module:function` reference returning `(model_class, init_dict, cpu_inputs)` and, + only if the model spec's head count doesn't divide 4, override `WORLD_SIZE`. `TP_ATOL` / `TP_RTOL` bound the + difference from the single-chip reference; override them only for a model whose TPU numerics depend on the shard + shapes. + + `WORLD_SIZE` defaults to 4 rather than an arbitrary rank count: `torch_tpu`'s per-generation topology table + (`torch_tpu._internal.utils.hardware`) only enumerates whole-pod-slice chip counts (1/4/8 for v6e, for example), + not arbitrary sub-slices of a larger single host. A rank count with no matching whole-slice topology has + nothing to advertise and the PJRT client never completes its start-session barrier — the test would hang for + the barrier's full multi-minute timeout instead of failing. 4 is the smallest whole-slice count every current + TPU generation defines (see `_V4_TOPOLOGY` / `_V5E_TOPOLOGY` / `_V6E_TOPOLOGY` / `_V7_TOPOLOGY` in + `torch_tpu._internal.utils.hardware`). `skip_if_unsupported` below still checks the actual host up front and + skips fast instead of hanging when it doesn't have exactly that many chips. + + Requires `TORCH_TPU_TOPOLOGY` and `TORCH_TPU_SLICEBUILDER_ADDRESSES` to be set. Source them via:: + + eval $(python -m torch_tpu._internal.distributed.launchers.singlehost_wrapper | sed 's/^/export /') + """ + + WORLD_SIZE = 4 + # The worker itself only needs a couple of minutes even from a cold XLA compile; this generously bounds the + # subprocess so a real hang (e.g. a barrier timeout this skip failed to catch) fails the test loudly instead of + # stalling the run. + TIMEOUT_S = 900 + TP_SPEC: str = "" + TP_ATOL = 1e-3 + TP_RTOL = 1e-3 + + def skip_if_unsupported(self): + """Skip unless the host has exactly `WORLD_SIZE` TPU chips. + + A topology *string* existing for a chip count (`hardware.get_tpu_topology`) isn't enough to guarantee the + PJRT client can actually form that session: a sub-slice of a larger single host (e.g. claiming 2 of a + 4-chip v6e-4's chips via `TORCH_TPU_TOPOLOGY`/`TORCH_TPU_SLICEBUILDER_ADDRESSES`) can still fail with a + low-level `START_SESSION` GRPC error, since the runtime's session setup is tied to the host's actual + provisioned slice, not just a topology label. The only combination verified to work is running with exactly + as many ranks as the host has chips. + """ + from torch_tpu._internal.utils import hardware + + try: + device_count = hardware.get_tpu_device_count() + except Exception as e: # pragma: no cover - defensive, hardware detection is best-effort + pytest.skip(f"Could not determine local TPU chip count: {e}") + return + + if device_count != self.WORLD_SIZE: + pytest.skip( + f"This host exposes {device_count} TPU chip(s), but this test requires exactly " + f"{self.WORLD_SIZE} (a TPU single-host tensor-parallel job must use all chips on the host; " + f"sub-slicing a larger host is not reliably supported by the runtime). Run this test on a host " + f"with exactly {self.WORLD_SIZE} TPU chips." + ) + + def test_tensor_parallel_tpu_inference(self): + self.skip_if_unsupported() + _run_tp_worker_subprocess( + "_tpu_tp_worker.py", + self.TP_SPEC, + world_size=self.WORLD_SIZE, + timeout_s=self.TIMEOUT_S, + extra_args=[f"--atol={self.TP_ATOL}", f"--rtol={self.TP_RTOL}"], + ) + + @is_context_parallel @require_torch_multi_accelerator class ContextParallelTesterMixin: diff --git a/tests/models/transformers/_tpu_tp_worker.py b/tests/models/transformers/_tpu_tp_worker.py new file mode 100644 index 000000000000..c9df01c4846d --- /dev/null +++ b/tests/models/transformers/_tpu_tp_worker.py @@ -0,0 +1,119 @@ +# coding=utf-8 +# Copyright 2026 HuggingFace Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Generic torchrun worker: assert a model's TPU tensor-parallel output matches its single-chip reference. + +Model-agnostic. The model under test is supplied as a `module:function` spec reference on the command line; the +referenced factory returns `(model_class, init_dict, inputs)` with CPU tensors, so all model-specific test data lives +with the launching test rather than here. + +Launched as a subprocess by `TensorParallelTPUTesterMixin` (and runnable directly for debugging):: + + eval $(python -m torch_tpu._internal.distributed.launchers.singlehost_wrapper | sed 's/^/export /') + torchrun --nproc_per_node=4 _tpu_tp_worker.py \\ + tests.models.transformers.test_models_transformer_flux2:make_tpu_tp_spec + +Exit code 0 means the TP path is numerically equivalent to the unsharded model; non-zero means failure. +""" + +import argparse +import copy +import importlib +import os +import sys +import traceback + + +# Make the in-repo `diffusers` and `tests` packages importable when run via torchrun from an arbitrary CWD. +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "..", "src")) +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "..")) + +import torch +import torch.distributed as dist +import torch_tpu # noqa: F401 — registers "tpu" device and "tpu_dist" backend +from torch.distributed.device_mesh import DeviceMesh +from torch_tpu._internal import sync as tpu_sync + +from diffusers import TensorParallelConfig + + +def _synchronize(): + tpu_sync.synchronize(None, wait=True) + + +def main(): + parser = argparse.ArgumentParser(description="TPU tensor-parallel correctness worker.") + parser.add_argument( + "spec", + help="`module:function` reference returning (model_class, init_dict, cpu_inputs) for the model under test.", + ) + parser.add_argument("--atol", type=float, default=1e-3) + parser.add_argument("--rtol", type=float, default=1e-3) + args = parser.parse_args() + module_name, _, fn_name = args.spec.partition(":") + model_class, init_dict, inputs = getattr(importlib.import_module(module_name), fn_name)() + + dist.init_process_group(backend="tpu_dist") + rank = dist.get_rank() + tp_size = dist.get_world_size() + + # TPU runs fp32 matmuls in bf16 by default, which would put the reference and the TP output ~1e-2 apart and force a + # tolerance loose enough to hide a sharding bug. At the highest precision they agree to ~1e-7 for most models. + torch.set_float32_matmul_precision("highest") + + # Identical weights on every rank (same seed). + torch.manual_seed(0) + model = model_class(**init_dict).eval() + + # Single-chip (unsharded) reference on the TPU rather than the CPU, so the reference and the TP pass run the same + # kernels and only the sharding differs. + ref_model = copy.deepcopy(model).to("tpu") + inputs_on_device = {k: v.to("tpu") if isinstance(v, torch.Tensor) else v for k, v in inputs.items()} + with torch.no_grad(): + ref_output = ref_model(**inputs_on_device, return_dict=False)[0] + _synchronize() + ref_output = ref_output.float().cpu() + del ref_model + + model.enable_parallelism(config=TensorParallelConfig(mesh=DeviceMesh("tpu", list(range(tp_size))))) + model = model.to("tpu") + with torch.no_grad(): + tp_output = model(**inputs_on_device, return_dict=False)[0] + _synchronize() + tp_output = tp_output.float().cpu() + + if rank == 0: + assert tp_output.shape == ref_output.shape, f"shape mismatch: {tp_output.shape} vs {ref_output.shape}" + assert torch.isfinite(tp_output).all(), "TP output contains non-finite values" + max_abs = (tp_output - ref_output).abs().max().item() + denom = ref_output.abs().max().item() + 1e-6 + print( + f"[rank0] tp_size={tp_size} output_shape={tuple(tp_output.shape)} " + f"max_abs_diff={max_abs:.4e} max_rel_diff={max_abs / denom:.4e}" + ) + torch.testing.assert_close(tp_output, ref_output, atol=args.atol, rtol=args.rtol) + print("[rank0] PASS: TPU tensor-parallel output matches single-device reference.") + + dist.barrier() + dist.destroy_process_group() + + +if __name__ == "__main__": + try: + main() + except Exception: + traceback.print_exc() + # Ensure a non-zero exit so the launching pytest sees the failure. + os._exit(1) diff --git a/tests/models/transformers/test_models_transformer_flux.py b/tests/models/transformers/test_models_transformer_flux.py index be76f892fc4c..2703103f5cdf 100644 --- a/tests/models/transformers/test_models_transformer_flux.py +++ b/tests/models/transformers/test_models_transformer_flux.py @@ -53,6 +53,7 @@ SingleFileTesterMixin, TaylorSeerCacheTesterMixin, TensorParallelTesterMixin, + TensorParallelTPUTesterMixin, TorchAoCompileTesterMixin, TorchAoTesterMixin, TorchCompileTesterMixin, @@ -276,6 +277,19 @@ class TestFluxTransformerTensorParallel(FluxTransformerTesterConfig, TensorParal """Tensor Parallel inference tests for Flux Transformer (CUDA/XPU multi-accelerator).""" +def make_tpu_tp_spec(): + """Model spec for `_tpu_tp_worker.py`, with 4 heads so the 4 TPU ranks of `TensorParallelTPUTesterMixin` divide them.""" + config = FluxTransformerTesterConfig() + init_dict = {**config.get_init_dict(), "num_attention_heads": 4} + return FluxTransformer2DModel, init_dict, config.get_dummy_inputs(device="cpu") + + +class TestFluxTransformerTensorParallelTPU(TensorParallelTPUTesterMixin): + """Tensor Parallel inference test for Flux Transformer on TPU.""" + + TP_SPEC = "tests.models.transformers.test_models_transformer_flux:make_tpu_tp_spec" + + def make_neuron_tp_spec(): """Model spec consumed by the generic Neuron TP worker (`_neuron_tp_worker.py`). diff --git a/tests/models/transformers/test_models_transformer_flux2.py b/tests/models/transformers/test_models_transformer_flux2.py index 3263ce68202c..3df62b6020c6 100644 --- a/tests/models/transformers/test_models_transformer_flux2.py +++ b/tests/models/transformers/test_models_transformer_flux2.py @@ -42,6 +42,7 @@ ModelTesterMixin, SingleFileTesterMixin, TensorParallelTesterMixin, + TensorParallelTPUTesterMixin, TorchAoCompileTesterMixin, TorchAoTesterMixin, TorchCompileTesterMixin, @@ -176,6 +177,19 @@ def make_neuron_tp_spec(): return Flux2Transformer2DModel, config.get_init_dict(), config.get_dummy_inputs(device="cpu") +def make_tpu_tp_spec(): + """Model spec for `_tpu_tp_worker.py`, with 4 heads so the 4 TPU ranks of `TensorParallelTPUTesterMixin` divide them.""" + config = Flux2TransformerTesterConfig() + init_dict = {**config.get_init_dict(), "num_attention_heads": 4} + return Flux2Transformer2DModel, init_dict, config.get_dummy_inputs(device="cpu") + + +class TestFlux2TransformerTensorParallelTPU(TensorParallelTPUTesterMixin): + """Tensor Parallel inference test for Flux2 Transformer on TPU.""" + + TP_SPEC = "tests.models.transformers.test_models_transformer_flux2:make_tpu_tp_spec" + + @is_tensor_parallel @require_torch_neuron class TestFlux2TransformerTensorParallelNeuron: diff --git a/tests/models/transformers/test_models_transformer_qwenimage.py b/tests/models/transformers/test_models_transformer_qwenimage.py index 5fcf37f6ff3f..a322d686fa11 100644 --- a/tests/models/transformers/test_models_transformer_qwenimage.py +++ b/tests/models/transformers/test_models_transformer_qwenimage.py @@ -37,6 +37,7 @@ MemoryTesterMixin, ModelTesterMixin, TensorParallelTesterMixin, + TensorParallelTPUTesterMixin, TorchAoTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, @@ -307,6 +308,22 @@ class TestQwenImageTransformerTensorParallel(QwenImageTransformerTesterConfig, T """Tensor Parallel inference tests for QwenImage Transformer (CUDA/XPU multi-accelerator).""" +def make_tpu_tp_spec(): + """Model spec for `_tpu_tp_worker.py`; the shared config's 4 heads already divide the 4 TPU ranks.""" + config = QwenImageTransformerTesterConfig() + return QwenImageTransformer2DModel, config.get_init_dict(), config.get_dummy_inputs(device="cpu") + + +class TestQwenImageTransformerTensorParallelTPU(TensorParallelTPUTesterMixin): + """Tensor Parallel inference test for QwenImage Transformer on TPU.""" + + TP_SPEC = "tests.models.transformers.test_models_transformer_qwenimage:make_tpu_tp_spec" + # On TPU the sharded and unsharded QwenImage outputs differ by ~1e-2 even at the highest matmul precision. The + # same comparison on CPU agrees to ~1e-7, so the gap is TPU numerics depending on the shard shapes, not the plan. + TP_ATOL = 2e-2 + TP_RTOL = 2e-2 + + def make_neuron_tp_spec(): """Model spec consumed by the generic Neuron TP worker (``_neuron_tp_worker.py``). diff --git a/tests/testing_utils.py b/tests/testing_utils.py index cce6f15325b4..a22da6de08e5 100644 --- a/tests/testing_utils.py +++ b/tests/testing_utils.py @@ -46,6 +46,7 @@ is_timm_available, is_torch_available, is_torch_neuronx_available, + is_torch_tpu_available, is_torch_version, is_torchao_available, is_torchsde_available, @@ -566,6 +567,14 @@ def require_torch_neuron(test_case): )(test_case) +def require_torch_tpu(test_case): + """Decorator marking a test that requires a TPU device (torch_tpu).""" + return pytest.mark.skipif( + not is_torch_tpu_available(), + reason="test requires TPU device (torch_tpu)", + )(test_case) + + def require_torch_multi_gpu(test_case): """ Decorator marking a test that requires a multi-GPU setup (in PyTorch). These tests are skipped on a machine without