Library / SDK
pytorch/xla avatar
pytorch/xla

PyTorch/XLA: Running PyTorch on TPUs, and What the TorchTPU Announcement Means for It

Enabling PyTorch on XLA Devices (e.g. Google TPU)

2,801 stars572 forksC++NOASSERTION

At a glance

What is it?
PyTorch/XLA is the Python package that connects PyTorch to Cloud TPUs through the XLA compiler. It works, it ships regular releases, and Google has said TorchTPU will replace it once public, which changes how you should plan around it.
Who is it for?
Adopt PyTorch/XLA if you already have PyTorch training code and TPU capacity, and you accept the migration risk that the README's own notice creates. Do not adopt it for GPU or CPU work, where plain PyTorch and CUDA are the shorter path, and do not start a long-lived platform project on it without reading issue #9684 first.
Can I use it commercially?
Check first. The repository uses a licence we do not classify automatically, so read its LICENSE file before any commercial use.
Is it still maintained?
Yes. The repository last received commits 125 days ago.
What is it written in?
Mainly C++, according to GitHub's language statistics.

Answers come from the project's GitHub data, last synced on September 29, 2026, and from our analysis. They are not legal advice.

Editorial analysis

What PyTorch/XLA is for, and who actually needs it

PyTorch/XLA exists for one situation: you have a PyTorch model and you want it to run on Google Cloud TPUs. The README describes the package as using the XLA deep learning compiler to connect the PyTorch framework and Cloud TPUs. That sentence is the whole scope. It is not a general accelerator abstraction, and it is not a way to make PyTorch faster on hardware you already own.

The audience is narrow and specific. You are likely already training in PyTorch, you have access to TPU capacity (a Cloud TPU VM, or a free single TPU VM through Kaggle, which the README points at), and rewriting the model in JAX is not something you want to do. PyTorch/XLA is the bridge that lets you keep the training loop you wrote.

The repository is not archived, and the last push was on 2026-05-27. Releases have been regular: v2.7.0 on 2025-04-24, v2.8.0 on 2025-08-13, v2.9.0 on 2025-11-17. That cadence is real. But the README opens with a note dated 4/22/2026 stating that once TorchTPU is public it will replace PyTorch/XLA, and a second note from 10/2025 pointing to an RFC at issue #9684 proposing a more native direction for PyTorch on TPU. Anyone evaluating this project has to read those two notes before reading anything else, because they determine whether the work you put in survives.

The lazy tensor mechanism and why tracing is the performance boundary

The core idea is that PyTorch tensors are not executed eagerly on the TPU. Operations are recorded, and the recorded graph is handed to the XLA compiler, which produces a program for the device. This is the lazy tensor model, and it explains most of the project's behaviour and most of its failure modes.

Because the graph is compiled rather than interpreted, the host CPU spends time tracing the model before the TPU does any work. The README describes exactly this symptom: if your model is tracing bound, you see the host CPU busy tracing while the TPUs sit idle. That is not a bug you can tune away with a flag; it is the cost of the compilation step.

This is also why the project ships two C++ ABI flavours. Starting from the 2.7 release, C++11 ABI builds are the default, and the README states that pre-C++11 ABI wheels are no longer provided. In 2.6 both flavours existed, and the README gives a reason: C++11 ABI wheels and docker images have better lazy tensor tracing performance. The README reports Mixtral 8x7B numbers on v5p-256 at global batch size 1024, with pre-C++11 ABI at 33% MFU and C++ ABI at 39%. Those are the project's own figures for one model on one slice, not a general promise, but they show the size of the effect when tracing is your bottleneck.

If your model is compute bound rather than tracing bound, none of this helps you. The ABI choice is a tracing optimisation, and it only pays off where tracing is what the host is waiting on.

Installing PyTorch/XLA on a TPU VM and training a first step

The README gives the stable install as a single pip command, and it pins both torch and torch_xla to the same version. Builds are available for Python 3.8 to 3.11 according to the installation section, while a later note says that from the 2.8 release onward nightly and release wheels are provided for Python 3.11 to 3.13. Those two statements overlap awkwardly, so check which Python your VM has before you start.

bash
pip install torch==2.8.0 'torch_xla[tpu]==2.8.0'

The README suggests creating an isolated environment first, either a venv or a conda env, and shows python3.11 as the example. If you use custom kernels, there is a separate optional install for pallas dependencies, which pulls from the JAX release index and the libtpu wheels index. That optional line is in the README and is not needed for ordinary training.

For a nightly build instead of a release, the README installs torch and torchvision from the nightly CPU index, then installs a torch_xla wheel by direct URL, editing the cp312-cp312 tag to match your Python version. That is a manual step and an easy place to get the tag wrong.

The code change to an existing training loop is small. The README shows it as a diff: import torch_xla, wrap the step body in a context manager, move inputs and labels to the xla device, move the model parameters to xla before training, and call torch_xla.sync() after the loop. The README's claim is that these changes should get your model to train on the TPU. Note the word should. There is no statement that every model works without further change, and the docs directory structure suggests troubleshooting material exists for the cases where it does not.

Single process, multi process, and the SPMD mode the tutorial skips

PyTorch/XLA has three execution shapes, and the README's getting-started guide covers only two of them. Single process means one Python interpreter controlling one TPU. Multi process means N interpreters, one per TPU on the system. SPMD means one interpreter controlling all N TPUs.

The README is explicit that multi processing is more complex and is not compatible with SPMD, and that the tutorial does not cover SPMD, pointing instead at a separate SPMD guide. That is an honest boundary. It also means the shortest path in the documentation leads you to the model that does not scale across the whole slice without reading further.

This matters for planning. If you follow only the single-process example, you are using one TPU and the rest of the slice is idle. Moving to SPMD is a different programming model, not a configuration flag, and the docs spread across docs/source/perf and docs/source/features are where that work lives. Budget for reading them.

The repository layout supports that reading. There are separate documentation trees for learn, accelerators, perf, features and contribute, plus example directories for data_parallel, fsdp, scan, host_offloading and flash_attention. The examples directory is the practical map of what the project supports, and its breadth is a reasonable proxy for how much of the distributed story is actually exercised.

Where PyTorch/XLA is the wrong tool

The clearest case against it is hardware. If you are on NVIDIA GPUs, PyTorch with CUDA is the default path and PyTorch/XLA adds a compiler between you and the device for no benefit. The README frames the package around Cloud TPUs throughout; nothing in it positions XLA as a faster route on GPU. There is a CPU PJRT plugin under plugins/cpu, but the README links it as a plugin README rather than as a recommended deployment target.

The second case is version lock. torch and torch_xla are pinned to matching versions in every install command shown, and the nightly path requires you to hand-edit a wheel URL tag. If your team needs to track PyTorch releases quickly, this coupling is friction. You are not upgrading one library; you are moving two in lockstep.

The third case is the one the README itself raises. A project that has announced its own replacement is a project where long-lived infrastructure investment carries a migration cost that is not yet quantified. The README does not document a migration path to TorchTPU, does not give a timeline beyond once TorchTPU is public, and does not say what happens to the PyTorch/XLA API surface. If you cannot tolerate that uncertainty, this is the wrong foundation regardless of how well it works today.

PyTorch/XLA compared with JAX on the same hardware

The natural alternative on TPUs is JAX, and the difference is not performance, it is where the compilation boundary sits. JAX is built around XLA from the start: jit, vmap and pmap are the programming model, and the compiler is not something you opt into. PyTorch/XLA goes the other way. You write ordinary PyTorch, and the bridge records your operations and hands them to XLA. The README's diff-based tutorial is the clearest illustration of the difference: the selling point is that your existing loop keeps working with a handful of edits.

That is a real trade. Keeping the PyTorch loop means keeping PyTorch semantics and the PyTorch ecosystem, but it also means the tracing boundary is something you have to reason about, and the README's own note about tracing-bound models is the cost of that arrangement. A JAX user does not see that note because the model is compiled by construction.

