SnakeScan as a scan planner

This commit is contained in:
Joe Knapper 2026-02-06 11:53:36 +00:00
parent 663a18a38f
commit b4957e0420
2 changed files with 303 additions and 1 deletions

View file

@ -11,7 +11,7 @@ from __future__ import annotations
import logging import logging
from copy import copy from copy import copy
from typing import Any, Optional, TypeAlias from typing import Any, Literal, Optional, TypeAlias
import numpy as np import numpy as np
@ -605,6 +605,152 @@ class SmartSpiral(ScanPlanner):
return np.max(np.abs(displacement_in_moves)) 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( def distance_between(
current_pos: XYPos | np.ndarray | FutureScanLocation, current_pos: XYPos | np.ndarray | FutureScanLocation,
next_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") next_pos = np.array(next_pos, dtype="float64")
current_pos = np.array(current_pos, dtype="float64") current_pos = np.array(current_pos, dtype="float64")
return float(np.linalg.norm(next_pos - current_pos)) 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

View file

@ -308,3 +308,113 @@ def test_example_smart_spiral():
assert planner.path_history == expected_planner.path_history assert planner.path_history == expected_planner.path_history
assert planner.imaged_locations == expected_planner.imaged_locations 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),
]