Calculate simulation images downsampled for speed

This commit is contained in:
Julian Stirling 2025-12-26 12:15:22 +00:00
parent 1e36811796
commit fa249a0f24

View file

@ -13,7 +13,7 @@ import logging
import time import time
from threading import Thread from threading import Thread
from types import TracebackType from types import TracebackType
from typing import Literal, Optional, Self from typing import Literal, Optional, Self, overload
import numpy as np import numpy as np
from PIL import Image, ImageFilter from PIL import Image, ImageFilter
@ -43,6 +43,26 @@ BG_COLOR = [220, 215, 217]
# Random Number Generator # Random Number Generator
RNG = np.random.default_rng() RNG = np.random.default_rng()
DOWNSAMPLE = 2
@overload
def _downsample_shape(shape: tuple[int, int]) -> tuple[int, int]: ...
@overload
def _downsample_shape(shape: tuple[int, int, int]) -> tuple[int, int, int]: ...
def _downsample_shape(
shape: tuple[int, int] | tuple[int, int, int],
) -> tuple[int, int] | tuple[int, int, int]:
if len(shape) == 2:
return (shape[0] // DOWNSAMPLE, shape[1] // DOWNSAMPLE)
if len(shape) == 3:
return (shape[0] // DOWNSAMPLE, shape[1] // DOWNSAMPLE, shape[2])
raise ValueError("A shape should be a 2 or 3 element tuple.")
class SimulatedCamera(BaseCamera): class SimulatedCamera(BaseCamera):
"""A Thing that simulates a camera for testing.""" """A Thing that simulates a camera for testing."""
@ -75,11 +95,11 @@ class SimulatedCamera(BaseCamera):
""" """
super().__init__(thing_server_interface) super().__init__(thing_server_interface)
self.shape = shape self.shape = shape
self.glyph_shape = glyph_shape self.ds_shape = _downsample_shape(shape)
self.canvas_shape = canvas_shape self.glyph_shape = _downsample_shape(glyph_shape)
self.sample_limits = ( self.canvas_shape = _downsample_shape(canvas_shape)
canvas_shape[:2] if sample_limits is None else sample_limits sample_limits = canvas_shape[:2] if sample_limits is None else sample_limits
) self.sample_limits = _downsample_shape(sample_limits)
self.frame_interval = frame_interval self.frame_interval = frame_interval
self._capture_thread: Optional[Thread] = None self._capture_thread: Optional[Thread] = None
self._capture_enabled = False self._capture_enabled = False
@ -108,6 +128,7 @@ class SimulatedCamera(BaseCamera):
def generate_sprites(self) -> None: def generate_sprites(self) -> None:
"""Generate sprites to populate the image.""" """Generate sprites to populate the image."""
sprite_sizes = [10, 21, 36, 40, 50] sprite_sizes = [10, 21, 36, 40, 50]
sprite_sizes = [s / DOWNSAMPLE for s in sprite_sizes]
self.sprites = [] self.sprites = []
channel_block = np.zeros(self.glyph_shape[0:2]) channel_block = np.zeros(self.glyph_shape[0:2])
@ -202,9 +223,13 @@ class SimulatedCamera(BaseCamera):
:param pos: a 3-item tuple containing the x,y,z coordinates of the 'stage' :param pos: a 3-item tuple containing the x,y,z coordinates of the 'stage'
""" """
canvas_width, canvas_height, _ = self.canvas_shape canvas_width, canvas_height, _ = self.canvas_shape
image_width, image_height, _ = self.shape image_width, image_height, _ = self.ds_shape
# Scale position by RATIO to get position in base image. im_pos = (
im_pos = tuple(x * ratio for x, ratio in zip(pos, RATIO, strict=True)) pos[0] * RATIO[0] / DOWNSAMPLE,
pos[1] * RATIO[1] / DOWNSAMPLE,
pos[2] * RATIO[2],
)
top_left = ( top_left = (
int(im_pos[0]) - image_width // 2 + self.sample_limits[0] // 2, int(im_pos[0]) - image_width // 2 + self.sample_limits[0] // 2,
int(im_pos[1]) - image_height // 2 + self.sample_limits[1] // 2, int(im_pos[1]) - image_height // 2 + self.sample_limits[1] // 2,
@ -213,23 +238,24 @@ class SimulatedCamera(BaseCamera):
# canvas edge. # canvas edge.
x_indices = (np.arange(top_left[0], top_left[0] + image_width)) % canvas_width x_indices = (np.arange(top_left[0], top_left[0] + image_width)) % canvas_width
y_indices = (np.arange(top_left[1], top_left[1] + image_height)) % canvas_height y_indices = (np.arange(top_left[1], top_left[1] + image_height)) % canvas_height
z_indices = np.arange(self.shape[2]) z_indices = np.arange(self.ds_shape[2])
canvas = self.canvas if self._show_sample else self.blank_canvas canvas = self.canvas if self._show_sample else self.blank_canvas
# Use npx to make each 1d index list 3D # Use npx to make each 1d index list 3D
focused_image = canvas[np.ix_(x_indices, y_indices, z_indices)] focused_image = canvas[np.ix_(x_indices, y_indices, z_indices)]
image = fast_pil_blur(focused_image, sigma=np.abs(im_pos[2]) / 5) image = fast_pil_blur(focused_image, sigma=np.abs(im_pos[2]) / 5)
if image.shape != self.shape: if image.shape != self.ds_shape:
raise ValueError( raise ValueError(
f"Image shape {image.shape} does not match intended shape {self.shape}" f"Image shape {image.shape} does not match intended shape {self.ds_shape}"
) )
# Add noise and convert to uint8 # Add noise and convert to uint8
image += RNG.normal(scale=self.noise_level, size=self.shape).astype("int16") image += RNG.normal(scale=self.noise_level, size=self.ds_shape).astype("int16")
image[image < 0] = 0 image[image < 0] = 0
image[image > 255] = 255 image[image > 255] = 255
return Image.fromarray(image.astype("uint8")) image = Image.fromarray(image.astype("uint8"))
return image.resize((self.shape[1], self.shape[0]), Image.BILINEAR)
def generate_frame(self) -> Image.Image: def generate_frame(self) -> Image.Image:
"""Generate a frame with blobs based on the stage coordinates.""" """Generate a frame with blobs based on the stage coordinates."""