Skip to content

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))
(False, True, 'complex128')

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
False

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)
float64

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

set_precision(float64: bool = True) -> None

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
def set_precision(float64: bool = True) -> None:
    """
    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.
    """
    global _explicit
    jax.config.update("jax_enable_x64", float64)
    _explicit = True

ensure_x64

ensure_x64() -> None

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
def ensure_x64() -> None:
    """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."""
    if not _explicit:
        jax.config.update("jax_enable_x64", True)

with_x64

with_x64(fn)

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.

Source code in dense_evolution/config.py
def with_x64(fn):
    """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."""
    @functools.wraps(fn)
    def wrapper(*args, **kwargs):
        ensure_x64()
        return fn(*args, **kwargs)
    return wrapper