diff --git a/_additional_platforms/torch_tpu.json b/_additional_platforms/torch_tpu.json new file mode 100644 index 000000000000..de55b5fbdf3c --- /dev/null +++ b/_additional_platforms/torch_tpu.json @@ -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" + } +} \ No newline at end of file diff --git a/_get_started/additional_platforms/torch_tpu.md b/_get_started/additional_platforms/torch_tpu.md new file mode 100644 index 000000000000..afc23e782b49 --- /dev/null +++ b/_get_started/additional_platforms/torch_tpu.md @@ -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") + +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/)