diff --git a/benchmarks/precision.py b/benchmarks/precision.py new file mode 100644 index 00000000..f15471c2 --- /dev/null +++ b/benchmarks/precision.py @@ -0,0 +1,69 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import pytest +import sympy as sp + +from tensorwaves import configure +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: 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 = _create_data(backend) + result = benchmark(lambda: _evaluate(function, data, backend)) + _assert_precision(backend, precision, data, result) + + +def _configure_backend(backend: Backend, precision: Precision) -> None: + if backend == "jax": + configure(jax_precision=precision) + else: + configure(tensorflow_precision=precision) + + +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 diff --git a/benchmarks/unbinned_nll.py b/benchmarks/unbinned_nll.py index 295df212..c5ed4906 100644 --- a/benchmarks/unbinned_nll.py +++ b/benchmarks/unbinned_nll.py @@ -11,6 +11,7 @@ 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 @@ -203,8 +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]: - jax.config.update("jax_enable_x64", True) - + 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() @@ -223,8 +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]]: - jax.config.update("jax_enable_x64", True) - + 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 6edd27a7..f3904faf 100644 --- a/docs/install.md +++ b/docs/install.md @@ -60,6 +60,25 @@ pip install tensorwaves[jax,scipy,tf] pip install tensorwaves[all] # all runtime dependencies ``` +## Backend precision + +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_precision="float32", + tensorflow_precision="float32", +) +``` + +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 ::::{dropdown} **GPU support** diff --git a/src/tensorwaves/__init__.py b/src/tensorwaves/__init__.py index 6a270d9e..62fe936d 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 00000000..98a5ced3 --- /dev/null +++ b/src/tensorwaves/config.py @@ -0,0 +1,119 @@ +"""Configure optional computational backends.""" + +from __future__ import annotations + +import os +import sys +from dataclasses import dataclass +from importlib import import_module +from typing import TYPE_CHECKING, Literal, get_args + +if TYPE_CHECKING: + from collections.abc import Callable + from types import ModuleType + +Precision = Literal["float32", "float64"] + + +def configure( + *, + jax_precision: Precision | None = None, + tensorflow_precision: Precision | 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_precision`` is not specified. + """ + _validate_precision("jax_precision", jax_precision) + _validate_precision("tensorflow_precision", tensorflow_precision) + if jax_precision is not None: + _jax_config.precision = jax_precision + _configure_imported_module("jax", _set_jax_precision, jax_precision) + if tensorflow_precision is not None: + _tensorflow_config.precision = 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: + 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 + + +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: + _set_tensorflow_precision(tf, _tensorflow_precision()) + _tensorflow_config.initialized = True + return tf + + +def _set_tensorflow_precision(tensorflow: ModuleType, precision: Precision) -> None: + tensorflow.experimental.numpy.experimental_enable_numpy_behavior( + prefer_float32=precision == "float32" + ) + + +def _tensorflow_precision() -> Precision: + return _tensorflow_config.precision or "float64" + + +def _validate_precision(name: str, precision: object) -> None: + if precision is not None and precision not in get_args(Precision): + 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/src/tensorwaves/data/phasespace.py b/src/tensorwaves/data/phasespace.py index 10483f13..cc823824 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 946e4b7e..3054ef71 100644 --- a/src/tensorwaves/data/rng.py +++ b/src/tensorwaves/data/rng.py @@ -6,6 +6,7 @@ import numpy as np +from tensorwaves.config import _tensorflow_precision from tensorwaves.function._backend import raise_missing_module_error from tensorwaves.interface import RealNumberGenerator @@ -41,11 +42,11 @@ class TFUniformRealNumberGenerator(RealNumberGenerator): def __init__(self, seed: int | None = None) -> None: try: - from tensorflow import float64 # ruff:ignore[import-outside-top-level] + import tensorflow as tf # ruff:ignore[import-outside-top-level] except ImportError: # pragma: no cover raise_missing_module_error("tensorflow", extras_require="tf") self.seed = seed - self.dtype = float64 # ty:ignore[possibly-unresolved-reference] + 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 @@ -78,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 18d4167c..3e3ef68e 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,11 +81,10 @@ 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] + 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 3d74f953..c1668b0e 100644 --- a/src/tensorwaves/function/_backend.py +++ b/src/tensorwaves/function/_backend.py @@ -6,9 +6,11 @@ from typing import TYPE_CHECKING from warnings import warn +from tensorwaves.config import _initialize_jax, _initialize_tensorflow + 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") @@ -40,13 +42,12 @@ 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] + return jnp, jsp.special if backend in {"numpy", "numba"}: import numpy as np @@ -54,13 +55,11 @@ 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 tnp.__dict__, tf return backend @@ -85,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) @@ -103,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 cb140a05..28bcad82 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 22dcf66e..88e16d5a 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 8bbd9f7d..915ea2b2 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/test_config.py b/tests/test_config.py new file mode 100644 index 00000000..e6246811 --- /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