Skip to content

[TPU] TorchTPU backend integration - eager / torch.compile / tp - #14039

Open
JingyaHuang wants to merge 66 commits into
huggingface:mainfrom
JingyaHuang:add-torchtpu-support
Open

JingyaHuang wants to merge 66 commits into
huggingface:mainfrom
JingyaHuang:add-torchtpu-support

Conversation

@JingyaHuang

@JingyaHuang JingyaHuang commented Jun 22, 2026 •

Copy link
Copy Markdown
Contributor

What does this PR do?

Need the fix #14739 and perferrably merge the sharding improvement PR #14544 first.

This is a preparation based on TorchTPU beta before the official release.

Who can review?

Anyone in the community is free to review the PR once the tests have passed. Feel free to tag
members/contributors who may be interested in your PR.

@github-actions github-actions Bot added documentation Improvements or additions to documentation models utils pipelines size/L PR with diff > 200 LOC labels Jun 22, 2026
@JingyaHuang JingyaHuang changed the title [TPU] Initial TorchTPU backend integration (eager + torch.compile) [TPU] TorchTPU backend integration - eager / torch.compile / tp Sep 7, 2026
@JingyaHuang
JingyaHuang marked this pull request as ready for review September 7, 2026 16:23
…ster tpu.md; drop redundant execution_device check

- Propagate the text_encoder.device-based fix (introduced for TPU CPU-offload
  support) from FluxPipeline/Flux2KleinPipeline/WanPipeline into their
  `# Copied from` copies (flux/*, flux2_klein_inpaint, visualcloze, anyflow,
  chronoedit, lucy_edit, skyreels_v2/*). SDXL-family copies of
  StableDiffusionXLPipeline.encode_prompt are intentionally left untouched;
  they'll be handled in a follow-up PR that fixes device placement for every
  pipeline component (not just text encoders).
- Register docs/source/en/optimization/tpu.md in _toctree.yml (was breaking
  the docs build: "not present in the table of contents").
- Remove the redundant "prefer non-CPU, non-meta component" loop from
  DiffusionPipeline._execution_device: PR huggingface#14383 already fixed this in
  DiffusionPipeline.device, which _execution_device falls back to. Verified
  on TPU hardware that _execution_device still resolves correctly for a
  split-placement pipeline after the removal.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
@github-actions github-actions Bot removed the pipelines label Oct 1, 2026
@JingyaHuang
JingyaHuang requested a review from sayakpaul October 2, 2026 12:33
@sayakpaul sayakpaul added this to the Release 0.42.0 milestone Oct 5, 2026

@sayakpaul sayakpaul left a comment •

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nice, thanks! The changes are minimal and we should be able to merge soon.

(Ran the updated tests on 2 A10Gs and they worked: https://huggingface.co/jobs/sayakpaul/6ac30cecfbc85ba6823a4556)

Comment on lines +19 to +22
| 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` |

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we present a faithful comparison between their latency numbers?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Not at this stage but maybe in the future, need to ask TPU team.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Sorry but why do we need to ask the TPU team about presenting compilation numbers?

Comment thread docs/source/en/optimization/tpu.md
Comment thread docs/source/en/optimization/tpu.md Outdated
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),

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think we should provide a working example here. Or we could use transformers' support for TP?

Also, the same argument (a model to be placed on a TPU chip is too large) applies to the DiT. I think we could make an entirely separate section to talk about TP, etc.

Or, we could include a statement like "For details on how to shard the denoising module, refer to the "## Tensor parallelism" section."

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yeah, it would be better just show a simple example with everything ran on tpu w/o tp in this section, and have an example with TP on the TP dedicated section.

Comment thread docs/source/en/optimization/tpu.md Outdated
Comment thread docs/source/en/optimization/tpu.md Outdated
Comment on lines +102 to +103
pipe.transformer = torch.compile(pipe.transformer, backend="tpu", fullgraph=True, dynamic=False)
pipe.vae = torch.compile(pipe.vae, backend="tpu", fullgraph=True, dynamic=False)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Not a necessity but:

  1. We can simplify this by calling the module-level compile method like transformer.compile().
  2. Have we verified if compile_repeated_blocks() works on the transformer? It is related to regional compilation.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

transformer.compile() with offload fails when setting fullgraph=True, because accelerate's CpuOffload.pre_forward is wrapped in torch.compiler.disable.
https://github.com/huggingface/accelerate/blob/7d8824f10f8be1ea405d611c98d51ebfbf4821cd/src/accelerate/hooks.py#L748-L757

And yes, compile_repeated_blocks() has been verified on the TPU. All components run on tpu:0, and the output is correct.

Comment thread docs/source/en/optimization/tpu.md Outdated
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)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is so good! Thanks for the work, here!

Comment thread tests/models/transformers/_tpu_tp_worker.py Outdated
Comment thread tests/models/testing_utils/parallelism.py Outdated
Comment thread tests/models/testing_utils/parallelism.py Outdated
@JingyaHuang
JingyaHuang requested a review from sayakpaul October 6, 2026 13:48
Comment on lines +62 to +64
> 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.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

the code example below compiles only the transformer's repeated blocks, so changing num_inference_steps shouldn't trigger a recompile?

Comment thread docs/source/en/optimization/tpu.md Outdated

repo_id = "black-forest-labs/FLUX.2-dev"
text_encoder = Mistral3ForConditionalGeneration.from_pretrained(
repo_id, subfolder="text_encoder", dtype=torch.bfloat16, tp_plan="auto", device_mesh=mesh

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

i think passing tp_plan in from_pretrained is going to be removed soon, so it may be better to use DistributedConfig here instead

@sayakpaul sayakpaul left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks! I left some further comments.

Comment on lines +19 to +22
| 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` |

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Sorry but why do we need to ask the TPU team about presenting compilation numbers?


def test_tensor_parallel_tpu_inference(self):
from torch_tpu._internal.distributed.launchers.singlehost_wrapper import prepare_tpu_environment
from torch_tpu._internal.utils import hardware

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we not also skip the tests, if the underlying model class doesn't have any _tp_plan defined?

Comment on lines +507 to +508
tp_atol = 1e-3
tp_rtol = 1e-3

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should these be configurable from test_tensor_parallel_tpu_inference? Easy to override no?

sayakpaul and others added 2 commits October 7, 2026 11:55
Co-authored-by: Steven Liu <59462357+stevhliu@users.noreply.github.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

documentation Improvements or additions to documentation hooks models size/L PR with diff > 200 LOC tests utils

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants