Skip to content

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

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

JingyaHuang wants to merge 64 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.

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

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