$npx skillfedfor your agent

torchax

torchax is a library for running Jax and PyTorch together

With conditionsPyPI Software DevelopmentReleased Jun 2026164.6K downloads / mopermissive licensePure Python

Decision gist · record as of 2026-08-14

pure-Python wheel — torchax-0.0.13-py3-none-any.whl
v0.0.13 · released 2026-06-17 · Python >=3.11

Yes, if you need to run PyTorch on TPUs or want deep PyTorch–JAX interoperability. The package is actively maintained, has no known vulnerabilities, and installs easily. However, it is in alpha (Development Status :: 3), so expect incomplete operation coverage and potential API changes. Verify that your specific operations are supported before committing to production use.AI-flagged interpretation of the facts on this page — verify before relying

Before you install

  • You must install PyTorch (CPU version) and JAX for your accelerator (TPU, GPU, or CPU) before installing torchax; torchax requires Python 3.11 or later.
  • Installation is straightforward with no runtime dependencies; the package itself is a pure-Python wheel.
  • However, you must pre-install PyTorch (CPU version) and JAX separately for your target accelerator, which adds setup complexity despite low install friction for torchax itself.

License · maintenance · safety

permissive license (permissive) — Licensed under Apache License 2.0, a permissive license that allows commercial use, modification, and distribution with minimal restrictions—suitable for most use cases.

last release 2026-06-17 (58 days) · last repo commit 2026-08-07 · 237 stars

0 known vulnerabilities (OSV.dev, 2026-08-14) · 164,579 downloads/mo, #10,542 on PyPI

Verify before relying

pip install torchax

import torchax
torchax.enable_globally()

# Create model and tensors on 'jax' device
m = MyModel().to('jax')
inputs = torch.randn(3, 3, 28, 28, device='jax')
res = m(inputs)
print(res.jax())  # access underlying jax.Array
  • Extent of PyTorch operation coverage—which ops are currently implemented vs. which will fall back or error.
  • Performance overhead of the dispatch mechanism compared to native execution.
  • Stability and maturity of TPU execution path given the alpha status (Development Status :: 3).
Same gist for agents: .md · .json

What it is and what it does

torchax bridges PyTorch and JAX by implementing a custom PyTorch backend that intercepts PyTorch operations and executes them via JAX. This allows you to write standard PyTorch code but run it on TPUs and other JAX-supported accelerators, and to mix PyTorch and JAX code in the same program—calling JAX functions with tensors and vice versa.

The package works by subclassing torch.Tensor into torchax.tensor.Tensor and overriding __torch_dispatch__ to redirect operations to JAX equivalents. You enable it globally with torchax.enable_globally(), then create tensors and models on the 'jax' device. Operations run in eager mode by default, but you can wrap functions with jax.jit for compiled execution. The package also provides checkpoint save/load utilities and a JittableModule helper for easier JIT compilation of models.

Use it for

  • Run existing PyTorch models on Google Cloud TPUs without rewriting them for JAX.
  • Use JAX optimization libraries and gradient transformations to train PyTorch models.
  • Call JAX functions with compiled performance from within PyTorch code, passing jax.Arrays directly.
  • Combine a PyTorch feature extractor with a JAX model in the same training loop.
  • Leverage JAX features like GSPMD for distributed training of PyTorch models.

Worth the install?

AI-flagged interpretation of the facts on this page. Verify before relying on it.

With conditions

Yes, if you need to run PyTorch on TPUs or want deep PyTorch–JAX interoperability.

The package is actively maintained, has no known vulnerabilities, and installs easily. However, it is in alpha (Development Status :: 3), so expect incomplete operation coverage and potential API changes. Verify that your specific operations are supported before committing to production use.

Install

torchax on PyPI

Before you install

Installation is straightforward with no runtime dependencies; the package itself is a pure-Python wheel. However, you must pre-install PyTorch (CPU version) and JAX separately for your target accelerator, which adds setup complexity despite low install friction for torchax itself.

You must install PyTorch (CPU version) and JAX for your accelerator (TPU, GPU, or CPU) before installing torchax; torchax requires Python 3.11 or later.

License in practice

Licensed under Apache License 2.0, a permissive license that allows commercial use, modification, and distribution with minimal restrictions—suitable for most use cases.

Quickstart

pip install torchax

import torchax
torchax.enable_globally()

