Performance Optimization with Unitful Quantities#

In this guide, we’ll explore how to think about performance optimization when working with coordinax objects in JAX. The key insight is understanding where the overhead lives and when it matters.

Key Concepts#

  1. Wrapper overhead: Operations on objects have overhead compared to raw JAX arrays.

  2. JIT removes overhead: JAX’s JIT compiler can eliminate much of this wrapper overhead by tracing through the code.

  3. Pytree complexity: Objects are JAX pytrees, which adds cost when crossing JIT boundaries (converting between traced and non-traced values).

  4. Strategy: The secret to performance is to minimize pytree conversions at the boundary between traced and non-traced code.

Let’s explore this last point a little more. The details will become clearer in the examples below, but the general idea is that for optimal performance you want to structure your code so that the JIT-compiled functions take raw arrays as input and output, and you convert to/from coordinax objects inside the JIT-compiled function. This way, the overhead of working with coordinax objects is paid only once per JIT compilation, rather than on every call. Also as a note, coordinax is not “special” in this regard – any PyTree object will have this same overhead at the boundary, so this is a general principle for working with JAX and PyTrees effectively.

The pseudo-code below illustrates the idea:

@jax.jit
def function_that_takes_pytrees(*objects: PyTree):
    # This function can take and return PyTrees,
    # but it will be slower due to wrapper overhead and pytree complexity.
    ...

@jax.jit
def optimized_function(*arrays: Array):
    # This function takes and returns raw arrays,
    # so it can be optimized by JIT without overhead.
    # Inside this function, we can convert to/from PyTrees as needed,
    # but this conversion will be compiled away by JIT.
    ...

What Eager Costs, and Why#

The guide above teaches jitting. It is worth saying plainly what you pay if you do not.

One pt_map from sph3d to cart3d, at a single point:

call

eager

jitted, reused

dict[str, Quantity]

686 us

21.2 us

dict[str, Array] + usys=

170 us

13.8 us

That quantity column used to read 3650 us. The difference is not a faster unxt: it is that the transition bodies no longer do their arithmetic on Quantity operands.

Why that mattered so much. quax builds and evaluates a separate jaxpr for every primitive whose operand is a Quantity. Eagerly, one primitive costs ~4 us on a raw array against 120-820 us on a Quantity, so a transition’s cost was simply the sum of its operations – Quantity ** 2 120 us, + 198 us, / 229 us, == 0 444 us, atan2 617 us. Taking the reverse map as a worked example, since it was the most expensive: Cart3D -> Spherical3D spent ~2900 us of its ~3500 us there.

The bodies now resolve their units once on the way in, compute on raw arrays, and re-attach once on the way out. Same numbers – the conversions are checked point-by-point against the previous implementation – and, on the Quantity route, a fraction of the dispatches. The raw-array route rises slightly instead: strip probes unit_of even on values that turn out to have none. Counting calls to plum-dispatched functions for one eager call, before and after:

route

calls, before

now

of which pt_map

dict[str, Quantity]

178

45

2

dict[str, Array] + usys=

17

19

2

Two things follow.

Raw arrays are still cheaper eagerly, but by ~4x rather than the 20x this guide once reported. Both routes compute identical numbers, bit for bit, so a pipeline already carrying bare arrays should still hand them straight to pt_map with an explicit usys=. It is no longer worth restructuring a pipeline around.

What remains is mostly not coordinax’s to remove. Of the 45 dispatches left, 2 are pt_map; the rest are unxt resolving convert, ustrip and unit_of at the boundary – resolving the units, not computing with them. Pre-resolving with plum’s Function.invoke is therefore not the lever it looks like: it removes 2 of 45 and measures 1.01x. It pays where a hot loop repeatedly re-resolves its own call, which is why norm keeps a module-level array_norm = norm.invoke(Array, Array).

Dispatch is per call, not per element. Batched and jitted, the routes converge – at 10,000 points they land within a few percent of each other:

route, N = 10,000, jitted

per point

dict[str, Array], broadcast

22.3 ns

dict[str, Quantity], broadcast

23.7 ns

dict[str, Array], jit(vmap(...))

22.1 ns

So the eager figure is a statement about many small calls, not about throughput. Batch, and it disappears; stay eager and per-point, and the route is worth a few-fold.

Note

The remaining eager cost is dominated by plum resolutions on unxt’s boundary functions – ustrip, uconvert, dimension_of – which carry signatures that are not faithful (Literal, Mapping[...], type[...]) and so run with their method cache disabled. plum#290 proposes separating is_cacheable from is_faithful, which would let exactly those signatures be cached. coordinax calls those functions rather than reaching past them for this reason: the improvement arrives without any change here. Expect the Quantity column to narrow further when it lands; none of the advice above depends on it.

Coordinate Changes#

Let’s start by importing the libraries we’ll need and setting up some test data.

import functools as ft
from dataclasses import replace

from jaxtyping import Array
import jax
import jax.numpy as jnp
import quax

from jaxmore import vmap
import unxt as u

import coordinax as cx

We’ll define this function now:

usys = u.unitsystems.si

_c2s_cx = cx.pt_map(cx.cart3d, cx.sph3d, usys=usys)
c2s_cx = jax.jit(vmap(_c2s_cx))

Note

This one coordinax-backed function will be able to transform ALL of the object types we work with without modification, thanks to the way coordinax objects are designed to work with JAX and Quax.

Array#

Basic JAX#

@jax.jit
@vmap
def c2s_arr(x: Array, /) -> Array:
    r = jnp.linalg.norm(x, axis=-1)
    theta = jnp.acos(x[..., 2] / r)
    phi = jnp.arctan2(x[..., 1], x[..., 0])
    return jnp.stack((r, theta, phi), axis=-1)

xarr = jnp.array([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]])

%time jax.block_until_ready(c2s_arr(xarr))
%timeit jax.block_until_ready(c2s_arr(xarr))

c2s_arr(xarr)
CPU times: user 70.6 ms, sys: 6.09 ms, total: 76.7 ms
Wall time: 57.1 ms
10.2 μs ± 23.6 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)
Array([[3.7416575 , 0.64052236, 1.1071488 ],
       [8.774964  , 0.81788856, 0.8960554 ]], dtype=float32)

This is a good baseline, but it’s applied to a raw JAX array, which in coordinax is assumed to be in Cartesian coordinates. coordinax allows us to work with coordinates in many different types of charts. Let’s see how the performance compares when we use coordinax objects, especially after JIT compilation.

Coordinax#

Now let’s work with coordinax and perform the same coordinate change. We’ll see how the performance compares, especially after JIT compilation.

The key to performance is to close over the static quantities (like the charts and unit system) so that JAX can optimize the computation effectively. This means we want to avoid passing coordinax objects directly into JIT-compiled functions if they contain static information that can’t be traced.

Let’s time this function with raw arrays first, to see the baseline overhead of using coordinax objects.

%time jax.block_until_ready(c2s_cx(xarr))
%timeit jax.block_until_ready(c2s_cx(xarr))

c2s_cx(xarr)
CPU times: user 68.5 ms, sys: 5.07 ms, total: 73.6 ms
Wall time: 72.4 ms
9.04 μs ± 27.4 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)
Array([[3.7416573 , 0.64052236, 1.1071488 ],
       [8.774964  , 0.81788856, 0.8960554 ]], dtype=float32)

There is NONE! The performance is the same as the hardcoded version. This is because we closed over the static arguments (the charts and unit system) so that JAX can optimize the computation effectively. The coordinax objects are only used inside the JIT-compiled function, so there is no overhead from pytrees or wrappers at the boundary.

But only if the closure itself is built once. _c2s_cx and c2s_cx above are each built exactly once, at module scope, and every timing reuses that same object. jax.jit caches on the identity of the Python function it wraps, not on argument equality – so a fresh cx.pt_map(...) closure re-wrapped in a fresh jax.jit(...) on every call is a cache miss every time, not a cheap re-dispatch:

def rebuild_each_call(x):
    fn = jax.jit(cx.pt_map(cx.cart3d, cx.sph3d, usys=usys))
    return jax.block_until_ready(fn(x))

%timeit rebuild_each_call(xarr)
%timeit jax.block_until_ready(c2s_cx(xarr))
45.1 ms ± 463 μs per loop (mean ± std. dev. of 7 runs, 10 loops each)
8.95 μs ± 54 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)

Rebuilding pays a full retrace-and-compile on every call – tens of milliseconds instead of low microseconds, roughly a 1000-4000x difference depending on the function. Build the pt_map/jit/vmap stack once, at module scope or in __init__, and call that same object every time; if several call sites need the same conversion, define it once in a shared helper module and import the built closure rather than re-deriving it at each site.

