Skip to main content

metal-pjrt-plugin: run JAX on your Mac's GPU

An open-source JAX plugin that runs JAX programs on Apple Silicon GPUs through Metal. It's for people who write JAX and want their Mac's GPU to do the work: training small models, running LLM inference, or scientific code that's happy in float32.

Under the hood it treats Metal as one more XLA:GPU platform, next to CUDA and ROCm. XLA's own GPU compiler fuses your program and generates the kernels, which the plugin translates to Metal Shading Language.

The JAX platform name is "mtl" (Metal's own prefix, as in MTLDevice), so it doesn't collide with Apple's closed-source jax-metal plugin, which uses "metal". This project isn't affiliated with or endorsed by Apple or Google.

Status

Early, but it works. Float32, float16 and bfloat16 programs run end to end on one GPU: training loops, sorting, FFTs, host callbacks and float32 linear algebra. Performance is broadly comparable to MLX and PyTorch's MPS backend on the workloads we've tried, with some slow spots (see docs/performance.md). Everything so far has been tested on a single machine (an M3 with 8 GB, macOS 26.2), so expect rough edges elsewhere, and please report them.

Requirements

  • An Apple Silicon Mac. Intel Macs aren't supported.
  • macOS 26 or later.
  • Python 3.12+ with jax and jaxlib 0.10.0 or later (tested: 0.10.0 through 0.11.2 and a 0.12 nightly). Older versions load with a warning.

There's no PyPI release yet, so you build from source. That needs the Xcode command-line tools (xcode-select --install; full Xcode isn't needed), plus bazelisk and uv (brew install bazelisk uv).

Install

git clone https://github.com/dfm/metal-pjrt-plugin.git
cd metal-pjrt-plugin
scripts/install_dev.sh     # creates .venv, builds the plugin, installs it
bazel shutdown             # frees the build server's memory

The first build compiles XLA and takes a while: about 1.5-2 hours on an 8 GB M3 (95 minutes from a clean clone), with ~8 GB of build output plus a disk cache in ~/.cache/metal-pjrt-plugin/ (~4.5 GB after the first build, growing with rebuilds). Every checkout shares that cache, so later builds and fresh clones take minutes. Running the plugin needs no developer tools.

To check that it works, run the smoke tests:

.venv/bin/python -m pytest tests/test_smoke.py

To install into another environment, build the wheels with scripts/build_wheel.sh and pip install both files it puts in dist/: metal-pjrt-plugin, the Python package JAX finds, and metal-pjrt-core, the compiled library it loads. The library talks to JAX through stable interfaces, so one build works across the supported JAX versions.

Quick start

import jax
jax.config.update("jax_platforms", "mtl,cpu")  # before using any device
import jax.numpy as jnp

x = jnp.arange(4.0)
y = jax.jit(lambda v: v * 2)(x)
print(y, y.devices())  # [0. 2. 4. 6.] {MtlDevice(id=0)}

JAX prints a warning that the "mtl" platform is experimental; that's expected. You can also pick the platform from the shell:

JAX_PLATFORMS=mtl,cpu python my_script.py

The plugin is opt-in: installing it doesn't change JAX's default (CPU). To keep CPU as the default and send only some work to the GPU, place it explicitly with jax.device_put(x, jax.devices("mtl")[0]).

Before running anything heavy, read Is it safe for my GPU?: a kernel that runs too long makes macOS reset the GPU for the whole machine. The FAQ also covers what works, accuracy, memory and the compilation cache, and docs/troubleshooting.md maps the plugin's error messages to what to do.

Examples

Learn more

Contributing

Issues and pull requests are welcome. CONTRIBUTING.md says what makes a useful report, and docs/development.md covers building and testing. Please run GPU tests under scripts/device_lock.py. For security issues, see SECURITY.md.

License

Apache-2.0 (LICENSE). The plugin includes third-party code: matmul, convolution and FFT kernels ported from MLX (MIT), Apple's metal-cpp headers (Apache-2.0), and a statically linked XLA. Their notices are in THIRD_PARTY_NOTICES, which ships in both wheels.

Metadata

Release files for metal-pjrt-plugin 0.0.1

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Built distribution (wheel)

Table of built distributions (wheels) for metal-pjrt-plugin 0.0.1
File Interpreter ABI Platform
metal_pjrt_plugin-0.0.1-py3-none-any.whl Python 3 none any Details

Release files / metal_pjrt_plugin-0.0.1-py3-none-any.whl

Download URL metal_pjrt_plugin-0.0.1-py3-none-any.whl
Size 59.9 kB
Tags Python 3
SHA-256 checksum
How to use checksums
a8172c439b73d3f0a9938d39281cd19947fb90a9f2ecc042361583749055caa7
BLAKE2b-256 checksum
How to use checksums
73b0eb9caa27260db2bab9385a6611191f2d6a2568d3e19bfc5c742c19a35374
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Oct 3, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

0.0.1 This release

1 release file

0.0.0

1 release file

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page