Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .cspell.json
Original file line number Diff line number Diff line change
Expand Up @@ -242,6 +242,7 @@
"vectorize",
"venv",
"vmap",
"vmapped",
"weisskopf",
"wirtinger",
"xcode",
Expand Down
6 changes: 3 additions & 3 deletions benchmarks/unbinned_nll.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
50 changes: 32 additions & 18 deletions src/tensorwaves/estimator.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
DataSample,
DataTransformer,
Estimator,
ParameterType,
ParameterValue,
ParametrizedFunction,
)
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -220,27 +230,27 @@ 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(
self, parameters: Mapping[str, ParameterValue]
) -> 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,
Expand Down Expand Up @@ -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,
Expand All @@ -315,24 +325,28 @@ 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(
self, parameters: Mapping[str, ParameterValue]
) -> 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,
Expand Down
20 changes: 10 additions & 10 deletions src/tensorwaves/function/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from tensorwaves.interface import (
DataSample,
Function,
ParameterType,
ParameterValue,
ParametrizedFunction,
)
Expand Down Expand Up @@ -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]:
Expand All @@ -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 "
Expand All @@ -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:
Expand Down
27 changes: 21 additions & 6 deletions src/tensorwaves/interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
<https://numpy.org/doc/stable/user/basics.broadcasting.html>`_ 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]):
Expand All @@ -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
Expand All @@ -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.
Expand All @@ -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(
Expand Down
2 changes: 1 addition & 1 deletion src/tensorwaves/optimizer/minuit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down
4 changes: 2 additions & 2 deletions src/tensorwaves/optimizer/scipy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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,
),
Expand Down
6 changes: 3 additions & 3 deletions tests/optimizer/test_minuit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)

Expand All @@ -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)
Expand Down
6 changes: 3 additions & 3 deletions tests/optimizer/test_scipy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)

Expand All @@ -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)
Expand Down
Loading
Loading