"""Matplotlib canvases and the ROI primitive.""" import matplotlib as mpl import numpy as np from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg from matplotlib.colors import BoundaryNorm, ListedColormap from matplotlib.figure import Figure from matplotlib.patches import Polygon from matplotlib.path import Path as MplPath from PyQt6.QtCore import Qt, pyqtSignal from PyQt6.QtGui import QKeyEvent from PyQt6.QtWidgets import QSizePolicy from sras_format import CH1_IDX, CH3_IDX, CH4_IDX, CH_NAMES, SrasFile, adc_to_mv def count_colormap(n_angles: int): """(cmap, norm, ticks) for an integer "how many angles cover this pixel" image, 0..n_angles. Discrete, not continuous: the judgement the wizard's stack view exists for is "is this a plateau at N, or a fan of partial overlaps", so a region covered by one angle too few has to read as its own band rather than a slightly darker shade. Count 0 is fully transparent so uncovered canvas cannot be mistaken for a low count. Shared by both wizard pages that draw this image — they use different canvas classes, and the same number must not change colour between them. """ n = max(1, int(n_angles)) base = mpl.colormaps["viridis"].resampled(n) colors = [(0.0, 0.0, 0.0, 0.0)] + [base(i) for i in range(n)] return (ListedColormap(colors), BoundaryNorm(np.arange(-0.5, n + 1), len(colors)), np.arange(0, n + 1)) # --------------------------------------------------------------------------- # ROI (free quadrilateral in data coordinates) # --------------------------------------------------------------------------- class RoiQuad: """Free quadrilateral defined in data coordinates (mm). Stored as 4 corner points (shape (4, 2)) in CCW order: BL, BR, TR, TL. Each corner can be positioned independently, allowing skewed / non-orthogonal regions of interest. Because it lives in scan/data coords it persists unchanged when the displayed channel/mode switches. """ def __init__(self, pts: np.ndarray): """pts : array-like, shape (4, 2).""" self._pts = np.asarray(pts, dtype=np.float64).reshape(4, 2).copy() @classmethod def from_bbox(cls, x0: float, y0: float, x1: float, y1: float) -> "RoiQuad": """Create an axis-aligned rectangle from two opposite corners.""" lx, rx = min(x0, x1), max(x0, x1) by, ty = min(y0, y1), max(y0, y1) return cls(np.array([[lx, by], [rx, by], [rx, ty], [lx, ty]])) def copy(self) -> "RoiQuad": return RoiQuad(self._pts.copy()) def corners(self) -> np.ndarray: """World-coord corners, shape (4, 2), CCW: BL, BR, TR, TL.""" return self._pts.copy() def centroid(self) -> np.ndarray: return self._pts.mean(axis=0) def bbox_size(self) -> np.ndarray: """Width and height of the axis-aligned bounding box, shape (2,).""" return self._pts.max(axis=0) - self._pts.min(axis=0) def contains(self, x: float, y: float) -> bool: return bool(MplPath(self._pts).contains_point((x, y))) def mask_for_grid(self, x_axis: np.ndarray, y_axis: np.ndarray) -> np.ndarray: """Boolean mask (n_rows, n_frames) of pixels whose centres lie inside the quadrilateral. Only the quad's axis-aligned bounding box is tested — meshgrid and contains_points over the *whole* grid would be tens of millions of point-in-polygon tests (and hundreds of MB of float64 temporaries) on a large scan, on every ROI edit. """ x = np.asarray(x_axis, dtype=np.float64) y = np.asarray(y_axis, dtype=np.float64) mask = np.zeros((y.size, x.size), dtype=bool) (x0, y0), (x1, y1) = self._pts.min(axis=0), self._pts.max(axis=0) cols = np.nonzero((x >= x0) & (x <= x1))[0] rows = np.nonzero((y >= y0) & (y <= y1))[0] if cols.size == 0 or rows.size == 0: return mask c0, c1 = int(cols[0]), int(cols[-1]) + 1 r0, r1 = int(rows[0]), int(rows[-1]) + 1 X, Y = np.meshgrid(x[c0:c1], y[r0:r1]) inside = MplPath(self._pts).contains_points( np.column_stack([X.ravel(), Y.ravel()])) mask[r0:r1, c0:c1] = inside.reshape(X.shape) return mask # --------------------------------------------------------------------------- # Matplotlib canvases # --------------------------------------------------------------------------- class ImageCanvas(FigureCanvasQTAgg): pixel_clicked = pyqtSignal(int, int) # row_idx, frame_idx roi_changed = pyqtSignal() # ROI created / edited / cleared draw_mode_changed = pyqtSignal(bool) # "draw new ROI" arm toggled # Interaction state values _IDLE = "idle" _DRAW_NEW = "draw_new" _MOVE = "move" _DRAG_CORNER = "drag_corner" # Hit tolerance (display pixels) for handles. _HANDLE_PX = 12 _CLICK_THRESH_PX = 4 # releases within this of press count as a click def __init__(self, parent=None, *, rect_only: bool = False): """*rect_only* constrains the ROI to an axis-aligned rectangle. Used by the alignment wizard's crop page, where a free quadrilateral would be actively misleading: v6 geometry can only express an axis-aligned rectangle, so anything else the user drew would have to be squared off behind their back. Default off, so the main window's free-quad ROI is unaffected. """ fig = Figure(figsize=(7, 5), tight_layout=True) self.ax = fig.add_subplot(111) super().__init__(fig) self.setParent(parent) self.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding) self._extent = None self._img_shape = None self._rect_only = rect_only # ROI state self._roi: RoiQuad | None = None self._roi_artists: list = [] self._state = self._IDLE self._draw_mode = False # Per-interaction snapshots / anchors self._press_xy: tuple[float, float] | None = None self._press_pixel: tuple[float, float] | None = None self._press_button = None self._snapshot: RoiQuad | None = None self._drag_corner_idx: int = -1 self._move_anchor = None # press-point in world coords self._draw_previous: RoiQuad | None = None self.mpl_connect("button_press_event", self._on_press) self.mpl_connect("motion_notify_event", self._on_motion) self.mpl_connect("button_release_event", self._on_release) # ------------------------------------------------------------------ # Public API # ------------------------------------------------------------------ def show_image(self, img: np.ndarray, extent: list[float], cmap, vmin: float, vmax: float, xlabel: str, ylabel: str, title: str, colorbar_label: str = "", cb_ticks=None, norm=None, bad_color=None): """*cmap* may be a name or a Colormap instance. *norm* (which overrides vmin/vmax) and *cb_ticks* let a caller draw a discrete integer image — the wizard's overlap-count view — with whole-number colorbar bands instead of a continuous shade. *bad_color*, if given, is the fill for NaN pixels — a copy of *cmap* is made so the shared, registered instance is never mutated.""" self.figure.clf() self.ax = self.figure.add_subplot(111) # Patches and lines are destroyed by figure.clf(); drop stale refs. self._roi_artists = [] self._extent = extent self._img_shape = img.shape if bad_color is not None: cmap = (cmap if hasattr(cmap, "with_extremes") else mpl.colormaps[cmap]).with_extremes(bad=bad_color) kw = ({"norm": norm} if norm is not None else {"vmin": vmin, "vmax": vmax}) im = self.ax.imshow( img, aspect="auto", origin="upper", extent=extent, cmap=cmap, interpolation="nearest", **kw, ) cb = self.figure.colorbar(im, ax=self.ax, fraction=0.046, pad=0.04, ticks=cb_ticks) if colorbar_label: cb.set_label(colorbar_label) self.ax.set_xlabel(xlabel) self.ax.set_ylabel(ylabel) self.ax.set_title(title) # Re-draw the ROI (if any) on top of the fresh image so it persists # unchanged across mode / angle / channel switches. self._draw_roi() self.draw() def get_roi(self) -> RoiQuad | None: return self._roi def set_roi(self, roi: RoiQuad | None): self._roi = roi.copy() if roi is not None else None self._draw_roi() self.draw_idle() self.roi_changed.emit() def clear_roi(self): self._roi = None self._remove_roi_artists() self.draw_idle() self.roi_changed.emit() def start_drawing(self): """Arm the next click+drag on the image to create a new ROI, replacing any existing one.""" self._draw_mode = True self.setCursor(Qt.CursorShape.CrossCursor) self.draw_mode_changed.emit(True) def cancel_drawing(self): if self._draw_mode: self._draw_mode = False self.setCursor(Qt.CursorShape.ArrowCursor) self.draw_mode_changed.emit(False) # ------------------------------------------------------------------ # Rendering # ------------------------------------------------------------------ def _remove_roi_artists(self): for a in self._roi_artists: try: a.remove() except (ValueError, AttributeError, NotImplementedError): pass self._roi_artists = [] def _draw_roi(self): self._remove_roi_artists() if self._roi is None or self.ax is None: return corners = self._roi.corners() # Filled quad, then a sharp unfilled edge for visibility over bright # images, then draggable corner handles. for kwargs in ( dict(fill=True, facecolor="#ffd93a", edgecolor="#e53935", alpha=0.22, linewidth=2.0, zorder=10), dict(fill=False, edgecolor="#e53935", linewidth=1.8, zorder=11), ): patch = Polygon(corners, closed=True, **kwargs) self.ax.add_patch(patch) self._roi_artists.append(patch) self._roi_artists.append(self.ax.scatter( corners[:, 0], corners[:, 1], s=60, c="white", edgecolors="#e53935", linewidths=1.6, zorder=13)) # ------------------------------------------------------------------ # Hit testing (display pixels for handles, data coords for "inside") # ------------------------------------------------------------------ def _hit_test(self, event) -> tuple[str, int | None] | None: if self._roi is None or self.ax is None: return None if event.x is None or event.y is None: return None corners_disp = self.ax.transData.transform(self._roi.corners()) click = np.array([event.x, event.y]) for i in range(4): if np.hypot(*(corners_disp[i] - click)) <= self._HANDLE_PX: return ("corner", i) if event.xdata is not None and event.ydata is not None: if self._roi.contains(event.xdata, event.ydata): return ("inside", None) return None # ------------------------------------------------------------------ # Mouse event handlers # ------------------------------------------------------------------ def _on_press(self, event): if event.inaxes is not self.ax or self._extent is None: return if event.button != 1: # only left mouse button return # If the matplotlib toolbar is in pan / zoom mode, let it handle # the interaction instead of starting a ROI manipulation. tb = getattr(self, "toolbar", None) if tb is not None and getattr(tb, "mode", ""): return self._press_xy = (event.xdata, event.ydata) self._press_pixel = (event.x, event.y) self._press_button = event.button if self._draw_mode: self._draw_previous = self._roi.copy() if self._roi else None self._roi = RoiQuad.from_bbox(event.xdata, event.ydata, event.xdata, event.ydata) self._state = self._DRAW_NEW self._draw_roi() self.draw_idle() return hit = self._hit_test(event) if hit is None: self._state = self._IDLE return kind, idx = hit self._snapshot = self._roi.copy() if kind == "corner": self._state = self._DRAG_CORNER self._drag_corner_idx = idx else: self._state = self._MOVE self._move_anchor = (event.xdata, event.ydata) def _on_motion(self, event): if self._state == self._IDLE: return if event.xdata is None or event.ydata is None: return if event.inaxes is not self.ax: return if self._state == self._DRAW_NEW: x0, y0 = self._press_xy self._roi = RoiQuad.from_bbox(x0, y0, event.xdata, event.ydata) elif self._state == self._MOVE: delta = np.array([event.xdata - self._move_anchor[0], event.ydata - self._move_anchor[1]]) self._roi._pts = self._snapshot.corners() + delta elif self._state == self._DRAG_CORNER: self._roi._pts[self._drag_corner_idx] = [event.xdata, event.ydata] if self._rect_only: self._rectify_corner(self._drag_corner_idx) self._draw_roi() self.draw_idle() def _rectify_corner(self, idx: int): """Re-square the quad after a corner drag, anchored on the *opposite* corner. Anchoring on the diagonal opposite (idx ^ 2, since corners run BL, BR, TR, TL) rather than taking the bbox of all four points is what lets the rectangle shrink: a bbox over the three stale corners plus the new one is the union of the old rectangle and the new point, so dragging inward would never make it smaller. """ pts = self._roi.corners() ax_, ay = pts[idx ^ 2] bx, by = pts[idx] self._roi._pts = RoiQuad.from_bbox(min(ax_, bx), min(ay, by), max(ax_, bx), max(ay, by)).corners() def _on_release(self, event): if event.button != 1 and self._press_button != 1: return prev_state = self._state self._state = self._IDLE try: if prev_state == self._DRAW_NEW: self._finish_draw() elif prev_state in (self._MOVE, self._DRAG_CORNER): self._draw_roi() self.draw_idle() self.roi_changed.emit() else: self._maybe_emit_pixel_click(event) finally: self._press_xy = self._press_pixel = None self._press_button = None def _finish_draw(self): """Commit (or reject) a freshly-dragged quad.""" if self._extent is not None: x0, x1, y_bot, y_top = self._extent min_w = abs(x1 - x0) * 0.01 # minimum: 1% of each axis range min_h = abs(y_bot - y_top) * 0.01 else: min_w = min_h = 1e-6 if self._roi is None: too_small = True else: bbox = self._roi.bbox_size() too_small = bbox[0] < min_w or bbox[1] < min_h if too_small: self._roi = self._draw_previous self._draw_previous = None self.cancel_drawing() self._draw_roi() self.draw_idle() self.roi_changed.emit() def _maybe_emit_pixel_click(self, event): """A release close enough to its press counts as a pixel click.""" if (self._press_pixel is None or event.x is None or event.y is None or self._extent is None or event.inaxes is not self.ax or event.xdata is None): return dx_px = event.x - self._press_pixel[0] dy_px = event.y - self._press_pixel[1] if dx_px * dx_px + dy_px * dy_px > self._CLICK_THRESH_PX ** 2: return x0, x1, y_bot, y_top = self._extent n_rows, n_frames = self._img_shape col = int((event.xdata - x0) / (x1 - x0) * n_frames) row = int((event.ydata - y_top) / (y_bot - y_top) * n_rows) self.pixel_clicked.emit(max(0, min(row, n_rows - 1)), max(0, min(col, n_frames - 1))) class WaveformCanvas(FigureCanvasQTAgg): def __init__(self, parent=None): fig = Figure(figsize=(8, 3), tight_layout=True) self.ax_wave = fig.add_subplot(121) self.ax_right = fig.add_subplot(122) super().__init__(fig) self.setParent(parent) self.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding) def show_rf_waveform(self, sras: SrasFile, angle_idx: int, row_idx: int, frame_idx: int, apply_bg_sub: bool = True): """CH1 RF: time-domain + FFT spectrum. If apply_bg_sub is True and sras.background is not None, the background waveform is overlaid on the time-domain plot and the FFT is computed on the subtracted signal. The unsubtracted FFT is also shown faintly for comparison. """ data = sras.data[angle_idx] waveform = data[row_idx, CH1_IDX, frame_idx, :].astype(np.float32) t_ns = sras.time_axis_ns() f_mhz = sras.freq_axis_mhz() dc3_val = data[row_idx, CH3_IDX, frame_idx, :].astype(np.float32).mean() dc4_val = data[row_idx, CH4_IDX, frame_idx, :].astype(np.float32).mean() bg = sras.background if (apply_bg_sub and sras.background is not None) else None waveform_plot = waveform - bg if bg is not None else waveform self.ax_wave.cla() self.ax_right.cla() if bg is not None: self.ax_wave.plot(t_ns, waveform, linewidth=0.5, color="#aaaaaa", label="raw", zorder=1) self.ax_wave.plot(t_ns, bg, linewidth=0.5, color="#e07030", linestyle="--", label="background", zorder=2) self.ax_wave.plot(t_ns, waveform_plot, linewidth=0.7, color="#4488cc", label="subtracted", zorder=3) self.ax_wave.legend(fontsize=7, loc="upper right") else: self.ax_wave.plot(t_ns, waveform, linewidth=0.7, color="#4488cc") self.ax_wave.set_xlabel("Time (ns)") self.ax_wave.set_ylabel("ADC counts") bg_tag = " [bg sub]" if bg is not None else "" dc3_mv = adc_to_mv(dc3_val, *sras.cal(CH3_IDX)) dc4_mv = adc_to_mv(dc4_val, *sras.cal(CH4_IDX)) self.ax_wave.set_title( f"CH1 RF row={row_idx} frame={frame_idx}{bg_tag}\n" f"CH3={dc3_val:.1f} CH4={dc4_val:.1f} " f"({dc3_mv:.2f} / {dc4_mv:.2f} mV)", fontsize=8, ) # FFT of the (possibly subtracted) waveform power_sub = np.abs(np.fft.rfft(waveform_plot)) ** 2 power_sub[0] = 0.0 peak_mhz = f_mhz[int(np.argmax(power_sub))] if bg is not None: # Also show the unsubtracted FFT for reference power_raw = np.abs(np.fft.rfft(waveform)) ** 2 power_raw[0] = 0.0 self.ax_right.plot(f_mhz, power_raw, linewidth=0.5, color="#aaaaaa", label="raw FFT", zorder=1) self.ax_right.plot(f_mhz, power_sub, linewidth=0.7, color="#4488cc", label="subtracted FFT" if bg is not None else None, zorder=2) self.ax_right.axvline(peak_mhz, color="tomato", linestyle="--", linewidth=1.2, label=f"peak = {peak_mhz:.1f} MHz") self.ax_right.set_xlabel("Frequency (MHz)") self.ax_right.set_ylabel("Power (arb.)") self.ax_right.set_title("FFT Power Spectrum") self.ax_right.set_xlim(0, 500) self.ax_right.legend(fontsize=8) self.draw() def show_dc_waveform(self, sras: SrasFile, angle_idx: int, ch_idx: int, row_idx: int, frame_idx: int): """CH3 or CH4 DC: time-domain + mean annotation.""" waveform = sras.data[angle_idx][row_idx, ch_idx, frame_idx, :].astype(np.float32) mean_val = float(waveform.mean()) mean_mv = adc_to_mv(mean_val, *sras.cal(ch_idx)) self.ax_wave.cla() self.ax_right.cla() self.ax_wave.plot(sras.time_axis_ns(), waveform, linewidth=0.7, color="#4488cc") self.ax_wave.axhline(mean_val, color="tomato", linestyle="--", linewidth=1.2, label=f"mean = {mean_val:.2f} ADC") self.ax_wave.set_xlabel("Time (ns)") self.ax_wave.set_ylabel("ADC counts") self.ax_wave.set_title( f"{CH_NAMES[ch_idx]} DC row={row_idx} frame={frame_idx}") self.ax_wave.legend(fontsize=8) self.ax_right.text( 0.5, 0.5, f"DC mode\n\nmean = {mean_val:.3f} ADC\n = {mean_mv:.3f} mV", ha="center", va="center", transform=self.ax_right.transAxes, fontsize=11, ) self.ax_right.set_axis_off() self.draw() class AlignOverlayCanvas(FigureCanvasQTAgg): """Renders the alignment wizard's multi-angle mask views and turns keyboard input into translate/rotate nudge requests for whichever angle is active. Two views of the same reprojected masks, because they answer different questions. show_counts colours each pixel by *how many* angles cover it, which is the at-a-glance verdict on a correlation run: a good alignment is one saturated plateau, a bad one is a fringe of low-count halos. show_overlay gives each angle its own colour, which is what you need while nudging a specific angle by hand. A pure input+render widget — it holds no alignment state and never touches SrasFile itself; the wizard page owns all of that and decides, from these signals, whether a cheap single-layer refresh or a full preview-canvas rebuild is needed. FigureCanvasQTAgg is a real QWidget, so keyPressEvent works like on any other widget, but Qt only ever delivers key events to whichever widget currently has focus — StrongFocus, plus grabbing focus on click and once right after the dialog is shown, are both required or arrow keys silently do nothing. Rotate keys are letters (Q/E), not punctuation (comma/period or brackets): Shift+letter still reports the same Qt.Key on every platform, whereas Shift+comma/bracket can report a different virtual key (Key_Less / Key_BraceLeft) depending on platform and keyboard layout — which would silently break the "Shift = coarse step" modifier for rotation specifically. Arrow keys have no such hazard. """ nudge_translate = pyqtSignal(int, int, bool) # dir_x, dir_y in {-1,0,1}; coarse nudge_rotate = pyqtSignal(int, bool) # dir in {-1,1} (CCW/CW); coarse _TRANSLATE_KEYS = { Qt.Key.Key_Left: (-1, 0), Qt.Key.Key_Right: (1, 0), Qt.Key.Key_Up: (0, -1), Qt.Key.Key_Down: (0, 1), } _ROTATE_KEYS = {Qt.Key.Key_Q: 1, Qt.Key.Key_E: -1} # CCW, CW def __init__(self, parent=None): fig = Figure(figsize=(6, 6), tight_layout=True) self.ax = fig.add_subplot(111) super().__init__(fig) self.setParent(parent) self.setFocusPolicy(Qt.FocusPolicy.StrongFocus) self.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding) self.mpl_connect("button_press_event", lambda _e: self.setFocus()) def show_overlay(self, rgba: np.ndarray, extent: list[float], title: str): self.figure.clf() self.ax = self.figure.add_subplot(111) self.ax.imshow(rgba, extent=extent, origin="upper", aspect="auto") self._finish(title) def show_counts(self, counts: np.ndarray, n_angles: int, extent: list[float], title: str): """The mask stack coloured by how many angles cover each pixel. A discrete colormap with integer-ticked colorbar rather than a continuous one: the judgement being made is "is this a plateau at N, or a fan of partial overlaps", and a region covered by one angle too few has to read as its own band rather than a slightly darker shade. Uncovered pixels are transparent so they cannot be mistaken for a low count. """ self.figure.clf() self.ax = self.figure.add_subplot(111) cmap, norm, ticks = count_colormap(n_angles) im = self.ax.imshow( np.asarray(counts), extent=extent, origin="upper", aspect="auto", interpolation="nearest", cmap=cmap, norm=norm) cb = self.figure.colorbar(im, ax=self.ax, fraction=0.046, pad=0.04, ticks=ticks) cb.set_label("angles overlapping") self._finish(title) def _finish(self, title: str): self.ax.set_xlabel("X (mm)") self.ax.set_ylabel("Y (mm)") self.ax.set_title(title) self.draw_idle() # coalesces rapid redraws — matters for key-repeat. def keyPressEvent(self, event: QKeyEvent): key = event.key() coarse = bool(event.modifiers() & Qt.KeyboardModifier.ShiftModifier) if key in self._TRANSLATE_KEYS: dx, dy = self._TRANSLATE_KEYS[key] self.nudge_translate.emit(dx, dy, coarse) event.accept() elif key in self._ROTATE_KEYS: self.nudge_rotate.emit(self._ROTATE_KEYS[key], coarse) event.accept() else: super().keyPressEvent(event)