diff --git a/src/openflexure_microscope_server/scan_planners.py b/src/openflexure_microscope_server/scan_planners.py index 7387e70d..f3ace6f1 100644 --- a/src/openflexure_microscope_server/scan_planners.py +++ b/src/openflexure_microscope_server/scan_planners.py @@ -5,7 +5,11 @@ is only one type the SmartSpiral. More can be added using by subclassing the ScanPlanner """ -from typing import TypeAlias, Optional +# Future annotations needed for typhinting same class in __eq__ method. Other option +# would be to import Union and use a string. +from __future__ import annotations + +from typing import TypeAlias, Optional, Any import logging from copy import copy @@ -51,6 +55,101 @@ def enforce_xyz_tuple(value: XYZPos) -> XYZPos: return value +# Disable PLW1641, this warns against hash not being set when equals is. But the +# class shouldn't be hashable. +class FutureScanLocation: # noqa PLW1641 + """Data for information on future locations to scan. + + This object is only used for internal calculation and data storage it shouldn't + be an input or output of public methods. + """ + + def __init__(self, xy_pos: XYPos, **kwargs: Any) -> None: + """Initialise FutureScanLocation with an xy position. + + :param xy_pos: The (x, y) position to scan. + :param kwargs: Any other information about the location. This will be passed + through to the VisitedScanLocation object. + """ + self._xy_pos = enforce_xy_tuple(xy_pos) + self.planner_data = copy(kwargs) + + @property + def xy_tuple(self) -> XYPos: + """The xy position tuple.""" + return self._xy_pos + + def __eq__(self, other: XYPos | FutureScanLocation) -> bool: + """Check for equality, only checks the xy position. + + Will check against tuple or other FutureScanLocation object. + """ + if isinstance(other, FutureScanLocation): + return self._xy_pos == other.xy_tuple + if isinstance(other, tuple) and len(other) == 2: + return self._xy_pos[:2] == other + return NotImplemented + + +# Disable PLW1641, this warns against hash not being set when equals is. But the +# class shouldn't be hashable. +class VisitedScanLocation: # noqa PLW1641 + """Data for information on locations already visited during a scan. + + This object is only used for internal calculation and data storage it shouldn't + be an input or output of public methods. + """ + + def __init__( + self, xyz_pos: XYZPos, imaged: bool, focused: bool, **kwargs: Any + ) -> None: + """Initialise VisitedScanLocation with an xyz-position, and whether imaged/focused. + + :param xyz_pos: The (x, y, z) position visited. + :param imaged: True if an image was taken, False if not (due to background + detect) + :param focused: True if autofocus completed successfully + :param kwargs: Any other information about the location. This should be passed + through to from the FutureScanLocation object that requested the location + to be scanned. + """ + self._xyz_pos = enforce_xyz_tuple(xyz_pos) + self.imaged = imaged + self.focused = focused + self.planner_data = copy(kwargs) + + @property + def xyz_tuple(self) -> XYZPos: + """The xyz position tuple.""" + return self._xyz_pos + + @property + def xy_tuple(self) -> XYPos: + """The xy position tuple.""" + return self._xyz_pos[:2] + + def __eq__( + self, other: XYPos | XYZPos | FutureScanLocation | VisitedScanLocation + ) -> bool: + """Check for equality, only checks the xyz-position or xy-position if z isn't available. + + Will check xyz-position against 3-value tuples and other VisitedScanLocation + objects. + + Will check xy-position against 2-value tuples and FutureScanLocation objects. + """ + if isinstance(other, VisitedScanLocation): + return self._xyz_pos == other.xyz_tuple + if isinstance(other, FutureScanLocation): + return self._xyz_pos[:2] == other.xy_tuple + if isinstance(other, tuple): + if len(other) == 2: + return self._xyz_pos[:2] == other + if len(other) == 3: + return self._xyz_pos == other + return NotImplemented + + class ScanPlanner: """A base class for a scan planner. @@ -79,25 +178,10 @@ class ScanPlanner: self._initial_position = enforce_xy_tuple(initial_position) self._parse(planner_settings) - # The remaining (x,y) locations to scan - # (This was `path` before refactoring from the long `sample_scan` code) - self._remaining_locations: XYPosList = self._initial_location_list() - - # This holds a list of all (x,y,z) locations where images were taken - # this may not be equivalent to the x,y positions ins self._path_history - # if background detect is used - # (This was not used in the `sample_scan` code) - self._imaged_locations: XYZPosList = [] - - # This holds a list of all (x,y,z) locations where autofocus was successful - # (This was `focused_path` before refactoring from the long `sample_scan` code) - self._focused_locations: XYZPosList = [] - - # This holds a list of all x,y locations visited in order since the start - # (This was `true_path` before refactoring from the long `sample_scan` code - # previously it had z set, but if we don't take an image, not z is needed and - # it slows other checks) - self._path_history: XYPosList = [] + self._remaining_locations: list[FutureScanLocation] = ( + self._initial_location_list() + ) + self._path_history: list[VisitedScanLocation] = [] @property def scan_complete(self) -> bool: @@ -107,48 +191,47 @@ class ScanPlanner: @property def remaining_locations(self) -> XYPosList: """Property to access a copy of the remaining_locations.""" - return copy(self._remaining_locations) + return [loc.xy_tuple for loc in self._remaining_locations] @property def imaged_locations(self) -> XYZPosList: """Property to access a copy of the imaged_locations.""" - return copy(self._imaged_locations) + return [loc.xyz_tuple for loc in self._path_history if loc.imaged] @property def focused_locations(self) -> XYZPosList: """Property to access a copy of the focused_locations.""" - return copy(self._focused_locations) + return [loc.xyz_tuple for loc in self._path_history if loc.focused] @property def path_history(self) -> XYPosList: """Property to access a copy of the path_history.""" - return copy(self._path_history) + return [loc.xy_tuple for loc in self._path_history] def _parse(self, planner_settings: Optional[dict] = None) -> None: """Parse any settings sent to this planner and store them if needed.""" raise NotImplementedError("Did you call the ScanPlanner base class?") - def _initial_location_list(self) -> XYPosList: + def _initial_location_list(self) -> list[FutureScanLocation]: """Set the initial list of locations for this scan planner. This is called on initialisation. For a simple grid scan/snake scan this would be all locations to move to. - Note for implementation that this _must_ contain (x,y) tuples, not [x, y] - lists or matching errors could occur. + :return: A list of FutureScanLocation objects with all planned locations. """ raise NotImplementedError("Did you call the ScanPlanner base class?") - def position_visited(self, position: XYPos) -> bool: - """Return True if input xy position has been visited before.""" + def position_visited(self, position: FutureScanLocation) -> bool: + """Return True if input scan position has been visited before.""" # Ensure tuple for correct matching! - return tuple(position) in self._path_history + return position in self._path_history - def position_planned(self, position: XYPos) -> bool: - """Return True if input xy position is planned.""" + def position_planned(self, position: FutureScanLocation) -> bool: + """Return True if input scan position position is planned.""" # Ensure tuple for correct matching! - return tuple(position) in self._remaining_locations + return position in self._remaining_locations def get_next_location_and_z_estimate(self) -> tuple[XYPos, Optional[int]]: """Return the next location to scan and its estimated z-position. @@ -159,7 +242,7 @@ class ScanPlanner: if self.scan_complete: raise RuntimeError("Can't get next position, scan is complete") - next_location = self._remaining_locations[0] + next_location = self._remaining_locations[0].xy_tuple # If focussed locations exist return closest location, favouring most recent closest_pos = self.closest_focus_site(next_location) @@ -177,12 +260,14 @@ class ScanPlanner: Returns None if there if no focussed locations are present """ - if not self._focused_locations: + # save to variable rather than search for focussed sites each time. + focused_locations = self.focused_locations + if not focused_locations: return None # must be float64 (double precision) to deal with the huge numbers involved! current_pos = np.array(xy_pos, dtype="float64") - path_pos = np.array(self._focused_locations, dtype="float64")[:, :2] + path_pos = np.array(focused_locations, dtype="float64")[:, :2] # Use linalg.norm to calculate the direct distance bweween the points # Note linalg.norm always uses float64 @@ -193,36 +278,35 @@ class ScanPlanner: indices = np.where(dists == np.min(dists))[0] # The last index is most recent - return self._focused_locations[indices[-1]] + return focused_locations[indices[-1]] def mark_location_visited( self, xyz_pos: XYZPos, imaged: bool, focused: bool ) -> None: """Mark the location as visited. - Args: - xyz_pos: the x_y_z position - imaged: true if an image was taken, false if not (due to background detect) - focused: true if autofocus completed successfully - + :param xyz_pos: the x_y_z position + :param imaged: true if an image was taken, false if not (due to background detect) + :param focused: true if autofocus completed successfully """ # ensure is tuple! xyz_pos = enforce_xyz_tuple(xyz_pos) - xy_pos = xyz_pos[:2] # Remove the expected position from the remaining locations list # and check it's correct - expected_pos = tuple(self._remaining_locations.pop(0)) - if xy_pos != expected_pos: + expected_pos = self._remaining_locations.pop(0) + if xyz_pos[:2] != expected_pos.xy_tuple: raise RuntimeError("Wrong scan location visited!") # Append xy position for path_history - self._path_history.append(xy_pos) - # And full x,y,z for imaged and foucsed if appropriate - if imaged: - self._imaged_locations.append(xyz_pos) - if focused: - self._focused_locations.append(xyz_pos) + self._path_history.append( + VisitedScanLocation( + xyz_pos=xyz_pos, + imaged=imaged, + focused=focused, + **expected_pos.planner_data, + ) + ) class SmartSpiral(ScanPlanner): @@ -275,25 +359,23 @@ class SmartSpiral(ScanPlanner): self._dy = int(planner_settings["dy"]) self._max_dist = int(planner_settings["max_dist"]) - def _initial_location_list(self) -> XYPosList: + def _initial_location_list(self) -> list[FutureScanLocation]: """Set the initial list of locations for this scan planner. This is salled on initialisation. For smart spiral this is just the first point """ - return [self._initial_position] + return [FutureScanLocation(self._initial_position)] def mark_location_visited( self, xyz_pos: XYZPos, imaged: bool = True, focused: bool = True ) -> None: - """Mark the location as visited. Adjust extra positions accordingly. - - Args: - xyz_pos: the x_y_z position - imaged: true if an image was taken, false if not (due to background detect) - focused: true if autofocus completed successfully + """Mark the location as visited. + :param xyz_pos: the x_y_z position + :param imaged: true if an image was taken, false if not (due to background detect) + :param focused: true if autofocus completed successfully """ # First call the base class to update the positions super().mark_location_visited(xyz_pos, imaged, focused) @@ -314,10 +396,10 @@ class SmartSpiral(ScanPlanner): * already visited """ new_positions = [ - (xy_pos[0] - self._dx, xy_pos[1]), - (xy_pos[0] + self._dx, xy_pos[1]), - (xy_pos[0], xy_pos[1] - self._dy), - (xy_pos[0], xy_pos[1] + self._dy), + FutureScanLocation((xy_pos[0] - self._dx, xy_pos[1])), + FutureScanLocation((xy_pos[0] + self._dx, xy_pos[1])), + FutureScanLocation((xy_pos[0], xy_pos[1] - self._dy)), + FutureScanLocation((xy_pos[0], xy_pos[1] + self._dy)), ] for new_pos in new_positions: @@ -356,7 +438,7 @@ class SmartSpiral(ScanPlanner): if self.scan_complete: raise RuntimeError("Can't get next position, scan is complete") - next_location = self._remaining_locations[0] + next_location = self._remaining_locations[0].xy_tuple # If focused locations exist, return the neighbour with the lowest z position closest_pos = self.select_nearby_focus_site(next_location) @@ -376,12 +458,14 @@ class SmartSpiral(ScanPlanner): Returns None if no focused locations are present """ - if not self._focused_locations: + # save to variable rather than search for focussed sites each time. + focused_locations = self.focused_locations + if not focused_locations: return None # must be float64 (double precision) to deal with the huge numbers involved! current_pos = np.array(xy_pos, dtype="float64") - path_pos = np.array(self._focused_locations, dtype="float64")[:, :2] + path_pos = np.array(focused_locations, dtype="float64")[:, :2] # Use linalg.norm to calculate the direct distance between the points # Note linalg.norm always uses float64 @@ -399,7 +483,7 @@ class SmartSpiral(ScanPlanner): indices = np.where(dists <= distance_cutoff)[0] # Turning into an array allows slicing based on a list - focused_locations_array = np.array(self._focused_locations) + focused_locations_array = np.array(focused_locations) # Choose the lowest (smallest z) of the neighbouring sites. Smart stack works best # if started too low, so the lowest z will perform best @@ -410,8 +494,8 @@ class SmartSpiral(ScanPlanner): def moves_between( self, - starting_pos: XYPos | np.ndarray, - ending_pos: XYPos | np.ndarray, + starting_pos: XYPos | np.ndarray | FutureScanLocation, + ending_pos: XYPos | np.ndarray | FutureScanLocation, ) -> float: """Return the larger of x moves or y moves between two xy positions. @@ -419,6 +503,10 @@ class SmartSpiral(ScanPlanner): :param ending_pos: the position to measure to """ + if isinstance(starting_pos, FutureScanLocation): + starting_pos = starting_pos.xy_tuple + if isinstance(ending_pos, FutureScanLocation): + ending_pos = ending_pos.xy_tuple move_size = np.array([self._dx, self._dy]) starting_pos = np.array(starting_pos, dtype="float64") @@ -430,12 +518,17 @@ class SmartSpiral(ScanPlanner): def distance_between( - current_pos: XYPos | np.ndarray, next_pos: XYPos | np.ndarray + current_pos: XYPos | np.ndarray | FutureScanLocation, + next_pos: XYPos | np.ndarray | FutureScanLocation, ) -> float: """Calculate the distance between the two xy positions. This was previously called ``distance_to_site`` """ + if isinstance(current_pos, FutureScanLocation): + current_pos = current_pos.xy_tuple + if isinstance(next_pos, FutureScanLocation): + next_pos = next_pos.xy_tuple next_pos = np.array(next_pos, dtype="float64") current_pos = np.array(current_pos, dtype="float64") return float(np.linalg.norm(next_pos - current_pos)) diff --git a/tests/test_scan_planners.py b/tests/test_scan_planners.py index 89996945..940257d0 100644 --- a/tests/test_scan_planners.py +++ b/tests/test_scan_planners.py @@ -254,17 +254,27 @@ def test_closest_focus_with_large_numbers(): planner = scan_planners.SmartSpiral( initial_position=initial_position, planner_settings=planner_settings ) - # Directly overwrite the private focussed locations list for test + # Directly overwrite the private path history locations list for test # For two points 1m points away it should choose the last as they are equal - planner._focused_locations = [(1000000, 0, 0), (0, 1000000, 0)] + + planner._path_history = [ + scan_planners.VisitedScanLocation((1000000, 0, 0), imaged=True, focused=True), + scan_planners.VisitedScanLocation((0, 1000000, 0), imaged=True, focused=True), + ] assert planner.closest_focus_site((0, 0)) == (0, 1000000, 0) # Try similar - planner._focused_locations = [(1234567, 0, 0), (-1234567, 0, 0)] + planner._path_history = [ + scan_planners.VisitedScanLocation((1234567, 0, 0), imaged=True, focused=True), + scan_planners.VisitedScanLocation((-1234567, 0, 0), imaged=True, focused=True), + ] assert planner.closest_focus_site((0, 0)) == (-1234567, 0, 0) # Make the first point 1 step closer - planner._focused_locations = [(1234566, 0, 0), (-1234567, 0, 0)] + planner._path_history = [ + scan_planners.VisitedScanLocation((1234566, 0, 0), imaged=True, focused=True), + scan_planners.VisitedScanLocation((-1234567, 0, 0), imaged=True, focused=True), + ] assert planner.closest_focus_site((0, 0)) == (1234566, 0, 0) @@ -292,5 +302,14 @@ def test_example_smart_spiral(): expected_planner = ( scan_test_helpers.get_expected_result_for_example_smart_spiral(sample_name) ) - assert planner.path_history == expected_planner.path_history - assert planner.imaged_locations == expected_planner.imaged_locations + # Use hasattr to check if the saved test data is the old object where imaged + # locations were saved directly as a list of tuples rather than being generated + # from a history of VisitedScanLocation objects. + # This "if" section of this if-else block can be deleted once we are confident + # of the new format and the pickles are updated + if hasattr(expected_planner, "_imaged_locations"): + assert planner.path_history == expected_planner._path_history + assert planner.imaged_locations == expected_planner._imaged_locations + else: + assert planner.path_history == expected_planner.path_history + assert planner.imaged_locations == expected_planner.imaged_locations diff --git a/tests/utilities/scan_test_helpers.py b/tests/utilities/scan_test_helpers.py index 288996ac..cec99249 100644 --- a/tests/utilities/scan_test_helpers.py +++ b/tests/utilities/scan_test_helpers.py @@ -60,8 +60,8 @@ def visualise_scan(sample: FakeSample, planner: scan_planners.ScanPlanner) -> Fi """For a given sample and scanner object return a matplotlib figure of the scan.""" fig, ax = plt.subplots(figsize=(8, 8)) ax.add_artist(sample.patch) - xh, yh = zip(*planner._path_history, strict=True) - xi, yi, _ = zip(*planner._imaged_locations, strict=True) + xh, yh = zip(*planner.path_history, strict=True) + xi, yi, _zi = zip(*planner.imaged_locations, strict=True) # convert history to numpy array so can calculate quiver arrows xh = np.array(xh) @@ -99,7 +99,7 @@ def interp_closed_path(xy_points: list[tuple[int, int]], n_points: int) -> MatPa # fit splines to x=f(u) and y=g(u), treating both as periodic. also note that s=0 # is needed in order to force the spline fit to pass through all the input points. - spline_data, _ = interpolate.splprep([x, y], s=0, per=True) + spline_data, *_extra = interpolate.splprep([x, y], s=0, per=True) # evaluate the spline xi, yi = interpolate.splev(np.linspace(0, 1, n_points), spline_data)