Skip to content

Follow-up on #42 / #55 : reproducibility notes for the JAX/Optax mechanical backend #56

Description

@spital

Hi maintainers,

This report comes from my local work on a PR attempt for the JAX/Optax mechanical backend during UnitaryHack 2026 / #42. I was not able to submit that PR before #42 was closed, and I understand #55 is already the active PR for that work.
I only wanted to document checks I can reproduce against the current default branch and PR #55 while the work is still easy to verify.

References:

Checked on 2026-06-26:

1. Tracking state: #42 is closed, but current master still has the old mechanical backend

This may simply be because #55 is still waiting to be merged. I am listing it because users following #42 as completed will currently get the pre-port implementation from a fresh clone of the default branch.

Reproduction:

tmp="$(mktemp -d)"
git clone https://github.com/OpenQuantumDesign/oqd-trical.git "$tmp/oqd-trical"
cd "$tmp/oqd-trical"
git fetch origin pull/55/head:pr-55

printf 'origin/master: '; git rev-parse --short origin/master
printf 'PR #55 head:   '; git rev-parse --short pr-55
git merge-base --is-ancestor pr-55 origin/master
printf 'is PR #55 an ancestor of origin/master? exit=%s\n' "$?"
git log --date=short --format='origin/master commit: %h %cd %s' -n 1 origin/master

git grep -n 'autograd\|scipy.optimize\|optax\|jax' origin/master -- \
  pyproject.toml \
  src/oqd_trical/mechanical \
  src/oqd_trical/misc/optimize.py \
  src/oqd_trical/misc/polynomial.py

Expected if #42 is complete on the default branch:

  • The accepted JAX/Optax implementation, or an equivalent implementation, is reachable from origin/master.
  • autograd is removed from the mechanical backend dependencies/imports.
  • The mechanical optimizer uses JAX/Optax rather than the old autograd/scipy path.

Actual result I see:

origin/master: e63a49f
PR #55 head:   63ea8da
is PR #55 an ancestor of origin/master? exit=1
origin/master commit: e63a49f 2026-06-09 Merge pull request #53 from OpenQuantumDesign/yhteoh-patch-1

origin/master:pyproject.toml:30:    "autograd",
origin/master:src/oqd_trical/mechanical/potential.py:18:import autograd as ag
origin/master:src/oqd_trical/mechanical/potential.py:820:        expr (Callable): function of the potential that is defined using the numpy submodule of autograd package.
origin/master:src/oqd_trical/mechanical/potential.py:882:        intensity_expr (Callable): function of the expression for intensity of the optical potential that is defined using the numpy submodule of autograd package.
origin/master:src/oqd_trical/misc/polynomial.py:22:from autograd import numpy as agnp

2. PR #55 tests/notebooks enable JAX x64, but normal package import leaves users in float32 mode

PR #55's tests and notebook explicitly enable jax_enable_x64, but importing oqd_trical normally leaves JAX in default float32 mode. That means tests are checking a different precision mode from what downstream users get by default.

Reproduction against PR #55:

tmp="$(mktemp -d)"
git clone https://github.com/OpenQuantumDesign/oqd-trical.git "$tmp/oqd-trical"
cd "$tmp/oqd-trical"
git fetch origin pull/55/head:pr-55
git switch --detach pr-55
uv sync --frozen --extra tests

.venv/bin/python - <<'PY'
import jax
import oqd_trical

print("jax_enable_x64 after oqd_trical import:", jax.config.jax_enable_x64)
PY

git grep -n 'jax_enable_x64' HEAD -- . ':!uv.lock'

Expected:

  • If the mechanical solver requires float64 precision, the package should enable it during initialization, or the public API/docs should make the requirement explicit.
  • Tests should not silently exercise a higher precision mode than normal users get.

Actual:

jax_enable_x64 after oqd_trical import: False

HEAD:examples/mechanical/1_examples_mechanical.ipynb:35:    "jax.config.update(\"jax_enable_x64\", True)\n",
HEAD:examples/mechanical/1_examples_mechanical.ipynb:626:    "jax.config.update(\"jax_enable_x64\", True)\n",
HEAD:tests/test_mechanical.py:25:jax.config.update("jax_enable_x64", True)

On a small N=8 harmonic-chain calculation, the same PR #55 code gives different dtypes and a visible numerical difference between default import and JAX_ENABLE_X64=1:

default False float32 float32
forced True float64 float64
max_abs_x_diff 2.171780261907958e-08
max_rel_w_diff 0.0019270162117970603
worst default/x64 frequencies 1222438.125 1224798.3312404824

For mechanical calculations using tight tolerances, float32-by-default may be fine if it is intentional, but it seems worth making explicit.

3. PR #55's frozen gpu extra does not install a CUDA backend package

PR #55 declares:

gpu = ["jax[cuda13]"]

