Skip to content
Merged
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
6 changes: 6 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
26 changes: 11 additions & 15 deletions attune/_discrete_tune.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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()):

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

reversed iteration insures earlier items will have precedence.

out[(ind_value >= imin) & (ind_value <= imax)] = key
return out

def __eq__(self, other):
return self.ranges == other.ranges and self.default == other.default
Expand Down
2 changes: 1 addition & 1 deletion attune/_tune.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
18 changes: 9 additions & 9 deletions docs/structure.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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")



Expand Down
13 changes: 6 additions & 7 deletions tests/discrete_tune.py
Original file line number Diff line number Diff line change
@@ -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)])
2 changes: 1 addition & 1 deletion tests/instrument/test_update.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"


Expand Down