# JAX's three transformations, the two packages that version apart, and a research project warning

> JAX is built around grad, jit and vmap, and composing them is where its power and its friction both come from. Underneath, jax and jaxlib carry separate version numbers, the test suite turns warnings into errors, and scaling is three different programming models rather than one.

**jax-ml/jax** — Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more.

- Repository: https://github.com/jax-ml/jax
- Website: https://docs.jax.dev
- Stars: 36,338 · Forks: 3,809
- Language: Python
- License: Apache-2.0
- Published: 2026-08-08 · Updated: 2026-08-18 · Language: en
- Canonical page: https://hysenlabs.com/projects/jax-ml-jax

## grad re-evaluates the function, so a branch inside it runs twice

The absolute value example makes the mechanism visible, because the comment on the second line is the whole warning:

```python
def abs_val(x):
  if x > 0:
    return x
  else:
    return -x

abs_val_grad = jax.grad(abs_val)
print(abs_val_grad(1.0))   # prints 1.0
print(abs_val_grad(-1.0))  # prints 1.0 (abs_val is re-evaluated)
```

At x equal to minus one the derivative is plus one, not minus one, because the function ran again and took the other branch. The consequence for a reader is that a function you differentiate is a function that gets called again, so anything with a side effect, a print, a file write, or an expensive step inside it will run more than once per gradient. The same property is why JAX can differentiate through loops, branches, recursion and closures, and why derivatives of derivatives are just `jax.grad(jax.grad(jax.grad(tanh)))`.

## jit narrows the Python control flow your function is allowed to use

Automatic differentiation accepting arbitrary Python is not the same thing as compilation accepting it. The section on jit says plainly that using `jax.jit` constrains the kind of Python control flow the function can use, and the remedy offered is a tutorial on control flow and logical operators with JIT rather than a config flag. The tension is structural: grad traces your Python, jit compiles it, and the second pass can only keep control flow it can trace. A branch on a traced value is the case that changes shape between the two, because the compiler has to commit to a path. The consequence for a reader is that code which differentiates perfectly can fail the moment you wrap it in jit, and the error will come from the compiler rather than from your logic. Composition is the point, and it is also where the sharp edges accumulate.

## vmap pushes the loop into the primitives instead of wrapping them

The stated difference from writing a loop yourself is that vmap does not loop over function applications, it pushes the loop down onto the function's primitive operations, turning matrix-vector multiplies into matrix-matrix multiplies. That is where the performance story comes from, and it is also where the semantic risk sits, because the function no longer runs once per element in the way you wrote it. The example guards against exactly this:

```python
def l1_distance(x, y):
  assert x.ndim == y.ndim == 1  # only works on 1D inputs
  return jnp.sum(jnp.abs(x - y))
```

The consequence is that a shape assumption written as an assert only holds for the shape you tested. Stack vmap on it the way the next lines do and the inner function sees a different rank than the comment promised, and nothing in the assert fires. Composed with grad and jit, `per_example_grads = jax.jit(jax.vmap(jax.grad(loss), in_axes=(None, 0, 0)))` is the shape of code you would have to write by hand otherwise.

## jax and jaxlib carry separate version numbers and an rc0 suffix

The packaging is where friction moves from theory into your installer. `setup.py` holds four version values: the project name is `jax`, the current jaxlib version is `0.11.2`, the latest jaxlib version on PyPI is `0.11.2`, and the libtpu version is pinned as `0.0.48.*`. A minimum jaxlib version is read from the version module, and if the build is a pre-release, the string `rc0` is appended to both the minimum and the current jaxlib version. Both live in this one repository, with `jax/` and `jaxlib/` as sibling directories, and the released tags are named for the jax side, such as jax-v0.11.2 published 2026-09-17. The consequence is that installing jax is a two package negotiation, and a release candidate is resolved by string suffix rather than by a separate channel, so a resolver that ignores the rc0 convention can pair a jax with a jaxlib it was never tested against.

## The test suite turns warnings into errors and keeps a deprecation ledger

`pyproject.toml` starts its pytest configuration with `filterwarnings` set to `error`, and everything after it is a named exception. The exceptions name the things being retired: an ignore for `jax.cloud_tpu_init was deprecated` with a note to remove it once the symbol is gone, an ignore for transparent hugepages until a specific change lands, exceptions for array_api_tests warnings that are not yet stable, and a long escape for the tensorflow protobuf import inside jax.profiler, which is documented as failing with python 3.12 and some protobuf versions. That last comment also records that protobuf versions in TF releases can lag the ones in code. The consequence is two sided. A green run says a lot about the library, and says nothing about which deprecations your own code is about to hit, since those are exactly the warnings the project suppresses for itself.