Let’s see what happens if we don’t close over the static arguments:

c2s_bad = jax.jit(
  vmap(cx.pt_map, in_axes=(0, None, None), in_kw={"usys": None}),
  static_argnums=(1, 2), static_argnames=("usys"),
)

usys = u.unitsystems.si

%time jax.block_until_ready(c2s_bad(xarr, cx.cart3d, cx.sph3d, usys=usys))
%timeit jax.block_until_ready(c2s_bad(xarr, cx.cart3d, cx.sph3d, usys=usys))

c2s_bad(xarr, cx.cart3d, cx.sph3d, usys=usys)
CPU times: user 46.4 ms, sys: 1.99 ms, total: 48.3 ms
Wall time: 47.7 ms
22.3 μs ± 440 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)
Array([[3.7416573 , 0.64052236, 1.1071488 ],
       [8.774964  , 0.81788856, 0.8960554 ]], dtype=float32)

Note

The vmap we are using is from jaxmore, which is a thin wrapper around jax.vmap that adds some extra features, in particular support for keyword arguments.

As expected, this is much slower than the hard-coded version. This is all the wrapper overhead and pytree complexity in action. Even though the argument xarr is a raw JAX array the vmap-with-kwargs has kwarg-related overhead and the JIT has to deal with the static arguments (chart_from,chart_to,usys).

Dict[str, Array]#

Basic JAX#

def c2s_dict(x: dict[str, Array], /) -> dict[str, Array]:
    r = jnp.sqrt(x["x"] **2 + x["y"]** 2 + x["z"] ** 2)
    theta = jnp.acos(x["z"] / r)
    phi = jnp.arctan2(x["y"], x["x"])
    return {"r": r, "theta": theta, "phi": phi}

vec_c2s_dict = jax.jit(vmap(c2s_dict))

xdict = {"x": jnp.array([1.0, 4.0]), "y": jnp.array([2.0, 5.0]),
     "z": jnp.array([3.0, 6.0])}

%time jax.block_until_ready(vec_c2s_dict(xdict))
%timeit jax.block_until_ready(vec_c2s_dict(xdict))

vec_c2s_dict(xdict)
CPU times: user 86.6 ms, sys: 3.01 ms, total: 89.6 ms
Wall time: 57.7 ms
16.6 μs ± 44.4 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)
{'phi': Array([1.1071488, 0.8960554], dtype=float32),
 'r': Array([3.7416575, 8.774964 ], dtype=float32),
 'theta': Array([0.64052236, 0.81788856], dtype=float32)}

This is around 50% slower than the raw array version, which is expected due to the overhead of using dictionaries and the way JAX handles them as pytrees. However, this is still quite fast and may be acceptable depending on the use case.

If we want to achieve the same performance as the raw array version, we can shift the pytree conversion inside the JIT boundary. This way, the overhead of converting between pytrees and arrays is also compiled away.

@jax.jit
@vmap
def c2s_dict_comp(x: Array, /) -> Array:
    d = {"x": x[..., 0], "y": x[..., 1], "z": x[..., 2]}
    r = jnp.sqrt(d["x"] **2 + d["y"]** 2 + d["z"] ** 2)
    theta = jnp.acos(d["z"] / r)
    phi = jnp.arctan2(d["y"], d["x"])
    return jnp.stack((r, theta, phi), axis=-1)

%time jax.block_until_ready(c2s_dict_comp(xarr))
%timeit jax.block_until_ready(c2s_dict_comp(xarr))
CPU times: user 60.6 ms, sys: 2.01 ms, total: 62.6 ms
Wall time: 47 ms
9.54 μs ± 26.7 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)

This is now as fast as the raw array version, because we have minimized the overhead at the JIT boundary.

Coordinax#

We can apply the same function to coordinax objects.

c2s_cx(xdict)
{'phi': Array([1.1071488, 0.8960554], dtype=float32),
 'r': Array([3.7416573, 8.774964 ], dtype=float32),
 'theta': Array([0.64052236, 0.81788856], dtype=float32)}

To achieve the same performance as the raw array version, we can shift the pytree conversion inside the JIT boundary. This way, the overhead of converting between pytrees and arrays is also compiled away.

from jaxmore import structured

