Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions _additional_platforms/torch_tpu.json
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
{
"name": "TorchTPU",
"support_channel": "https://github.com/google-pytorch/torch_tpu/issues",
"stable": {
"linux": "pip3 install torch --index-url https://download.pytorch.org/whl/cpu && pip3 install torch-tpu"
}
}
57 changes: 57 additions & 0 deletions _get_started/additional_platforms/torch_tpu.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
# Installing on Google Cloud TPU (TorchTPU)

**TorchTPU (`torch_tpu`)** is a native PyTorch backend built for Google Cloud Tensor Processing Units (TPUs). It enables Google Cloud TPUs to run PyTorch workloads natively by integrating as a `PrivateUse1` backend using the dispatch key `"tpu"`. TorchTPU supports eager execution, `torch.compile()`, distributed training (`torch.distributed`, `DTensor`, `FSDP2`), and custom kernels (`Pallas`, `Helion`).

## Prerequisites

### Hardware Requirements

* A provisioned Google Cloud TPU VM (TPU v5e, v5p, v6e, or v7x) on Google Compute Engine (GCE), access to Google Kubernetes Engine (GKE) cluster, or a TPU runtime in Google Colab.

### Software Requirements

* Linux (Ubuntu 22.04+ recommended)
* Python >= 3.10
* `libtpu` runtime library (automatically installed with `torch-tpu`)

## Installation

### pip

```bash
pip3 install torch --index-url https://download.pytorch.org/whl/cpu && pip3 install torch-tpu
```

Use the `pip` package manager to install the CPU build of PyTorch alongside `torch-tpu`. Select your preferred options in the selector above to get the installation command.

## Verification

To ensure that PyTorch was installed correctly with TorchTPU support, run the following code:

```python
import torch
device = torch.device("tpu")

Comment thread
ejmartinezm marked this conversation as resolved.
x = torch.randn(2, 2, device="tpu")
y = torch.randn(2, 2, device="tpu")
z = x.mm(y)

print(z)
```

The following, or a similar output, indicates successful installation:

```bash
tensor([[-0.4218, 0.8912],
[ 0.1534, -1.1045]], device='tpu:0')
```

## Documentation

For more information, please visit:

* [TorchTPU Official Documentation & User Guide](https://google-pytorch.github.io/torch_tpu/)
* [Supported vs. Unsupported Feature Matrix](https://github.com/google-pytorch/torch_tpu/blob/main/docs/features_matrix.md)
* [Google Cloud TPU Documentation](https://cloud.google.com/tpu/docs)
* [GitHub Repository](https://github.com/google-pytorch/torch_tpu)
* [PyPI](https://pypi.org/project/torch-tpu/)