diff --git a/.cspell.json b/.cspell.json index 3dd99ec6..5b971deb 100644 --- a/.cspell.json +++ b/.cspell.json @@ -242,6 +242,7 @@ "vectorize", "venv", "vmap", + "vmapped", "weisskopf", "wirtinger", "xcode", diff --git a/benchmarks/unbinned_nll.py b/benchmarks/unbinned_nll.py index c5ed4906..2b632222 100644 --- a/benchmarks/unbinned_nll.py +++ b/benchmarks/unbinned_nll.py @@ -277,16 +277,16 @@ def _compute_estimator_reference( def _benchmark_estimator_numpy( - benchmark: Callable[[Callable[[], float]], float], + benchmark: Callable[[Callable[[], float | np.ndarray]], float | np.ndarray], backend: str, data: dict[str, np.ndarray], phsp: dict[str, np.ndarray], parameters: dict[str, float], -) -> float: +) -> float | np.ndarray: estimator = _create_estimator(backend, data, phsp) estimator(parameters) - def run() -> float: + def run() -> float | np.ndarray: return estimator(parameters) return benchmark(run) diff --git a/src/tensorwaves/estimator.py b/src/tensorwaves/estimator.py index 07470c3d..d1d08a76 100644 --- a/src/tensorwaves/estimator.py +++ b/src/tensorwaves/estimator.py @@ -15,6 +15,7 @@ DataSample, DataTransformer, Estimator, + ParameterType, ParameterValue, ParametrizedFunction, ) @@ -85,14 +86,23 @@ def _determine_backend(function: ParametrizedFunction, backend: str | None) -> s def _coerce_parameter_types( - parameters: Mapping[str, ParameterValue], -) -> dict[str, ParameterValue]: - # normalize to float/complex so that JIT compilers see stable input types - # (an int value would otherwise trigger a re-trace once it becomes a float) - return { - name: complex(value) if isinstance(value, complex) else float(value) - for name, value in parameters.items() - } + parameters: Mapping[str, ParameterType], +) -> dict[str, ParameterType]: + return {name: _coerce_parameter_value(value) for name, value in parameters.items()} + + +def _coerce_parameter_value(value: ParameterType) -> ParameterType: + # normalize scalars to float/complex so that JIT compilers see stable input + # types (an int value would otherwise trigger a re-trace once it becomes a + # float) and give parameter arrays a new trailing axis, so that they + # broadcast against the event axis of the data samples + if isinstance(value, complex): + return complex(value) + if isinstance(value, (int, float)): + return float(value) + if getattr(value, "ndim", 0) >= 1: + return value[..., None] + return value def _import_jax(): # ruff: ignore[missing-return-type-private-function] @@ -220,19 +230,19 @@ def __init__( sum_function = find_function("sum", backend) def estimator( - parameters: Mapping[str, ParameterValue], + parameters: Mapping[str, ParameterType], domain: DataSample, observed_values: np.ndarray, weights: np.ndarray, ) -> float: computed_values = function(domain, parameters) chi_squared = weights * (computed_values - observed_values) ** 2 - return sum_function(chi_squared) + return sum_function(chi_squared, axis=-1) self.__estimator = _jit_estimator_core(estimator, backend) self.__gradient = _create_core_gradient(estimator, backend) - def __call__(self, parameters: Mapping[str, ParameterValue]) -> float: + def __call__(self, parameters: Mapping[str, ParameterType]) -> float | np.ndarray: return self.__estimator(*self.__estimator_args(parameters)) def gradient( @@ -240,7 +250,7 @@ def gradient( ) -> dict[str, ParameterValue]: return self.__gradient(*self.__estimator_args(parameters)) - def __estimator_args(self, parameters: Mapping[str, ParameterValue]) -> tuple: + def __estimator_args(self, parameters: Mapping[str, ParameterType]) -> tuple: return ( _coerce_parameter_types(parameters), self.__domain, @@ -306,7 +316,7 @@ def __init__( log_function = find_function("log", backend) def estimator( - parameters: Mapping[str, ParameterValue], + parameters: Mapping[str, ParameterType], data: DataSample, phsp: DataSample, phsp_weights: np.ndarray | None, @@ -315,16 +325,20 @@ def estimator( phsp_intensities = function(phsp, parameters) if phsp_weights is not None: phsp_intensities *= phsp_weights - normalization_integral = phsp_volume * mean_function(phsp_intensities) - log_normalization = len(bare_intensities) * log_function( + normalization_integral = phsp_volume * mean_function( + phsp_intensities, axis=-1 + ) + log_normalization = bare_intensities.shape[-1] * log_function( normalization_integral ) - return log_normalization - sum_function(log_function(bare_intensities)) + return log_normalization - sum_function( + log_function(bare_intensities), axis=-1 + ) self.__estimator = _jit_estimator_core(estimator, backend) self.__gradient = _create_core_gradient(estimator, backend) - def __call__(self, parameters: Mapping[str, ParameterValue]) -> float: + def __call__(self, parameters: Mapping[str, ParameterType]) -> float | np.ndarray: return self.__estimator(*self.__estimator_args(parameters)) def gradient( @@ -332,7 +346,7 @@ def gradient( ) -> dict[str, ParameterValue]: return self.__gradient(*self.__estimator_args(parameters)) - def __estimator_args(self, parameters: Mapping[str, ParameterValue]) -> tuple: + def __estimator_args(self, parameters: Mapping[str, ParameterType]) -> tuple: return ( _coerce_parameter_types(parameters), self.__data, diff --git a/src/tensorwaves/function/__init__.py b/src/tensorwaves/function/__init__.py index e8f1ad9b..cd2ef0d8 100644 --- a/src/tensorwaves/function/__init__.py +++ b/src/tensorwaves/function/__init__.py @@ -12,6 +12,7 @@ from tensorwaves.interface import ( DataSample, Function, + ParameterType, ParameterValue, ParametrizedFunction, ) @@ -146,10 +147,13 @@ def __init__( def __call__( self, data: DataSample, - parameters: Mapping[str, ParameterValue] | None = None, + parameters: Mapping[str, ParameterType] | None = None, ) -> np.ndarray: - extended_data = {**data, **self.__merge_parameters(parameters)} - return self.__function(extended_data) # ty:ignore[invalid-argument-type] + extended_data: dict = {**data, **self.__parameters} + if parameters is not None: + self.__validate_parameters(parameters) + extended_data.update(parameters) + return self.__function(extended_data) @property def function(self) -> Callable[..., np.ndarray]: @@ -170,18 +174,15 @@ def parameters(self) -> dict[str, ParameterValue]: def with_parameters( self, parameters: Mapping[str, ParameterValue] ) -> ParametrizedBackendFunction: + self.__validate_parameters(parameters) return ParametrizedBackendFunction( function=self.function, argument_order=self.argument_order, - parameters=self.__merge_parameters(parameters), + parameters={**self.__parameters, **parameters}, backend=self.backend, ) - def __merge_parameters( - self, parameters: Mapping[str, ParameterValue] | None - ) -> dict[str, ParameterValue]: - if parameters is None: - return self.__parameters + def __validate_parameters(self, parameters: Mapping[str, ParameterType]) -> None: over_defined = set(parameters) - set(self.__parameters) if over_defined: sep = "\n " @@ -191,7 +192,6 @@ def __merge_parameters( f" Expecting one of:{sep}{parameter_listing}" ) raise ValueError(msg) - return {**self.__parameters, **parameters} def get_source_code(function: Function) -> str: diff --git a/src/tensorwaves/interface.py b/src/tensorwaves/interface.py index 85c5e740..9db00655 100644 --- a/src/tensorwaves/interface.py +++ b/src/tensorwaves/interface.py @@ -39,7 +39,16 @@ def __call__(self, data: InputType) -> OutputType: ... DataSample = dict[str, np.ndarray] """Mapping of variable names to a sequence of data points, used by `Function`.""" ParameterValue = complex | float -"""Allowed types for parameter values.""" +"""Allowed types for scalar parameter values.""" +ParameterType = ParameterValue | np.ndarray +"""Types for parameter values in an evaluation, including arrays of values. + +An array of parameter values represents several parameter points that are evaluated in +one call through `broadcasting +`_ against the event axis +of a `.DataSample`. This can be used to propagate fit uncertainties by evaluating over +e.g. bootstrapped parameter samples in a single, backend-parallelized call. +""" class ParametrizedFunction(Function[InputType, OutputType]): @@ -62,12 +71,13 @@ class ParametrizedFunction(Function[InputType, OutputType]): def __call__( self, data: InputType, - parameters: Mapping[str, ParameterValue] | None = None, + parameters: Mapping[str, ParameterType] | None = None, ) -> OutputType: """Evaluate the function over :code:`data` for these parameter values. Given parameter values are merged with the defaults in :attr:`parameters` for - this evaluation only. + this evaluation only. Parameter values may be arrays, as long as they + broadcast against the event arrays in :code:`data` (see `.ParameterType`). """ @property @@ -90,7 +100,7 @@ class DataTransformer(Function[DataSample, DataSample]): """ -class Estimator(Function[Mapping[str, ParameterValue], float]): +class Estimator(Function[Mapping[str, ParameterType], float | np.ndarray]): """Estimator for discrepancy model and data. See the :mod:`.estimator` module for different implementations of this interface. @@ -99,8 +109,13 @@ class Estimator(Function[Mapping[str, ParameterValue], float]): """ @abstractmethod - def __call__(self, parameters: Mapping[str, ParameterValue]) -> float: # ty:ignore[invalid-method-override] - """Compute estimator value for this combination of parameter values.""" + def __call__(self, parameters: Mapping[str, ParameterType]) -> float | np.ndarray: # ty:ignore[invalid-method-override] + """Compute estimator value for this combination of parameter values. + + Parameter values may be one-dimensional arrays of shape :code:`(p,)`, in which + case the estimator returns an array of :code:`p` estimator values, one for + each parameter point (see `.ParameterType`). + """ @abstractmethod def gradient( diff --git a/src/tensorwaves/optimizer/minuit.py b/src/tensorwaves/optimizer/minuit.py index 9bb01cd8..41a27c33 100644 --- a/src/tensorwaves/optimizer/minuit.py +++ b/src/tensorwaves/optimizer/minuit.py @@ -70,7 +70,7 @@ def optimize( logs=_create_log( optimizer=type(self), estimator_type=type(estimator), - estimator_value=estimator(parameters), + estimator_value=float(estimator(parameters)), function_call=n_function_calls, parameters=parameters, ) diff --git a/src/tensorwaves/optimizer/scipy.py b/src/tensorwaves/optimizer/scipy.py index 915ea2b2..a546f536 100644 --- a/src/tensorwaves/optimizer/scipy.py +++ b/src/tensorwaves/optimizer/scipy.py @@ -62,7 +62,7 @@ def optimize( # ruff:ignore[complex-structure] logs=_create_log( optimizer=type(self), estimator_type=type(estimator), - estimator_value=estimator(parameters), + estimator_value=float(estimator(parameters)), function_call=n_function_calls, parameters=parameters, ) @@ -91,7 +91,7 @@ def wrapped_function(pars: list) -> float: logs=_create_log( optimizer=type(self), estimator_type=type(estimator), - estimator_value=estimator(parameters), + estimator_value=float(estimator(parameters)), function_call=n_function_calls, parameters=parameters, ), diff --git a/tests/optimizer/test_minuit.py b/tests/optimizer/test_minuit.py index c6413e5f..40f2b98d 100644 --- a/tests/optimizer/test_minuit.py +++ b/tests/optimizer/test_minuit.py @@ -4,7 +4,7 @@ import pytest -from tensorwaves.interface import Estimator, ParameterValue +from tensorwaves.interface import Estimator, ParameterType, ParameterValue from tensorwaves.optimizer.minuit import Minuit2 from . import CallbackMock, assert_invocations @@ -19,7 +19,7 @@ class Polynomial1DMinimaEstimator(Estimator): def __init__(self, polynomial: Callable) -> None: self.__polynomial = polynomial - def __call__(self, parameters: Mapping[str, ParameterValue]) -> float: + def __call__(self, parameters: Mapping[str, ParameterType]) -> float: x = parameters["x"] return self.__polynomial(x) @@ -33,7 +33,7 @@ class Polynomial2DMinimaEstimator(Estimator): def __init__(self, polynomial: Callable) -> None: self.__polynomial = polynomial - def __call__(self, parameters: Mapping[str, ParameterValue]) -> float: + def __call__(self, parameters: Mapping[str, ParameterType]) -> float: x = parameters["x"] y = parameters["y"] return self.__polynomial(x, y) diff --git a/tests/optimizer/test_scipy.py b/tests/optimizer/test_scipy.py index 03f6ec88..12759d90 100644 --- a/tests/optimizer/test_scipy.py +++ b/tests/optimizer/test_scipy.py @@ -4,7 +4,7 @@ import pytest -from tensorwaves.interface import Estimator, ParameterValue +from tensorwaves.interface import Estimator, ParameterType, ParameterValue from tensorwaves.optimizer.scipy import ScipyMinimizer from . import CallbackMock, assert_invocations @@ -19,7 +19,7 @@ class Polynomial1DMinimaEstimator(Estimator): def __init__(self, polynomial: Callable) -> None: self.__polynomial = polynomial - def __call__(self, parameters: Mapping[str, ParameterValue]) -> float: + def __call__(self, parameters: Mapping[str, ParameterType]) -> float: x = parameters["x"] return self.__polynomial(x) @@ -33,7 +33,7 @@ class Polynomial2DMinimaEstimator(Estimator): def __init__(self, polynomial: Callable) -> None: self.__polynomial = polynomial - def __call__(self, parameters: Mapping[str, ParameterValue]) -> float: + def __call__(self, parameters: Mapping[str, ParameterType]) -> float: x = parameters["x"] y = parameters["y"] return self.__polynomial(x, y) diff --git a/tests/test_estimator.py b/tests/test_estimator.py index 1b44ddb9..cb25e073 100644 --- a/tests/test_estimator.py +++ b/tests/test_estimator.py @@ -42,6 +42,18 @@ def test_call(self, backend): ) assert estimator({"a": 0, "b": 2}) == 2.5 + def test_array_valued_parameters(self): + x_data = {"x": np.array([0.0, 1.0, 2.0])} + y_data = np.array([0.0, 1.0, 2.0]) + function = ParametrizedBackendFunction( + function=lambda a, b, x: a + b * x, + argument_order=("a", "b", "x"), + parameters={"a": 0, "b": 1}, + ) + estimator = ChiSquared(function, x_data, y_data) + b_values = np.array([1.0, 2.0]) + np.testing.assert_allclose(estimator({"b": b_values}), [0.0, 5.0]) + def test_jit_compiled_once(self): trace_count = 0 @@ -173,6 +185,43 @@ def test_create_cached_function(backend): np.testing.assert_allclose(intensities, cached_intensities) +@pytest.mark.parametrize("backend", ["jax", "numba", "numpy", "tf"]) +def test_unbinned_nll_with_array_valued_parameters(backend: str): + x, mu, sigma = sp.symbols("x mu sigma") + function = create_parametrized_function( + expression=sp.exp(-(((x - mu) / sigma) ** 2) / 2), + parameters={mu: 0.5, sigma: 0.1}, + backend=backend, + ) + rng = np.random.default_rng(seed=0) + data = {"x": rng.normal(0.5, 0.1, size=2_000)} + phsp = {"x": rng.uniform(-2.0, 5.0, size=5_000)} + estimator = UnbinnedNLL(function, data, phsp, phsp_volume=7.0) + mu_values = np.array([0.4, 0.5, 0.6]) + batched_output = np.asarray(estimator({"mu": mu_values})) + scalar_outputs = [float(estimator({"mu": value})) for value in mu_values] + assert batched_output.shape == mu_values.shape + np.testing.assert_allclose(batched_output, scalar_outputs, rtol=1e-8) + + +def test_unbinned_nll_batched_evaluation_equals_jax_vmap(): + jax = pytest.importorskip("jax") + x, mu, sigma = sp.symbols("x mu sigma") + function = create_parametrized_function( + expression=sp.exp(-(((x - mu) / sigma) ** 2) / 2), + parameters={mu: 0.5, sigma: 0.1}, + backend="jax", + ) + rng = np.random.default_rng(seed=0) + data = {"x": rng.normal(0.5, 0.1, size=2_000)} + phsp = {"x": rng.uniform(-2.0, 5.0, size=5_000)} + estimator = UnbinnedNLL(function, data, phsp, phsp_volume=7.0) + mu_values = np.array([0.4, 0.5, 0.6]) + batched_output = np.asarray(estimator({"mu": mu_values})) + vmapped_output = jax.vmap(lambda value: estimator({"mu": value}))(mu_values) + np.testing.assert_allclose(batched_output, np.asarray(vmapped_output), rtol=1e-8) + + NUMPY_RNG = np.random.default_rng(12345)