structurer = structured(lambda x: cx.cdict(x, cx.cart3d),
                        lambda x:cx.carray(x, cx.sph3d.components, usys).value)
c2s_cx_dict = jax.jit(vmap(structurer(_c2s_cx)))

%time jax.block_until_ready(c2s_cx_dict(xarr))
%timeit jax.block_until_ready(c2s_cx_dict(xarr))

c2s_cx_dict(xarr)
CPU times: user 53 ms, sys: 2 ms, total: 55 ms
Wall time: 54.8 ms
9.08 μs ± 44.4 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)
Array([[3.7416573 , 0.64052236, 1.1071488 ],
       [8.774964  , 0.81788856, 0.8960554 ]], dtype=float32)

Based on this, we can see that the actual cost of coordinax is paid at the JIT boundary. If we can minimize the number of times we cross the JIT boundary with coordinax objects, we can achieve performance that is close to raw arrays. This is a key insight for working with coordinax in performance-critical code: minimize pytree conversions at the boundary.

Dict[str, Quantity]#

Basic JAX#

@jax.jit
@quax.quaxify  # enables Quantity support!
@vmap
def c2s_qdict(x: dict[str, u.Q], /) -> dict[str, u.Q]:
    r = jnp.sqrt(x["x"] **2 + x["y"]** 2 + x["z"] ** 2)
    theta = jnp.acos(x["z"] / r)
    phi = jnp.arctan2(x["y"], x["x"])
    return {"r": r, "theta": theta, "phi": phi}

xqdict = {k: u.Q(v, "m") for k, v in xdict.items()}

%time jax.block_until_ready(c2s_qdict(xqdict))
%timeit jax.block_until_ready(c2s_qdict(xqdict))

c2s_qdict(xqdict)
CPU times: user 108 ms, sys: 2 ms, total: 110 ms
Wall time: 80.9 ms
48.9 μs ± 91 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)
{'phi': Q([1.1071488, 0.8960554], 'rad'),
 'r': Q([3.7416575, 8.774964 ], 'm'),
 'theta': Q([0.64052236, 0.81788856], 'rad')}

This is around 2-3x slower than the raw array version, which is expected due to the additional overhead of handling Quantity objects. Let’s see if we can improve this by shifting the pytree conversion inside the JIT boundary, just like we did with the dict of arrays.

@jax.jit
@vmap
def c2s_qdict_comp(x: Array, /) -> Array:
    d = {"x": x[..., 0], "y": x[..., 1], "z": x[..., 2]}
    r = jnp.sqrt(d["x"] **2 + d["y"]** 2 + d["z"] ** 2)
    theta = jnp.acos(d["z"] / r)
    phi = jnp.arctan2(d["y"], d["x"])
    return jnp.stack((r, theta, phi), axis=-1)

%time jax.block_until_ready(c2s_qdict_comp(xarr))
%timeit jax.block_until_ready(c2s_qdict_comp(xarr))
CPU times: user 64.3 ms, sys: 2 μs, total: 64.3 ms
Wall time: 47 ms
9.68 μs ± 42.4 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)

The speed is now much closer to the raw array version, but now we can see the problem that was implicit in the previous transformation – you have to manually manage units. There’s a good way around this. If we stick to a particular unit system (e.g. SI), we can close over the static unit system inside the JIT-compiled function, so that you don’t have to manage units at all:

import quaxed.numpy as qnp

usys = u.unitsystems.si

@jax.jit
@jax.vmap
def c2s_qdict_comp2(x: Array, /) -> Array:
    d = cx.cdict(x, usys["length"], cx.cart3d)
    r = qnp.sqrt(d["x"] **2 + d["y"]** 2 + d["z"] ** 2)
    theta = qnp.acos(d["z"] / r)
    phi = qnp.arctan2(d["y"], d["x"])
    return jnp.stack([u.ustrip(usys, r), u.ustrip(usys, theta), u.ustrip(usys, phi)], axis=-1)

%time jax.block_until_ready(c2s_qdict_comp2(xarr))
%timeit jax.block_until_ready(c2s_qdict_comp2(xarr))
CPU times: user 71.1 ms, sys: 1.99 ms, total: 73.1 ms
Wall time: 57.6 ms
9.69 μs ± 73.8 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)

This is now very close to the raw array version, and the user doesn’t have to manage units at all, so long as they stick to the predefined unit system. On the other hand, this function is not a bit of a mess! Let’s do better…

