You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
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.
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.
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.
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.
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:
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.
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.
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:
Should the x64 behavior be made consistent between tests and normal package imports, or documented as a caller responsibility?
Should the gpu extra be corrected so the frozen environment installs a CUDA-capable backend?
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:
origin/master:e63a49f(Merge pull request #53 from OpenQuantumDesign/yhteoh-patch-1)63ea8da1. Tracking state: #42 is closed, but current
masterstill has the old mechanical backendThis 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:
Expected if #42 is complete on the default branch:
origin/master.autogradis removed from the mechanical backend dependencies/imports.Actual result I see:
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 importingoqd_tricalnormally 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:
Expected:
Actual:
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: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
gpuextra does not install a CUDA backend packagePR #55 declares:
but the frozen environment I get from
uv sync --frozen --extra gpuinstalls CPU-onlyjaxliband no CUDA plugin packages.Reproduction against PR #55:
Expected on a CUDA host:
uv sync --frozen --extra gpuinstalls a CUDA-capable JAX backend, or the extra is renamed/removed until it works.JAX_PLATFORMS=cuda,cpu, JAX should initialize a CUDA device on a CUDA host.Actual result on my CUDA-capable host:
Strict CUDA initialization then fails with:
4. PR #55's
normal_modes()path appears much slower than currentmasteron a small chainThis 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 currentmasterfor this small example.Reproduction comparing current
masterand PR #55:Expected:
normal_modes()calls for a small N=8 chain should ideally be close to currentmaster, or the slower behavior should be documented/benchmarked before relying on this path for larger ion chains.Actual result on my machine:
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:
gpuextra be corrected so the frozen environment installs a CUDA-capable backend?normal_modes()performance be benchmarked before Replace scipy optimization with optax #55 is merged?