Precision (config)¶
JAX computes in 32-bit floats unless it is told otherwise: about 7 correct
digits instead of about 16. A quantum simulation needs the 64-bit version
(complex128 amplitudes), so Dense-Evolution switches JAX to 64 bits for you,
at the moment it is needed and never at import time.
Step 1. It happens automatically¶
import jax
import dense_evolution as de
before = jax.config.jax_enable_x64
sim = de.DenseSVSimulator(2)
sim.run_circuit_jit(de.QASMParser().parse('OPENQASM 2.0; include "qelib1.inc"; qreg q[2]; h q[0]; cx q[0],q[1];').to_tuples())
(before, jax.config.jax_enable_x64, str(sim.get_statevector().dtype))
import dense_evolution alone leaves JAX as it was (False). Creating the
simulator calls ensure_x64(), which turns on JAX's 64-bit mode for the whole
process, and the Bell state comes back as complex128. The mitigation
functions, arithmetic, postselect and spectral do the same on entry.
Step 2. Choosing the precision yourself¶
import jax
import dense_evolution as de
de.set_precision(False)
sim = de.DenseSVSimulator(2)
jax.config.jax_enable_x64
set_precision is for the case where another JAX library in the same process
must stay in 32 bits. Once you call it, your choice sticks: ensure_x64() no
longer turns 64 bits back on, so the simulator above runs in single precision.
Call it before creating anything.
Step 3. The same guard on your own functions¶
import jax.numpy as jnp
from dense_evolution.config import with_x64
@with_x64
def norm(v):
return jnp.linalg.norm(jnp.asarray(v))
str(norm([0.6, 0.8]).dtype)
with_x64 runs ensure_x64() before every call of the function it wraps.
Without it, the same jnp.linalg.norm(jnp.asarray([0.6, 0.8])) in a fresh
process returns float32: JAX silently truncates input it builds before
anything has enabled 64 bits.
Details¶
jax_enable_x64 is one flag for the whole Python process, shared by every
library that uses JAX. Earlier versions set it at import time in three
modules, which silently overrode a precision chosen by unrelated code running
in the same process; config.py is now the single place that sets it, lazily.
config ¶
Centralized JAX precision control for dense_evolution.
jax_enable_x64 is a process-wide JAX flag: once set, it affects every
JAX array in the process, not just dense_evolution's own. Before this
module existed, registry.py, compiler.py and statevector.py each called
jax.config.update("jax_enable_x64", True) unconditionally at MODULE
IMPORT time -- so merely import dense_evolution, even without ever
constructing a simulator, silently overrode a precision a caller had
already configured for unrelated JAX code running earlier in the same
process. mps.py and chunk.py never did this; their own docstrings
already documented the intended convention (see mps.py): dense_evolution
enables x64 lazily, only when something that actually needs complex128
precision is used, not as an import-time side effect.
This module is that single point. ensure_x64() is idempotent and safe to call from every float64 entry point (DenseSVSimulator.init, QuantumHardwareRegistry.init, circuit_to_energy_fn); set_precision() is the public opt-out for a caller who wants to choose precision explicitly, before constructing anything, and have that choice stick.
set_precision ¶
Explicitly configure JAX's process-wide numeric precision.
Call this yourself, before constructing anything, if you need control over exactly when/whether dense_evolution enables x64 -- for example because another float32-only JAX library must be initialized first. Once called, dense_evolution's own lazy ensure_x64() (used internally by DenseSVSimulator, etc.) no longer forces float64 back on, so an explicit set_precision(False) sticks.
Source code in dense_evolution/config.py
ensure_x64 ¶
Enable jax_enable_x64 unless a caller has already explicitly configured precision via set_precision(). Called internally, lazily, by dense_evolution components that need complex128 precision -- never at import time.
Source code in dense_evolution/config.py
with_x64 ¶
Wrap fn so ensure_x64() runs before every call -- for public entry
points that build JAX arrays from user input, which JAX would otherwise
silently truncate to complex64/float32 if nothing else enabled x64 first.