Coordinax#

With coordinax all we need is:

c2s_cx(xqdict)
{'phi': Angle([1.1071488, 0.8960554], 'rad'),
 'r': Q([3.7416573, 8.774964 ], 'm'),
 'theta': Angle([0.64052236, 0.81788856], 'rad')}

Just like the basic JAX version, this is around 2-3x slower than the raw array version due to the overhead of handling Quantity objects. However, as with the dict of arrays, we can shift the pytree conversion inside the JIT boundary to minimize this overhead.

# Pre-specialize `c2s_cx`, pushing pytrees into a JIT-optimizable closure.
structurer = structured(lambda x: cx.cdict(x, usys["length"], cx.cart3d),
                        lambda x: cx.carray(x, cx.sph3d.components, usys).value)
c2s_cx_qdict = jax.jit(jax.vmap(structurer(_c2s_cx)))

%time jax.block_until_ready(c2s_cx_qdict(xarr))
%timeit jax.block_until_ready(c2s_cx_qdict(xarr))

c2s_cx_qdict(xarr)
CPU times: user 50 ms, sys: 1.01 ms, total: 51 ms
Wall time: 49.9 ms
8.96 μs ± 12.4 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)
Array([[3.7416573 , 0.64052236, 1.1071488 ],
       [8.774964  , 0.81788856, 0.8960554 ]], dtype=float32)

There’s no overhead, and we got to use the same function to transform coordinax objects with Quantity data, without having to manually manage units at all. This is the power of coordinax.

Point#

Basic JAX#

@jax.jit
@vmap
def c2s_vec(x: cx.Point, /) -> cx.Point:
    r = quax.quaxify(jnp.sqrt)(x["x"] **2 + x["y"]** 2 + x["z"] ** 2)
    theta = quax.quaxify(jnp.acos)(x["z"] / r)
    phi = quax.quaxify(jnp.arctan2)(x["y"], x["x"])
    return replace(x, data={"r": r, "theta": theta, "phi": phi}, chart=cx.sph3d)

xvec = cx.Point.from_(xqdict, cx.cart3d)

%time jax.block_until_ready(c2s_vec(xvec))
%timeit jax.block_until_ready(c2s_vec(xvec))

c2s_vec(xvec)
CPU times: user 88.6 ms, sys: 1.01 ms, total: 89.6 ms
Wall time: 58.9 ms
63.8 μs ± 266 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)
Point(
  {
    'phi': Q([1.1071488, 0.8960554], 'rad'),
    'r': Q([3.7416575, 8.774964 ], 'm'),
    'theta': Q([0.64052236, 0.81788856], 'rad')
  },
  chart=Spherical3D(M=Rn(3))
)

This works, but again it’s not particularly optimized. Manually optimizing is similar to the cases above.

Coordinax#

With coordinax all we need is:

c2s_cx(xvec)
Point(
  {
    'phi': Angle([1.1071488, 0.8960554], 'rad'),
    'r': Q([3.7416573, 8.774964 ], 'm'),
    'theta': Angle([0.64052236, 0.81788856], 'rad')
  },
  chart=Spherical3D(M=Rn(3))
)

With coordinax, optimizing is very easy.

# Pre-specialize `c2s_cx`, pushing pytrees into a JIT-optimizable closure.
structurer = structured(lambda x: cx.Point.from_(u.Q(x, usys["length"]), cx.cart3d),
                        lambda x: cx.carray(x.data, cx.sph3d.components, usys).value)
c2s_cx_vec = jax.jit(vmap(structurer(_c2s_cx)))

%time jax.block_until_ready(c2s_cx_vec(xarr))
%timeit jax.block_until_ready(c2s_cx_vec(xarr))

c2s_cx_vec(xarr)
CPU times: user 52.2 ms, sys: 1 ms, total: 53.2 ms
Wall time: 51.9 ms
8.94 μs ± 32.4 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)
Array([[3.7416573 , 0.64052236, 1.1071488 ],
       [8.774964  , 0.81788856, 0.8960554 ]], dtype=float32)


Jacobian of Point Maps#

jax.jacfwd computes the forward-mode Jacobian \(J^j{}_i = \partial \phi^j / \partial q^i\) of the chart transition map. We compare three approaches evaluated at a single 3-D base point.

import coordinax.charts as cxc

at = jnp.array([1.0, 0.0, 0.0])

Raw JAX (baseline)#

c2s_arr defined earlier is vmapped, so for differentiation we write the scalar version directly:

def _c2s_arr_scalar(x: Array, /) -> Array:
    r = jnp.linalg.norm(x)
    theta = jnp.acos(x[2] / r)
    phi = jnp.arctan2(x[1], x[0])
    return jnp.stack([r, theta, phi])

jac_c2s_arr = jax.jit(jax.jacfwd(_c2s_arr_scalar))

%time jac_c2s_arr(at).block_until_ready()
%timeit jac_c2s_arr(at).block_until_ready()
jac_c2s_arr(at)
CPU times: user 167 ms, sys: 5.02 ms, total: 172 ms
Wall time: 115 ms
8.15 μs ± 33.3 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)
Array([[ 1.,  0.,  0.],
       [-0., -0., -1.],
       [ 0.,  1.,  0.]], dtype=float32)

pt_map with jax.jacfwd#

_c2s_cx (defined above as cx.pt_map(cx.cart3d, cx.sph3d, usys=usys)) is already a scalar Array -> Array callable that closes over the chart pair and unit system. Passing it to jax.jacfwd gives the same Jacobian with no extra per-call cost:

jac_pt_map_fn = jax.jit(jax.jacfwd(_c2s_cx))

%time jac_pt_map_fn(at).block_until_ready()
%timeit jac_pt_map_fn(at).block_until_ready()
jac_pt_map_fn(at)
CPU times: user 384 ms, sys: 12 ms, total: 396 ms
Wall time: 258 ms
7.9 μs ± 18.5 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)
Array([[ 1.,  0.,  0.],
       [ 0.,  0., -1.],
       [ 0.,  1.,  0.]], dtype=float32)

There is a small one-time compilation overhead from the coordinax dispatch layer, but the per-call runtime is identical to the raw-JAX baseline.

jac_pt_map (idiomatic API)#

cxc.jac_pt_map(from_chart, to_chart, usys=usys) is the high-level API. The curried form returns a callable that wraps the full jacfwd pipeline, including unit handling. Wrapping it in jax.jit gives the same performance:

jac_fn = cxc.jac_pt_map(cxc.cart3d, cxc.sph3d, usys=usys)
jac_jit = jax.jit(jac_fn)

%time jac_jit(at).block_until_ready()
%timeit jac_jit(at).block_until_ready()
jac_jit(at)
CPU times: user 106 ms, sys: 2.99 ms, total: 109 ms
Wall time: 78 ms
7.26 μs ± 32 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)
Array([[ 1.,  0.,  0.],
       [ 0.,  0., -1.],
       [-0.,  1.,  0.]], dtype=float32)

Same runtime as the baseline. The idiomatic form also accepts quantity-valued dicts directly, without any manual unit management.

Eager Jacobians#

Everything above is jitted. Eager is a different story — and it is not only an eager-path concern, because tracing a jitted function runs this code too.

jax.jacfwd builds and evaluates a jaxpr on every call. Jitting hoists that into a one-time compilation; eagerly you pay it per call, and it dominates everything else:

at_cyl = {"rho": jnp.asarray(2.0), "phi": jnp.asarray(0.7), "z": jnp.asarray(3.0)}

jitted = jax.jit(lambda a: cxc.jac_pt_map(a, cxc.cyl3d, cxc.sph3d, usys=usys))
jax.block_until_ready(jitted(at_cyl))  # compile
jax.block_until_ready(cxc.jac_pt_map(at_cyl, cxc.cyl3d, cxc.sph3d, usys=usys))  # warm up

print("eager: ", end="")
%timeit jax.block_until_ready(cxc.jac_pt_map(at_cyl, cxc.cyl3d, cxc.sph3d, usys=usys))
print("jitted:", end="")
%timeit jax.block_until_ready(jitted(at_cyl))
eager: 
2.15 ms ± 5.34 μs per loop (mean ± std. dev. of 7 runs, 100 loops each)
jitted:
12 μs ± 44.9 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)

jax.block_until_ready on both: JAX dispatches asynchronously, so timing without it measures how fast Python can queue the work rather than how long it takes.

The cost is per call, not per point#

A chart map is pointwise, so a batch of points is a batch of independent Jacobians — but the trace happens once for the whole batch. Going from one point to ten thousand costs well under twice the time:

for n in (1, 100, 10_000):
    batch = {k: jnp.full((n,), v) for k, v in (("rho", 2.0), ("phi", 0.7), ("z", 3.0))}
    jax.block_until_ready(cxc.jac_pt_map(batch, cxc.cyl3d, cxc.sph3d, usys=usys))  # warm up
    print(f"N = {n:<6d}", end=" ")
    %timeit -r 3 jax.block_until_ready(cxc.jac_pt_map(batch, cxc.cyl3d, cxc.sph3d, usys=usys))
N = 1      
6.39 ms ± 36 μs per loop (mean ± std. dev. of 3 runs, 100 loops each)
N = 100    
6.34 ms ± 19.3 μs per loop (mean ± std. dev. of 3 runs, 100 loops each)
N = 10000  
7.03 ms ± 5.68 μs per loop (mean ± std. dev. of 3 runs, 100 loops each)

(A batch of one is not the same as a scalar point: any leading axis takes the vmap route, so N = 1 need not match the scalar call timed above.)

So batch your points rather than looping: N points cost one trace, not N.

Tracing pays it too#

jit does not avoid this work, it moves it: tracing your jitted function runs jac_pt_map with tracers, so whatever an eager call costs, tracing costs about the same. Measuring that needs care, because JAX caches traces — time a warm one and you get microseconds and no signal:

import time


def cold_trace_ms(chart_from, chart_to, at):
    """A fresh function each time, so JAX's trace cache cannot answer for us."""
    times = []
    for _ in range(5):
        fn = lambda a, f=chart_from, t=chart_to: cxc.jac_pt_map(a, f, t, usys=usys)
        start = time.perf_counter()
        jax.make_jaxpr(fn)(at)
        times.append(time.perf_counter() - start)
    return min(times) * 1e3


def first_call_ms(chart_from, chart_to, at):
    """Trace *and* XLA compilation: what a jitted function's first call costs."""
    times = []
    for _ in range(5):
        fn = jax.jit(lambda a, f=chart_from, t=chart_to: cxc.jac_pt_map(a, f, t, usys=usys))
        start = time.perf_counter()
        jax.block_until_ready(fn(at))
        times.append(time.perf_counter() - start)
    return min(times) * 1e3


print(f"cold trace:     {cold_trace_ms(cxc.cyl3d, cxc.sph3d, at_cyl):6.2f} ms")
print(f"first jit call: {first_call_ms(cxc.cyl3d, cxc.sph3d, at_cyl):6.2f} ms")
cold trace:       3.27 ms
first jit call:  58.29 ms

The cold trace matches the eager timing above — that is the same work, run once with tracers. But note the second number: tracing is only a small part of a first jit call, which is dominated by XLA compiling the jaxpr. So a faster eager path shortens tracing, not compilation, and is not a way to make your first call cheap.

Where it does pay is anywhere the trace itself is repeated or is the whole cost: eager loops, and code that re-traces — a new input shape, a fresh jitted closure per call, grad of something not yet cached. That is also why some chart pairs have a closed-form Jacobian written out rather than differentiated; the dispatch picks one where it exists.

A closed form is a hand-written matrix, so the fast path is not a less-trusted number: every one is tested against jax.jacfwd of the very map it claims — never against a restatement of its own formula, which would agree with a wrong derivation — at a generic point in both rad and deg, and in SI across magnitudes from \(10^{-16}\) to \(10^{17}\). Where two of them are the same matrix rearranged, as cart3d -> lonlat_sph3d is cart3d -> sph3d with its rows reordered and one flipped, they are also pinned to each other, so a correction to one cannot quietly miss the other.



TimeDep Transforms and JIT#

Time-dependent transforms are expressed with TimeDep(builder), where builder is usually an equinox.Module whose numeric fields (angular frequency, boost rate, …) are pytree leaves. This changes how jit caching behaves compared to the closure-based pattern it replaces.

A builder is an ordinary pytree: jax.jit (and eqx.filter_jit) key their compilation cache on structure — the pytree treedef — not on the identity of the Python object. Two TimeDep operators built from the same builder type retrace only once, no matter how many different parameter values are passed through:

import equinox as eqx
import jax.numpy as jnp
import unxt as u
import coordinax.charts as cxc
import coordinax.representations as cxr
import coordinax.transforms as cxfm

axis = jnp.array([0.0, 0.0, 1.0])
x = {"x": u.Q(1.0, "m"), "y": u.Q(0.0, "m"), "z": u.Q(0.0, "m")}

traces = []


@jax.jit
def apply_at(op, tau):
    traces.append(1)
    return cxfm.act(op, tau, x, cxc.cart3d, cxr.point)["y"].ustrip("m")


op_a = cxfm.TimeDep(cxfm.builders.RotationAboutAxis(u.Q(30.0, "deg/s"), axis=axis))
op_b = cxfm.TimeDep(cxfm.builders.RotationAboutAxis(u.Q(60.0, "deg/s"), axis=axis))
apply_at(op_a, u.Q(1.0, "s"))
apply_at(op_b, u.Q(1.0, "s"))  # same builder structure -> no retrace

print("structural retraces:", len(traces))
structural retraces: 1

Contrast this with TimeDep.from_(fn) for a bare function, which cannot be a pytree leaf and so is stored in a static field: a fresh closure is a fresh object identity, and every fresh identity is a cache miss:

traces.clear()


def make_fn(omega_deg):
    def build(tau):
        theta = jnp.deg2rad(omega_deg) * tau.ustrip("s")
        ct, st = jnp.cos(theta), jnp.sin(theta)
        R = jnp.array([[ct, -st, 0.0], [st, ct, 0.0], [0.0, 0.0, 1.0]])
        return cxfm.Rotate(R)

    return build


op_c = cxfm.TimeDep.from_(make_fn(30.0))
op_d = cxfm.TimeDep.from_(make_fn(60.0))  # a fresh closure -- different identity
apply_at(op_c, u.Q(1.0, "s"))
apply_at(op_d, u.Q(1.0, "s"))

print("closure retraces:", len(traces))
closure retraces: 2

Binding the varying parameter with eqx.Partial instead of closing over it restores structural caching without writing a Module: from_ leaves an already-pytree callable unwrapped, so the bound value stays a dynamic leaf. Since a Partial also carries the function as a leaf, apply it under eqx.filter_jit rather than plain jax.jit:

traces.clear()


def build_at(omega_deg, tau):
    theta = jnp.deg2rad(omega_deg) * tau.ustrip("s")
    ct, st = jnp.cos(theta), jnp.sin(theta)
    R = jnp.array([[ct, -st, 0.0], [st, ct, 0.0], [0.0, 0.0, 1.0]])
    return cxfm.Rotate(R)


@eqx.filter_jit
def apply_at_filtered(op, tau):
    traces.append(1)
    return cxfm.act(op, tau, x, cxc.cart3d, cxr.point)["y"].ustrip("m")


op_e = cxfm.TimeDep.from_(eqx.Partial(build_at, jnp.asarray(30.0)))
op_f = cxfm.TimeDep.from_(eqx.Partial(build_at, jnp.asarray(60.0)))
apply_at_filtered(op_e, u.Q(1.0, "s"))
apply_at_filtered(op_f, u.Q(1.0, "s"))

print("eqx.Partial retraces:", len(traces))
eqx.Partial retraces: 1

The array-leaf hashing cliff. Structural caching is a property of how op is passed into the jitted function — as a traced argument, whose leaves JAX inspects at trace time. It breaks down if a builder carrying array leaves is instead treated as static, which happens whenever it becomes part of the jitted function’s identity rather than one of its arguments — most commonly by jitting a bound method directly. A builder whose fields are plain Python floats survives this because Python floats are hashable; a builder whose fields are jax.Arrays or Quantitys is not:

class Builder(eqx.Module):
    omega: float  # or a jax.Array

    def __call__(self, tau):
        return self.omega * tau


b_float = Builder(1.0)  # plain Python float -- hashable
print(jax.jit(b_float)(2.0))  # plain jax.jit works

b_array = Builder(jnp.asarray(1.0))  # jax array leaf
try:
    jax.jit(b_array)(2.0)
except TypeError as e:
    print(f"TypeError: {e}")

print(eqx.filter_jit(b_array)(2.0))  # eqx.filter_jit partitions leaves correctly
2.0
TypeError: unhashable type: 'jaxlib._jax.ArrayImpl'
2.0

The fix is eqx.filter_jit (or eqx.partition/eqx.combine by hand) wherever a builder or transform carrying array leaves might be hashed as a static argument — it is a safe default for any code path that builds TimeDep operators, since it behaves identically to jax.jit when every field happens to be hashable.