From b4957e042013813bcb640c77cb9864f4b7f8fcf6 Mon Sep 17 00:00:00 2001 From: Joe Knapper Date: Fri, 6 Feb 2026 11:53:36 +0000 Subject: [PATCH] SnakeScan as a scan planner --- .../scan_planners.py | 194 +++++++++++++++++- tests/unit_tests/test_scan_planners.py | 110 ++++++++++ 2 files changed, 303 insertions(+), 1 deletion(-) diff --git a/src/openflexure_microscope_server/scan_planners.py b/src/openflexure_microscope_server/scan_planners.py index 454c0660..d2af3cad 100644 --- a/src/openflexure_microscope_server/scan_planners.py +++ b/src/openflexure_microscope_server/scan_planners.py @@ -11,7 +11,7 @@ from __future__ import annotations import logging from copy import copy -from typing import Any, Optional, TypeAlias +from typing import Any, Literal, Optional, TypeAlias import numpy as np @@ -605,6 +605,152 @@ class SmartSpiral(ScanPlanner): return np.max(np.abs(displacement_in_moves)) +class SnakeScan(ScanPlanner): + """A scan planner that performs a snake scan, right and down from a corner. + + This planner starts at the corner of the region to scan, snaking back and forth, + starting moving right and down (assuming positive dx and dy.) + """ + + _dx: int = 0 + _dy: int = 0 + _x_count: int = 0 + _y_count: int = 0 + + def __init__( + self, initial_position: XYPos, planner_settings: Optional[dict] = None + ) -> None: + """Set up the lists inherited from ScanPlanner, plus a distance cutoff. + + Use the supplied _dx and _dy to set a distance cutoff for an image to be + considered neighbouring another + """ + super().__init__(initial_position, planner_settings) + self._distance_cutoff: float = max([self._dx, self._dy]) * 1.1 + + def _parse(self, planner_settings: Optional[dict] = None) -> None: + """Parse SnakeScan Settings dictionary. + + * ``dx`` - the movement size in x + * ``dy`` - the movement size in y + * ``x_count`` - The number of columns in the scan. + * ``y_count`` - The number of rows in the scan. + """ + expected_keys = ["x_count", "y_count", "dx", "dy"] + invalid_msg = "SnakeScan requires a planner_settings dictionary with keys: " + if not planner_settings: + raise ValueError(invalid_msg + ",".join(expected_keys)) + if not all(keys in planner_settings for keys in expected_keys): + raise KeyError(invalid_msg + ",".join(expected_keys)) + + self._dx = int(planner_settings["dx"]) + self._dy = int(planner_settings["dy"]) + self._x_count = int(planner_settings["x_count"]) + self._y_count = int(planner_settings["y_count"]) + + def _initial_location_list(self) -> list[FutureScanLocation]: + """Set the initial list of locations for this scan planner. + + This is called on initialisation. + + For snake scan, this is the full grid, and none will be added during scanning. + """ + grid = create_rectangular_scan_path( + starting_pos=self._initial_position, + x_count=self._x_count, + y_count=self._y_count, + dx=self._dx, + dy=self._dy, + style="snake", + ) + + # create_rectangular_scan_path provides a nested list, which is flattened here + path = [] + for line in grid: + path += line + + return [FutureScanLocation(location) for location in path] + + def mark_location_visited( + self, xyz_pos: XYZPos, imaged: bool = True, focused: bool = True + ) -> None: + """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 + """ + # Only call the base class to update the positions + super().mark_location_visited(xyz_pos, imaged, focused) + + def get_next_location_and_z_estimate(self) -> tuple[XYPos, Optional[int]]: + """Return the next location to scan and its estimated z-position. + + This overrides the default behaviour of ScanPlanner to take the lowest value of + nearest neighbours as this works best for smart stack. + + Note z-position may be None! This indicates that the current z position + should be used. + """ + if self.scan_complete: + raise RuntimeError("Can't get next position, scan is complete") + + next_site = self._remaining_locations[0] + next_location = next_site.xy_tuple + + # If focused locations exist, return the neighbour with the lowest z position + closest_pos = self.select_nearby_focus_site(next_location) + z = None if closest_pos is None else closest_pos[2] + + return next_location, z + + def select_nearby_focus_site(self, xy_pos: XYPos) -> Optional[XYZPos]: + """Return the xyz position of the nearby site with the lowest z position. + + Lowest position is best, as starting too high causes smart stacking to + autofocus and restart. Starting too low just requires extra movements in +z. + Nearby is defined as within 1.1 times the larger of the x and y scan offsets. + + If no focused sites are within this range, use the height of the nearest + focused site. + + Returns None if no focused locations are present + """ + # 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(focused_locations, dtype="float64")[:, :2] + + # Use linalg.norm to calculate the direct distance between the points + # Note linalg.norm always uses float64 + dists = np.linalg.norm((path_pos - current_pos), axis=1) + + # Get indices of all focused sites within distance_cutoff. + # Note np.where always returns a tuple of arrays, hence the trailing [0] + indices = np.where(dists <= self._distance_cutoff)[0] + + # Handle the case that no focused positions are within this range, and + # instead use the nearest focused position. This will always return a + # height, due to the check that self._focused_locations exists. + if len(indices) == 0: + distance_cutoff = min(dists) + indices = np.where(dists <= distance_cutoff)[0] + + # Turning into an array allows slicing based on a list + 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 + chosen_focused_site = min(focused_locations_array[indices], key=lambda x: x[-1]) + + # Convert back into list so values are of type int instead of np.int32 + return tuple(chosen_focused_site.tolist()) + + def distance_between( current_pos: XYPos | np.ndarray | FutureScanLocation, next_pos: XYPos | np.ndarray | FutureScanLocation, @@ -620,3 +766,49 @@ def distance_between( 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)) + + +def create_rectangular_scan_path( + starting_pos: XYPos, + x_count: int, + y_count: int, + dx: int, + dy: int, + style: Literal["snake", "raster"], +) -> list: + """Generate a 2D grid of (x, y) coordinates representing a rectangular scan path. + + The grid is generated from starting_pos, and expanded in the + positive x and y directions using the provided step sizes. The scan order + can be either raster (left-to-right for every row) or snake (alternating + left-to-right and right-to-left per row). + + :param starting_pos: Starting (x, y) position for the scan grid. + :param x_count: Number of points in the x-direction (columns). + :param y_count: Number of points in the y-direction (rows). + :param dx: Step size between points in the x-direction. + :param dy: Step size between points in the y-direction. + :param style: Scan pattern style. Either raster or snake. + :return: Nested list of (x, y) coordinates arranged by row. + """ + coords = [] + + # Populate grid with coordinates in a regular grid + for y_index in range(y_count): # rows + coords.append([]) + for x_index in range(x_count): # cols + # Create a coordinate tuple + coord = ( + starting_pos[0] + [x_index, y_index][0] * dx, + starting_pos[1] + [x_index, y_index][1] * dy, + ) + # Append coordinate array to position grid + coords[y_index].append(coord) + + # If style is snake, reverse every second row + if style == "snake": + for i, line in enumerate(coords): + if i % 2 != 0: + # Reverse the list of coordinates + line.reverse() + return coords diff --git a/tests/unit_tests/test_scan_planners.py b/tests/unit_tests/test_scan_planners.py index bbaf061a..fcbd9dd7 100644 --- a/tests/unit_tests/test_scan_planners.py +++ b/tests/unit_tests/test_scan_planners.py @@ -308,3 +308,113 @@ def test_example_smart_spiral(): assert planner.path_history == expected_planner.path_history assert planner.imaged_locations == expected_planner.imaged_locations + + +def test_snake_scan_basic_grid(): + """Check that SnakeScan generates a single point for a 1x1 scan.""" + initial_position = (100, 50) + planner_settings = {"dx": 100, "dy": 100, "x_count": 1, "y_count": 1} + + planner = scan_planners.SnakeScan( + initial_position=initial_position, + planner_settings=planner_settings, + ) + + assert not planner.scan_complete + # When we start it should want to stay in the initial pos and have + # no z_estimate + xy_pos, z_pos = planner.get_next_location_and_z_estimate() + assert xy_pos == initial_position + assert z_pos is None + + # Try to mark location as imaged with only xy_position + with pytest.raises(ValueError, match="3 value tuple expected"): + planner.mark_location_visited(xy_pos, imaged=False, focused=False) + # scan still not complete + assert not planner.scan_complete + # if we mark this position as visited but not imaged + planner.mark_location_visited( + (xy_pos[0], xy_pos[1], 10), imaged=False, focused=False + ) + # scan is now complete + assert planner.scan_complete + + # if scan is complete, asking for the next location returns an error + with pytest.raises(RuntimeError): + planner.get_next_location_and_z_estimate() + + +def test_snake_scan_basic_length(): + """SnakeScan should generate the correct number of locations.""" + initial_position = (100, 50) + planner_settings = {"dx": 100, "dy": 100, "x_count": 3, "y_count": 4} + + planner = scan_planners.SnakeScan( + initial_position=initial_position, + planner_settings=planner_settings, + ) + + coords = planner.remaining_locations + + assert len(coords) == 3 * 4 + + +def test_snake_scan_ordering(): + """Test that snake scan returns a path in the right order.""" + initial_position = (0, 0) + planner_settings = {"dx": 10, "dy": 10, "x_count": 4, "y_count": 3} + + planner = scan_planners.SnakeScan( + initial_position=initial_position, + planner_settings=planner_settings, + ) + + coords = planner.remaining_locations + + expected = [ + (0, 0), + (10, 0), + (20, 0), + (30, 0), + (30, 10), + (20, 10), + (10, 10), + (0, 10), + (0, 20), + (10, 20), + (20, 20), + (30, 20), + ] + + assert coords == expected + + +def test_snake_scan_single_row(): + """Test edge case of a single row scan.""" + initial_position = (0, 0) + planner_settings = {"dx": 5, "dy": 5, "x_count": 4, "y_count": 1} + + planner = scan_planners.SnakeScan( + initial_position=initial_position, + planner_settings=planner_settings, + ) + + assert planner.remaining_locations == [(0, 0), (5, 0), (10, 0), (15, 0)] + + +def test_snake_scan_single_column(): + """Test edge case of a single column scan.""" + initial_position = (0, 0) + planner_settings = {"dx": 5, "dy": 5, "x_count": 1, "y_count": 4} + + planner = scan_planners.SnakeScan( + initial_position=initial_position, + planner_settings=planner_settings, + ) + + assert planner.remaining_locations == [ + (0, 0), + (0, 5), + (0, 10), + (0, 15), + ]