[TPU] TorchTPU backend integration - eager / torch.compile / tp - #14039
JingyaHuang wants to merge 64 commits into
Conversation
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
…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>
Co-authored-by: Steven Liu <59462357+stevhliu@users.noreply.github.com>
…rs into add-torchtpu-support
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
There was a problem hiding this comment.
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)
| | 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` | |
There was a problem hiding this comment.
Should we present a faithful comparison between their latency numbers?
There was a problem hiding this comment.
Not at this stage but maybe in the future, need to ask TPU team.
| 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), |
There was a problem hiding this comment.
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."
There was a problem hiding this comment.
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.
| pipe.transformer = torch.compile(pipe.transformer, backend="tpu", fullgraph=True, dynamic=False) | ||
| pipe.vae = torch.compile(pipe.vae, backend="tpu", fullgraph=True, dynamic=False) |
There was a problem hiding this comment.
Not a necessity but:
- We can simplify this by calling the module-level compile method like
transformer.compile(). - Have we verified if
compile_repeated_blocks()works on thetransformer? It is related to regional compilation.
There was a problem hiding this comment.
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.
| 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) |
There was a problem hiding this comment.
This is so good! Thanks for the work, here!
What does this PR do?
Need the fix #14739 and perferrably merge the sharding improvement PR #14544 first.
torch.compileThis 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.