## Scaling is three different programming models, not one switch

The scaling section offers three ways to spread a computation over thousands of devices, and a table states what each one costs you:

| Mode | View? | Explicit sharding? | Explicit Collectives? |
|---|---|---|---|
| Auto | Global | No | No |
| Explicit | Global | Yes | No |
| Manual | Per-device | Yes | Yes |

Compiler-based automatic parallelization has you program as if on a single global machine and lets the compiler choose how to shard, with some constraints you supply. Explicit sharding keeps the global view but puts shardings into the JAX types, inspectable with `jax.typeof`. Manual per-device programming via shard_map gives you a per-device view and explicit collectives. The consequence is that choosing a scaling mode is choosing what you are allowed to know, and the example shows the cost concretely: with a mesh of one data axis, `jax.typeof(W)` reports `f32[512@data,512]` while `jax.typeof(b)` reports `f32[512]`, so the same loop prints parameters at two different shardings.

## Building runs through Bazel, installing does not

The repository carries two build systems and they are for different people. Contributors and CI use Bazel: there is a .bazelrc, a pinned .bazelversion, BUILD.bazel, MODULE.bazel, bazel_downloader.cfg, a test_shard_count.bzl, a third_party/ directory, and a ci/ directory. Everyone installing from an index uses setuptools instead, with setup.py, build_wheel.py and a pyproject.toml whose build-system requires setuptools and wheel. pytest markers in that same file confirm the shape of the test run, since one marker declares that a test can use and may require multiple accelerators and another marks a test for Slurm multinode GPU nightly CI. The consequence is that a bug report or a patch will be evaluated against a Bazel build you never install locally, and the configuration that governs how tests shard across accelerators lives in files a pip user never sees.

## The examples are mnist and infrastructure glue, not model training

The `examples/` directory says what the project considers a worked example: mnist_classifier.py, mnist_classifier_fromscratch.py, mnist_vae.py, spmd_mnist_classifier_fromscratch.py, advi.py, datasets.py and differentially_private_sgd.py, alongside infrastructure directories for ffi, jax_cpp, k8s and a converter named onnx2xla.py. That is a deliberately small surface, and the consequence is specific: nothing here shows a large model trained across many devices, so the sharding and pmap-style work you would actually need at scale is left for you to write. The `jax_plugins/` directory is the other clue to how this grows, since a plugin namespace is where backends live outside the core package. Contributors should also look at benchmarks/ and cloud_tpu_colabs/ for how work is measured and run.

## Conclusion

JAX suits people who want to compose transformations rather than call a framework API, who are comfortable reading a compiler error, and who need reverse-mode gradients through control flow the derivative would not otherwise survive. It does not suit anyone looking for a stable supported surface, because the project calls itself a research project rather than an official Google product and tells you to expect sharp edges. Before you commit, pin both jax and jaxlib, since they version separately and pre-release wheels rely on an rc0 suffix, and read the Common Gotchas page the README links rather than discovering the rules from tracebacks. The last push is dated 2026-09-25, and v0.11.2 shipped 2026-09-17.

## FAQ

### What is JAX in machine learning?

JAX is a Python library for accelerator-oriented array computation and program transformation, built around three transformations: jax.grad for reverse-mode gradients, jax.jit for compiling with XLA, and jax.vmap for mapping functions along array axes. The project describes it as a research project rather than an official Google product.

### Will JAX replace PyTorch?

The repository makes no such claim. It positions JAX as an extensible system for composable function transformations at scale, and the only PyTorch-adjacent thing in it is a converter example, onnx2xla.py, alongside the general framing of large-scale machine learning workloads.

### Why is JAX called JAX?

The repository does not explain the name. It describes JAX as an extensible system for composable function transformations and points to the reference documentation at docs.jax.dev for anything the README leaves open.

### Which is better, TensorFlow or JAX?

No comparison exists in this repository. The one documented TensorFlow coupling is in the profiler, which imports tensorflow.python.profiler.trace internally, and pyproject.toml carries a warning filter for a protobuf version skew between TF releases and the code.

## Sources

- [Official documentation](https://docs.jax.dev)
- [Official README](https://github.com/jax-ml/jax#readme)
- [Project repository](https://github.com/jax-ml/jax)
- [Release notes](https://github.com/jax-ml/jax/releases)

---

Hysen Labs editorial analysis, written from the project's own repository and release notes. Cite the canonical page: https://hysenlabs.com/projects/jax-ml-jax