# Create model and tensors on 'jax' device
m = MyModel().to('jax')
inputs = torch.randn(3, 3, 28, 28, device='jax')
res = m(inputs)
print(res.jax())  # access underlying jax.Array

Verify before relying

  • Extent of PyTorch operation coverage—which ops are currently implemented vs. which will fall back or error.
  • Performance overhead of the dispatch mechanism compared to native execution.
  • Stability and maturity of TPU execution path given the alpha status (Development Status :: 3).

Package facts

Licensepermissive license permissive
Python supportSupports the current Python release >=3.11
Install frictionLow. Pure-Python wheel
Runtime dependenciesNone
MaintenanceActively maintained 58 days since the last release
Last repo commit
First released
Downloads164,579 / month, #10,542 on PyPI 30-day window, as of 2026-08-14
Known vulnerabilitiesNone known OSV.dev, checked 2026-08-14
Classifiers
Development Status :: 3 - AlphaIntended Audience :: DevelopersIntended Audience :: EducationIntended Audience :: Science/ResearchLicense :: OSI Approved :: Apache Software LicenseProgramming Language :: Python :: 3.11Programming Language :: Python :: 3.12Programming Language :: Python :: 3.13Topic :: Scientific/EngineeringTopic :: Scientific/Engineering :: Artificial IntelligenceTopic :: Scientific/Engineering :: MathematicsTopic :: Software DevelopmentTopic :: Software Development :: LibrariesTopic :: Software Development :: Libraries :: Python Modules

Evidence: torchax-0.0.13-py3-none-any.whl

Tags

Capabilities
pytorch on tpupytorch jax interoprun pytorch with jaxtpu backend for pytorchpytorch jax bridgejax array torch tensorpytorch tpu support
Topics
tpu-backendpytorch-jax-interopaccelerator-support

Let your AI agent find packages like this

Example. Real query, live index.

You found this page by searching. An agent finds it by wishing: SkillFed indexes 14,416 PyPI packages by what they can do, searchable in plain language.

wish › “pytorch jax interop”

  • torchaxtorchax is a PyTorch backend that runs PyTorch code on Google Cloud…
  • apache-tvm-ffiProvides a stable, minimal C ABI and FFI for machine learning systems…
  • kerasKeras 3 is a multi-backend deep learning framework supporting JAX,…

Give your agent the search over MCP, or paste the wish link into any chat.

More Software Development packages

typing-extensions Worth it
PyPI · Software Development · released Jul 2026

Provides backported and experimental type hints for Python 3.9+, allowing use of newer typing features on older Python versions and enabling early experimentation with type system PEPs before they enter the standard library.

PSF-2.0pure Python · 3.9+
1.9Bdownloads / mo
numpy Worth it
PyPI · Software Development · released Aug 2026

NumPy provides an N-dimensional array object and a comprehensive suite of mathematical, linear algebra, Fourier transform, and random number functions for scientific computing in Python.

BSD-3-Clause AND 0BSD AND MIT AND Zlib AND CC0-1.0compiled wheel · 3.12+
1.1Bdownloads / mo
fastapi Worth it
PyPI · Software Development · released Jul 2026

FastAPI is a Python web framework for building REST APIs using type hints, with automatic request validation, serialization, and interactive API documentation.

MITpure Python · 3.10+
568.6Mdownloads / mo
annotated-doc With conditions
PyPI · Software Development · released Jul 2026

Provides a way to document function parameters, class attributes, return types, and variables inline using Python's `Annotated` type hint syntax instead of traditional docstrings.

MITpure Python · 3.9+
456.2Mdownloads / mo
typer Worth it
PyPI · Software Development · released Aug 2026

Typer builds command-line applications from Python functions using type hints, automatically generating help text, argument parsing, and shell completion.

Install it if you are building CLIs in Python.

MITpure Python · 3.10+
369.3Mdownloads / mo
distlib With conditions
PyPI · Software Development · released Jun 2026

Distlib provides low-level packaging utilities for building, distributing, and managing Python software—including metadata handling, version specifiers, wheel support, script installation, and dependency resolution.

permissive licensepure Python
323.3Mdownloads / mo

See also jax-cuda13-pjrt · libtpu · jaxlib · jax-cuda12-plugin · tokamax · jax-cuda12-pjrt · jax-cuda13-plugin · jax · pathwaysutils · tpu-inference