If your team already knows PyTorch and the model is large enough that a rewrite is out of the question, the bridge is worth the tracing complexity. If you are starting fresh on TPUs with no PyTorch investment, the argument for going through a bridge is weaker, and the RFC at issue #9684 proposing a more native direction for PyTorch on TPU is essentially the project acknowledging the same tension.

Licence, maintenance and the cost of staying current

The repository reports NOASSERTION as its licence identifier, and pyproject.toml declares license = { file = "LICENSE" } with a classifier of "License :: OSI Approved :: BSD License". Those two signals do not agree, and the classifier is a declaration rather than the text. Read the LICENSE file in the repository root before you depend on it, and treat the classifier as a hint rather than an answer. This is not legal advice; it is a note that the metadata is inconsistent enough to warrant a look.

On maintenance, the facts are: not archived, last push on 2026-05-27, releases at a steady cadence through v2.9.0 on 2025-11-17. That is a project that is being worked on. It is also a project whose README carries a replacement notice, so the maintenance question is not whether commits land but how long they will keep landing under this name.

The upgrade cost is the pinned pair. Every install command in the README pins torch and torch_xla to the same version, and the C++ ABI default changed at 2.7, which means an upgrade across that boundary can change your tracing performance in either direction. If you are on pre-2.7 wheels, moving forward means moving to C++11 ABI, and the README's Mixtral figures suggest that is usually an improvement for tracing-bound models but is still a change you should measure rather than assume.

Editorial conclusion

Adopt PyTorch/XLA if you already have PyTorch training code and TPU capacity, and you accept the migration risk that the README's own notice creates. Do not adopt it for GPU or CPU work, where plain PyTorch and CUDA are the shorter path, and do not start a long-lived platform project on it without reading issue #9684 first. Before committing, verify three things on your own VM: that pip install torch==2.8.0 'torch_xla[tpu]==2.8.0' resolves on your Python version, that your training loop produces correct results under torch_xla.step() and torch_xla.sync(), and that the TorchTPU timeline has moved since the 4/22/2026 blog post.

Frequently asked questions

What is an XLA compiler?

XLA is the deep learning compiler that PyTorch/XLA uses to connect PyTorch to Cloud TPUs. It takes the operations recorded from your PyTorch model and produces a program for the device, which is why the host CPU spends time tracing before the TPU runs.

What is PyTorch and why is it used?

PyTorch is the deep learning framework you write the training loop in, and PyTorch/XLA exists so that loop can run on TPU hardware. The README's tutorial is written as a diff against an existing PyTorch training loop, which shows the intended workflow of keeping PyTorch code and adding the XLA bridge.

Is PyTorch just Python?

No. PyTorch/XLA is a Python package, but the repository is primarily C++ and builds a compiled XLA client through bazel, with setup.py exposing build flags such as TPUVM_MODE and BUNDLE_LIBTPU. The Python layer sits on top of that compiled bridge.

What is PyTorch/XLA?

It is a Python package that uses the XLA deep learning compiler to connect the PyTorch framework and Cloud TPUs. The README describes it as the bridge that lets an existing PyTorch training loop run on TPU hardware with a small set of code changes.

What is the difference between PyTorch and XLA?

PyTorch is the deep learning framework you write the model in; XLA is the compiler that turns the recorded operations into a program for the device. PyTorch/XLA is the package that connects the two, so PyTorch is the interface and XLA is the execution layer underneath it.

Official sources

  1. Issues
  2. Project website
  3. pytorch/xla on GitHub
  4. README
  5. Releases
Add this badge to your README

If you maintain this project, the badge below links readers to this analysis and shows its maintenance status from the daily GitHub snapshot. Paste the markdown into your README; add ?metric=license or ?metric=stars to the image URL for a different field.

Add this badge to your README

markdown
[![Hysen Labs](https://hysenlabs.com/badge/pytorch-xla.svg)](https://hysenlabs.com/projects/pytorch-xla)