but the frozen environment I get from uv sync --frozen --extra gpu installs CPU-only jaxlib and no CUDA plugin packages.

Reproduction against PR #55:

tmp="$(mktemp -d)"
git clone https://github.com/OpenQuantumDesign/oqd-trical.git "$tmp/oqd-trical"
cd "$tmp/oqd-trical"
git fetch origin pull/55/head:pr-55
git switch --detach pr-55

uv sync --frozen --extra tests --extra gpu

.venv/bin/python - <<'PY'
import importlib.metadata as m
import jax
import oqd_trical

print("jax_enable_x64 after oqd_trical import:", jax.config.jax_enable_x64)
for pkg in ["jax", "jaxlib", "jax-cuda12-plugin", "jax-cuda13-plugin", "jax-cuda13-pjrt"]:
    try:
        print(pkg, m.version(pkg))
    except m.PackageNotFoundError:
        print(pkg, "not installed")
print("devices", jax.devices())
PY

JAX_PLATFORMS=cuda,cpu .venv/bin/python - <<'PY'
import jax
print(jax.devices())
PY

Expected on a CUDA host:

  • uv sync --frozen --extra gpu installs a CUDA-capable JAX backend, or the extra is renamed/removed until it works.
  • With JAX_PLATFORMS=cuda,cpu, JAX should initialize a CUDA device on a CUDA host.

Actual result on my CUDA-capable host:

jax_enable_x64 after oqd_trical import: False
jax 0.5.2
jaxlib 0.5.1
jax-cuda12-plugin not installed
jax-cuda13-plugin not installed
jax-cuda13-pjrt not installed
devices [CpuDevice(id=0)]

Strict CUDA initialization then fails with:

RuntimeError: Unable to initialize backend 'cuda':  (set JAX_PLATFORMS='' to automatically choose an available backend)

4. PR #55's normal_modes() path appears much slower than current master on a small chain

This is not a correctness failure, and first-call JAX compilation cost is expected. The surprising part is that repeated normal_modes() calls are still much slower than current master for this small example.

Reproduction comparing current master and PR #55:

tmp="$(mktemp -d)"
git clone https://github.com/OpenQuantumDesign/oqd-trical.git "$tmp/oqd-trical"
cd "$tmp/oqd-trical"
git fetch origin pull/55/head:pr-55

for ref in origin/master pr-55; do
  echo "REF $ref"
  git switch --detach --quiet "$ref"
  uv sync --frozen --extra tests >/dev/null
  .venv/bin/python - <<'PY'
import time
import numpy as np
try:
    import jax
    print("jax_enable_x64", jax.config.jax_enable_x64)
except Exception:
    print("jax not imported")
import oqd_trical
import oqd_trical.misc.constants as cst

N = 8
mass = 171 * cst.m_u
omega_x = 2 * np.pi * 0.4e6
omega_y = 2 * np.pi * 0.36e6
omega_z = 2 * np.pi * 0.08e6

alpha = np.zeros((3, 3, 3))
alpha[2, 0, 0] = mass * omega_x**2 / 2
alpha[0, 2, 0] = mass * omega_y**2 / 2
alpha[0, 0, 2] = mass * omega_z**2 / 2

pp = oqd_trical.mechanical.PolynomialPotential(alpha, N=N)
ti = oqd_trical.mechanical.TrappedIons(N, pp, m=mass)

s = time.perf_counter()
ti.equilibrium_position()
e = time.perf_counter()
print("equilibrium %.3fs" % (e - s))

for i in range(3):
    s = time.perf_counter()
    ti.normal_modes()
    e = time.perf_counter()
    print("normal_modes call%d %.3fs" % (i + 1, e - s))
PY
done

Expected:

  • PR Replace scipy optimization with optax #55 may pay some one-time JAX compilation cost.
  • Repeated normal_modes() calls for a small N=8 chain should ideally be close to current master, or the slower behavior should be documented/benchmarked before relying on this path for larger ion chains.

Actual result on my machine:

REF origin/master
jax_enable_x64 False
equilibrium 0.056s
normal_modes call1 0.017s
normal_modes call2 0.017s
normal_modes call3 0.017s

REF pr-55
jax_enable_x64 False
equilibrium 0.928s
normal_modes call1 4.478s
normal_modes call2 2.841s
normal_modes call3 2.989s

From a quick look, this may be related to many small jitted Hessian helpers with static ion/axis arguments, but that is only my read of the current implementation. The main point is the reproducible timing difference above.

Summary

I think there are a few separate tracking questions:

  1. Should the x64 behavior be made consistent between tests and normal package imports, or documented as a caller responsibility?
  2. Should the gpu extra be corrected so the frozen environment installs a CUDA-capable backend?
  3. Should normal_modes() performance be benchmarked before Replace scipy optimization with optax #55 is merged?

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions