From 7dab7f2e52afe1a352a51e8eb783aac46428dac4 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:16:04 +0200 Subject: [PATCH 01/11] DX: add benchmark for different NLL implementations --- benchmarks/unbinned_nll.py | 65 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 65 insertions(+) create mode 100644 benchmarks/unbinned_nll.py diff --git a/benchmarks/unbinned_nll.py b/benchmarks/unbinned_nll.py new file mode 100644 index 00000000..b880aecc --- /dev/null +++ b/benchmarks/unbinned_nll.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import numpy as np +import pytest + +if TYPE_CHECKING: + from collections.abc import Callable + + +def _original_unbinned_nll( + bare_intensities: np.ndarray, + phsp_intensities: np.ndarray, +) -> float: + normalization_factor = 1.0 / np.mean(phsp_intensities) + likelihoods = normalization_factor * bare_intensities + return -np.sum(np.log(likelihoods)) + + +def _optimized_unbinned_nll( + bare_intensities: np.ndarray, + phsp_intensities: np.ndarray, +) -> float: + normalization_integral = np.mean(phsp_intensities) + return len(bare_intensities) * np.log(normalization_integral) - np.sum( + np.log(bare_intensities) + ) + + +_IMPLEMENTATIONS: dict[ + str, + Callable[[np.ndarray, np.ndarray], float], +] = { + "original": _original_unbinned_nll, + "optimized": _optimized_unbinned_nll, +} + + +@pytest.fixture(scope="module") +def intensities() -> tuple[np.ndarray, np.ndarray]: + rng = np.random.default_rng(seed=0) + bare_intensities = rng.uniform(low=0.1, high=10.0, size=5_000_000) + phsp_intensities = rng.uniform(low=0.1, high=10.0, size=1_000_000) + return bare_intensities, phsp_intensities + + +@pytest.mark.benchmark(group="unbinned-nll-normalization") +@pytest.mark.parametrize("implementation", _IMPLEMENTATIONS) +def test_unbinned_nll_normalization_formula( + benchmark, + implementation: str, + intensities: tuple[np.ndarray, np.ndarray], +) -> None: + bare_intensities, phsp_intensities = intensities + reference = _original_unbinned_nll( + bare_intensities, + phsp_intensities, + ) + result = benchmark( + _IMPLEMENTATIONS[implementation], + bare_intensities, + phsp_intensities, + ) + assert result == pytest.approx(reference) From 9cdfd873d2044c8fc62caee488f2cba5ff4ccf37 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:16:05 +0200 Subject: [PATCH 02/11] MAINT: rename `bare` to `data` --- benchmarks/unbinned_nll.py | 20 ++++++++++---------- src/tensorwaves/estimator.py | 4 ++-- 2 files changed, 12 insertions(+), 12 deletions(-) diff --git a/benchmarks/unbinned_nll.py b/benchmarks/unbinned_nll.py index b880aecc..f6c4ef6e 100644 --- a/benchmarks/unbinned_nll.py +++ b/benchmarks/unbinned_nll.py @@ -10,21 +10,21 @@ def _original_unbinned_nll( - bare_intensities: np.ndarray, + data_intensities: np.ndarray, phsp_intensities: np.ndarray, ) -> float: normalization_factor = 1.0 / np.mean(phsp_intensities) - likelihoods = normalization_factor * bare_intensities + likelihoods = normalization_factor * data_intensities return -np.sum(np.log(likelihoods)) def _optimized_unbinned_nll( - bare_intensities: np.ndarray, + data_intensities: np.ndarray, phsp_intensities: np.ndarray, ) -> float: normalization_integral = np.mean(phsp_intensities) - return len(bare_intensities) * np.log(normalization_integral) - np.sum( - np.log(bare_intensities) + return len(data_intensities) * np.log(normalization_integral) - np.sum( + np.log(data_intensities) ) @@ -40,9 +40,9 @@ def _optimized_unbinned_nll( @pytest.fixture(scope="module") def intensities() -> tuple[np.ndarray, np.ndarray]: rng = np.random.default_rng(seed=0) - bare_intensities = rng.uniform(low=0.1, high=10.0, size=5_000_000) + data_intensities = rng.uniform(low=0.1, high=10.0, size=5_000_000) phsp_intensities = rng.uniform(low=0.1, high=10.0, size=1_000_000) - return bare_intensities, phsp_intensities + return data_intensities, phsp_intensities @pytest.mark.benchmark(group="unbinned-nll-normalization") @@ -52,14 +52,14 @@ def test_unbinned_nll_normalization_formula( implementation: str, intensities: tuple[np.ndarray, np.ndarray], ) -> None: - bare_intensities, phsp_intensities = intensities + data_intensities, phsp_intensities = intensities reference = _original_unbinned_nll( - bare_intensities, + data_intensities, phsp_intensities, ) result = benchmark( _IMPLEMENTATIONS[implementation], - bare_intensities, + data_intensities, phsp_intensities, ) assert result == pytest.approx(reference) diff --git a/src/tensorwaves/estimator.py b/src/tensorwaves/estimator.py index 60d61b2e..3b07ac1c 100644 --- a/src/tensorwaves/estimator.py +++ b/src/tensorwaves/estimator.py @@ -204,14 +204,14 @@ def __init__( def __call__(self, parameters: Mapping[str, ParameterValue]) -> float: self.__function.update_parameters(parameters) - bare_intensities = self.__function(self.__data) + data_intensities = self.__function(self.__data) phsp_intensities = self.__function(self.__phsp) if self.__phsp_weights is not None: phsp_intensities *= self.__phsp_weights normalization_factor = 1.0 / ( self.__phsp_volume * self.__mean_function(phsp_intensities) ) - likelihoods = normalization_factor * bare_intensities + likelihoods = normalization_factor * data_intensities return -self.__sum_function(self.__log_function(likelihoods)) def gradient( From adce51e89042e4f46573ebc5525d075e8c8abbbb Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:16:06 +0200 Subject: [PATCH 03/11] DX: allow passing paths to `poe benchmark` --- pyproject.toml | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index c1168fb4..32d797e0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -267,7 +267,7 @@ heading = "Testing" [tool.poe.groups.test.tasks.benchmark] cmd = """ -pytest benchmarks \ +pytest ${paths} \ --benchmark-autosave \ --benchmark-json benchmarks/output.json \ --durations=0 \ @@ -276,6 +276,12 @@ pytest benchmarks \ executor = { extra = ["jax", "numba", "pwa"], group = "test" } help = "Run benchmark tests and visualize performance" +[[tool.poe.groups.test.tasks.benchmark.args]] +default = "" +multiple = true +name = "paths" +positional = true + [tool.poe.groups.test.tasks.cov] cmd = """ pytest \ From 63aec5f9ed75851ab853da9eee9de593062f26b6 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:16:07 +0200 Subject: [PATCH 04/11] DX: benchmark JAX arrays --- benchmarks/unbinned_nll.py | 99 ++++++++++++++++++++++++++++++++++---- 1 file changed, 90 insertions(+), 9 deletions(-) diff --git a/benchmarks/unbinned_nll.py b/benchmarks/unbinned_nll.py index f6c4ef6e..c2fa7df3 100644 --- a/benchmarks/unbinned_nll.py +++ b/benchmarks/unbinned_nll.py @@ -2,6 +2,8 @@ from typing import TYPE_CHECKING +import jax +import jax.numpy as jnp import numpy as np import pytest @@ -28,6 +30,35 @@ def _optimized_unbinned_nll( ) +def _create_jax_implementations() -> dict[ + str, + Callable[[jax.Array, jax.Array], jax.Array], +]: + @jax.jit + def original_unbinned_nll( + data_intensities: jax.Array, + phsp_intensities: jax.Array, + ) -> jax.Array: + normalization_factor = 1.0 / jnp.mean(phsp_intensities) + likelihoods = normalization_factor * data_intensities + return -jnp.sum(jnp.log(likelihoods)) + + @jax.jit + def optimized_unbinned_nll( + data_intensities: jax.Array, + phsp_intensities: jax.Array, + ) -> jax.Array: + normalization_integral = jnp.mean(phsp_intensities) + return len(data_intensities) * jnp.log(normalization_integral) - jnp.sum( + jnp.log(data_intensities) + ) + + return { + "original": original_unbinned_nll, + "optimized": optimized_unbinned_nll, + } + + _IMPLEMENTATIONS: dict[ str, Callable[[np.ndarray, np.ndarray], float], @@ -45,21 +76,71 @@ def intensities() -> tuple[np.ndarray, np.ndarray]: return data_intensities, phsp_intensities +@pytest.fixture(scope="module") +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() + return jax_data_intensities, jax_phsp_intensities + + +def _benchmark_numpy_implementation( + benchmark: Callable[[Callable[[], float]], float], + implementation: str, + intensities: tuple[np.ndarray, np.ndarray], +) -> float: + data_intensities, phsp_intensities = intensities + function = _IMPLEMENTATIONS[implementation] + + def run() -> float: + return function(data_intensities, phsp_intensities) + + return benchmark(run) + + +def _benchmark_jax_implementation( + benchmark: Callable[[Callable[[], jax.Array]], jax.Array], + implementation: str, + intensities: tuple[jax.Array, jax.Array], +) -> jax.Array: + data_intensities, phsp_intensities = intensities + function = _create_jax_implementations()[implementation] + function(data_intensities, phsp_intensities).block_until_ready() + + def run() -> jax.Array: + return function(data_intensities, phsp_intensities).block_until_ready() + + return benchmark(run) + + @pytest.mark.benchmark(group="unbinned-nll-normalization") +@pytest.mark.parametrize("backend", ["numpy", "jax"]) @pytest.mark.parametrize("implementation", _IMPLEMENTATIONS) def test_unbinned_nll_normalization_formula( benchmark, + backend: str, implementation: str, intensities: tuple[np.ndarray, np.ndarray], + request: pytest.FixtureRequest, ) -> None: - data_intensities, phsp_intensities = intensities reference = _original_unbinned_nll( - data_intensities, - phsp_intensities, - ) - result = benchmark( - _IMPLEMENTATIONS[implementation], - data_intensities, - phsp_intensities, + *intensities, ) - assert result == pytest.approx(reference) + if backend == "jax": + result = _benchmark_jax_implementation( + benchmark, + implementation, + request.getfixturevalue("jax_intensities"), + ) + else: + result = _benchmark_numpy_implementation( + benchmark, + implementation, + intensities, + ) + + assert float(np.asarray(result)) == pytest.approx(reference) From 0b3d032c9ed14aa0fade467f4ace0f466a05a3ab Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:16:07 +0200 Subject: [PATCH 05/11] DX: benchmark tensorflow --- benchmarks/unbinned_nll.py | 63 +++++++++++++++++++++++++++++++++++++- 1 file changed, 62 insertions(+), 1 deletion(-) diff --git a/benchmarks/unbinned_nll.py b/benchmarks/unbinned_nll.py index c2fa7df3..82ff98f5 100644 --- a/benchmarks/unbinned_nll.py +++ b/benchmarks/unbinned_nll.py @@ -6,6 +6,8 @@ import jax.numpy as jnp import numpy as np import pytest +import tensorflow as tf +import tensorflow.experimental.numpy as tnp # ty: ignore[unresolved-import] if TYPE_CHECKING: from collections.abc import Callable @@ -59,6 +61,36 @@ def optimized_unbinned_nll( } +def _create_tensorflow_implementations() -> dict[ + str, + Callable[[tf.Tensor, tf.Tensor], tf.Tensor], +]: + @tf.function + def original_unbinned_nll( + data_intensities: tf.Tensor, + phsp_intensities: tf.Tensor, + ) -> tf.Tensor: + normalization_factor = 1.0 / tnp.mean(phsp_intensities) + likelihoods = normalization_factor * data_intensities + return -tnp.sum(tnp.log(likelihoods)) + + @tf.function + def optimized_unbinned_nll( + data_intensities: tf.Tensor, + phsp_intensities: tf.Tensor, + ) -> tf.Tensor: + normalization_integral = tnp.mean(phsp_intensities) + n_events = data_intensities.shape[0] + return n_events * tnp.log(normalization_integral) - tnp.sum( + tnp.log(data_intensities) + ) + + return { + "original": original_unbinned_nll, + "optimized": optimized_unbinned_nll, + } + + _IMPLEMENTATIONS: dict[ str, Callable[[np.ndarray, np.ndarray], float], @@ -88,6 +120,14 @@ def jax_intensities( return jax_data_intensities, jax_phsp_intensities +@pytest.fixture(scope="module") +def tensorflow_intensities( + intensities: tuple[np.ndarray, np.ndarray], +) -> tuple[tf.Tensor, tf.Tensor]: + data_intensities, phsp_intensities = intensities + return tnp.asarray(data_intensities), tnp.asarray(phsp_intensities) + + def _benchmark_numpy_implementation( benchmark: Callable[[Callable[[], float]], float], implementation: str, @@ -117,8 +157,23 @@ def run() -> jax.Array: return benchmark(run) +def _benchmark_tensorflow_implementation( + benchmark: Callable[[Callable[[], np.ndarray]], np.ndarray], + implementation: str, + intensities: tuple[tf.Tensor, tf.Tensor], +) -> np.ndarray: + data_intensities, phsp_intensities = intensities + function = _create_tensorflow_implementations()[implementation] + function(data_intensities, phsp_intensities).numpy() + + def run() -> np.ndarray: + return function(data_intensities, phsp_intensities).numpy() + + return benchmark(run) + + @pytest.mark.benchmark(group="unbinned-nll-normalization") -@pytest.mark.parametrize("backend", ["numpy", "jax"]) +@pytest.mark.parametrize("backend", ["numpy", "jax", "tensorflow"]) @pytest.mark.parametrize("implementation", _IMPLEMENTATIONS) def test_unbinned_nll_normalization_formula( benchmark, @@ -136,6 +191,12 @@ def test_unbinned_nll_normalization_formula( implementation, request.getfixturevalue("jax_intensities"), ) + elif backend == "tensorflow": + result = _benchmark_tensorflow_implementation( + benchmark, + implementation, + request.getfixturevalue("tensorflow_intensities"), + ) else: result = _benchmark_numpy_implementation( benchmark, From db05c4491f447994fb66dc9d5ee4b61dbd6591ea Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:16:08 +0200 Subject: [PATCH 06/11] DX: test Numba with loops --- .cspell.json | 2 + benchmarks/unbinned_nll.py | 75 +++++++++++++++++++++++++++++++++++++- 2 files changed, 75 insertions(+), 2 deletions(-) diff --git a/.cspell.json b/.cspell.json index a6cb1f53..c180ced9 100644 --- a/.cspell.json +++ b/.cspell.json @@ -124,6 +124,7 @@ "ncalls", "ncols", "ndarray", + "njit", "noqa", "noreply", "nrows", @@ -133,6 +134,7 @@ "pcolormesh", "phasespace", "phsp", + "prange", "precommit", "prefactor", "preorder", diff --git a/benchmarks/unbinned_nll.py b/benchmarks/unbinned_nll.py index 82ff98f5..791fb6d9 100644 --- a/benchmarks/unbinned_nll.py +++ b/benchmarks/unbinned_nll.py @@ -1,9 +1,11 @@ from __future__ import annotations -from typing import TYPE_CHECKING +import math +from typing import TYPE_CHECKING, cast import jax import jax.numpy as jnp +import numba import numpy as np import pytest import tensorflow as tf @@ -12,6 +14,11 @@ if TYPE_CHECKING: from collections.abc import Callable + def prange(stop: int) -> range: ... + +else: + from numba import prange + def _original_unbinned_nll( data_intensities: np.ndarray, @@ -61,6 +68,49 @@ def optimized_unbinned_nll( } +def _create_numba_implementations() -> dict[ + str, + Callable[[np.ndarray, np.ndarray], float], +]: + @numba.njit(parallel=True) + def original_unbinned_nll( + data_intensities: np.ndarray, + phsp_intensities: np.ndarray, + ) -> float: + phsp_sum = 0.0 + for i in prange(len(phsp_intensities)): + phsp_sum += phsp_intensities[i] + normalization_factor = len(phsp_intensities) / phsp_sum + + log_likelihood = 0.0 + for i in prange(len(data_intensities)): + log_likelihood += math.log(normalization_factor * data_intensities[i]) + return -log_likelihood + + @numba.njit(parallel=True) + def optimized_unbinned_nll( + data_intensities: np.ndarray, + phsp_intensities: np.ndarray, + ) -> float: + phsp_sum = 0.0 + for i in prange(len(phsp_intensities)): + phsp_sum += phsp_intensities[i] + normalization_integral = phsp_sum / len(phsp_intensities) + + log_sum = 0.0 + for i in prange(len(data_intensities)): + log_sum += math.log(data_intensities[i]) + return len(data_intensities) * math.log(normalization_integral) - log_sum + + return cast( + "dict[str, Callable[[np.ndarray, np.ndarray], float]]", + { + "original": original_unbinned_nll, + "optimized": optimized_unbinned_nll, + }, + ) + + def _create_tensorflow_implementations() -> dict[ str, Callable[[tf.Tensor, tf.Tensor], tf.Tensor], @@ -142,6 +192,21 @@ def run() -> float: return benchmark(run) +def _benchmark_numba_implementation( + benchmark: Callable[[Callable[[], float]], float], + implementation: str, + intensities: tuple[np.ndarray, np.ndarray], +) -> float: + data_intensities, phsp_intensities = intensities + function = _create_numba_implementations()[implementation] + function(data_intensities, phsp_intensities) + + def run() -> float: + return function(data_intensities, phsp_intensities) + + return benchmark(run) + + def _benchmark_jax_implementation( benchmark: Callable[[Callable[[], jax.Array]], jax.Array], implementation: str, @@ -173,7 +238,7 @@ def run() -> np.ndarray: @pytest.mark.benchmark(group="unbinned-nll-normalization") -@pytest.mark.parametrize("backend", ["numpy", "jax", "tensorflow"]) +@pytest.mark.parametrize("backend", ["numpy", "numba", "jax", "tensorflow"]) @pytest.mark.parametrize("implementation", _IMPLEMENTATIONS) def test_unbinned_nll_normalization_formula( benchmark, @@ -191,6 +256,12 @@ def test_unbinned_nll_normalization_formula( implementation, request.getfixturevalue("jax_intensities"), ) + elif backend == "numba": + result = _benchmark_numba_implementation( + benchmark, + implementation, + intensities, + ) elif backend == "tensorflow": result = _benchmark_tensorflow_implementation( benchmark, From 38ea591ee0ae76955237b6a3935b2bc4e7a1bbf6 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:16:09 +0200 Subject: [PATCH 07/11] MAINT: remove `_function` from back-end agnostic attributes --- src/tensorwaves/estimator.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/src/tensorwaves/estimator.py b/src/tensorwaves/estimator.py index 3b07ac1c..33114aef 100644 --- a/src/tensorwaves/estimator.py +++ b/src/tensorwaves/estimator.py @@ -196,9 +196,9 @@ def __init__( self.__function = function self.__gradient = gradient_creator(self.__call__, backend) - self.__mean_function = find_function("mean", backend) - self.__sum_function = find_function("sum", backend) - self.__log_function = find_function("log", backend) + self.__mean = find_function("mean", backend) + self.__sum = find_function("sum", backend) + self.__log = find_function("log", backend) self.__phsp_volume = phsp_volume @@ -209,10 +209,10 @@ def __call__(self, parameters: Mapping[str, ParameterValue]) -> float: if self.__phsp_weights is not None: phsp_intensities *= self.__phsp_weights normalization_factor = 1.0 / ( - self.__phsp_volume * self.__mean_function(phsp_intensities) + self.__phsp_volume * self.__mean(phsp_intensities) ) likelihoods = normalization_factor * data_intensities - return -self.__sum_function(self.__log_function(likelihoods)) + return -self.__sum(self.__log(likelihoods)) def gradient( self, parameters: Mapping[str, ParameterValue] From 5af4e9e60c0c7a6e5375561fb0168fbde7ac2857 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:16:10 +0200 Subject: [PATCH 08/11] DX: enforce positive intensities in benchmark --- benchmarks/expression.py | 16 +++++++++------- 1 file changed, 9 insertions(+), 7 deletions(-) diff --git a/benchmarks/expression.py b/benchmarks/expression.py index cb28f8a9..796b9c9a 100644 --- a/benchmarks/expression.py +++ b/benchmarks/expression.py @@ -26,14 +26,16 @@ def poisson(x: sp.Symbol, k) -> sp.Expr: symbols = sp.symbols("x y (a:c) mu_(:2) sigma_(:2) omega") x, y, a, b, c, mu1, mu2, sigma1, sigma2, omega = symbols expression = ( - a * gaussian(x, mu1, sigma1) + b * gaussian(x, mu2, sigma2) + c * poisson(x, k=2) + a**2 * gaussian(x, mu1, sigma1) + + b**2 * gaussian(x, mu2, sigma2) + + c**2 * poisson(x, k=2) ) * sp.cos(y * omega) ** 2 domain_boundaries = {"x": (0, 5), "y": (-np.pi, +np.pi)} parameter_defaults = { - a: 0.15, - b: 0.05, - c: 0.3, + a: np.sqrt(0.15), + b: np.sqrt(0.05), + c: np.sqrt(0.3), mu1: 1.0, mu2: 2.7, omega: 0.5, @@ -41,9 +43,9 @@ def poisson(x: sp.Symbol, k) -> sp.Expr: sigma2: 0.5, } initial_parameters = { - "a": 0.2, - "b": 0.1, - "c": 0.2, + "a": np.sqrt(0.2), + "b": np.sqrt(0.1), + "c": np.sqrt(0.2), "mu_0": 0.9, "sigma_0": 0.4, "sigma_1": 0.4, From 42e8e03e5810bacd10e89e67da34a1dce49e2530 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:16:11 +0200 Subject: [PATCH 09/11] DX: benchmark `UnbinnedNLL` implementation --- benchmarks/unbinned_nll.py | 173 +++++++++++++++++++++++++++++++++++++ 1 file changed, 173 insertions(+) diff --git a/benchmarks/unbinned_nll.py b/benchmarks/unbinned_nll.py index 791fb6d9..295df212 100644 --- a/benchmarks/unbinned_nll.py +++ b/benchmarks/unbinned_nll.py @@ -11,6 +11,9 @@ import tensorflow as tf import tensorflow.experimental.numpy as tnp # ty: ignore[unresolved-import] +from tensorwaves.estimator import UnbinnedNLL +from tensorwaves.function import ParametrizedBackendFunction + if TYPE_CHECKING: from collections.abc import Callable @@ -150,6 +153,36 @@ def optimized_unbinned_nll( } +def _numpy_intensity(x: np.ndarray, center: float) -> np.ndarray: + return 1.0 + (x - center) ** 2 + + +@numba.njit(parallel=True) +def _numba_intensity(x: np.ndarray, center: float) -> np.ndarray: + intensities = np.empty_like(x) + for i in prange(len(x)): + intensities[i] = 1.0 + (x[i] - center) ** 2 + return intensities + + +@jax.jit +def _jax_intensity(x: jax.Array, center: float) -> jax.Array: + return 1.0 + (x - center) ** 2 + + +@tf.function +def _tensorflow_intensity(x: tf.Tensor, center: float) -> tf.Tensor: + return 1.0 + (x - center) ** 2 + + +_ESTIMATOR_FUNCTIONS = { + "numpy": _numpy_intensity, + "numba": _numba_intensity, + "jax": _jax_intensity, + "tensorflow": _tensorflow_intensity, +} + + @pytest.fixture(scope="module") def intensities() -> tuple[np.ndarray, np.ndarray]: rng = np.random.default_rng(seed=0) @@ -158,6 +191,14 @@ def intensities() -> tuple[np.ndarray, np.ndarray]: return data_intensities, phsp_intensities +@pytest.fixture(scope="module") +def estimator_samples() -> tuple[dict[str, np.ndarray], dict[str, np.ndarray]]: + rng = np.random.default_rng(seed=0) + data = {"x": rng.uniform(low=-2.0, high=2.0, size=1_000_000)} + phsp = {"x": rng.uniform(low=-2.0, high=2.0, size=1_000_000)} + return data, phsp + + @pytest.fixture(scope="module") def jax_intensities( intensities: tuple[np.ndarray, np.ndarray], @@ -178,6 +219,27 @@ def tensorflow_intensities( return tnp.asarray(data_intensities), tnp.asarray(phsp_intensities) +@pytest.fixture(scope="module") +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()}, + {"x": jnp.asarray(phsp["x"]).block_until_ready()}, + ) + + +@pytest.fixture(scope="module") +def tensorflow_estimator_samples( + estimator_samples: tuple[dict[str, np.ndarray], dict[str, np.ndarray]], +) -> tuple[dict[str, tf.Tensor], dict[str, tf.Tensor]]: + data, phsp = estimator_samples + return {"x": tnp.asarray(data["x"])}, {"x": tnp.asarray(phsp["x"])} + + def _benchmark_numpy_implementation( benchmark: Callable[[Callable[[], float]], float], implementation: str, @@ -192,6 +254,75 @@ def run() -> float: return benchmark(run) +def _create_estimator( + backend: str, + data: dict, + phsp: dict, +) -> UnbinnedNLL: + function = ParametrizedBackendFunction( + function=_ESTIMATOR_FUNCTIONS[backend], + argument_order=("x", "center"), + parameters={"center": 0.0}, + ) + return UnbinnedNLL(function, data, phsp, backend=backend) + + +def _compute_estimator_reference( + data: dict[str, np.ndarray], + phsp: dict[str, np.ndarray], + center: float, +) -> float: + data_intensities = _numpy_intensity(data["x"], center) + phsp_intensities = _numpy_intensity(phsp["x"], center) + return _original_unbinned_nll(data_intensities, phsp_intensities) + + +def _benchmark_estimator_numpy( + benchmark: Callable[[Callable[[], float]], float], + backend: str, + data: dict[str, np.ndarray], + phsp: dict[str, np.ndarray], + parameters: dict[str, float], +) -> float: + estimator = _create_estimator(backend, data, phsp) + estimator(parameters) + + def run() -> float: + return estimator(parameters) + + return benchmark(run) + + +def _benchmark_estimator_jax( + benchmark: Callable[[Callable[[], jax.Array]], jax.Array], + data: dict[str, jax.Array], + phsp: dict[str, jax.Array], + parameters: dict[str, float], +) -> jax.Array: + estimator = _create_estimator("jax", data, phsp) + estimator(parameters).block_until_ready() # ty: ignore[unresolved-attribute] + + def run() -> jax.Array: + return estimator(parameters).block_until_ready() # ty: ignore[unresolved-attribute] + + return benchmark(run) + + +def _benchmark_estimator_tensorflow( + benchmark: Callable[[Callable[[], np.ndarray]], np.ndarray], + data: dict[str, tf.Tensor], + phsp: dict[str, tf.Tensor], + parameters: dict[str, float], +) -> np.ndarray: + estimator = _create_estimator("tensorflow", data, phsp) + estimator(parameters).numpy() # ty: ignore[unresolved-attribute] + + def run() -> np.ndarray: + return estimator(parameters).numpy() # ty: ignore[unresolved-attribute] + + return benchmark(run) + + def _benchmark_numba_implementation( benchmark: Callable[[Callable[[], float]], float], implementation: str, @@ -276,3 +407,45 @@ def test_unbinned_nll_normalization_formula( ) assert float(np.asarray(result)) == pytest.approx(reference) + + +@pytest.mark.benchmark(group="unbinned-nll-estimator") +@pytest.mark.parametrize("backend", ["numpy", "numba", "jax", "tensorflow"]) +def test_unbinned_nll_estimator( + benchmark, + backend: str, + estimator_samples: tuple[dict[str, np.ndarray], dict[str, np.ndarray]], + request: pytest.FixtureRequest, +) -> None: + data, phsp = estimator_samples + parameters = {"center": 0.3} + reference = _compute_estimator_reference(data, phsp, parameters["center"]) + + if backend == "jax": + jax_data, jax_phsp = request.getfixturevalue("jax_estimator_samples") + result = _benchmark_estimator_jax( + benchmark, + jax_data, + jax_phsp, + parameters, + ) + elif backend == "tensorflow": + tensorflow_data, tensorflow_phsp = request.getfixturevalue( + "tensorflow_estimator_samples" + ) + result = _benchmark_estimator_tensorflow( + benchmark, + tensorflow_data, + tensorflow_phsp, + parameters, + ) + else: + result = _benchmark_estimator_numpy( + benchmark, + backend, + data, + phsp, + parameters, + ) + + assert float(np.asarray(result)) == pytest.approx(reference) From ddf6d6391da18c08370a2197814d484a9444a5b4 Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:16:11 +0200 Subject: [PATCH 10/11] ENH: speed up `UnbinnedNLL` implementation for NumPy --- src/tensorwaves/estimator.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/src/tensorwaves/estimator.py b/src/tensorwaves/estimator.py index 33114aef..e58113f6 100644 --- a/src/tensorwaves/estimator.py +++ b/src/tensorwaves/estimator.py @@ -208,10 +208,8 @@ def __call__(self, parameters: Mapping[str, ParameterValue]) -> float: phsp_intensities = self.__function(self.__phsp) if self.__phsp_weights is not None: phsp_intensities *= self.__phsp_weights - normalization_factor = 1.0 / ( - self.__phsp_volume * self.__mean(phsp_intensities) - ) - likelihoods = normalization_factor * data_intensities + normalization_integral = self.__phsp_volume * self.__mean(phsp_intensities) + likelihoods = data_intensities / normalization_integral return -self.__sum(self.__log(likelihoods)) def gradient( From 983ecd0932e827eb5813027f4f2d894f66b2400f Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:44:53 +0200 Subject: [PATCH 11/11] BEHAVIOR: use subtraction in `UnbinnedNLL` --- src/tensorwaves/estimator.py | 4 +-- tests/optimizer/test_fit_simple_model.py | 44 ++++++++++++------------ 2 files changed, 24 insertions(+), 24 deletions(-) diff --git a/src/tensorwaves/estimator.py b/src/tensorwaves/estimator.py index e58113f6..7f23393f 100644 --- a/src/tensorwaves/estimator.py +++ b/src/tensorwaves/estimator.py @@ -209,8 +209,8 @@ def __call__(self, parameters: Mapping[str, ParameterValue]) -> float: if self.__phsp_weights is not None: phsp_intensities *= self.__phsp_weights normalization_integral = self.__phsp_volume * self.__mean(phsp_intensities) - likelihoods = data_intensities / normalization_integral - return -self.__sum(self.__log(likelihoods)) + log_normalization = len(data_intensities) * self.__log(normalization_integral) + return log_normalization - self.__sum(self.__log(data_intensities)) def gradient( self, parameters: Mapping[str, ParameterValue] diff --git a/tests/optimizer/test_fit_simple_model.py b/tests/optimizer/test_fit_simple_model.py index 07c590b9..97c14576 100644 --- a/tests/optimizer/test_fit_simple_model.py +++ b/tests/optimizer/test_fit_simple_model.py @@ -74,14 +74,14 @@ def expression_and_parameters() -> tuple[sp.Expr, dict[sp.Symbol, float]]: symbols: tuple[sp.Symbol, ...] = sp.symbols("x y (a:c) mu_(:2) sigma_(:2) omega") x, y, a, b, c, mu1, mu2, sigma1, sigma2, omega = symbols expression = ( - a * gaussian(x, mu1, sigma1) - + b * gaussian(x, mu2, sigma2) - + c * poisson(x, k=2) + a**2 * gaussian(x, mu1, sigma1) + + b**2 * gaussian(x, mu2, sigma2) + + c**2 * poisson(x, k=2) ) * sp.cos(y * omega) ** 2 parameter_defaults = { - a: 0.15, - b: 0.05, - c: 0.3, + a: np.sqrt(0.15), + b: np.sqrt(0.05), + c: np.sqrt(0.3), mu1: 1.0, mu2: 2.7, omega: 0.5, @@ -173,27 +173,27 @@ def test_optimize_all_parameters( ( 0.1, # iminuit default tolerance { - "a": 0.15679884056468815, - "b": 0.051281396032855225, - "c": 0.26265501744837677, - "mu_0": 0.9871104323476636, - "mu_1": 2.6947038781339754, - "omega": 0.4982768824682492, - "sigma_0": 0.3075629925771585, - "sigma_1": 0.5768191611084318, + "a": 0.3970518449186512, + "b": 0.22706305032403012, + "c": 0.5138637509750306, + "mu_0": 0.9871165126781332, + "mu_1": 2.694601996919253, + "omega": 0.49827248627364273, + "sigma_0": 0.3075724216579571, + "sigma_1": 0.5770019226843762, }, ), ( 2.0, { - "a": 0.15676480837709061, - "b": 0.051242383278715484, - "c": 0.26300711305648883, - "mu_0": 0.9870594159578658, - "mu_1": 2.694891339245927, - "omega": 0.4982886866357587, - "sigma_0": 0.30740850647117296, - "sigma_1": 0.5760794793407019, + "a": 0.39702016316833466, + "b": 0.2269834228175954, + "c": 0.5140806019278283, + "mu_0": 0.9869234898869502, + "mu_1": 2.6947811738260827, + "omega": 0.49830861095377976, + "sigma_0": 0.3073670759992788, + "sigma_1": 0.5765008475572346, }, ), ],