From c9f958a4246d8fa038a68a54b947dea1818e4fd7 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Thu, 6 Aug 2026 15:56:03 +0200 Subject: [PATCH 1/8] ENH: set JAX precision through configuration --- benchmarks/unbinned_nll.py | 7 ++-- docs/install.md | 14 +++++++ src/tensorwaves/__init__.py | 2 + src/tensorwaves/config.py | 52 ++++++++++++++++++++++++++ src/tensorwaves/estimator.py | 4 +- src/tensorwaves/function/_backend.py | 5 ++- tests/function/test_backend.py | 55 ++++++++++++++++++++++++++++ 7 files changed, 131 insertions(+), 8 deletions(-) create mode 100644 src/tensorwaves/config.py diff --git a/benchmarks/unbinned_nll.py b/benchmarks/unbinned_nll.py index 295df2129..a1836325d 100644 --- a/benchmarks/unbinned_nll.py +++ b/benchmarks/unbinned_nll.py @@ -11,9 +11,12 @@ import tensorflow as tf import tensorflow.experimental.numpy as tnp # ty: ignore[unresolved-import] +from tensorwaves import configure from tensorwaves.estimator import UnbinnedNLL from tensorwaves.function import ParametrizedBackendFunction +configure(jax_enable_x64=True) + if TYPE_CHECKING: from collections.abc import Callable @@ -203,8 +206,6 @@ def estimator_samples() -> tuple[dict[str, np.ndarray], dict[str, np.ndarray]]: def jax_intensities( intensities: tuple[np.ndarray, np.ndarray], ) -> tuple[jax.Array, jax.Array]: - jax.config.update("jax_enable_x64", True) - data_intensities, phsp_intensities = intensities jax_data_intensities = jnp.asarray(data_intensities).block_until_ready() jax_phsp_intensities = jnp.asarray(phsp_intensities).block_until_ready() @@ -223,8 +224,6 @@ def tensorflow_intensities( def jax_estimator_samples( estimator_samples: tuple[dict[str, np.ndarray], dict[str, np.ndarray]], ) -> tuple[dict[str, jax.Array], dict[str, jax.Array]]: - jax.config.update("jax_enable_x64", True) - data, phsp = estimator_samples return ( {"x": jnp.asarray(data["x"]).block_until_ready()}, diff --git a/docs/install.md b/docs/install.md index 6edd27a78..a63dca5e4 100644 --- a/docs/install.md +++ b/docs/install.md @@ -60,6 +60,20 @@ pip install tensorwaves[jax,scipy,tf] pip install tensorwaves[all] # all runtime dependencies ``` +## JAX precision + +TensorWaves uses 64-bit precision for JAX by default. Configure 32-bit precision +before creating JAX arrays or TensorWaves functions with: + +```python +import tensorwaves + +tensorwaves.configure(jax_enable_x64=False) +``` + +TensorWaves also respects JAX's `JAX_ENABLE_X64` environment variable. An explicit +call to `tensorwaves.configure()` takes precedence over the environment variable. + :::::{container} full-width ::::{dropdown} **GPU support** diff --git a/src/tensorwaves/__init__.py b/src/tensorwaves/__init__.py index 6a270d9e2..62fe936d0 100644 --- a/src/tensorwaves/__init__.py +++ b/src/tensorwaves/__init__.py @@ -20,6 +20,7 @@ """ __all__ = [ + "configure", "data", "estimator", "function", @@ -27,3 +28,4 @@ ] from . import data, estimator, function, optimizer +from .config import configure diff --git a/src/tensorwaves/config.py b/src/tensorwaves/config.py new file mode 100644 index 000000000..216c46b00 --- /dev/null +++ b/src/tensorwaves/config.py @@ -0,0 +1,52 @@ +"""Configure optional computational backends.""" + +from __future__ import annotations + +import os +from dataclasses import dataclass +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from types import ModuleType + + +@dataclass +class _JaxConfig: + enable_x64: bool | None = None + initialized: bool = False + + +_jax_config = _JaxConfig() + + +def configure(*, jax_enable_x64: bool) -> None: + """Set the precision used by the JAX backend. + + Call this function before creating JAX arrays or TensorWaves functions. If it is + not called, TensorWaves respects ``JAX_ENABLE_X64`` and otherwise enables 64-bit + precision. + """ + if not isinstance(jax_enable_x64, bool): + msg = "jax_enable_x64 must be a bool" + raise TypeError(msg) + _jax_config.enable_x64 = jax_enable_x64 + if _jax_config.initialized: + import jax # ruff: ignore[import-outside-top-level] + + _set_jax_precision(jax, jax_enable_x64) + + +def _initialize_jax() -> ModuleType: + import jax # ruff: ignore[import-outside-top-level] + + if not _jax_config.initialized: + if _jax_config.enable_x64 is not None: + _set_jax_precision(jax, _jax_config.enable_x64) + elif "JAX_ENABLE_X64" not in os.environ: + _set_jax_precision(jax, True) + _jax_config.initialized = True + return jax + + +def _set_jax_precision(jax: ModuleType, enable_x64: bool) -> None: + jax.config.update("jax_enable_x64", enable_x64) diff --git a/src/tensorwaves/estimator.py b/src/tensorwaves/estimator.py index 18d4167cd..267737be3 100644 --- a/src/tensorwaves/estimator.py +++ b/src/tensorwaves/estimator.py @@ -7,6 +7,7 @@ from typing import TYPE_CHECKING +from tensorwaves.config import _initialize_jax from tensorwaves.data.transform import SympyDataTransformer from tensorwaves.function._backend import find_function, raise_missing_module_error from tensorwaves.function.sympy import create_parametrized_function, prepare_caching @@ -80,10 +81,9 @@ def gradient_creator( ) -> Callable[[Mapping[str, ParameterValue]], dict[str, ParameterValue]]: if backend == "jax": try: - import jax # ruff:ignore[import-outside-top-level] + jax = _initialize_jax() except ImportError: # pragma: no cover raise_missing_module_error("jax", extras_require="jax") - jax.config.update("jax_enable_x64", True) # ty:ignore[possibly-unresolved-reference] gradient = jax.grad(function) # ty:ignore[possibly-unresolved-reference] def conjugated_gradient( diff --git a/src/tensorwaves/function/_backend.py b/src/tensorwaves/function/_backend.py index 3d74f9539..96e4144b7 100644 --- a/src/tensorwaves/function/_backend.py +++ b/src/tensorwaves/function/_backend.py @@ -6,6 +6,8 @@ from typing import TYPE_CHECKING from warnings import warn +from tensorwaves.config import _initialize_jax + if TYPE_CHECKING: from collections.abc import Callable from typing import ParamSpec, TypeVar @@ -40,12 +42,11 @@ def get_backend_modules(backend: str | tuple | dict) -> str | tuple | dict: if isinstance(backend, str): if backend == "jax": try: - import jax + _initialize_jax() import jax.numpy as jnp import jax.scipy as jsp except ImportError: # pragma: no cover raise_missing_module_error("jax", extras_require="jax") - jax.config.update("jax_enable_x64", True) # ty:ignore[possibly-unresolved-reference] return jnp, jsp.special # ty:ignore[possibly-unresolved-reference] if backend in {"numpy", "numba"}: import numpy as np diff --git a/tests/function/test_backend.py b/tests/function/test_backend.py index 6bb150e96..e8edaabf5 100644 --- a/tests/function/test_backend.py +++ b/tests/function/test_backend.py @@ -1,4 +1,59 @@ +# ruff: file-ignore[suspicious-subprocess-import, subprocess-without-shell-equals-true] + +import os +import subprocess +import sys +from typing import Any + +import pytest + +from tensorwaves import configure +from tensorwaves.function._backend import find_function + + +@pytest.mark.parametrize( + ("environment_value", "configuration_value", "expected"), + [ + (None, None, True), + ("false", None, False), + ("true", False, False), + ], +) +def test_jax_precision_configuration( + environment_value: str | None, + configuration_value: bool | None, + expected: bool, +): + environment = os.environ.copy() + if environment_value is None: + environment.pop("JAX_ENABLE_X64", None) + else: + environment["JAX_ENABLE_X64"] = environment_value + configuration = ( + "" + if configuration_value is None + else f"configure(jax_enable_x64={configuration_value})" + ) + code = f""" +from tensorwaves import configure from tensorwaves.function._backend import find_function +{configuration} +find_function("array", backend="jax") +import jax +print(jax.config.x64_enabled) +""" + output = subprocess.check_output( + [sys.executable, "-c", code], + env=environment, + text=True, + ) + assert output.strip() == str(expected) + + +def test_configure_jax_precision_type(): + invalid_value: Any = 1 + with pytest.raises(TypeError, match="jax_enable_x64 must be a bool"): + configure(jax_enable_x64=invalid_value) def test_find_function_jax(): From b013057ce438be8331b4e3dd0a562f0e0e25e0fd Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Thu, 6 Aug 2026 15:58:16 +0200 Subject: [PATCH 2/8] FEAT: set TensorFlow precision --- docs/install.md | 15 +++--- src/tensorwaves/config.py | 74 ++++++++++++++++++++++------ src/tensorwaves/data/rng.py | 8 +-- src/tensorwaves/function/_backend.py | 8 ++- tests/function/test_backend.py | 47 ++++++++++++++++-- 5 files changed, 120 insertions(+), 32 deletions(-) diff --git a/docs/install.md b/docs/install.md index a63dca5e4..30fee986c 100644 --- a/docs/install.md +++ b/docs/install.md @@ -60,19 +60,22 @@ pip install tensorwaves[jax,scipy,tf] pip install tensorwaves[all] # all runtime dependencies ``` -## JAX precision +## Backend precision -TensorWaves uses 64-bit precision for JAX by default. Configure 32-bit precision -before creating JAX arrays or TensorWaves functions with: +TensorWaves uses 64-bit precision for JAX and TensorFlow by default. Configure 32-bit precision before creating backend arrays or TensorWaves functions with: ```python import tensorwaves -tensorwaves.configure(jax_enable_x64=False) +tensorwaves.configure( + jax_enable_x64=False, + tensorflow_prefer_float32=True, +) ``` -TensorWaves also respects JAX's `JAX_ENABLE_X64` environment variable. An explicit -call to `tensorwaves.configure()` takes precedence over the environment variable. +The two options are independent and can be specified separately. TensorWaves also +respects JAX's `JAX_ENABLE_X64` environment variable. An explicit call to +`tensorwaves.configure()` takes precedence over the environment variable. :::::{container} full-width diff --git a/src/tensorwaves/config.py b/src/tensorwaves/config.py index 216c46b00..993a2a313 100644 --- a/src/tensorwaves/config.py +++ b/src/tensorwaves/config.py @@ -4,6 +4,7 @@ import os from dataclasses import dataclass +from importlib import import_module from typing import TYPE_CHECKING if TYPE_CHECKING: @@ -19,25 +20,46 @@ class _JaxConfig: _jax_config = _JaxConfig() -def configure(*, jax_enable_x64: bool) -> None: - """Set the precision used by the JAX backend. +@dataclass +class _TensorFlowConfig: + prefer_float32: bool | None = None + initialized: bool = False - Call this function before creating JAX arrays or TensorWaves functions. If it is - not called, TensorWaves respects ``JAX_ENABLE_X64`` and otherwise enables 64-bit - precision. - """ - if not isinstance(jax_enable_x64, bool): - msg = "jax_enable_x64 must be a bool" - raise TypeError(msg) - _jax_config.enable_x64 = jax_enable_x64 - if _jax_config.initialized: - import jax # ruff: ignore[import-outside-top-level] - _set_jax_precision(jax, jax_enable_x64) +_tensorflow_config = _TensorFlowConfig() + + +def configure( + *, + jax_enable_x64: bool | None = None, + tensorflow_prefer_float32: bool | None = None, +) -> None: + """Set the precision used by computational backends. + + Call this function before creating backend arrays or TensorWaves functions. By + default, TensorWaves uses 64-bit precision for JAX and TensorFlow. TensorWaves + respects ``JAX_ENABLE_X64`` if ``jax_enable_x64`` is not specified. + """ + if jax_enable_x64 is not None: + if not isinstance(jax_enable_x64, bool): + msg = "jax_enable_x64 must be a bool or None" + raise TypeError(msg) + _jax_config.enable_x64 = jax_enable_x64 + if _jax_config.initialized: + jax = import_module("jax") + _set_jax_precision(jax, jax_enable_x64) + if tensorflow_prefer_float32 is not None: + if not isinstance(tensorflow_prefer_float32, bool): + msg = "tensorflow_prefer_float32 must be a bool or None" + raise TypeError(msg) + _tensorflow_config.prefer_float32 = tensorflow_prefer_float32 + if _tensorflow_config.initialized: + tf = import_module("tensorflow") + _enable_tensorflow_numpy_behavior(tf, tensorflow_prefer_float32) def _initialize_jax() -> ModuleType: - import jax # ruff: ignore[import-outside-top-level] + jax = import_module("jax") if not _jax_config.initialized: if _jax_config.enable_x64 is not None: @@ -50,3 +72,27 @@ def _initialize_jax() -> ModuleType: def _set_jax_precision(jax: ModuleType, enable_x64: bool) -> None: jax.config.update("jax_enable_x64", enable_x64) + + +def _initialize_tensorflow() -> ModuleType: + tf = import_module("tensorflow") + + if not _tensorflow_config.initialized: + prefer_float32 = _tensorflow_config.prefer_float32 is True + _enable_tensorflow_numpy_behavior(tf, prefer_float32) + _tensorflow_config.initialized = True + return tf + + +def _enable_tensorflow_numpy_behavior( + tensorflow: ModuleType, prefer_float32: bool +) -> None: + tensorflow.experimental.numpy.experimental_enable_numpy_behavior( + prefer_float32=prefer_float32 + ) + + +def _tensorflow_float_dtype(tensorflow: ModuleType) -> object: + if _tensorflow_config.prefer_float32 is True: + return tensorflow.float32 + return tensorflow.float64 diff --git a/src/tensorwaves/data/rng.py b/src/tensorwaves/data/rng.py index 946e4b7e2..69c68adc8 100644 --- a/src/tensorwaves/data/rng.py +++ b/src/tensorwaves/data/rng.py @@ -2,10 +2,11 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, cast import numpy as np +from tensorwaves.config import _initialize_tensorflow, _tensorflow_float_dtype from tensorwaves.function._backend import raise_missing_module_error from tensorwaves.interface import RealNumberGenerator @@ -41,11 +42,12 @@ class TFUniformRealNumberGenerator(RealNumberGenerator): def __init__(self, seed: int | None = None) -> None: try: - from tensorflow import float64 # ruff:ignore[import-outside-top-level] + tf = _initialize_tensorflow() except ImportError: # pragma: no cover raise_missing_module_error("tensorflow", extras_require="tf") + raise self.seed = seed - self.dtype = float64 # ty:ignore[possibly-unresolved-reference] + self.dtype = cast("tf.DType", _tensorflow_float_dtype(tf)) def __call__( self, size: int, min_value: float = 0.0, max_value: float = 1.0 diff --git a/src/tensorwaves/function/_backend.py b/src/tensorwaves/function/_backend.py index 96e4144b7..e885c5b80 100644 --- a/src/tensorwaves/function/_backend.py +++ b/src/tensorwaves/function/_backend.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING from warnings import warn -from tensorwaves.config import _initialize_jax +from tensorwaves.config import _initialize_jax, _initialize_tensorflow if TYPE_CHECKING: from collections.abc import Callable @@ -55,12 +55,10 @@ def get_backend_modules(backend: str | tuple | dict) -> str | tuple | dict: # returning only np.__dict__ does not work well with conditionals if backend in {"tensorflow", "tf"}: try: - import tensorflow as tf - import tensorflow.experimental.numpy as tnp # ty:ignore[unresolved-import] - from tensorflow.python.ops.numpy_ops import np_config + tf = _initialize_tensorflow() + tnp = tf.experimental.numpy except ImportError: # pragma: no cover raise_missing_module_error("tensorflow", extras_require="tf") - np_config.enable_numpy_behavior() # ty:ignore[possibly-unresolved-reference] return tnp.__dict__, tf # ty:ignore[possibly-unresolved-reference] return backend diff --git a/tests/function/test_backend.py b/tests/function/test_backend.py index e8edaabf5..80e727a25 100644 --- a/tests/function/test_backend.py +++ b/tests/function/test_backend.py @@ -50,10 +50,49 @@ def test_jax_precision_configuration( assert output.strip() == str(expected) -def test_configure_jax_precision_type(): - invalid_value: Any = 1 - with pytest.raises(TypeError, match="jax_enable_x64 must be a bool"): - configure(jax_enable_x64=invalid_value) +@pytest.mark.parametrize( + ("prefer_float32", "expected"), + [ + (None, "float64"), + (True, "float32"), + ], +) +def test_tensorflow_precision_configuration(prefer_float32: bool | None, expected: str): + configuration = ( + "" + if prefer_float32 is None + else f"configure(tensorflow_prefer_float32={prefer_float32})" + ) + code = f""" +from tensorwaves import configure +from tensorwaves.data import TFUniformRealNumberGenerator +from tensorwaves.function._backend import find_function +{configuration} +asarray = find_function("asarray", backend="tensorflow") +array = asarray([1.0]) +random_values = TFUniformRealNumberGenerator(seed=0)(size=1) +print(array.dtype.name, random_values.dtype.name) +""" + output = subprocess.check_output( + [sys.executable, "-c", code], + text=True, + ) + assert output.strip() == f"{expected} {expected}" + + +@pytest.mark.parametrize( + ("argument", "message"), + [ + ({"jax_enable_x64": 1}, "jax_enable_x64 must be a bool or None"), + ( + {"tensorflow_prefer_float32": 1}, + "tensorflow_prefer_float32 must be a bool or None", + ), + ], +) +def test_configure_precision_type(argument: dict[str, Any], message: str): + with pytest.raises(TypeError, match=message): + configure(**argument) def test_find_function_jax(): From db8112ed6222e58be8d799b71b75f721e1b38474 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Thu, 6 Aug 2026 16:20:05 +0200 Subject: [PATCH 3/8] ENH: generalize keyword arguments --- benchmarks/unbinned_nll.py | 2 +- docs/install.md | 4 +-- src/tensorwaves/config.py | 66 +++++++++++++++++----------------- tests/function/test_backend.py | 29 +++++++-------- 4 files changed, 51 insertions(+), 50 deletions(-) diff --git a/benchmarks/unbinned_nll.py b/benchmarks/unbinned_nll.py index a1836325d..a7e8348c0 100644 --- a/benchmarks/unbinned_nll.py +++ b/benchmarks/unbinned_nll.py @@ -15,7 +15,7 @@ from tensorwaves.estimator import UnbinnedNLL from tensorwaves.function import ParametrizedBackendFunction -configure(jax_enable_x64=True) +configure(jax_precision="float64") if TYPE_CHECKING: from collections.abc import Callable diff --git a/docs/install.md b/docs/install.md index 30fee986c..1cbcaa9e0 100644 --- a/docs/install.md +++ b/docs/install.md @@ -68,8 +68,8 @@ TensorWaves uses 64-bit precision for JAX and TensorFlow by default. Configure 3 import tensorwaves tensorwaves.configure( - jax_enable_x64=False, - tensorflow_prefer_float32=True, + jax_precision="float32", + tensorflow_precision="float32", ) ``` diff --git a/src/tensorwaves/config.py b/src/tensorwaves/config.py index 993a2a313..92ba7e251 100644 --- a/src/tensorwaves/config.py +++ b/src/tensorwaves/config.py @@ -5,94 +5,94 @@ import os from dataclasses import dataclass from importlib import import_module -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Literal if TYPE_CHECKING: from types import ModuleType +Precision = Literal["float32", "float64"] + @dataclass class _JaxConfig: - enable_x64: bool | None = None + precision: Precision | None = None initialized: bool = False -_jax_config = _JaxConfig() - - @dataclass class _TensorFlowConfig: - prefer_float32: bool | None = None + precision: Precision | None = None initialized: bool = False +_jax_config = _JaxConfig() _tensorflow_config = _TensorFlowConfig() def configure( *, - jax_enable_x64: bool | None = None, - tensorflow_prefer_float32: bool | None = None, + jax_precision: Literal["float32", "float64"] | None = None, + tensorflow_precision: Literal["float32", "float64"] | None = None, ) -> None: """Set the precision used by computational backends. Call this function before creating backend arrays or TensorWaves functions. By default, TensorWaves uses 64-bit precision for JAX and TensorFlow. TensorWaves - respects ``JAX_ENABLE_X64`` if ``jax_enable_x64`` is not specified. + respects ``JAX_ENABLE_X64`` if ``jax_precision`` is not specified. """ - if jax_enable_x64 is not None: - if not isinstance(jax_enable_x64, bool): - msg = "jax_enable_x64 must be a bool or None" - raise TypeError(msg) - _jax_config.enable_x64 = jax_enable_x64 + if jax_precision is not None: + _validate_precision("jax_precision", jax_precision) + _jax_config.precision = jax_precision if _jax_config.initialized: jax = import_module("jax") - _set_jax_precision(jax, jax_enable_x64) - if tensorflow_prefer_float32 is not None: - if not isinstance(tensorflow_prefer_float32, bool): - msg = "tensorflow_prefer_float32 must be a bool or None" - raise TypeError(msg) - _tensorflow_config.prefer_float32 = tensorflow_prefer_float32 + _set_jax_precision(jax, jax_precision) + if tensorflow_precision is not None: + _validate_precision("tensorflow_precision", tensorflow_precision) + _tensorflow_config.precision = tensorflow_precision if _tensorflow_config.initialized: tf = import_module("tensorflow") - _enable_tensorflow_numpy_behavior(tf, tensorflow_prefer_float32) + _set_tensorflow_precision(tf, tensorflow_precision) def _initialize_jax() -> ModuleType: jax = import_module("jax") if not _jax_config.initialized: - if _jax_config.enable_x64 is not None: - _set_jax_precision(jax, _jax_config.enable_x64) + if _jax_config.precision is not None: + _set_jax_precision(jax, _jax_config.precision) elif "JAX_ENABLE_X64" not in os.environ: - _set_jax_precision(jax, True) + _set_jax_precision(jax, "float64") _jax_config.initialized = True return jax -def _set_jax_precision(jax: ModuleType, enable_x64: bool) -> None: - jax.config.update("jax_enable_x64", enable_x64) +def _set_jax_precision(jax: ModuleType, precision: Precision) -> None: + jax.config.update("jax_enable_x64", precision == "float64") def _initialize_tensorflow() -> ModuleType: tf = import_module("tensorflow") if not _tensorflow_config.initialized: - prefer_float32 = _tensorflow_config.prefer_float32 is True - _enable_tensorflow_numpy_behavior(tf, prefer_float32) + precision = _tensorflow_config.precision or "float64" + _set_tensorflow_precision(tf, precision) _tensorflow_config.initialized = True return tf -def _enable_tensorflow_numpy_behavior( - tensorflow: ModuleType, prefer_float32: bool -) -> None: +def _set_tensorflow_precision(tensorflow: ModuleType, precision: Precision) -> None: tensorflow.experimental.numpy.experimental_enable_numpy_behavior( - prefer_float32=prefer_float32 + prefer_float32=precision == "float32" ) def _tensorflow_float_dtype(tensorflow: ModuleType) -> object: - if _tensorflow_config.prefer_float32 is True: + if _tensorflow_config.precision == "float32": return tensorflow.float32 return tensorflow.float64 + + +def _validate_precision(name: str, precision: object) -> None: + if precision not in {"float32", "float64"}: + msg = f"{name} must be 'float32', 'float64', or None" + raise ValueError(msg) diff --git a/tests/function/test_backend.py b/tests/function/test_backend.py index 80e727a25..9a4d48b89 100644 --- a/tests/function/test_backend.py +++ b/tests/function/test_backend.py @@ -16,12 +16,12 @@ [ (None, None, True), ("false", None, False), - ("true", False, False), + ("true", "float32", False), ], ) def test_jax_precision_configuration( environment_value: str | None, - configuration_value: bool | None, + configuration_value: str | None, expected: bool, ): environment = os.environ.copy() @@ -32,7 +32,7 @@ def test_jax_precision_configuration( configuration = ( "" if configuration_value is None - else f"configure(jax_enable_x64={configuration_value})" + else f"configure(jax_precision={configuration_value!r})" ) code = f""" from tensorwaves import configure @@ -51,17 +51,15 @@ def test_jax_precision_configuration( @pytest.mark.parametrize( - ("prefer_float32", "expected"), + ("precision", "expected"), [ (None, "float64"), - (True, "float32"), + ("float32", "float32"), ], ) -def test_tensorflow_precision_configuration(prefer_float32: bool | None, expected: str): +def test_tensorflow_precision_configuration(precision: str | None, expected: str): configuration = ( - "" - if prefer_float32 is None - else f"configure(tensorflow_prefer_float32={prefer_float32})" + "" if precision is None else f"configure(tensorflow_precision={precision!r})" ) code = f""" from tensorwaves import configure @@ -83,15 +81,18 @@ def test_tensorflow_precision_configuration(prefer_float32: bool | None, expecte @pytest.mark.parametrize( ("argument", "message"), [ - ({"jax_enable_x64": 1}, "jax_enable_x64 must be a bool or None"), ( - {"tensorflow_prefer_float32": 1}, - "tensorflow_prefer_float32 must be a bool or None", + {"jax_precision": "float16"}, + "jax_precision must be 'float32', 'float64', or None", + ), + ( + {"tensorflow_precision": "float16"}, + "tensorflow_precision must be 'float32', 'float64', or None", ), ], ) -def test_configure_precision_type(argument: dict[str, Any], message: str): - with pytest.raises(TypeError, match=message): +def test_configure_precision_value(argument: dict[str, Any], message: str): + with pytest.raises(ValueError, match=message): configure(**argument) From 1acfa76fb2a53455a78743783682f1a6cdfc923e Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Thu, 6 Aug 2026 16:29:19 +0200 Subject: [PATCH 4/8] MAINT: improve function order --- src/tensorwaves/config.py | 32 ++++++++++++++++---------------- tests/function/test_backend.py | 16 ++++++++++------ 2 files changed, 26 insertions(+), 22 deletions(-) diff --git a/src/tensorwaves/config.py b/src/tensorwaves/config.py index 92ba7e251..f18d1a49c 100644 --- a/src/tensorwaves/config.py +++ b/src/tensorwaves/config.py @@ -13,22 +13,6 @@ Precision = Literal["float32", "float64"] -@dataclass -class _JaxConfig: - precision: Precision | None = None - initialized: bool = False - - -@dataclass -class _TensorFlowConfig: - precision: Precision | None = None - initialized: bool = False - - -_jax_config = _JaxConfig() -_tensorflow_config = _TensorFlowConfig() - - def configure( *, jax_precision: Literal["float32", "float64"] | None = None, @@ -96,3 +80,19 @@ def _validate_precision(name: str, precision: object) -> None: if precision not in {"float32", "float64"}: msg = f"{name} must be 'float32', 'float64', or None" raise ValueError(msg) + + +@dataclass +class _JaxConfig: + precision: Precision | None = None + initialized: bool = False + + +@dataclass +class _TensorFlowConfig: + precision: Precision | None = None + initialized: bool = False + + +_jax_config = _JaxConfig() +_tensorflow_config = _TensorFlowConfig() diff --git a/tests/function/test_backend.py b/tests/function/test_backend.py index 9a4d48b89..2c04ee83b 100644 --- a/tests/function/test_backend.py +++ b/tests/function/test_backend.py @@ -12,8 +12,12 @@ @pytest.mark.parametrize( - ("environment_value", "configuration_value", "expected"), - [ + argnames=( + "environment_value", + "configuration_value", + "expected", + ), + argvalues=[ (None, None, True), ("false", None, False), ("true", "float32", False), @@ -51,8 +55,8 @@ def test_jax_precision_configuration( @pytest.mark.parametrize( - ("precision", "expected"), - [ + argnames=("precision", "expected"), + argvalues=[ (None, "float64"), ("float32", "float32"), ], @@ -79,8 +83,8 @@ def test_tensorflow_precision_configuration(precision: str | None, expected: str @pytest.mark.parametrize( - ("argument", "message"), - [ + argnames=("argument", "message"), + argvalues=[ ( {"jax_precision": "float16"}, "jax_precision must be 'float32', 'float64', or None", From 2c97d7ce2659ac808a2edfb2d909ec0d218b17ec Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Thu, 6 Aug 2026 17:18:51 +0200 Subject: [PATCH 5/8] FIX: apply configured precision to computational backends * BEHAVIOR: stop enabling TensorFlow numpy behavior in `TFUniformRealNumberGenerator` * DOC: describe when backend precision takes effect * FIX: configure precision for `tuple` and `dict` backends * FIX: read the value of `JAX_ENABLE_X64` instead of testing for the key * FIX: validate all `configure()` arguments before applying any of them * MAINT: annotate `raise_missing_module_error()` as `NoReturn` * MAINT: confine benchmark precision configuration to the JAX fixtures * MAINT: move configuration tests to `tests/test_config.py` --- benchmarks/unbinned_nll.py | 4 +- docs/install.md | 8 +- src/tensorwaves/config.py | 65 +++++--- src/tensorwaves/data/phasespace.py | 2 +- src/tensorwaves/data/rng.py | 15 +- src/tensorwaves/estimator.py | 2 +- src/tensorwaves/function/_backend.py | 14 +- src/tensorwaves/function/sympy/__init__.py | 6 +- src/tensorwaves/optimizer/callbacks.py | 6 +- src/tensorwaves/optimizer/scipy.py | 2 +- tests/function/test_backend.py | 99 ------------ tests/test_config.py | 176 +++++++++++++++++++++ 12 files changed, 251 insertions(+), 148 deletions(-) create mode 100644 tests/test_config.py diff --git a/benchmarks/unbinned_nll.py b/benchmarks/unbinned_nll.py index a7e8348c0..c5ed4906c 100644 --- a/benchmarks/unbinned_nll.py +++ b/benchmarks/unbinned_nll.py @@ -15,8 +15,6 @@ from tensorwaves.estimator import UnbinnedNLL from tensorwaves.function import ParametrizedBackendFunction -configure(jax_precision="float64") - if TYPE_CHECKING: from collections.abc import Callable @@ -206,6 +204,7 @@ def estimator_samples() -> tuple[dict[str, np.ndarray], dict[str, np.ndarray]]: def jax_intensities( intensities: tuple[np.ndarray, np.ndarray], ) -> tuple[jax.Array, jax.Array]: + configure(jax_precision="float64") data_intensities, phsp_intensities = intensities jax_data_intensities = jnp.asarray(data_intensities).block_until_ready() jax_phsp_intensities = jnp.asarray(phsp_intensities).block_until_ready() @@ -224,6 +223,7 @@ def tensorflow_intensities( def jax_estimator_samples( estimator_samples: tuple[dict[str, np.ndarray], dict[str, np.ndarray]], ) -> tuple[dict[str, jax.Array], dict[str, jax.Array]]: + configure(jax_precision="float64") data, phsp = estimator_samples return ( {"x": jnp.asarray(data["x"]).block_until_ready()}, diff --git a/docs/install.md b/docs/install.md index 1cbcaa9e0..f3904faf4 100644 --- a/docs/install.md +++ b/docs/install.md @@ -73,9 +73,11 @@ tensorwaves.configure( ) ``` -The two options are independent and can be specified separately. TensorWaves also -respects JAX's `JAX_ENABLE_X64` environment variable. An explicit call to -`tensorwaves.configure()` takes precedence over the environment variable. +The two options are independent and can be specified separately. + +Precision is a property of the backend, not of TensorWaves, so it applies from the moment it is set: arrays that were created earlier keep the dtype they were created with. A backend that has already been imported is reconfigured immediately, and one that has not is configured when TensorWaves first uses it. Calling `tensorwaves.configure()` right after your imports therefore covers both cases. + +TensorWaves also respects JAX's `JAX_ENABLE_X64` environment variable. An explicit call to `tensorwaves.configure()` takes precedence over it. Note that JAX itself only reads that variable when it is imported, whereas TensorWaves reads it when it first uses JAX. TensorFlow has no equivalent environment variable, so `tensorflow_precision` is the only way to select its precision. :::::{container} full-width diff --git a/src/tensorwaves/config.py b/src/tensorwaves/config.py index f18d1a49c..98a5ced30 100644 --- a/src/tensorwaves/config.py +++ b/src/tensorwaves/config.py @@ -3,11 +3,13 @@ from __future__ import annotations import os +import sys from dataclasses import dataclass from importlib import import_module -from typing import TYPE_CHECKING, Literal +from typing import TYPE_CHECKING, Literal, get_args if TYPE_CHECKING: + from collections.abc import Callable from types import ModuleType Precision = Literal["float32", "float64"] @@ -15,8 +17,8 @@ def configure( *, - jax_precision: Literal["float32", "float64"] | None = None, - tensorflow_precision: Literal["float32", "float64"] | None = None, + jax_precision: Precision | None = None, + tensorflow_precision: Precision | None = None, ) -> None: """Set the precision used by computational backends. @@ -24,28 +26,41 @@ def configure( default, TensorWaves uses 64-bit precision for JAX and TensorFlow. TensorWaves respects ``JAX_ENABLE_X64`` if ``jax_precision`` is not specified. """ + _validate_precision("jax_precision", jax_precision) + _validate_precision("tensorflow_precision", tensorflow_precision) if jax_precision is not None: - _validate_precision("jax_precision", jax_precision) _jax_config.precision = jax_precision - if _jax_config.initialized: - jax = import_module("jax") - _set_jax_precision(jax, jax_precision) + _configure_imported_module("jax", _set_jax_precision, jax_precision) if tensorflow_precision is not None: - _validate_precision("tensorflow_precision", tensorflow_precision) _tensorflow_config.precision = tensorflow_precision - if _tensorflow_config.initialized: - tf = import_module("tensorflow") - _set_tensorflow_precision(tf, tensorflow_precision) + _configure_imported_module( + "tensorflow", _set_tensorflow_precision, tensorflow_precision + ) + + +def _configure_imported_module( + module_name: str, + set_precision: Callable[[ModuleType, Precision], None], + precision: Precision, +) -> None: + """Set the precision on a backend that has already been imported. + + A backend that has not been imported yet is configured on first use, so that + :func:`configure` never triggers a backend import itself. + """ + module = sys.modules.get(module_name) + if module is not None: + set_precision(module, precision) def _initialize_jax() -> ModuleType: jax = import_module("jax") if not _jax_config.initialized: - if _jax_config.precision is not None: - _set_jax_precision(jax, _jax_config.precision) - elif "JAX_ENABLE_X64" not in os.environ: - _set_jax_precision(jax, "float64") + precision = _jax_config.precision + if precision is None: + precision = _precision_from_flag(os.environ.get("JAX_ENABLE_X64", "1")) + _set_jax_precision(jax, precision) _jax_config.initialized = True return jax @@ -54,12 +69,20 @@ def _set_jax_precision(jax: ModuleType, precision: Precision) -> None: jax.config.update("jax_enable_x64", precision == "float64") +def _precision_from_flag(value: str) -> Precision: + """Interpret a JAX-style boolean environment variable value. + + >>> _precision_from_flag("1"), _precision_from_flag("false") + ('float64', 'float32') + """ + return "float64" if value.strip().lower() in {"1", "true", "yes"} else "float32" + + def _initialize_tensorflow() -> ModuleType: tf = import_module("tensorflow") if not _tensorflow_config.initialized: - precision = _tensorflow_config.precision or "float64" - _set_tensorflow_precision(tf, precision) + _set_tensorflow_precision(tf, _tensorflow_precision()) _tensorflow_config.initialized = True return tf @@ -70,14 +93,12 @@ def _set_tensorflow_precision(tensorflow: ModuleType, precision: Precision) -> N ) -def _tensorflow_float_dtype(tensorflow: ModuleType) -> object: - if _tensorflow_config.precision == "float32": - return tensorflow.float32 - return tensorflow.float64 +def _tensorflow_precision() -> Precision: + return _tensorflow_config.precision or "float64" def _validate_precision(name: str, precision: object) -> None: - if precision not in {"float32", "float64"}: + if precision is not None and precision not in get_args(Precision): msg = f"{name} must be 'float32', 'float64', or None" raise ValueError(msg) diff --git a/src/tensorwaves/data/phasespace.py b/src/tensorwaves/data/phasespace.py index 10483f13b..cc8238248 100644 --- a/src/tensorwaves/data/phasespace.py +++ b/src/tensorwaves/data/phasespace.py @@ -107,7 +107,7 @@ def __init__( except ImportError: # pragma: no cover raise_missing_module_error("phasespace", extras_require="phsp") sorted_ids = sorted(final_state_masses) - self.__phsp_gen = phasespace.nbody_decay( # ty:ignore[possibly-unresolved-reference] + self.__phsp_gen = phasespace.nbody_decay( mass_top=initial_state_mass, masses=[final_state_masses[i] for i in sorted_ids], names=list(map(str, sorted_ids)), diff --git a/src/tensorwaves/data/rng.py b/src/tensorwaves/data/rng.py index 69c68adc8..3054ef71d 100644 --- a/src/tensorwaves/data/rng.py +++ b/src/tensorwaves/data/rng.py @@ -2,11 +2,11 @@ from __future__ import annotations -from typing import TYPE_CHECKING, cast +from typing import TYPE_CHECKING import numpy as np -from tensorwaves.config import _initialize_tensorflow, _tensorflow_float_dtype +from tensorwaves.config import _tensorflow_precision from tensorwaves.function._backend import raise_missing_module_error from tensorwaves.interface import RealNumberGenerator @@ -42,12 +42,11 @@ class TFUniformRealNumberGenerator(RealNumberGenerator): def __init__(self, seed: int | None = None) -> None: try: - tf = _initialize_tensorflow() + import tensorflow as tf # ruff:ignore[import-outside-top-level] except ImportError: # pragma: no cover raise_missing_module_error("tensorflow", extras_require="tf") - raise self.seed = seed - self.dtype = cast("tf.DType", _tensorflow_float_dtype(tf)) + self.dtype = tf.float32 if _tensorflow_precision() == "float32" else tf.float64 def __call__( self, size: int, min_value: float = 0.0, max_value: float = 1.0 @@ -80,10 +79,10 @@ def _get_tensorflow_rng(seed: SeedLike | None = None) -> tf.random.Generator: raise_missing_module_error("tensorflow", extras_require="tf") if seed is None: - return tf.random.get_global_generator() # ty:ignore[possibly-unresolved-reference] + return tf.random.get_global_generator() if isinstance(seed, int): - return tf.random.Generator.from_seed(seed=seed) # ty:ignore[possibly-unresolved-reference] - if isinstance(seed, tf.random.Generator): # ty:ignore[possibly-unresolved-reference] + return tf.random.Generator.from_seed(seed=seed) + if isinstance(seed, tf.random.Generator): return seed msg = f"Cannot create a tf.random.Generator from a {type(seed).__name__}" raise TypeError(msg) diff --git a/src/tensorwaves/estimator.py b/src/tensorwaves/estimator.py index 267737be3..3e3ef68e3 100644 --- a/src/tensorwaves/estimator.py +++ b/src/tensorwaves/estimator.py @@ -84,7 +84,7 @@ def gradient_creator( jax = _initialize_jax() except ImportError: # pragma: no cover raise_missing_module_error("jax", extras_require="jax") - gradient = jax.grad(function) # ty:ignore[possibly-unresolved-reference] + gradient = jax.grad(function) def conjugated_gradient( parameters: Mapping[str, ParameterValue], diff --git a/src/tensorwaves/function/_backend.py b/src/tensorwaves/function/_backend.py index e885c5b80..c1668b0e2 100644 --- a/src/tensorwaves/function/_backend.py +++ b/src/tensorwaves/function/_backend.py @@ -10,7 +10,7 @@ if TYPE_CHECKING: from collections.abc import Callable - from typing import ParamSpec, TypeVar + from typing import NoReturn, ParamSpec, TypeVar P = ParamSpec("P") T = TypeVar("T") @@ -47,7 +47,7 @@ def get_backend_modules(backend: str | tuple | dict) -> str | tuple | dict: import jax.scipy as jsp except ImportError: # pragma: no cover raise_missing_module_error("jax", extras_require="jax") - return jnp, jsp.special # ty:ignore[possibly-unresolved-reference] + return jnp, jsp.special if backend in {"numpy", "numba"}: import numpy as np @@ -59,7 +59,7 @@ def get_backend_modules(backend: str | tuple | dict) -> str | tuple | dict: tnp = tf.experimental.numpy except ImportError: # pragma: no cover raise_missing_module_error("tensorflow", extras_require="tf") - return tnp.__dict__, tf # ty:ignore[possibly-unresolved-reference] + return tnp.__dict__, tf return backend @@ -84,14 +84,14 @@ def jit_compile(backend: str) -> Callable[[Callable[P, T]], Callable[P, T]]: import jax except ImportError: # pragma: no cover raise_missing_module_error("jax", extras_require="jax") - return jax.jit # ty:ignore[possibly-unresolved-reference] + return jax.jit if backend == "numba": try: import numba except ImportError: # pragma: no cover raise_missing_module_error("numba", extras_require="numba") - return partial(numba.jit, forceobj=True, parallel=True) # ty:ignore[possibly-unresolved-reference] + return partial(numba.jit, forceobj=True, parallel=True) msg = f"Backend {backend} does not yet support JIT compilation" warn(msg, category=UserWarning, stacklevel=3) @@ -102,7 +102,9 @@ def _do_not_compile(function: Callable[P, T]) -> Callable[P, T]: return function -def raise_missing_module_error(module_name: str, *, extras_require: str = "") -> None: +def raise_missing_module_error( + module_name: str, *, extras_require: str = "" +) -> NoReturn: """Raise an `ImportError` with install instructions. >>> raise_missing_module_error("missing") diff --git a/src/tensorwaves/function/sympy/__init__.py b/src/tensorwaves/function/sympy/__init__.py index cb140a05c..28bcad821 100644 --- a/src/tensorwaves/function/sympy/__init__.py +++ b/src/tensorwaves/function/sympy/__init__.py @@ -7,6 +7,7 @@ from tqdm.auto import tqdm +from tensorwaves.config import _initialize_jax, _initialize_tensorflow from tensorwaves.function import ParametrizedBackendFunction, PositionalArgumentFunction from tensorwaves.function._backend import ( get_backend_modules, @@ -245,6 +246,7 @@ def lambdify( # ruff:ignore[complex-structure, too-many-return-statements] """ def jax_lambdify() -> Callable: + _initialize_jax() from ._printer import JaxPrinter return _sympy_lambdify( @@ -265,7 +267,7 @@ def numba_lambdify() -> Callable: def tensorflow_lambdify() -> Callable: try: - import tensorflow.experimental.numpy as tnp # ty:ignore[unresolved-import] + tnp = _initialize_tensorflow().experimental.numpy except ImportError: # pragma: no cover raise_missing_module_error("tensorflow", extras_require="tf") from ._printer import TensorflowPrinter @@ -273,7 +275,7 @@ def tensorflow_lambdify() -> Callable: return _sympy_lambdify( expression, symbols, - modules=tnp, # ty:ignore[possibly-unresolved-reference] + modules=tnp, printer=TensorflowPrinter(), use_cse=use_cse, ) diff --git a/src/tensorwaves/optimizer/callbacks.py b/src/tensorwaves/optimizer/callbacks.py index 22dcf66e7..88e16d5ad 100644 --- a/src/tensorwaves/optimizer/callbacks.py +++ b/src/tensorwaves/optimizer/callbacks.py @@ -361,7 +361,7 @@ def on_optimize_start(self, logs: dict[str, Any] | None = None) -> None: output_dir = self.__logdir + "/" + datetime.now().strftime("%Y%m%d-%H%M%S") if self.__subdir is not None: output_dir += "/" + self.__subdir - self.__stream = tf.summary.create_file_writer(output_dir) # ty:ignore[possibly-unresolved-reference] + self.__stream = tf.summary.create_file_writer(output_dir) self.__stream.set_as_default() def on_optimize_end(self, logs: dict[str, Any] | None = None) -> None: @@ -387,10 +387,10 @@ def on_function_call_end( return parameters = logs["parameters"] for par_name, value in parameters.items(): - tf.summary.scalar(par_name, value, step=function_call) # ty:ignore[possibly-unresolved-reference] + tf.summary.scalar(par_name, value, step=function_call) estimator_value = logs.get("estimator", {}).get("value", None) if estimator_value is not None: - tf.summary.scalar("estimator", estimator_value, step=function_call) # ty:ignore[possibly-unresolved-reference] + tf.summary.scalar("estimator", estimator_value, step=function_call) if self.__stream is not None: self.__stream.flush() diff --git a/src/tensorwaves/optimizer/scipy.py b/src/tensorwaves/optimizer/scipy.py index 8bbd9f7d0..915ea2b2d 100644 --- a/src/tensorwaves/optimizer/scipy.py +++ b/src/tensorwaves/optimizer/scipy.py @@ -120,7 +120,7 @@ def wrapped_callback(pars: Iterable[float]) -> None: ) start_time = time.time() - fit_result = minimize( # ty:ignore[possibly-unresolved-reference] + fit_result = minimize( wrapped_function, list(flattened_parameters.values()), method=self.__method, diff --git a/tests/function/test_backend.py b/tests/function/test_backend.py index 2c04ee83b..6bb150e96 100644 --- a/tests/function/test_backend.py +++ b/tests/function/test_backend.py @@ -1,103 +1,4 @@ -# ruff: file-ignore[suspicious-subprocess-import, subprocess-without-shell-equals-true] - -import os -import subprocess -import sys -from typing import Any - -import pytest - -from tensorwaves import configure -from tensorwaves.function._backend import find_function - - -@pytest.mark.parametrize( - argnames=( - "environment_value", - "configuration_value", - "expected", - ), - argvalues=[ - (None, None, True), - ("false", None, False), - ("true", "float32", False), - ], -) -def test_jax_precision_configuration( - environment_value: str | None, - configuration_value: str | None, - expected: bool, -): - environment = os.environ.copy() - if environment_value is None: - environment.pop("JAX_ENABLE_X64", None) - else: - environment["JAX_ENABLE_X64"] = environment_value - configuration = ( - "" - if configuration_value is None - else f"configure(jax_precision={configuration_value!r})" - ) - code = f""" -from tensorwaves import configure from tensorwaves.function._backend import find_function -{configuration} -find_function("array", backend="jax") -import jax -print(jax.config.x64_enabled) -""" - output = subprocess.check_output( - [sys.executable, "-c", code], - env=environment, - text=True, - ) - assert output.strip() == str(expected) - - -@pytest.mark.parametrize( - argnames=("precision", "expected"), - argvalues=[ - (None, "float64"), - ("float32", "float32"), - ], -) -def test_tensorflow_precision_configuration(precision: str | None, expected: str): - configuration = ( - "" if precision is None else f"configure(tensorflow_precision={precision!r})" - ) - code = f""" -from tensorwaves import configure -from tensorwaves.data import TFUniformRealNumberGenerator -from tensorwaves.function._backend import find_function -{configuration} -asarray = find_function("asarray", backend="tensorflow") -array = asarray([1.0]) -random_values = TFUniformRealNumberGenerator(seed=0)(size=1) -print(array.dtype.name, random_values.dtype.name) -""" - output = subprocess.check_output( - [sys.executable, "-c", code], - text=True, - ) - assert output.strip() == f"{expected} {expected}" - - -@pytest.mark.parametrize( - argnames=("argument", "message"), - argvalues=[ - ( - {"jax_precision": "float16"}, - "jax_precision must be 'float32', 'float64', or None", - ), - ( - {"tensorflow_precision": "float16"}, - "tensorflow_precision must be 'float32', 'float64', or None", - ), - ], -) -def test_configure_precision_value(argument: dict[str, Any], message: str): - with pytest.raises(ValueError, match=message): - configure(**argument) def test_find_function_jax(): diff --git a/tests/test_config.py b/tests/test_config.py new file mode 100644 index 000000000..e62468117 --- /dev/null +++ b/tests/test_config.py @@ -0,0 +1,176 @@ +# ruff: file-ignore[suspicious-subprocess-import, subprocess-without-shell-equals-true] + +import os +import subprocess +import sys +from typing import Any + +import pytest + +from tensorwaves import configure +from tensorwaves.config import _jax_config, _tensorflow_precision + + +def _run(code: str, env: dict[str, str] | None = None) -> str: + result = subprocess.run( + [sys.executable, "-c", code], + capture_output=True, + check=False, + env=env, + text=True, + timeout=600, + ) + assert result.returncode == 0, result.stderr + return result.stdout.strip() + + +@pytest.fixture +def restore_jax_precision(): + import jax + + precision = _jax_config.precision + x64_enabled = jax.config.x64_enabled + yield + _jax_config.precision = precision + jax.config.update("jax_enable_x64", x64_enabled) + + +@pytest.mark.parametrize( + argnames=( + "environment_value", + "configuration_value", + "expected", + ), + argvalues=[ + (None, None, True), + ("false", None, False), + ("0", None, False), + ("1", None, True), + ("true", "float32", False), + ], +) +def test_jax_precision_configuration( + environment_value: str | None, + configuration_value: str | None, + expected: bool, +): + environment = os.environ.copy() + if environment_value is None: + environment.pop("JAX_ENABLE_X64", None) + else: + environment["JAX_ENABLE_X64"] = environment_value + configuration = ( + "" + if configuration_value is None + else f"configure(jax_precision={configuration_value!r})" + ) + code = f""" +from tensorwaves import configure +from tensorwaves.function._backend import find_function +{configuration} +find_function("array", backend="jax") +import jax +print(jax.config.x64_enabled) +""" + assert _run(code, env=environment) == str(expected) + + +def test_jax_reads_environment_variable_set_after_import(): + """JAX resolves ``JAX_ENABLE_X64`` on import, TensorWaves on first backend use.""" + environment = os.environ.copy() + environment.pop("JAX_ENABLE_X64", None) + code = """ +import jax +import os +os.environ["JAX_ENABLE_X64"] = "1" +from tensorwaves.function._backend import find_function +find_function("array", backend="jax") +print(jax.config.x64_enabled) +""" + assert _run(code, env=environment) == "True" + + +@pytest.mark.parametrize("precision", ["float32", "float64"]) +def test_configure_before_creating_arrays(precision: str): + code = f""" +import jax.numpy as jnp +from tensorwaves import configure +configure(jax_precision={precision!r}) +print(jnp.asarray([1.0]).dtype.name) +""" + assert _run(code) == precision + + +@pytest.mark.parametrize( + argnames=("precision", "expected"), + argvalues=[ + (None, "float64"), + ("float32", "float32"), + ], +) +def test_tensorflow_precision_configuration(precision: str | None, expected: str): + configuration = ( + "" if precision is None else f"configure(tensorflow_precision={precision!r})" + ) + code = f""" +from tensorwaves import configure +from tensorwaves.data import TFUniformRealNumberGenerator +from tensorwaves.function._backend import find_function +{configuration} +asarray = find_function("asarray", backend="tensorflow") +array = asarray([1.0]) +random_values = TFUniformRealNumberGenerator(seed=0)(size=1) +print(array.dtype.name, random_values.dtype.name) +""" + assert _run(code) == f"{expected} {expected}" + + +def test_tensorflow_rng_does_not_enable_numpy_behavior(): + code = """ +import tensorflow as tf +from tensorwaves.data import TFUniformRealNumberGenerator +TFUniformRealNumberGenerator(seed=0) +print(hasattr(tf.constant([1, 2]), "astype")) +""" + assert _run(code) == "False" + + +@pytest.mark.usefixtures("restore_jax_precision") +def test_configure_applies_to_imported_backend(): + import jax + + configure(jax_precision="float32") + assert not jax.config.x64_enabled + configure(jax_precision="float64") + assert jax.config.x64_enabled + + +@pytest.mark.parametrize( + argnames=("argument", "message"), + argvalues=[ + ( + {"jax_precision": "float16"}, + "jax_precision must be 'float32', 'float64', or None", + ), + ( + {"tensorflow_precision": "float16"}, + "tensorflow_precision must be 'float32', 'float64', or None", + ), + ], +) +def test_configure_precision_value(argument: dict[str, Any], message: str): + with pytest.raises(ValueError, match=message): + configure(**argument) + + +def test_configure_validates_before_applying(): + jax_precision = _jax_config.precision + tensorflow_precision = _tensorflow_precision() + arguments: dict[str, Any] = { + "jax_precision": "float32", + "tensorflow_precision": "float16", + } + with pytest.raises(ValueError, match="tensorflow_precision"): + configure(**arguments) + assert _jax_config.precision == jax_precision + assert _tensorflow_precision() == tensorflow_precision From 943829fb754e2264d9cfd3aa0a0d0c7d8bb502a8 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Thu, 6 Aug 2026 19:14:17 +0200 Subject: [PATCH 6/8] DX: add benchmark that tests precision --- benchmarks/precision.py | 39 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 39 insertions(+) create mode 100644 benchmarks/precision.py diff --git a/benchmarks/precision.py b/benchmarks/precision.py new file mode 100644 index 000000000..b409ee880 --- /dev/null +++ b/benchmarks/precision.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import numpy as np +import pytest +import sympy as sp + +from tensorwaves import configure +from tensorwaves.function.sympy import create_function + +if TYPE_CHECKING: + from tensorwaves.config import Precision + + +@pytest.mark.benchmark(group="precision") +@pytest.mark.parametrize("backend", ["jax", "tensorflow"]) +@pytest.mark.parametrize("precision", ["float32", "float64"]) +def test_precision(benchmark, backend: str, precision: Precision) -> None: + _configure_backend(backend, precision) + x = sp.Symbol("x") + function = create_function(sp.sin(x) ** 2 + sp.exp(-x), backend=backend) + data = {"x": np.linspace(0, 10, num=1_000_000, dtype=precision)} + result = benchmark(lambda: _evaluate(function, data, backend)) + assert result.dtype.name == precision + + +def _configure_backend(backend: str, precision: Precision) -> None: + if backend == "jax": + configure(jax_precision=precision) + else: + configure(tensorflow_precision=precision) + + +def _evaluate(function, data: dict[str, np.ndarray], backend: str): + result = function(data) + if backend == "jax": + result.block_until_ready() + return result From 698823192e6d23c68b2dbe8a8025e1be0be6c5c6 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Thu, 6 Aug 2026 20:11:23 +0200 Subject: [PATCH 7/8] FIX: assert that backend precision has indeed been set --- benchmarks/precision.py | 42 +++++++++++++++++++++++++++++++++++------ 1 file changed, 36 insertions(+), 6 deletions(-) diff --git a/benchmarks/precision.py b/benchmarks/precision.py index b409ee880..f15471c28 100644 --- a/benchmarks/precision.py +++ b/benchmarks/precision.py @@ -2,7 +2,6 @@ from typing import TYPE_CHECKING -import numpy as np import pytest import sympy as sp @@ -10,30 +9,61 @@ from tensorwaves.function.sympy import create_function if TYPE_CHECKING: + from typing import Literal + from tensorwaves.config import Precision + from tensorwaves.interface import DataSample + + Backend = Literal["jax", "tensorflow"] @pytest.mark.benchmark(group="precision") @pytest.mark.parametrize("backend", ["jax", "tensorflow"]) @pytest.mark.parametrize("precision", ["float32", "float64"]) -def test_precision(benchmark, backend: str, precision: Precision) -> None: +def test_precision(benchmark, backend: Backend, precision: Precision) -> None: _configure_backend(backend, precision) x = sp.Symbol("x") function = create_function(sp.sin(x) ** 2 + sp.exp(-x), backend=backend) - data = {"x": np.linspace(0, 10, num=1_000_000, dtype=precision)} + data = _create_data(backend) result = benchmark(lambda: _evaluate(function, data, backend)) - assert result.dtype.name == precision + _assert_precision(backend, precision, data, result) -def _configure_backend(backend: str, precision: Precision) -> None: +def _configure_backend(backend: Backend, precision: Precision) -> None: if backend == "jax": configure(jax_precision=precision) else: configure(tensorflow_precision=precision) -def _evaluate(function, data: dict[str, np.ndarray], backend: str): +def _create_data(backend: Backend): + if backend == "jax": + import jax.numpy as jnp + + return {"x": jnp.linspace(0, 10, num=1_000_000).block_until_ready()} + + import tensorflow.experimental.numpy as tnp # ty: ignore[unresolved-import] + + return {"x": tnp.linspace(0, 10, num=1_000_000)} + + +def _evaluate(function, data: DataSample, backend: Backend): result = function(data) if backend == "jax": result.block_until_ready() return result + + +def _assert_precision( + backend: Backend, precision: Precision, data: dict, result +) -> None: + assert data["x"].dtype.name == precision + assert result.dtype.name == precision + if backend == "jax": + import jax + + assert jax.config.x64_enabled == (precision == "float64") + else: + import tensorflow.experimental.numpy as tnp # ty: ignore[unresolved-import] + + assert tnp.asarray(1.0).dtype.name == precision From ad5a4cfaa9c9fa915087318c2b0c94c8ad3e5b72 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Fri, 7 Aug 2026 09:33:23 +0200 Subject: [PATCH 8/8] Kick CI