diff --git a/CHANGELOG.md b/CHANGELOG.md index a8ee4e7..755c23b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,12 @@ The format is based on [Keep a Changelog](https://keepachangelog.com/). ## [Unreleased] +## Changed +- DiscreteTune.__call__ will now always return a numpy.ndarray object, regardless of argument type + +## Fixed +- fixed bug where DiscreteTune did not respect order of identifiers when called with an array + ## [0.5.1] ### Added diff --git a/attune/_discrete_tune.py b/attune/_discrete_tune.py index 52ae447..74cb650 100644 --- a/attune/_discrete_tune.py +++ b/attune/_discrete_tune.py @@ -35,7 +35,7 @@ def __init__( def __repr__(self): return f"DiscreteTune({repr(self.ranges)}, {repr(self.default)})" - def __call__(self, ind_value, *, ind_units=None): + def __call__(self, ind_value, *, ind_units=None) -> np.ndarray: """Evaluate the DiscreteTune at specific independent value(s). Paramters @@ -54,20 +54,16 @@ def __call__(self, ind_value, *, ind_units=None): """ if ind_units is not None and self._ind_units is not None: ind_value = wt.units.convert(ind_value, ind_units, self._ind_units) - if isinstance(ind_value, np.ndarray): - out = np.full( - ind_value.shape, - self.default, - dtype=f"U{max([len(s) for s in self.ranges.keys()])}", - ) - for key, (imin, imax) in self.ranges.items(): - out[(ind_value >= imin) & (ind_value <= imax)] = key - return out - else: - for key, (imin, imax) in self.ranges.items(): - if imin <= ind_value <= imax: - return key - return self.default + ind_value = np.asarray(ind_value) + out = np.full( + ind_value.shape, + self.default, + dtype=f"U{max([len(s) for s in self.ranges.keys()])}", + ) + # work in reverse so that the first valid entry persists + for key, (imin, imax) in reversed(self.ranges.items()): + out[(ind_value >= imin) & (ind_value <= imax)] = key + return out def __eq__(self, other): return self.ranges == other.ranges and self.default == other.default diff --git a/attune/_tune.py b/attune/_tune.py index 6b4128e..4b2a5ad 100644 --- a/attune/_tune.py +++ b/attune/_tune.py @@ -55,7 +55,7 @@ def __repr__(self): return f"Tune({repr(self.independent)}, {repr(self.dependent)})" return f"Tune({repr(self.independent)}, {repr(self.dependent)}, dep_units={repr(self.dep_units)})" - def __call__(self, ind_value, *, ind_units=None, dep_units=None): + def __call__(self, ind_value, *, ind_units=None, dep_units=None) -> np.ndarray: if ind_units is not None and self._ind_units is not None: ind_value = wt.units.convert(ind_value, ind_units, self._ind_units) ret = self._interp(ind_value) diff --git a/docs/structure.rst b/docs/structure.rst index 3b43428..c4cdd7a 100644 --- a/docs/structure.rst +++ b/docs/structure.rst @@ -78,15 +78,15 @@ You can, however, place higher priority (earlier) ranges inside of other ranges .. code-block:: python dt = attune.DiscreteTune({"hi": (100, 200), "lo": (10, 20), "inner": (50, 60), "med": (20, 100)}, default="def") - dt(5) == "def" - dt(15) == "lo" - dt(20) == "lo" - dt(30) == "med" - dt(55) == "inner" - dt(70) == "med" - dt(100) == "hi" - dt(150) == "hi" - dt(500) == "def" + dt(5) == np.array("def") + dt(15) == np.array("lo") + dt(20) == np.array("lo") + dt(30) == np.array("med") + dt(55) == np.array("inner") + dt(70) == np.array("med") + dt(100) == np.array("hi") + dt(150) == np.array("hi") + dt(500) == np.array("def") diff --git a/tests/discrete_tune.py b/tests/discrete_tune.py index 042d3e0..1d6bff7 100644 --- a/tests/discrete_tune.py +++ b/tests/discrete_tune.py @@ -1,12 +1,11 @@ import attune +import numpy as np def test_discrete(): dt = attune.DiscreteTune({"hi": (100, 200), "lo": (10, 20), "med": (20, 100)}, default="def") - assert dt(150) == "hi" - assert dt(20) == "lo" - assert dt(15) == "lo" - assert dt(100) == "hi" - assert dt(70) == "med" - assert dt(5) == "def" - assert dt(500) == "def" + x = [150, 20, 15, 100, 70, 5, 500] + y = ["hi", "lo", "lo", "hi", "med", "def", "def"] + assert np.all(dt(x) == np.asarray(y)) + assert dt(20) == dt(np.array(20)) == "lo" # should choose first valid range + assert all([dt(xi).item() == yi for xi, yi in zip(x, y)]) diff --git a/tests/instrument/test_update.py b/tests/instrument/test_update.py index 06ce880..106c8af 100644 --- a/tests/instrument/test_update.py +++ b/tests/instrument/test_update.py @@ -46,7 +46,7 @@ def test_update_existing(): inst2 = attune.Instrument({"arr": arr_new}, {"tune": attune.Setable("tune")}) inst_new = attune.update_merge(inst, inst2) assert math.isclose(inst_new(0.5)["tune"], 1.5) - assert inst_new(0.5)["discrete"] == "med" + assert inst_new(0.5)["discrete"] == inst(0.5)["discrete"] assert inst_new.transition.type == "update_merge"