From d5028db445d5ddbb5ade8fcb00ceaa98893372c7 Mon Sep 17 00:00:00 2001 From: Thomas Ales Date: Thu, 6 Aug 2026 09:35:30 -0500 Subject: [PATCH 01/10] Add scan-editing CLI and alignment tests; extend Manual Alignment correlation Continues the Manual Alignment work: refines the FFT cross-correlation and mask handling, adds sras_edit_scans.py (drop/renumber bad angle scans), tools/test_alignment.py (registration ground-truth suite), and a rotating test fixture in make_test_sras.py. Co-Authored-By: Claude Fable 5 --- sras_compute.py | 1164 +++++++++++++++++++++------------- sras_edit_scans.py | 295 +++++++++ sras_viewer.py | 204 +++--- sras_viewer_requirements.txt | 8 + sras_workers.py | 61 +- tools/make_test_sras.py | 132 ++++ tools/test_alignment.py | 223 +++++++ tools/test_gui.py | 129 ++-- 8 files changed, 1619 insertions(+), 597 deletions(-) create mode 100644 sras_edit_scans.py create mode 100644 tools/test_alignment.py diff --git a/sras_compute.py b/sras_compute.py index ee34166..28c8b3d 100644 --- a/sras_compute.py +++ b/sras_compute.py @@ -364,28 +364,50 @@ def cache_file(path: str, mode: str, apply_bg_sub: bool, # Angle alignment (Fusion menu) # # Puts every angle's images onto one shared, zero-padded pixel grid using a -# rigid transform only (rotation + translation, never scale). Rotation for -# angle `a` is the *known* scan-angle delta relative to a reference angle — -# never searched. Only the residual translation is found, via FFT phase -# correlation of each angle's binarized CH4 ("dc-mask") image. +# rigid transform only — rotation + translation, never scale. # -# Rotation is done in physical mm space rather than on raw pixel indices: -# the x-pixel pitch (SrasFile.pixel_x_mm) is file-wide constant but the -# y-pixel pitch (row spacing) can differ from it, and for v6 files can even -# vary per angle. Rotating the raw index grid directly would implicitly -# assume square pixels and shear a non-square-pixel image — an unwanted -# effective anisotropic scale. Instead each angle gets one affine that maps -# shared-canvas pixel index -> mm -> undo rotation/shift -> that angle's own -# local mm -> that angle's own raw pixel index, matching the output->input -# convention scipy.ndimage.affine_transform expects. +# Angle 0 (the reference) is the sole coordinate authority: it is the only +# angle whose stage XY (x_start_mm / y_positions_mm) is ever read, and the +# shared canvas is literally an extension of angle 0's own pixel grid, so the +# aligned view carries angle 0's real X/Y axes. Every *other* angle is placed +# purely by content — its rotation and translation come from cross-correlating +# its CH4 image against angle 0's (register_angle_to_reference) — and its own +# stage XY is deliberately never consulted. That is not an oversight: the +# rotation stage moves the sample relative to the scan window, so where a +# window sat in stage coordinates says nothing about where the sample is, and +# an earlier design that pivoted each angle on a signal-weighted centroid of +# its own window put every angle on a ~20 mm circle around the optical center +# instead of stacking them into one shape. +# +# Only two coordinate frames exist here: +# +# local mm — one angle's own physical frame: origin at the *center of its own +# pixel array*, x along +column, y along +row, scaled by that angle's own +# pitches. Carries no stage position whatsoever. +# +# ref mm — the reference angle's local mm. A registration result +# (rotation_deg, shift_mm) is exactly the rigid map from an angle's local +# mm to ref mm: q = R(rotation_deg) @ l + shift_mm. Stage coordinates +# re-enter once, at the very end, when the canvas origin is converted to +# angle 0's stage mm (AlignmentResult.canvas_origin_mm). +# +# Rotation is done in mm, never on raw pixel indices: the x pitch +# (SrasFile.pixel_x_mm, 5 µm on a real scan) and the y/row pitch (50 µm) differ +# by 10x, so rotating the raw index grid would shear the image — an unwanted +# anisotropic scale. Registration runs on a resampled *isotropic* grid for the +# same reason, and every affine here maps shared-grid index -> mm -> undo +# rotation/shift -> that angle's own local mm -> that angle's own raw index, +# matching the output->input convention scipy.ndimage.affine_transform wants. # --------------------------------------------------------------------------- @dataclass class AngleTransform: rotation_deg: float - shift_mm: tuple[float, float] # (dx_mm, dy_mm) found by phase correlation + shift_mm: tuple[float, float] # (dx_mm, dy_mm) in ref mm matrix: np.ndarray # (2,2): canvas (row,col) -> this angle's raw (row,col) offset: np.ndarray # (2,) + score: float = 1.0 # registration NCC (1.0 = reference/manual) + source: str = "" # image that won registration: "signal"/"mask" @dataclass @@ -403,36 +425,90 @@ class AlignmentResult: class ManualAngleParams: """One angle's manual-alignment state, independent of any canvas. - rotation_deg/shift_mm are exactly AngleTransform's non-derived fields — - the pair a canvas-bound AngleTransform's matrix/offset get built from - once a canvas is decided (build_manual_alignment). Defaults to identity - (no rotation, no shift): a fresh angle with no prior alignment is shown - raw, exactly as scanned — the same "fully unaligned" state Clear - Alignment resets back to. + rotation_deg/shift_mm are exactly AngleTransform's non-derived fields — the + rigid map from this angle's local mm to ref mm, the pair a canvas-bound + AngleTransform's matrix/offset get built from once a canvas is decided + (build_manual_alignment). Defaults to identity: a fresh angle with no prior + alignment is shown centered on the reference with no rotation, which is the + same "fully unaligned" state Clear Alignment resets back to. """ rotation_deg: float = 0.0 shift_mm: tuple[float, float] = (0.0, 0.0) +# Side length of the grid the rotation refinement runs on. The whole cost of +# registration scales with it: ~325 ms per candidate rotation at 640, roughly +# quadrupling per doubling, against ~200 MB of live masked-correlation buffers +# (see _registration_workers). 640 puts a full-size scan at ~0.1 mm/px, which +# resolves rotation on a 20 mm sample to well under a tenth of a degree. +_DEFAULT_FINE_DIM = 640 + + +@dataclass +class RigidFit: + """What register_angle_to_reference found for one angle.""" + rotation_deg: float + shift_mm: tuple[float, float] + score: float # zero-mean NCC over the valid overlap + source: str # "signal", "mask", or "reference" + + +def _skimage_phase_cross_correlation(): + """Import skimage.registration lazily and cache it. + + Deliberately not a module-level import: this module is imported by every + multiprocessing child (see the module docstring), skimage costs ~0.6 s to + import, and no child ever registers anything — registration runs in GUI- + process threads. + """ + global _pcc + try: + return _pcc + except NameError: + from skimage.registration import phase_cross_correlation as _fn + _pcc = _fn + return _pcc + + +# ---- Geometry: local mm, ref mm, and the one affine builder --------------- + def _pixel_pitch_mm(sras: SrasFile, angle_idx: int) -> tuple[float, float]: """(dx, dy) mm/pixel for one angle: dx is the file-wide constant pixel_x_mm; dy is this angle's own row spacing (assumed uniform, the same - assumption _redraw_image makes when it builds the display extent).""" + assumption _redraw_image makes when it builds the display extent). dy keeps + its sign, so +row always means the same physical direction as +y.""" y = sras.y_positions_mm(angle_idx) return sras.pixel_x_mm, (float(y[1] - y[0]) if len(y) > 1 else 1.0) -def _bbox_center_mm(sras: SrasFile, angle_idx: int) -> tuple[float, float]: - x = sras.x_axis_mm(angle_idx) - y = sras.y_positions_mm(angle_idx) - return float((x[0] + x[-1]) / 2.0), float((y[0] + y[-1]) / 2.0) +def _center_idx(sras: SrasFile, angle_idx: int) -> np.ndarray: + """(row, col) index of this angle's array center — the origin of its local + mm frame. Purely geometric: it depends on the array shape and nothing + else, which is what keeps local mm free of stage position.""" + n_rows, n_frames = sras.image_shape(angle_idx) + return np.array([(n_rows - 1) / 2.0, (n_frames - 1) / 2.0]) -def _bbox_corners_mm(sras: SrasFile, angle_idx: int) -> np.ndarray: - """4 corners (x, y) of this angle's raw mm bounding box, shape (4, 2).""" - x = sras.x_axis_mm(angle_idx) - y = sras.y_positions_mm(angle_idx) - return np.array([[xx, yy] for xx in (x[0], x[-1]) for yy in (y[0], y[-1])]) +def ref_center_mm(sras: SrasFile, ref_angle_idx: int) -> np.ndarray: + """Stage mm of the reference angle's array center: ref mm + this == stage + mm. + + The single bridge between ref mm and stage mm, and the only place in the + whole alignment path where any angle's stage position is read at all — + which is why it takes the *reference* index by name rather than an + arbitrary angle. + """ + x = sras.x_axis_mm(ref_angle_idx) + y = sras.y_positions_mm(ref_angle_idx) + return np.array([(x[0] + x[-1]) / 2.0, (y[0] + y[-1]) / 2.0], dtype=np.float64) + + +def _local_half_extent_mm(sras: SrasFile, angle_idx: int) -> tuple[float, float]: + """(half width, half height) in mm from this angle's array center to the + center of its outermost pixel.""" + n_rows, n_frames = sras.image_shape(angle_idx) + dx, dy = _pixel_pitch_mm(sras, angle_idx) + return (n_frames - 1) / 2.0 * abs(dx), (n_rows - 1) / 2.0 * abs(dy) def _rotation_matrix(theta_deg: float) -> np.ndarray: @@ -441,152 +517,63 @@ def _rotation_matrix(theta_deg: float) -> np.ndarray: return np.array([[c, -s], [s, c]]) # CCW rotation acting on (x, y) -def _theta_deg(sras: SrasFile, angle_idx: int, ref_idx: int) -> float: - """CCW rotation, in degrees and in _rotation_matrix's convention, that - maps angle_idx's own local mm frame onto ref_idx's. +def _nominal_delta_deg(sras: SrasFile, angle_idx: int, ref_idx: int) -> float: + """The rotation-stage's own reported angle change between two angles. - This is the *negative* of the raw angles_deg delta: the GR rotation - stage's reported angle increases in the opposite rotational sense from - this module's math-positive (CCW, x toward y) convention in scan mm - space. Rotating by +(angles_deg[a] - angles_deg[ref]) therefore turns - misalignment the wrong way — confirmed empirically (Auto De-rotate made - real scans worse, not better, before this negation). + Used only to *seed* the rotation search, never as the answer: the stage's + sign convention relative to this module's math-positive (CCW, x toward y) + convention in scan mm is not knowable from the file, so + register_angle_to_reference scores both +this and -this and lets the image + content decide (see _rotation_candidates). """ - return -float(sras.angles_deg[angle_idx] - sras.angles_deg[ref_idx]) + return float(sras.angles_deg[angle_idx] - sras.angles_deg[ref_idx]) -def _signal_centroid_mm(sras: SrasFile, angle_idx: int, - dc4_mv: np.ndarray) -> tuple[float, float]: - """Intensity-weighted centroid (mean x, mean y, weighted by CH4 signal - after subtracting this angle's own minimum) in angle_idx's own local mm - frame — the alignment pivot used in place of the raw scan-window bbox - center (see compute_pivot_points_mm). - - Weighting by the continuous DC signal, rather than a binary >= - dc_threshold_mv mask, means the pivot never depends on how well one - shared threshold happens to suit this particular angle: real signal - levels vary scan to scan, so a threshold tuned for one angle can leave - another angle's binary mask empty — and a centroid of an empty mask has - nothing to fall back to *except* the raw bbox center, silently - reproducing the exact "aligned to the scan window, not the sample" - problem this pivot exists to avoid. Falls back to the bbox center only - in the fully-degenerate case of a perfectly flat signal (nothing to - weight by at all). - """ - weights = dc4_mv - dc4_mv.min() - total = float(weights.sum()) - if total <= 0.0: - return _bbox_center_mm(sras, angle_idx) - x = sras.x_axis_mm(angle_idx) - y = sras.y_positions_mm(angle_idx) - cx = float((x * weights.sum(axis=0)).sum() / total) - cy = float((y * weights.sum(axis=1)).sum() / total) - return cx, cy +def _footprint_corners_ref_mm(sras: SrasFile, angle_idx: int, + rotation_deg: float, + shift_mm: tuple[float, float]) -> np.ndarray: + """This angle's 4 footprint corners mapped into ref mm by its rigid + transform, shape (4, 2). Zero-padding cost, not content: it is the scan + window's corners, which is what the shared canvas has to cover.""" + hw, hh = _local_half_extent_mm(sras, angle_idx) + corners = np.array([[sx * hw, sy * hh] for sx in (-1.0, 1.0) for sy in (-1.0, 1.0)]) + R = _rotation_matrix(rotation_deg) + return corners @ R.T + np.asarray(shift_mm, dtype=np.float64) -def compute_pivot_points_mm(sras: SrasFile, - dc4_mv: dict[int, np.ndarray] | None = None - ) -> dict[int, tuple[float, float]]: - """Per-angle alignment pivot, in each angle's own local mm frame: the - CH4-signal-weighted centroid of its own footprint (see - _signal_centroid_mm), rather than the raw scan-window bbox center. - - Pivoting on each angle's own content — instead of on wherever its - scanned window happened to sit in microscope/global XY space — is what - makes alignment purely relative *between scans* rather than to a global - coordinate system: the rotation stage's true mechanical axis need not - coincide with the scan window's geometric center, and the sample need - not be perfectly centered on that axis either, so a bbox-center pivot - leaves a residual orbital motion between angles that a content-centroid - pivot does not. - - Deliberately independent of dc_threshold_mv (the RF-mask / overlay - threshold): that value is a display/masking choice and must never - silently change where alignment pivots. - - *dc4_mv* lets a caller that has already computed each angle's CH4 mV - image (compute_angle_alignment's Step 1, or ManualAlignmentDialog's own - cache) reuse it instead of recomputing; angles missing from it are - computed fresh via dc_image_mv, which prefers a stored v5/v7 cache over - recomputing from raw waveforms. - """ - dc4_mv = dc4_mv or {} - pivots: dict[int, tuple[float, float]] = {} - for a in range(sras.n_angles): - img = dc4_mv.get(a) - if img is None: - img = dc_image_mv(sras, a, CH4_IDX) - pivots[a] = _signal_centroid_mm(sras, a, img) - return pivots - - -def _corners_in_ref_frame(sras: SrasFile, angle_idx: int, ref_idx: int, - pivot_mm: dict[int, tuple[float, float]], - shift_mm=(0.0, 0.0), - theta_deg: float | None = None) -> np.ndarray: - """Angle *angle_idx*'s bbox corners, rotated about its own alignment - pivot (pivot_mm[angle_idx] — see compute_pivot_points_mm) into the - reference frame and translated by *shift_mm*. Shape (4, 2). - - theta_deg overrides the analytic angles_deg-derived rotation used by - default — the manual-alignment path (union_canvas_mm) passes a - user-chosen rotation here (which may differ from the known scan-angle - delta) without needing a parallel code path. - """ - theta = _theta_deg(sras, angle_idx, ref_idx) if theta_deg is None else theta_deg - R = _rotation_matrix(theta) - c_a = np.array(pivot_mm[angle_idx]) - c_ref = np.array(pivot_mm[ref_idx]) - shift = np.asarray(shift_mm, dtype=np.float64) - return np.array([R @ (corner - c_a) + c_ref + shift - for corner in _bbox_corners_mm(sras, angle_idx)]) - - -def _build_affine_canvas_to_raw(sras: SrasFile, angle_idx: int, ref_idx: int, - shift_mm: tuple[float, float], - canvas_dx: float, canvas_dy: float, - canvas_origin_mm: tuple[float, float], - pivot_mm: dict[int, tuple[float, float]], - theta_deg: float | None = None - ) -> tuple[np.ndarray, np.ndarray]: - """matrix, offset s.t. raw_index = matrix @ [row_out, col_out] + offset, +def _affine_out_to_src(*, out_pitch_mm: tuple[float, float], + out_origin_ref_mm, src_dx_mm: float, src_dy_mm: float, + src_center_idx, src_center_off_mm=(0.0, 0.0), + rotation_deg: float = 0.0, + shift_mm: tuple[float, float] = (0.0, 0.0) + ) -> tuple[np.ndarray, np.ndarray]: + """matrix, offset s.t. src_index = matrix @ [row_out, col_out] + offset, matching scipy.ndimage.affine_transform's output->input convention. - Pipeline (all mm unless noted): - [X;Y] = A_out @ [row_out;col_out] + b_out # canvas idx -> ref-frame mm - [lx;ly] = R(theta)^T @ ([X;Y]-c_ref-shift) + c_a # undo rotation+shift -> angle a's local mm - [row;col] = D @ ([lx;ly] - [x_start_a; y0_a]) # local mm -> angle a's raw idx + The single affine builder behind every resampling in this module — the + registration grid, the manual-alignment preview and the final canvas all + differ only in their arguments. - where theta is _theta_deg(angle_idx, ref_idx) unless *theta_deg* - overrides it (see _corners_in_ref_frame's note, used by the manual- - alignment path), c_ref/c_a are each angle's own alignment pivot - (pivot_mm — see compute_pivot_points_mm; the CH4-signal-weighted - centroid of its own footprint, not the raw scan-window bbox center — - this keeps the sample itself centered post-rotation, minimizing - required canvas padding and, more importantly, keeping alignment - relative to the sample rather than to wherever the scan window sat in - microscope/global XY space), and A_out/D are the index<->mm scaling - matrices for the canvas pitch and this angle's own native pitch - respectively. + Pipeline (mm unless noted): + [X;Y] = A_out @ [row_out;col_out] + out_origin_ref_mm # out idx -> ref mm + [lx;ly] = R(rotation)^T @ ([X;Y] - shift) # ref mm -> this angle's local mm + [row;col] = D @ ([lx;ly] - src_center_off) + src_center_idx + + *src_center_off_mm* is the local-mm position of the source array's own + center, nonzero only when the source has been block-mean downsampled (its + center can land up to half a block off the full-resolution center — see + _prepare_reg_image). Everything is expressed relative to array centers, so + no per-angle stage coordinate appears anywhere in here. """ - theta = _theta_deg(sras, angle_idx, ref_idx) if theta_deg is None else theta_deg - Rinv = _rotation_matrix(theta).T - cx_a, cy_a = pivot_mm[angle_idx] - cx_ref, cy_ref = pivot_mm[ref_idx] - dx_a, dy_a = _pixel_pitch_mm(sras, angle_idx) - x0_a = float(sras.x_start_mm[angle_idx]) - y0_a = float(sras.y_positions_mm(angle_idx)[0]) - - A_out = np.array([[0.0, canvas_dx], [canvas_dy, 0.0]]) # [row,col] -> [X,Y] - b_out = np.array(canvas_origin_mm, dtype=np.float64) - D = np.array([[0.0, 1.0 / dy_a], [1.0 / dx_a, 0.0]]) # [x,y] -> [row,col] - shift = np.array(shift_mm, dtype=np.float64) - c_ref_v = np.array([cx_ref, cy_ref]) - c_a_v = np.array([cx_a, cy_a]) - origin_a = np.array([x0_a, y0_a]) + Rinv = _rotation_matrix(rotation_deg).T + A_out = np.array([[0.0, out_pitch_mm[0]], [out_pitch_mm[1], 0.0]]) # [row,col] -> [x,y] + D = np.array([[0.0, 1.0 / src_dy_mm], [1.0 / src_dx_mm, 0.0]]) # [x,y] -> [row,col] + b_out = np.asarray(out_origin_ref_mm, dtype=np.float64) + shift = np.asarray(shift_mm, dtype=np.float64) + off = np.asarray(src_center_off_mm, dtype=np.float64) matrix = D @ Rinv @ A_out - offset = D @ Rinv @ (b_out - c_ref_v - shift) + D @ (c_a_v - origin_a) + offset = D @ (Rinv @ (b_out - shift) - off) + np.asarray(src_center_idx, dtype=np.float64) return matrix, offset @@ -596,7 +583,7 @@ def apply_alignment(result: AlignmentResult, angle_idx: int, img: np.ndarray, that angle's raw (n_rows, n_frames)) onto the shared alignment canvas. order=0 (nearest) avoids blending real data with zero-padding or with masked-out (0-valued) CH1/velocity pixels at mask edges. Channel- - agnostic: the same per-angle transform (found from the CH4 mask) works + agnostic: the same per-angle transform (found from the CH4 image) works for any channel's image of that angle.""" t = result.per_angle[angle_idx] return scipy_ndimage.affine_transform( @@ -605,122 +592,533 @@ def apply_alignment(result: AlignmentResult, angle_idx: int, img: np.ndarray, mode="constant", cval=0.0) -def _block_mean_downsample(img: np.ndarray, factor: int) -> np.ndarray: - if factor <= 1: +def _block_mean_2d(img: np.ndarray, fy: int, fx: int) -> np.ndarray: + """Block-mean by independent row/column factors. Independent factors matter + because the raw grid is strongly anisotropic (5 µm along x, 50 µm along y): + a single square factor would either alias along x or throw away rows.""" + fy, fx = max(1, int(fy)), max(1, int(fx)) + if fy == 1 and fx == 1: return img h, w = img.shape - h2, w2 = (h // factor) * factor, (w // factor) * factor + h2, w2 = (h // fy) * fy, (w // fx) * fx + if h2 == 0 or w2 == 0: + return img trimmed = img[:h2, :w2] - return trimmed.reshape(h2 // factor, factor, w2 // factor, factor).mean(axis=(1, 3)) + return trimmed.reshape(h2 // fy, fy, w2 // fx, fx).mean(axis=(1, 3)) -def _phase_correlate_shift(ref_img: np.ndarray, mov_img: np.ndarray) -> tuple[int, int]: - """FFT normalized cross-power-spectrum phase correlation. Returns the - integer (dr, dc) pixel shift of mov_img relative to ref_img; both must - be the same shape. Risk: if the true shift is near +/- half the array - size, wraparound can bias the peak — mitigated by the generous - margin_frac padding in _working_canvas_for_pair, which keeps the true - residual shift small relative to the correlation canvas.""" - F1 = scipy_fft.fft2(ref_img.astype(np.float64), workers=-1) - F2 = scipy_fft.fft2(mov_img.astype(np.float64), workers=-1) - R = F1 * np.conj(F2) - R /= np.maximum(np.abs(R), 1e-12) - corr = scipy_fft.ifft2(R, workers=-1).real - dr, dc = np.unravel_index(np.argmax(corr), corr.shape) - h, w = corr.shape - if dr > h // 2: - dr -= h - if dc > w // 2: - dc -= w - return int(dr), int(dc) +# ---- Registration: rigid (rotation + translation) fit against the reference +# +# Everything below works on _RegImage, an angle's image resampled to a shared +# *isotropic* grid centered on its own array center. Nothing in here can see a +# stage coordinate even in principle, which is the point: the fit is decided by +# image content alone. + +@dataclass +class _RegImage: + """One angle's image prepared for registration: block-mean downsampled to + roughly the registration pitch, carrying the physical pitch it ended up + with and the local-mm offset of its own array center from the + full-resolution array center (block-mean trims a partial trailing block, so + the two centers can differ by up to half a block).""" + img: np.ndarray + dx_mm: float + dy_mm: float + center_off_mm: tuple[float, float] -def _working_canvas_for_pair(sras: SrasFile, ref_idx: int, a_idx: int, - dx: float, dy: float, - pivot_mm: dict[int, tuple[float, float]], - margin_frac: float = 0.3 - ) -> tuple[tuple[float, float], tuple[int, int]]: - """Union of the reference's own raw bbox and angle a's raw bbox rotated - (about its own alignment pivot) into the ref frame with zero shift, - padded by margin_frac on each side — sized generously so the true - phase-correlation shift lands well inside the canvas (see - _phase_correlate_shift's wraparound note).""" - pts = np.vstack([_bbox_corners_mm(sras, ref_idx), - _corners_in_ref_frame(sras, a_idx, ref_idx, pivot_mm)]) - x_min, y_min = pts.min(axis=0) - x_max, y_max = pts.max(axis=0) - pad_x, pad_y = (x_max - x_min) * margin_frac, (y_max - y_min) * margin_frac - x_min, x_max = x_min - pad_x, x_max + pad_x - y_min, y_max = y_min - pad_y, y_max + pad_y - n_cols = int(np.ceil((x_max - x_min) / dx)) + 1 - n_rows = int(np.ceil((y_max - y_min) / abs(dy))) + 1 - origin = (x_min, y_min if dy > 0 else y_max) - return origin, (n_rows, n_cols) +def _block_center_off(n_full: int, n_small: int, factor: int, + pitch_mm: float) -> float: + """Local-mm offset of a block-mean-downsampled array's own center from the + full-resolution array's center along one axis. + + Small pixel k averages full pixels [k*f, k*f + f - 1], so its center sits + at full index k*f + (f-1)/2; block-mean also trims a partial trailing + block. Both effects together move the small array's center by up to half a + block, which has to be accounted for or a downsampled image registers (or + previews) at a systematically shifted position. + """ + return (((n_small - 1) / 2.0) * factor + (factor - 1) / 2.0 + - (n_full - 1) / 2.0) * pitch_mm -def correlate_translation_mm(sras: SrasFile, angle_idx: int, ref_angle_idx: int, - signal_mv: dict[int, np.ndarray], - pivot_mm: dict[int, tuple[float, float]], *, - use_mask: bool = False, - dc_threshold_mv: float = 0.0, - margin_frac: float = 0.3, - max_corr_dim: int = 1024) -> tuple[float, float]: - """FFT phase-correlation translation (shift_x_mm, shift_y_mm) that best - lines up angle_idx's CH4 content onto ref_angle_idx's, given each angle's - own alignment pivot (pivot_mm — see compute_pivot_points_mm) already - rotated about that pivot by the analytic scan-angle delta with zero - shift. Returns (0.0, 0.0) unconditionally for the reference angle. +def _prepare_reg_image(sras: SrasFile, angle_idx: int, img: np.ndarray, + pitch_mm: float) -> _RegImage: + """Block-mean an angle's image down to roughly *pitch_mm* before it is + resampled onto the registration grid. Pre-averaging matters: the raw grid + is 10x finer along x than along y, so sampling it directly at the (much + coarser) isotropic registration pitch would alias badly along x.""" + dx, dy = _pixel_pitch_mm(sras, angle_idx) + fx = max(1, int(pitch_mm / abs(dx))) + fy = max(1, int(pitch_mm / abs(dy))) + small = _block_mean_2d(np.asarray(img, dtype=np.float32), fy, fx) + n_rows, n_frames = img.shape + return _RegImage( + small, dx * fx, dy * fy, + (_block_center_off(n_frames, small.shape[1], fx, dx), + _block_center_off(n_rows, small.shape[0], fy, dy))) - By default (use_mask=False) correlates on each angle's own raw CH4 - signal minus its own minimum — the same weighting _signal_centroid_mm - uses — rather than a dc_threshold_mv binary mask: a shared threshold - that doesn't suit every angle's real signal level can empty or distort - one angle's mask (the same failure mode the alignment pivot was fixed - to avoid), and correlating on the raw signal also lets phase - correlation lock onto real internal sample structure rather than just - the scan window's silhouette. Subtracting each angle's own minimum - (rather than using the signal as-is) keeps the zero-padding surrounding - the rotated image from reading as a spurious high-contrast edge against - a nonzero DC baseline. use_mask=True switches to the binary - >= dc_threshold_mv mask instead (compute_angle_alignment's original - behavior), for cases where raw-signal correlation locks onto noise. + +def _embed(reg: _RegImage, pitch_mm: float, n: int, rotation_deg: float, + order: int) -> np.ndarray: + """Resample a _RegImage onto the shared n x n isotropic registration grid, + rotated by *rotation_deg* about the grid center and with no translation + (translation is what phase correlation then measures).""" + half = (n - 1) / 2.0 * pitch_mm + matrix, offset = _affine_out_to_src( + out_pitch_mm=(pitch_mm, pitch_mm), out_origin_ref_mm=(-half, -half), + src_dx_mm=reg.dx_mm, src_dy_mm=reg.dy_mm, + src_center_idx=((reg.img.shape[0] - 1) / 2.0, (reg.img.shape[1] - 1) / 2.0), + src_center_off_mm=reg.center_off_mm, rotation_deg=rotation_deg) + return scipy_ndimage.affine_transform( + reg.img, matrix, offset=offset, output_shape=(n, n), order=order, + mode="constant", cval=0.0) + + +def _embed_with_valid(reg: _RegImage, pitch_mm: float, n: int, + rotation_deg: float) -> tuple[np.ndarray, np.ndarray]: + """Embedded image plus the boolean mask of where that angle actually has + data. The valid mask is what lets registration ignore each angle's + differently-shaped scan window instead of locking onto its silhouette.""" + img = _embed(reg, pitch_mm, n, rotation_deg, order=1) + ones = _RegImage(np.ones_like(reg.img), reg.dx_mm, reg.dy_mm, reg.center_off_mm) + valid = _embed(ones, pitch_mm, n, rotation_deg, order=0) > 0.5 + return img, valid + + +def _shift_into(img: np.ndarray, dr: int, dc: int) -> np.ndarray: + """img translated by whole pixels with zero fill (never wrapping, unlike + np.roll — wrapped content would score as a spurious match).""" + out = np.zeros_like(img) + h, w = img.shape + sr0, sr1 = max(0, dr), min(h, h + dr) + sc0, sc1 = max(0, dc), min(w, w + dc) + if sr0 >= sr1 or sc0 >= sc1: + return out + out[sr0:sr1, sc0:sc1] = img[sr0 - dr:sr1 - dr, sc0 - dc:sc1 - dc] + return out + + +def _overlap_ncc(ref: np.ndarray, ref_valid: np.ndarray, + mov: np.ndarray, mov_valid: np.ndarray, + min_overlap_frac: float = 0.15) -> float: + """Zero-mean normalized cross-correlation over the two images' common valid + region — the score every rotation candidate is ranked by. + + Computed on the overlap only, and rejected outright (-1) when the overlap + is too small a fraction of the smaller footprint: without that floor a + candidate that slides the angles almost entirely apart can win on a handful + of coincidentally-similar pixels. + """ + both = ref_valid & mov_valid + n = int(both.sum()) + smaller = min(int(ref_valid.sum()), int(mov_valid.sum())) + if smaller == 0 or n < min_overlap_frac * smaller or n < 16: + return -1.0 + a = ref[both].astype(np.float64) + b = mov[both].astype(np.float64) + a -= a.mean() + b -= b.mean() + denom = np.sqrt((a * a).sum() * (b * b).sum()) + return float((a * b).sum() / denom) if denom > 0 else -1.0 + + +def _masked_shift(ref: np.ndarray, ref_valid: np.ndarray, + mov: np.ndarray, mov_valid: np.ndarray) -> tuple[int, int]: + """Integer (dr, dc) that best registers *mov* onto *ref*, from skimage's + masked FFT phase correlation (Padfield). The masked variant is the whole + reason scikit-image is a dependency: plain phase correlation on these + images locks onto the scan window's rectangular silhouette, which differs + per angle, instead of onto the sample.""" + pcc = _skimage_phase_cross_correlation() + result = pcc(ref, mov, reference_mask=ref_valid, moving_mask=mov_valid) + shift = result[0] if isinstance(result, tuple) else result + return int(round(float(shift[0]))), int(round(float(shift[1]))) + + +def _subpixel_residual(ref: np.ndarray, mov: np.ndarray, + both: np.ndarray) -> tuple[float, float]: + """Sub-pixel leftover shift between two already integer-aligned images, + from upsampled phase correlation over their common valid region. + + A separate pass because skimage's *masked* phase correlation has no + upsample_factor; here both images are zeroed outside the shared overlap and + mean-subtracted inside it, so the plain upsampled version is well posed. + Clamped to ±1 px: this only ever polishes an already-good integer fit, and + a larger "residual" means the peak was spurious. + + normalization=None (plain cross-correlation, not phase correlation) is + load-bearing. Whitening the spectrum is what makes phase correlation good + at finding a large unknown shift, but here the two images are already + aligned to within a pixel and the masked-off surroundings put a hard edge + in both: whitened, that edge and the high-frequency noise swamp the true + sub-pixel peak and the default returns a flat zero every time. + """ + if not both.any(): + return 0.0, 0.0 + a = np.zeros_like(ref) + b = np.zeros_like(mov) + a[both] = ref[both] - ref[both].mean() + b[both] = mov[both] - mov[both].mean() + if not (np.any(a) and np.any(b)): + return 0.0, 0.0 + pcc = _skimage_phase_cross_correlation() + result = pcc(a, b, upsample_factor=20, normalization=None) + shift = result[0] if isinstance(result, tuple) else result + dr, dc = float(shift[0]), float(shift[1]) + if abs(dr) > 1.0 or abs(dc) > 1.0: + return 0.0, 0.0 + return dr, dc + + +def _reg_pitch_and_size(sras: SrasFile, max_dim: int, + margin: float = 1.25) -> tuple[float, int]: + """Isotropic pitch (mm/px) and side length for the shared registration + grid: square, big enough for the largest angle's footprint at any rotation + (hence its diagonal) plus *margin* headroom for the translation search. + + The floor on pitch is the *geometric mean* of the two native pitches, not + the coarser of them. The raw grid is anisotropic (5 µm along x, 50 µm along + y): flooring at 50 µm would throw away all the extra x detail, and rotation + precision depends directly on it — a feature at radius r moves by r·δθ, so + at 50 µm a 20 mm-wide sample can only resolve rotation to a few tenths of a + degree. The geometric mean interpolates y up by ~3x rather than discarding + x, which costs a little memory and buys real angular precision. On a + full-size scan max_dim binds first and this floor never applies at all. + """ + diag = max(float(np.hypot(*(2 * v for v in _local_half_extent_mm(sras, a)))) + for a in range(sras.n_angles)) + span = diag * margin + native = float(np.sqrt( + abs(sras.pixel_x_mm) + * max(abs(_pixel_pitch_mm(sras, a)[1]) for a in range(sras.n_angles)))) + pitch = max(span / max_dim, native) + n = int(scipy_fft.next_fast_len(max(16, int(np.ceil(span / pitch))))) + return pitch, n + + +def _registration_workers(sras: SrasFile, fine_dim: int) -> int: + """How many angles may register concurrently. + + Not plan_angle_level: that budgets for waveform chunks, and registration + never touches a waveform — it works on already-computed DC images and a few + grid-sized arrays. The real limit is the masked phase correlation, which + pads to roughly twice the grid and holds several complex128 arrays of that + size live at once, so the count is derived from the same + _TOTAL_BYTES_BUDGET the rest of the module honours. The estimate below + comes out at 262 MB for the default fine_dim=640, against 199 MB measured — + deliberately on the pessimistic side, since overshooting the budget costs + swapping while undershooting only costs a little wall time. + """ + per_worker = 10 * (2 * fine_dim) ** 2 * 16 # ~10 complex128 grids + return int(max(1, min(_MAX_WORKERS, sras.n_angles, + _TOTAL_BYTES_BUDGET // max(1, per_worker)))) + + +def _rotation_candidates(nominal_deg: float, search_deg: float, + step_deg: float) -> list[float]: + """Coarse rotation candidates: a window around *both* signs of the stage's + reported angle change. Scoring both is what makes the stage's sign + convention a non-issue — the images decide which way the stage turns, and + a file whose stage reports the opposite sense registers just as well.""" + out: list[float] = [] + for center in (-nominal_deg, nominal_deg): + k = int(np.floor(search_deg / step_deg)) + for i in range(-k, k + 1): + out.append(center + i * step_deg) + # Dedupe (the two windows coincide when nominal_deg is 0) while keeping order. + seen: set[float] = set() + return [t for t in out if not (round(t, 6) in seen or seen.add(round(t, 6)))] + + +def _score_rotation(ref_img: np.ndarray, ref_valid: np.ndarray, + mov: _RegImage, pitch: float, n: int, theta: float, + subpixel: bool = False) -> tuple[float, tuple[float, float]]: + """Best score this rotation can reach, and the translation that reaches it: + rotate, phase-correlate for the shift, score the overlap. + + *subpixel* also removes the leftover sub-pixel translation before scoring. + That matters more than it sounds: without it every candidate is scored at + whole-pixel alignment, so rotations that differ by less than one pixel of + rim displacement are ranked by quantization noise rather than by fit, and + the refinement stalls a degree or so off. Left off for the coarse sweep, + which only has to pick a basin, and on for the refinement. + """ + img, valid = _embed_with_valid(mov, pitch, n, theta) + dr, dc = _masked_shift(ref_img, ref_valid, img, valid) + shifted = _shift_into(img, dr, dc) + shifted_valid = _shift_into(valid.astype(np.float32), dr, dc) > 0.5 + if subpixel: + sub_dr, sub_dc = _subpixel_residual( + ref_img, shifted, ref_valid & shifted_valid) + if sub_dr or sub_dc: + shifted = scipy_ndimage.shift(shifted, (sub_dr, sub_dc), order=1, + mode="constant", cval=0.0) + dr, dc = dr + sub_dr, dc + sub_dc + return (_overlap_ncc(ref_img, ref_valid, shifted, shifted_valid), + (float(dr), float(dc))) + + +def _refine_rotation(ref_img: np.ndarray, ref_valid: np.ndarray, + mov: _RegImage, pitch: float, n: int, + theta: float, score: float, shift: tuple[float, float], + step_deg: float, min_step_deg: float = 0.05, + max_evals: int = 80 + ) -> tuple[float, float, tuple[float, float]]: + """Hill-climb the rotation from the coarse winner: step out while the score + improves, halve the step when it doesn't, stop below *min_step_deg*. + + A walking search rather than a fixed grid around the coarse winner, because + the coarse stage ranks rotations at a coarse pitch where a degree can be + worth less than the translation quantization — its winner can legitimately + land a degree or two off, which a fixed ±½° refinement window could never + recover from. + + Every candidate here is scored with subpixel=True, matching how the + incoming *score* was measured. Mixing the two is silently fatal: the + sub-pixel-corrected score is strictly the higher of the two, so a + subpixel-scored start compared against un-corrected candidates can never be + beaten and the search sits still at whatever the coarse stage handed it. + """ + evals = 0 + while step_deg >= min_step_deg and evals < max_evals: + trials = [] + for cand in (theta - step_deg, theta + step_deg): + s, sh = _score_rotation(ref_img, ref_valid, mov, pitch, n, cand, + subpixel=True) + evals += 1 + trials.append((s, cand, sh)) + best_s, best_cand, best_sh = max(trials, key=lambda t: t[0]) + if best_s > score: + theta, score, shift = best_cand, best_s, best_sh + else: + step_deg /= 2.0 + return theta, score, shift + + +def _source_images(sras: SrasFile, angle_idx: int, ref_angle_idx: int, + signal_mv: dict[int, np.ndarray], dc_threshold_mv: float, + sources) -> list[tuple[str, np.ndarray, np.ndarray]]: + """(name, reference image, moving image) per requested registration source. + + "signal" is each angle's own CH4 image minus its own minimum — subtracting + per-angle rather than globally keeps a nonzero DC baseline from reading as + a high-contrast edge against the zero padding. "mask" is the binarized + >= dc_threshold_mv image, the same silhouette the overlay draws. A mask + that is empty or completely full for either angle carries no registration + information at all, so that source is dropped rather than scored. + """ + out = [] + for name in sources: + pair = [] + for a in (ref_angle_idx, angle_idx): + img = signal_mv[a] + if name == "mask": + m = img >= dc_threshold_mv + if not m.any() or m.all(): + pair = [] + break + pair.append(m.astype(np.float32)) + else: + pair.append((img - img.min()).astype(np.float32)) + if pair: + out.append((name, pair[0], pair[1])) + return out + + +def register_angle_to_reference( + sras: SrasFile, angle_idx: int, ref_angle_idx: int, + signal_mv: dict[int, np.ndarray], *, + dc_threshold_mv: float = 0.0, + sources: tuple[str, ...] = ("signal", "mask"), + coarse_dim: int = 256, fine_dim: int = _DEFAULT_FINE_DIM, + search_deg: float = 6.0, coarse_step_deg: float = 2.0) -> RigidFit: + """Rigid (rotation + translation, never scale) fit of *angle_idx* onto + *ref_angle_idx*, found entirely by cross-correlating image content. + + Two stages: + 1. Coarse sweep at *coarse_dim* over a ±*search_deg* window around both + signs of the stage's reported angle change (see _rotation_candidates), + for each requested source image, scored by _overlap_ncc. Only has to + pick the right basin. + 2. Hill-climbing refinement of the winning (source, rotation) at + *fine_dim* with sub-pixel translation folded into every score + (_refine_rotation, _score_rotation), down to 0.05°. + + The returned rotation_deg/shift_mm are the rigid map from this angle's + local mm to ref mm (q = R @ l + shift); score is the final NCC, which the + caller can surface so a bad scan is visible rather than silently fused in. + Returns identity for the reference angle itself. """ if angle_idx == ref_angle_idx: - return (0.0, 0.0) + return RigidFit(0.0, (0.0, 0.0), 1.0, "reference") - def corr_img(a: int) -> np.ndarray: - img = signal_mv[a] - if use_mask: - return (img >= dc_threshold_mv).astype(np.float32) - return (img - img.min()).astype(np.float32) + candidates = _source_images(sras, angle_idx, ref_angle_idx, signal_mv, + dc_threshold_mv, sources) + if not candidates: + return RigidFit(0.0, (0.0, 0.0), -1.0, "none") - img_a = corr_img(angle_idx) - img_ref = corr_img(ref_angle_idx) + nominal = _nominal_delta_deg(sras, angle_idx, ref_angle_idx) + thetas = _rotation_candidates(nominal, search_deg, coarse_step_deg) - dx_ref, dy_ref = _pixel_pitch_mm(sras, ref_angle_idx) - max_dim = max(img_a.shape + img_ref.shape) - factor = max(1, int(np.ceil(max_dim / max_corr_dim))) - dx_c, dy_c = dx_ref * factor, dy_ref * factor - small_a = _block_mean_downsample(img_a, factor) - small_ref = _block_mean_downsample(img_ref, factor) + # ---- Stage 1: coarse sweep, every source ------------------------------ + pitch_c, n_c = _reg_pitch_and_size(sras, coarse_dim) + best = (-2.0, 0.0, (0, 0), "none") # score, theta, (dr, dc), source + for name, ref_raw, mov_raw in candidates: + ref_reg = _prepare_reg_image(sras, ref_angle_idx, ref_raw, pitch_c) + mov_reg = _prepare_reg_image(sras, angle_idx, mov_raw, pitch_c) + ref_img, ref_valid = _embed_with_valid(ref_reg, pitch_c, n_c, 0.0) + for theta in thetas: + score, shift = _score_rotation(ref_img, ref_valid, mov_reg, + pitch_c, n_c, theta) + if score > best[0]: + best = (score, theta, shift, name) - work_origin, work_shape = _working_canvas_for_pair( - sras, ref_angle_idx, angle_idx, dx_c, dy_c, pivot_mm, margin_frac=margin_frac) + if best[3] == "none": + return RigidFit(0.0, (0.0, 0.0), -1.0, "none") - m_a, o_a = _build_affine_canvas_to_raw( - sras, angle_idx, ref_angle_idx, (0.0, 0.0), dx_c, dy_c, work_origin, pivot_mm) - m_ref, o_ref = _build_affine_canvas_to_raw( - sras, ref_angle_idx, ref_angle_idx, (0.0, 0.0), dx_c, dy_c, work_origin, pivot_mm) + # ---- Stage 2: refine the winner at full registration resolution ------- + name = best[3] + ref_raw, mov_raw = next((r, m) for n_, r, m in candidates if n_ == name) + pitch_f, n_f = _reg_pitch_and_size(sras, fine_dim) + ref_reg = _prepare_reg_image(sras, ref_angle_idx, ref_raw, pitch_f) + mov_reg = _prepare_reg_image(sras, angle_idx, mov_raw, pitch_f) + ref_img, ref_valid = _embed_with_valid(ref_reg, pitch_f, n_f, 0.0) - rotated_a = scipy_ndimage.affine_transform( - small_a, m_a / factor, offset=o_a / factor, - output_shape=work_shape, order=0, mode="constant", cval=0.0) - embedded_ref = scipy_ndimage.affine_transform( - small_ref, m_ref / factor, offset=o_ref / factor, - output_shape=work_shape, order=0, mode="constant", cval=0.0) + theta = best[1] + score, shift = _score_rotation(ref_img, ref_valid, mov_reg, pitch_f, n_f, + theta, subpixel=True) + theta, score, shift = _refine_rotation( + ref_img, ref_valid, mov_reg, pitch_f, n_f, theta, score, shift, + step_deg=coarse_step_deg) - dr, dc = _phase_correlate_shift(embedded_ref, rotated_a) - return (dc * dx_c, dr * dy_c) + dr, dc = shift + return RigidFit(float(theta), (float(dc * pitch_f), float(dr * pitch_f)), + float(score), name) + + +# ---- Shared canvas: angle 0's own pixel grid, extended -------------------- + +def canvas_for_params(sras: SrasFile, ref_angle_idx: int, + pitch_mm: tuple[float, float], + per_angle_params: dict[int, ManualAngleParams], + *, margin_frac: float = 0.0, snap: bool = True + ) -> tuple[tuple[float, float], tuple[int, int]]: + """Shared-canvas origin (stage mm) and (n_rows, n_cols) at *pitch_mm* that + contains every angle's footprint after its own rigid transform. Angles + missing from per_angle_params default to identity (e.g. a sidecar saved + before a rescan added more angles). + + snap=True aligns the canvas grid with the reference angle's own pixel grid, + so the reference lands on integer canvas pixels and is resampled by an + exact integer translation — the concrete meaning of "the canvas carries + angle 0's X/Y coordinates". It requires pitch_mm to be the reference's own + pitch; the manual-alignment preview passes a coarser pitch and snap=False. + + margin_frac pads the box on every side: 0 for a final canvas, nonzero for + ManualAlignmentDialog's preview canvas, which needs headroom so an ordinary + translation nudge never has to trigger a full canvas resize (an extreme + nudge can still push content past this padding; accepted, and cheap to + recover from by re-opening the dialog). + """ + dx, dy = pitch_mm + corners = np.vstack([ + _footprint_corners_ref_mm( + sras, a, + per_angle_params.get(a, ManualAngleParams()).rotation_deg, + per_angle_params.get(a, ManualAngleParams()).shift_mm) + for a in range(sras.n_angles)]) + x_min, y_min = corners.min(axis=0) + x_max, y_max = corners.max(axis=0) + if margin_frac: + pad_x, pad_y = (x_max - x_min) * margin_frac, (y_max - y_min) * margin_frac + x_min, x_max = x_min - pad_x, x_max + pad_x + y_min, y_max = y_min - pad_y, y_max + pad_y + + center = ref_center_mm(sras, ref_angle_idx) + if snap: + # Express the box in the reference's own pixel indices and grow it + # outward to whole pixels, so canvas index k lands exactly where the + # reference's own pixel (k + const) does. + cy, cx = _center_idx(sras, ref_angle_idx) + cols = sorted((x_min / dx + cx, x_max / dx + cx)) + rows = sorted((y_min / dy + cy, y_max / dy + cy)) + col0, col1 = int(np.floor(cols[0])), int(np.ceil(cols[1])) + row0, row1 = int(np.floor(rows[0])), int(np.ceil(rows[1])) + origin_ref = np.array([(col0 - cx) * dx, (row0 - cy) * dy]) + shape = (row1 - row0 + 1, col1 - col0 + 1) + else: + origin_ref = np.array([x_min, y_min if dy > 0 else y_max]) + shape = (int(np.ceil((y_max - y_min) / abs(dy))) + 1, + int(np.ceil((x_max - x_min) / dx)) + 1) + + origin_stage = origin_ref + center + return (float(origin_stage[0]), float(origin_stage[1])), shape + + +def build_canvas_affine(sras: SrasFile, angle_idx: int, ref_angle_idx: int, + rotation_deg: float, shift_mm: tuple[float, float], + pitch_mm: tuple[float, float], + canvas_origin_mm: tuple[float, float], + *, src_downsample: tuple[int, int] = (1, 1) + ) -> tuple[np.ndarray, np.ndarray]: + """canvas index -> this angle's raw index, for a canvas whose origin is + given in *stage* mm (the reference's frame). *src_downsample* is the + (rows, cols) block-mean factor already applied to the image the caller will + resample — 1:1 for the raw image, coarser for the manual-alignment + preview's downsampled masks.""" + dx_a, dy_a = _pixel_pitch_mm(sras, angle_idx) + n_rows, n_frames = sras.image_shape(angle_idx) + fy, fx = (max(1, int(v)) for v in src_downsample) + if (fy, fx) != (1, 1): + # Same block-mean bookkeeping _prepare_reg_image does: the downsampled + # array's own center can sit up to half a block off the full-resolution + # center, and that offset has to be undone here or every preview layer + # lands slightly (and inconsistently) off. + nr_s, nf_s = n_rows // fy, n_frames // fx + src_dx, src_dy = dx_a * fx, dy_a * fy + src_center = ((nr_s - 1) / 2.0, (nf_s - 1) / 2.0) + src_off = (_block_center_off(n_frames, nf_s, fx, dx_a), + _block_center_off(n_rows, nr_s, fy, dy_a)) + else: + src_dx, src_dy = dx_a, dy_a + src_center = _center_idx(sras, angle_idx) + src_off = (0.0, 0.0) + + origin_ref = np.asarray(canvas_origin_mm, dtype=np.float64) \ + - ref_center_mm(sras, ref_angle_idx) + return _affine_out_to_src( + out_pitch_mm=pitch_mm, out_origin_ref_mm=origin_ref, + src_dx_mm=src_dx, src_dy_mm=src_dy, src_center_idx=src_center, + src_center_off_mm=src_off, rotation_deg=rotation_deg, shift_mm=shift_mm) + + +def _result_from_params(sras: SrasFile, ref_angle_idx: int, + dc_threshold_mv: float, + params: dict[int, ManualAngleParams], + extra: dict[int, tuple[float, str]] | None = None + ) -> AlignmentResult: + """Assemble the final AlignmentResult from per-angle rigid parameters: pick + the shared canvas, then build each angle's canvas->raw affine. Pure matrix + and bbox math, so it is cheap enough to call synchronously on the GUI + thread on every manual edit.""" + pitch = _pixel_pitch_mm(sras, ref_angle_idx) + canvas_origin_mm, canvas_shape = canvas_for_params( + sras, ref_angle_idx, pitch, params, snap=True) + + extra = extra or {} + per_angle: dict[int, AngleTransform] = {} + for a in range(sras.n_angles): + p = params.get(a, ManualAngleParams()) + matrix, offset = build_canvas_affine( + sras, a, ref_angle_idx, p.rotation_deg, p.shift_mm, + pitch, canvas_origin_mm) + score, source = extra.get(a, (1.0, "")) + per_angle[a] = AngleTransform(p.rotation_deg, p.shift_mm, matrix, offset, + score, source) + + return AlignmentResult(ref_angle_idx, dc_threshold_mv, canvas_shape, + pitch[0], pitch[1], canvas_origin_mm, per_angle) def _parallel_map(fn, items, n_workers: int) -> list: @@ -733,14 +1131,17 @@ def _parallel_map(fn, items, n_workers: int) -> list: def compute_angle_alignment(sras: SrasFile, ref_angle_idx: int, - dc_threshold_mv: float, - progress_cb=None) -> AlignmentResult: - """Top-level alignment driver. Runs on a background thread (see - AngleAlignmentWorker) — deliberately recomputes CH4 DC images from - scratch rather than reading the GUI-thread _dc_cache dict, since - background-thread workers must not touch GUI-thread-owned caches.""" + dc_threshold_mv: float, progress_cb=None, + fine_dim: int = _DEFAULT_FINE_DIM) -> AlignmentResult: + """Top-level alignment driver: register every angle onto *ref_angle_idx* by + content, then lay them all out on that angle's own coordinate grid. + + Runs on a background thread (see AngleAlignmentWorker) — deliberately + recomputes CH4 DC images from scratch rather than reading the GUI-thread + _dc_cache dict, since background-thread workers must not touch + GUI-thread-owned caches. + """ n = sras.n_angles # already the *complete*-angle count for aborted v6 scans - dx_ref, dy_ref = _pixel_pitch_mm(sras, ref_angle_idx) # Parallel over angles, serial within each — see plan_angle_level. n_workers, angle_budget = plan_angle_level(sras) @@ -753,7 +1154,7 @@ def compute_angle_alignment(sras: SrasFile, ref_angle_idx: int, _ticks.append(1) progress_cb(base + int(min(len(_ticks), n) / n * span)) - # ---- Step 1: CH4 DC image + binarized mask per angle, native grid ----- + # ---- Step 1: CH4 DC image per angle, native grid ---------------------- def dc4_for(a: int) -> np.ndarray: dc4 = adc_to_mv( compute_dc_image(sras, a, CH4_IDX, max_workers=1, budget=angle_budget), @@ -763,51 +1164,28 @@ def compute_angle_alignment(sras: SrasFile, ref_angle_idx: int, dc4_mv = dict(enumerate(_parallel_map(dc4_for, range(n), n_workers))) - # Each angle's own alignment pivot — the CH4-signal-weighted centroid, - # not the raw scan-window bbox center (see compute_pivot_points_mm). - # Reuses the dc4_mv images just computed above, so this is free, and is - # deliberately independent of dc_threshold_mv (see that function's - # docstring) so a threshold that happens to leave some angle's binary - # mask empty can't silently degrade the pivot back to the bbox center. - pivot_mm = compute_pivot_points_mm(sras, dc4_mv=dc4_mv) + # ---- Step 2: rigid registration of every angle against the reference -- + # Its own worker count: registration is bounded by grid-sized FFT buffers, + # not by the waveform chunking plan_angle_level budgets for. + _ticks.clear() - # ---- Step 2: coarse translation via FFT phase correlation, on each - # angle's binarized CH4 mask against the reference's (see - # correlate_translation_mm — also the engine behind ManualAlignmentDialog's - # Auto Cross-Correlate button, there defaulting to the raw signal instead - # of a mask). - def shift_for(a: int) -> tuple[float, float]: - shift = correlate_translation_mm( - sras, a, ref_angle_idx, dc4_mv, pivot_mm, - use_mask=True, dc_threshold_mv=dc_threshold_mv) - tick(25, 50) - return shift + def fit_for(a: int) -> RigidFit: + fit = register_angle_to_reference( + sras, a, ref_angle_idx, dc4_mv, dc_threshold_mv=dc_threshold_mv, + fine_dim=fine_dim) + tick(25, 65) + return fit - shifts_mm = dict(enumerate(_parallel_map(shift_for, range(n), n_workers))) + fits = dict(enumerate(_parallel_map( + fit_for, range(n), _registration_workers(sras, fine_dim)))) - # ---- Step 3: union bounding box over all angles (rotation+shift applied) - corners = np.vstack([ - _corners_in_ref_frame(sras, a, ref_angle_idx, pivot_mm, shifts_mm[a]) - for a in range(n)]) - x_min, y_min = corners.min(axis=0) - x_max, y_max = corners.max(axis=0) - n_cols = int(np.ceil((x_max - x_min) / dx_ref)) + 1 - n_rows = int(np.ceil((y_max - y_min) / abs(dy_ref))) + 1 - canvas_origin_mm = (float(x_min), float(y_min if dy_ref > 0 else y_max)) - - # ---- Step 4: final per-angle full-resolution affine (canvas -> raw idx) - per_angle: dict[int, AngleTransform] = {} - for a in range(n): - matrix, offset = _build_affine_canvas_to_raw( - sras, a, ref_angle_idx, shifts_mm[a], dx_ref, dy_ref, canvas_origin_mm, - pivot_mm) - per_angle[a] = AngleTransform( - _theta_deg(sras, a, ref_angle_idx), shifts_mm[a], matrix, offset) - if progress_cb: - progress_cb(75 + int((a + 1) / n * 25)) - - return AlignmentResult(ref_angle_idx, dc_threshold_mv, (n_rows, n_cols), - dx_ref, dy_ref, canvas_origin_mm, per_angle) + # ---- Step 3/4: shared canvas on the reference's grid, per-angle affines + params = {a: ManualAngleParams(f.rotation_deg, f.shift_mm) for a, f in fits.items()} + extra = {a: (f.score, f.source) for a, f in fits.items()} + result = _result_from_params(sras, ref_angle_idx, dc_threshold_mv, params, extra) + if progress_cb: + progress_cb(100) + return result # Back-compat alias for the pre-split private name (used by tooling). @@ -817,80 +1195,36 @@ _compute_angle_alignment = compute_angle_alignment # --------------------------------------------------------------------------- # Manual alignment (Fusion menu -> Manual Alignment... dialog) # -# Skips compute_angle_alignment's mask + phase-correlation search entirely: -# every angle's rotation_deg/shift_mm is supplied directly by the caller -# (nudged by eye against a live multi-angle mask overlay, or pre-seeded from -# a previous compute_angle_alignment run or a saved sidecar). Building the -# final AlignmentResult from already-known per-angle parameters is pure -# closed-form matrix math (union_canvas_mm + _build_affine_canvas_to_raw) -# with no per-pixel image work at all, so build_manual_alignment is cheap -# enough to call synchronously on the GUI thread on every edit. The only -# genuinely expensive per-pixel operation anywhere in this flow is -# reproject_mask, and only ManualAlignmentDialog's own downsampled preview -# calls that per keystroke — see that class's docstring for how it limits -# each nudge to reprojecting only the actively-edited angle. +# Skips register_angle_to_reference's search entirely: every angle's +# rotation_deg/shift_mm is supplied directly by the caller (nudged by eye +# against a live multi-angle mask overlay, or pre-seeded from a registration +# run or a saved sidecar). Building the final AlignmentResult from already-known +# per-angle parameters is pure closed-form matrix math (_result_from_params) +# with no per-pixel image work at all, so build_manual_alignment is cheap enough +# to call synchronously on the GUI thread on every edit. The only genuinely +# expensive per-pixel operation anywhere in this flow is reproject_mask, and +# only ManualAlignmentDialog's own downsampled preview calls that per keystroke +# — see that class's docstring for how it limits each nudge to reprojecting only +# the actively-edited angle. # --------------------------------------------------------------------------- -def union_canvas_mm(sras: SrasFile, ref_angle_idx: int, dx: float, dy: float, - per_angle_params: dict[int, ManualAngleParams], - pivot_mm: dict[int, tuple[float, float]], - margin_frac: float = 0.0 - ) -> tuple[tuple[float, float], tuple[int, int]]: - """Shared-canvas origin (mm) and (n_rows, n_cols) at pitch (dx, dy) that - contains every angle's footprint after applying its own rotation+shift — - the generalisation of compute_angle_alignment's Step 3 to arbitrary (not - just _theta_deg-analytic) per-angle rotation. Angles missing from - per_angle_params default to identity (e.g. a sidecar saved before a - rescan added more angles). - - margin_frac pads the box on every side: 0 for a final canvas (this then - reproduces compute_angle_alignment's own Step-3 math exactly, when every - angle's rotation_deg equals the analytic delta and shift_mm matches); - nonzero for ManualAlignmentDialog's downsampled preview canvas, which - needs headroom so an ordinary translation nudge of the active angle never - has to trigger a full canvas resize (see that class's docstring — an - extreme nudge can still, in principle, push content past this padding; - accepted as a known edge case, same as _phase_correlate_shift's - wraparound risk note). - """ - n = sras.n_angles - corners = np.vstack([ - _corners_in_ref_frame( - sras, a, ref_angle_idx, pivot_mm, - per_angle_params.get(a, ManualAngleParams()).shift_mm, - theta_deg=per_angle_params.get(a, ManualAngleParams()).rotation_deg) - for a in range(n)]) - x_min, y_min = corners.min(axis=0) - x_max, y_max = corners.max(axis=0) - if margin_frac: - pad_x, pad_y = (x_max - x_min) * margin_frac, (y_max - y_min) * margin_frac - x_min, x_max = x_min - pad_x, x_max + pad_x - y_min, y_max = y_min - pad_y, y_max + pad_y - n_cols = int(np.ceil((x_max - x_min) / dx)) + 1 - n_rows = int(np.ceil((y_max - y_min) / abs(dy))) + 1 - origin = (float(x_min), float(y_min if dy > 0 else y_max)) - return origin, (n_rows, n_cols) - - def reproject_mask(sras: SrasFile, angle_idx: int, ref_angle_idx: int, mask: np.ndarray, rotation_deg: float, shift_mm: tuple[float, float], - canvas_dx: float, canvas_dy: float, + canvas_pitch_mm: tuple[float, float], canvas_origin_mm: tuple[float, float], canvas_shape: tuple[int, int], - pivot_mm: dict[int, tuple[float, float]]) -> np.ndarray: + *, src_downsample: tuple[int, int] = (1, 1)) -> np.ndarray: """Resample one angle's binary/float mask onto an arbitrary canvas via an - explicit rotation+shift — the single building block - ManualAlignmentDialog's live preview repeatedly calls (once per - keystroke, for only the actively-nudged angle), since it bypasses - compute_angle_alignment's phase-correlation search entirely and just - takes rotation_deg/shift_mm as given. order=0 (nearest) matches - apply_alignment's own reasoning: a binary mask must never be blended with - zero-padding. - """ - matrix, offset = _build_affine_canvas_to_raw( - sras, angle_idx, ref_angle_idx, shift_mm, canvas_dx, canvas_dy, - canvas_origin_mm, pivot_mm, theta_deg=rotation_deg) + explicit rotation+shift — the single building block ManualAlignmentDialog's + live preview repeatedly calls (once per keystroke, for only the + actively-nudged angle). *src_downsample* must match the (rows, cols) + block-mean factor already applied to *mask*, or the reprojection lands at + the wrong scale. order=0 (nearest) matches apply_alignment's own reasoning: + a binary mask must never be blended with zero-padding.""" + matrix, offset = build_canvas_affine( + sras, angle_idx, ref_angle_idx, rotation_deg, shift_mm, + canvas_pitch_mm, canvas_origin_mm, src_downsample=src_downsample) return scipy_ndimage.affine_transform( mask.astype(np.float32, copy=False), matrix, offset=offset, output_shape=canvas_shape, order=0, mode="constant", cval=0.0) @@ -898,48 +1232,24 @@ def reproject_mask(sras: SrasFile, angle_idx: int, ref_angle_idx: int, def build_manual_alignment(sras: SrasFile, ref_angle_idx: int, dc_threshold_mv: float, - per_angle_params: dict[int, ManualAngleParams], - pivot_mm: dict[int, tuple[float, float]] | None = None + per_angle_params: dict[int, ManualAngleParams] ) -> AlignmentResult: """Build a full, full-resolution AlignmentResult from user-supplied per-angle rotation+shift — the Manual Alignment counterpart to - compute_angle_alignment, skipping its mask/phase-correlation search - entirely (every angle's transform here is exactly what the caller - supplied). Pure matrix/bbox math once pivot_mm is known, so it is cheap - enough to call synchronously on the GUI thread. The reference angle's - params are always forced to identity, regardless of what - per_angle_params holds for it — it defines the shared origin and must - never be transformed. + compute_angle_alignment, skipping its registration search entirely (every + angle's transform here is exactly what the caller supplied). The reference + angle's params are always forced to identity, regardless of what + per_angle_params holds for it — it defines the shared origin and must never + be transformed. - Pass an already-computed *pivot_mm* (e.g. ManualAlignmentDialog's own - cache, built once from its live CH4 images) to skip recomputing every - angle's DC image here; otherwise it's computed fresh via - compute_pivot_points_mm, which is still fine for a one-off Save or - sidecar restore (just not free on a very large, not-yet-cached scan). - dc_threshold_mv itself plays no part in the pivot (see that function's - docstring) — it's stored on the returned AlignmentResult purely as a - record of the RF-mask threshold in effect at the time. + dc_threshold_mv plays no part in the geometry; it's stored on the returned + AlignmentResult purely as a record of the RF-mask threshold in effect at + the time. """ - n = sras.n_angles - dx_ref, dy_ref = _pixel_pitch_mm(sras, ref_angle_idx) - params = {a: per_angle_params.get(a, ManualAngleParams()) for a in range(n)} + params = {a: per_angle_params.get(a, ManualAngleParams()) + for a in range(sras.n_angles)} params[ref_angle_idx] = ManualAngleParams() - if pivot_mm is None: - pivot_mm = compute_pivot_points_mm(sras) - - canvas_origin_mm, canvas_shape = union_canvas_mm( - sras, ref_angle_idx, dx_ref, dy_ref, params, pivot_mm, margin_frac=0.0) - - per_angle: dict[int, AngleTransform] = {} - for a in range(n): - p = params[a] - matrix, offset = _build_affine_canvas_to_raw( - sras, a, ref_angle_idx, p.shift_mm, dx_ref, dy_ref, - canvas_origin_mm, pivot_mm, theta_deg=p.rotation_deg) - per_angle[a] = AngleTransform(p.rotation_deg, p.shift_mm, matrix, offset) - - return AlignmentResult(ref_angle_idx, dc_threshold_mv, canvas_shape, - dx_ref, dy_ref, canvas_origin_mm, per_angle) + return _result_from_params(sras, ref_angle_idx, dc_threshold_mv, params) # ---- Sidecar persistence (.sras.align.json) ------------------------- @@ -967,15 +1277,19 @@ def sidecar_path(sras_path) -> Path: return p.with_name(p.name + ".align.json") -# Bumped from 1 -> 2 when the rotation pivot changed from the raw scan- -# window bbox center to a content-derived centroid, *and* _theta_deg's sign -# convention was corrected — either change alone makes a version-1 file's -# stored rotation_deg/shift_mm numbers describe a different (and, for the -# pivot bug, actively wrong) transform than they would today. Loading one -# unchanged would silently reproduce exactly the "scans show up everywhere" -# symptom these fixes address, so version-1 sidecars are treated as absent -# rather than migrated. -_SIDECAR_SCHEMA_VERSION = 2 +# The stored rotation_deg/shift_mm are meaningless without the frame they were +# measured in, so this is bumped whenever that frame changes. Each bump makes +# older files describe a different (and, for the bugs each bump fixed, actively +# wrong) transform than the same numbers would today, and loading one unchanged +# would silently reproduce the very "scans show up everywhere" symptom the bump +# fixed — so older sidecars are treated as absent rather than migrated. +# 1 -> 2 pivot moved from the scan-window bbox center to a content-derived +# centroid, and the rotation sign convention was corrected. +# 2 -> 3 the content centroid was abandoned entirely: rotation is now about +# each angle's own array center, mapped onto the reference's array +# center, with shift_mm in the reference's local mm frame. No angle +# but the reference contributes stage coordinates any more. +_SIDECAR_SCHEMA_VERSION = 3 def save_manual_alignment(sras: SrasFile, ref_angle_idx: int, @@ -984,9 +1298,9 @@ def save_manual_alignment(sras: SrasFile, ref_angle_idx: int, """Write the sidecar JSON for sras.path (overwriting any existing one) and return the path written. - Schema (schema_version 2): + Schema (schema_version 3): { - "schema_version": 2, + "schema_version": 3, "ref_angle_idx": , "dc_threshold_mv": , "per_angle": { diff --git a/sras_edit_scans.py b/sras_edit_scans.py new file mode 100644 index 0000000..8774daf --- /dev/null +++ b/sras_edit_scans.py @@ -0,0 +1,295 @@ +#!/usr/bin/env python3 +""" +sras_edit_scans.py — Remove one or more angle scans from a .sras file. + +A .sras file holds one or more "angles" (rotation positions); the viewer +cross-correlates each non-reference angle against the reference to align +them. If one angle's acquisition went wrong (stage glitch, bad trigger, +laser dropout, ...) it throws off that alignment for the whole file. This +tool drops the bad angle(s) and renumbers the rest, writing a new .sras file +with everything else — waveform samples, calibration preambles, background +waveform, row/geometry tables — carried over byte-for-byte. + +Handles v2-v7. Any precomputed FFT/DC cache (v5 PREC tail, v7 CACH tail) is +dropped on write, since it's indexed by angle and would be stale/misaligned +after renumbering; the viewer just recomputes it next time the file opens. + +Usage: + python sras_edit_scans.py input.sras --list + python sras_edit_scans.py input.sras output.sras --drop 2,5 + python sras_edit_scans.py input.sras output.sras --keep 0,1,3,4,6 +""" + +import argparse +import struct +import sys +from pathlib import Path + +from sras_format import ( + GEO_FMT_V6, GEO_SIZE_V6, HDR_FMT, HDR_FMT_V6, HDR_SIZE, HDR_SIZE_V6, + SrasFile, +) + +_LEGACY_VERSIONS = (2, 3, 4, 5) +_V6_VERSIONS = (6, 7) + + +def parse_args(): + p = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + p.add_argument("input", help="Input .sras file") + p.add_argument("output", nargs="?", help="Output .sras file (omit with --list)") + p.add_argument("--list", action="store_true", + help="Print each angle's index/degrees/geometry and exit") + g = p.add_mutually_exclusive_group() + g.add_argument("--drop", metavar="I,J,...", + help="Comma-separated angle indices to remove") + g.add_argument("--keep", metavar="I,J,...", + help="Comma-separated angle indices to keep (all others dropped)") + return p.parse_args() + + +def _parse_index_list(s: str, n_angles: int) -> set[int]: + out = set() + for piece in s.split(","): + piece = piece.strip() + if not piece: + continue + i = int(piece) + if not (0 <= i < n_angles): + raise ValueError(f"angle index {i} out of range [0, {n_angles - 1}]") + out.add(i) + return out + + +def print_listing(sras: SrasFile): + print(f"\n{'idx':>4} {'angle_deg':>10} {'x_start_mm':>11} {'rows':>6} {'frames':>7}") + for a in range(sras.n_angles): + print(f"{a:>4} {sras.angles_deg[a]:>10.4f} {sras.x_start_mm[a]:>11.4f} " + f"{int(sras.n_rows[a]):>6} {int(sras.n_frames[a]):>7}") + + +def _copy_range(fin, fout, offset: int, nbytes: int, chunk: int = 64 * 1024 * 1024): + """Stream *nbytes* raw bytes from *fin* at *offset* into *fout*, without + ever holding more than one chunk in memory (waveform blocks can be + hundreds of MB to low GB each).""" + fin.seek(offset) + remaining = nbytes + while remaining: + buf = fin.read(min(chunk, remaining)) + if not buf: + raise IOError("unexpected EOF while copying waveform data") + fout.write(buf) + remaining -= len(buf) + + +# --------------------------------------------------------------------------- +# Legacy (v2-v5): uniform geometry across angles, one flat waveform block +# --------------------------------------------------------------------------- + +def _write_legacy(sras: SrasFile, keep: list[int], out_path: Path): + n_rows = int(sras.n_rows[0]) + n_frames = int(sras.n_frames[0]) # uniform across angles for v2-v5 + n_ch = sras.n_channels + spf = sras.samples_per_frame + bps = sras.bytes_per_sample + + header = struct.pack( + HDR_FMT, b"SRAS", sras.version, len(keep), n_rows, + float(sras.x_start_mm[0]), float(sras.x_delta_mm), + sras.velocity_mm_s, sras.laser_freq_hz, + n_frames, spf, sras.sample_rate_hz, bps, n_ch, + ) + + # Row table + preambles + background sit right after the angle table and + # don't vary per angle — copy that whole span through unmodified. + angle_table_size = sras.n_angles * 4 + with open(sras.path, "rb") as f: + f.seek(HDR_SIZE + angle_table_size) + shared_mid = f.read(sras._data_offset - (HDR_SIZE + angle_table_size)) + + angle_bytes = n_rows * n_ch * n_frames * spf * bps + + with open(sras.path, "rb") as fin, open(out_path, "wb") as fout: + fout.write(header) + fout.write(sras.angles_deg[keep].astype(">f4").tobytes()) + fout.write(shared_mid) + for a in keep: + _copy_range(fin, fout, sras._data_offset + a * angle_bytes, angle_bytes) + + +# --------------------------------------------------------------------------- +# v6/v7: per-angle geometry, ragged waveform blocks +# --------------------------------------------------------------------------- + +def _read_v6_sections(path: Path) -> dict: + """Raw, low-level read of everything before the waveform data. + + SrasFile._parse_v6 reads and discards each angle's x_delta (it's not + part of the display-facing geometry it exposes), so a round-trip through + SrasFile would silently drop that field. Re-parsing here keeps every + byte of the Per-Angle Geometry Table intact. + """ + with open(path, "rb") as f: + hdr_raw = f.read(HDR_SIZE_V6) + (magic, ver, n_angles_declared, x_start_nom, y_start_nom, x_delta_nom, + y_delta_nom, row_spacing, vel, freq, spf, sr, bps, + n_ch) = struct.unpack(HDR_FMT_V6, hdr_raw) + + angles_deg = list(struct.unpack(f">{n_angles_declared}f", + f.read(n_angles_declared * 4))) + + geo = [struct.unpack(GEO_FMT_V6, f.read(GEO_SIZE_V6)) + for _ in range(n_angles_declared)] # (x_start, x_delta, n_frames, n_rows) + + row_table = [f.read(geo[a][3] * 4) for a in range(n_angles_declared)] + + preamble_start = f.tell() + for _ in range(n_ch): + (length,) = struct.unpack(">H", f.read(2)) + f.read(length) + preambles_raw = _reread_span(f, preamble_start) + + bg_start = f.tell() + (n_bg,) = struct.unpack(">I", f.read(4)) + f.read(n_bg) + background_raw = _reread_span(f, bg_start) + + data_offset = f.tell() + + return { + "version": ver, "n_angles_declared": n_angles_declared, + "x_start_nom": x_start_nom, "y_start_nom": y_start_nom, + "x_delta_nom": x_delta_nom, "y_delta_nom": y_delta_nom, + "row_spacing": row_spacing, "vel": vel, "freq": freq, + "spf": spf, "sr": sr, "bps": bps, "n_ch": n_ch, + "angles_deg": angles_deg, "geo": geo, "row_table": row_table, + "preambles_raw": preambles_raw, "background_raw": background_raw, + "data_offset": data_offset, + } + + +def _reread_span(f, start: int) -> bytes: + end = f.tell() + f.seek(start) + span = f.read(end - start) + f.seek(end) + return span + + +def _v6_angle_offsets(sections: dict, file_size: int) -> list[tuple[int, int]]: + """(offset, nbytes) of each declared angle's waveform block, stopping at + the first angle whose data isn't fully on disk (aborted scan).""" + n_ch, spf, bps = sections["n_ch"], sections["spf"], sections["bps"] + offsets = [] + offset = sections["data_offset"] + for xs, xd, nf, nr in sections["geo"]: + nbytes = nr * n_ch * nf * spf * bps + if offset + nbytes > file_size: + break + offsets.append((offset, nbytes)) + offset += nbytes + return offsets + + +def _write_v6(in_path: Path, sections: dict, keep: list[int], out_path: Path): + file_size = in_path.stat().st_size + offsets = _v6_angle_offsets(sections, file_size) + geo = sections["geo"] + + header = struct.pack( + HDR_FMT_V6, b"SRAS", sections["version"], len(keep), + sections["x_start_nom"], sections["y_start_nom"], + sections["x_delta_nom"], sections["y_delta_nom"], + sections["row_spacing"], sections["vel"], sections["freq"], + sections["spf"], sections["sr"], sections["bps"], sections["n_ch"], + ) + + with open(in_path, "rb") as fin, open(out_path, "wb") as fout: + fout.write(header) + fout.write(struct.pack(f">{len(keep)}f", + *[sections["angles_deg"][i] for i in keep])) + for i in keep: + fout.write(struct.pack(GEO_FMT_V6, *geo[i])) + for i in keep: + fout.write(sections["row_table"][i]) + fout.write(sections["preambles_raw"]) + fout.write(sections["background_raw"]) + for i in keep: + off, nbytes = offsets[i] + _copy_range(fin, fout, off, nbytes) + + +def main(): + args = parse_args() + in_path = Path(args.input) + if not in_path.exists(): + print(f"Error: input file not found: {in_path}", file=sys.stderr) + sys.exit(1) + + print(f"Reading {in_path} ...", flush=True) + try: + sras = SrasFile(str(in_path)) + except ValueError as e: + print(f"Error: {e}", file=sys.stderr) + sys.exit(1) + + if sras.version not in (*_LEGACY_VERSIONS, *_V6_VERSIONS): + print(f"Error: unsupported .sras version: {sras.version}", file=sys.stderr) + sys.exit(1) + + aborted_note = " (scan aborted; trailing angle(s) already excluded)" if sras.scan_aborted else "" + print(f" Version : v{sras.version}", flush=True) + print(f" Angles : {sras.n_angles}{aborted_note}", flush=True) + + if args.list: + print_listing(sras) + return + + if not args.output: + print("Error: output path required unless --list is given.", file=sys.stderr) + sys.exit(1) + if not (args.drop or args.keep): + print("Error: specify --drop or --keep (see --list for indices).", file=sys.stderr) + sys.exit(1) + + out_path = Path(args.output) + if out_path.resolve() == in_path.resolve(): + print("Error: output path must differ from input path.", file=sys.stderr) + sys.exit(1) + + try: + if args.drop: + drop = _parse_index_list(args.drop, sras.n_angles) + keep = [a for a in range(sras.n_angles) if a not in drop] + else: + keep = sorted(_parse_index_list(args.keep, sras.n_angles)) + except ValueError as e: + print(f"Error: {e}", file=sys.stderr) + sys.exit(1) + + if not keep: + print("Error: at least one angle must remain.", file=sys.stderr) + sys.exit(1) + + dropped = [a for a in range(sras.n_angles) if a not in keep] + print(f"\nDropping angle(s): {dropped}") + print(f"Keeping angle(s) : {keep} ({len(keep)} of {sras.n_angles})") + print(f"\nWriting {out_path} ...", flush=True) + + if sras.version in _LEGACY_VERSIONS: + _write_legacy(sras, keep, out_path) + else: + _write_v6(in_path, _read_v6_sections(in_path), keep, out_path) + + in_mb = in_path.stat().st_size / 1024**2 + out_mb = out_path.stat().st_size / 1024**2 + print(f" Input size : {in_mb:.1f} MB") + print(f" Output size: {out_mb:.1f} MB") + print("Done.") + print("Note: any precomputed FFT/DC cache was dropped (it's indexed by " + "angle); the viewer will recompute it next time this file opens.") + + +if __name__ == "__main__": + main() diff --git a/sras_viewer.py b/sras_viewer.py index fb7e8b5..54272c0 100644 --- a/sras_viewer.py +++ b/sras_viewer.py @@ -821,18 +821,18 @@ class ManualAlignmentDialog(QDialog): misalignment is visible by eye. Reference angle (always index 0) is ground truth and never moves; every other angle is aligned to it. The user picks an "active" angle and nudges its rotation+translation with - the keyboard; Auto De-rotate sets every non-reference angle's rotation to - the known, analytic scan-angle delta without touching any translation; - Auto Cross-Correlate does the same rotation and additionally sets - translation to the FFT-phase-correlation best fit against the reference - (see compute.correlate_translation_mm) — meant to get every angle roughly - stacked on top of each other so keyboard nudging only has to make small - corrections, not find a coarse alignment from scratch. Save writes a - JSON sidecar next to the .sras file and hands a freshly-built, full- - resolution AlignmentResult back to the main window — the exact same - object shape compute_angle_alignment produces, so every existing - Aligned-View code path (apply_alignment, _aligned_canvas_axes, the - pixel-inspector inverse-transform) works completely unmodified. + the keyboard; Auto Cross-Correlate finds every non-reference angle's + rotation *and* translation by registering its image against the + reference's (see compute.register_angle_to_reference) — meant to get every + angle stacked on top of each other so keyboard nudging only has to make + small corrections, not find an alignment from scratch; Auto De-rotate is + the weaker fallback that just seeds rotation from the stage's reported + angle, leaving translation alone. Save writes a JSON sidecar next to the + .sras file and hands a freshly-built, full-resolution AlignmentResult back + to the main window — the exact same object shape compute_angle_alignment + produces, so every existing Aligned-View code path (apply_alignment, + _aligned_canvas_axes, the pixel-inspector inverse-transform) works + completely unmodified. Non-modal by design (shown via .show(), never .exec() or setModal(True)) so the user can still interact with the main window. Talks back to @@ -853,6 +853,16 @@ class ManualAlignmentDialog(QDialog): _ACTIVE_ALPHA = 0.75 _MAX_PREVIEW_DIM = 1024 + # (label, sources passed to compute.register_angle_to_reference). "Both" + # registers on each and keeps whichever scores higher per angle, which + # costs roughly double but removes the failure mode where the single + # chosen source is the one that happens to be uninformative for one angle. + _CORRELATE_SOURCES = ( + ("Both, keep best (recommended)", ("signal", "mask")), + ("Raw signal", ("signal",)), + ("Thresholded mask", ("mask",)), + ) + def __init__(self, parent: "SrasViewerWindow", sras: SrasFile, *, ref_angle_idx: int, dc_threshold_mv: float, seed_per_angle: dict[int, ManualAngleParams] | None, @@ -861,15 +871,16 @@ class ManualAlignmentDialog(QDialog): self._parent = parent self._sras = sras self._ref_angle_idx = ref_angle_idx - self._downsample_factor = 1 + self._downsample = (1, 1) # (rows, cols) block-mean factors self._dc4_mv: dict[int, np.ndarray] = {} self._masks_small: dict[int, np.ndarray] = {} - self._pivot_mm: dict[int, tuple[float, float]] = {} self._preview_layers: dict[int, np.ndarray] = {} self._preview_origin_mm = (0.0, 0.0) self._preview_shape = (1, 1) - self._preview_dx_mm = self._preview_dy_mm = 1.0 + self._preview_pitch_mm = (1.0, 1.0) self._masks_ready = False + self._fit_notes: dict[int, tuple[float, str]] = {} + self._derotate_sign_flipped = False self.setWindowTitle(f"Manual Alignment — {sras.path.name}") self.resize(1150, 760) @@ -1006,25 +1017,27 @@ class ManualAlignmentDialog(QDialog): self.grp_correlate, cl = _group("Cross-Correlate (FFT)") cform = _form() self.combo_correlate_source = QComboBox() - self.combo_correlate_source.addItems( - ["Raw signal (recommended)", "Thresholded mask"]) + for label, sources in self._CORRELATE_SOURCES: + self.combo_correlate_source.addItem(label, sources) cform.addRow("Correlate on:", self.combo_correlate_source) - self.spin_correlate_margin = QDoubleSpinBox() - self.spin_correlate_margin.setRange(0.05, 2.0) - self.spin_correlate_margin.setSingleStep(0.05) - self.spin_correlate_margin.setDecimals(2) - self.spin_correlate_margin.setValue(0.30) - self.spin_correlate_margin.setMinimumWidth(_SPIN_MIN_W) - cform.addRow("Search margin (× extent):", self.spin_correlate_margin) + self.spin_correlate_search_deg = QDoubleSpinBox() + self.spin_correlate_search_deg.setRange(0.0, 180.0) + self.spin_correlate_search_deg.setSingleStep(1.0) + self.spin_correlate_search_deg.setDecimals(1) + self.spin_correlate_search_deg.setSuffix(" °") + self.spin_correlate_search_deg.setValue(6.0) + self.spin_correlate_search_deg.setMinimumWidth(_SPIN_MIN_W) + cform.addRow("Rotation search (±):", self.spin_correlate_search_deg) cl.addLayout(cform) self.btn_auto_correlate = QPushButton("Auto Cross-Correlate (vs Reference)") cl.addWidget(self.btn_auto_correlate) cl.addWidget(_wrap_label( - "Sets rotation to the known scan angle and translation to the " - "FFT-correlated best fit for every non-reference angle. Run this " - "first, then use manual nudging only for small corrections.", - _CSS_HINT)) + "Finds each non-reference angle's rotation *and* translation by " + "cross-correlating its image against the reference's — the stage's " + "reported angle is only the starting point of the search, and both " + "of its signs are tried. Run this first, then nudge only for small " + "corrections.", _CSS_HINT)) panel_l.addWidget(self.grp_correlate) # ---- Actions ------------------------------------------------------ @@ -1091,15 +1104,16 @@ class ManualAlignmentDialog(QDialog): def _finish_mask_prep(self): if len(self._dc4_mv) < self._sras.n_angles: return # a mask-worker error left some angles unfetched - max_dim = max(max(img.shape) for img in self._dc4_mv.values()) - self._downsample_factor = max(1, int(np.ceil(max_dim / self._MAX_PREVIEW_DIM))) + # Rows and columns get their own factor. A real scan is ~7500 frames + # wide but only ~750 rows tall, so one shared factor sized for the + # frames would throw away 8x more row detail than the preview needs and + # leave the overlay too coarse in y to judge alignment by eye. + max_rows = max(img.shape[0] for img in self._dc4_mv.values()) + max_cols = max(img.shape[1] for img in self._dc4_mv.values()) + self._downsample = ( + max(1, int(np.ceil(max_rows / self._MAX_PREVIEW_DIM))), + max(1, int(np.ceil(max_cols / self._MAX_PREVIEW_DIM)))) self._recompute_masks_small() - # Alignment pivot: the CH4-signal-weighted centroid of each angle's - # own footprint (see compute.compute_pivot_points_mm) — computed once - # from the full-res CH4 images and deliberately independent of the - # mask threshold, so it never needs recomputing when that changes - # (unlike _masks_small, which is purely for the overlay's visuals). - self._pivot_mm = compute.compute_pivot_points_mm(self._sras, self._dc4_mv) self._rebuild_preview_canvas() self._set_controls_enabled(True) self.lbl_status.setText("Ready.") @@ -1108,13 +1122,12 @@ class ManualAlignmentDialog(QDialog): """Threshold + downsample every angle's already-in-memory full-res CH4 mV image. Cheap (a compare + block-mean), so this re-runs in full whenever the mask-threshold spin box changes — no re-fetch. - Purely for the overlay's visuals — the alignment pivot does not - depend on this threshold (see _pivot_mm / compute_pivot_points_mm).""" + Purely for the overlay's visuals: no alignment geometry depends on this + threshold, only which pixels the overlay paints.""" threshold = self.spin_mask_threshold_mv.value() - factor = self._downsample_factor + fy, fx = self._downsample self._masks_small = { - a: compute._block_mean_downsample( - (img >= threshold).astype(np.float32), factor) + a: compute._block_mean_2d((img >= threshold).astype(np.float32), fy, fx) for a, img in self._dc4_mv.items() } @@ -1131,32 +1144,35 @@ class ManualAlignmentDialog(QDialog): active angle. NOT triggered by a translation-only nudge — see _refresh_active_preview_layer.""" dx_ref, dy_ref = compute._pixel_pitch_mm(self._sras, self._ref_angle_idx) - factor = self._downsample_factor - dx_c, dy_c = dx_ref * factor, dy_ref * factor - origin, shape = compute.union_canvas_mm( - self._sras, self._ref_angle_idx, dx_c, dy_c, self._angle_params, - self._pivot_mm, margin_frac=self._PREVIEW_MARGIN_FRAC) + fy, fx = self._downsample + pitch = (dx_ref * fx, dy_ref * fy) + origin, shape = compute.canvas_for_params( + self._sras, self._ref_angle_idx, pitch, self._angle_params, + margin_frac=self._PREVIEW_MARGIN_FRAC, snap=False) self._preview_origin_mm, self._preview_shape = origin, shape - self._preview_dx_mm, self._preview_dy_mm = dx_c, dy_c + self._preview_pitch_mm = pitch self._preview_layers = { - a: compute.reproject_mask( - self._sras, a, self._ref_angle_idx, self._masks_small[a], - self._angle_params[a].rotation_deg, self._angle_params[a].shift_mm, - dx_c, dy_c, origin, shape, self._pivot_mm) - for a in range(self._sras.n_angles) + a: self._reproject(a) for a in range(self._sras.n_angles) } self._redraw_overlay() + def _reproject(self, angle_idx: int) -> np.ndarray: + """One angle's downsampled mask on the current preview canvas. + src_downsample must match _masks_small's block-mean factors, or the + layer lands magnified and offset instead of where the alignment + actually puts it.""" + p = self._angle_params[angle_idx] + return compute.reproject_mask( + self._sras, angle_idx, self._ref_angle_idx, + self._masks_small[angle_idx], p.rotation_deg, p.shift_mm, + self._preview_pitch_mm, self._preview_origin_mm, self._preview_shape, + src_downsample=self._downsample) + def _refresh_active_preview_layer(self): """Cheap path for a translation-only nudge/edit of the active angle: reproject just that one angle's downsampled mask onto the *existing* preview canvas — every other angle's cached layer is untouched.""" - a = self._active_angle - self._preview_layers[a] = compute.reproject_mask( - self._sras, a, self._ref_angle_idx, self._masks_small[a], - self._angle_params[a].rotation_deg, self._angle_params[a].shift_mm, - self._preview_dx_mm, self._preview_dy_mm, - self._preview_origin_mm, self._preview_shape, self._pivot_mm) + self._preview_layers[self._active_angle] = self._reproject(self._active_angle) self._redraw_overlay() def _redraw_overlay(self): @@ -1182,7 +1198,7 @@ class ManualAlignmentDialog(QDialog): rgba[..., 3] = fg_a + rgba[..., 3] * (1 - fg_a) x0, y0 = self._preview_origin_mm - dx, dy = self._preview_dx_mm, self._preview_dy_mm + dx, dy = self._preview_pitch_mm x_axis = x0 + np.arange(n_cols) * dx y_axis = y0 + np.arange(n_rows) * dy extent = [x_axis[0] - dx / 2, x_axis[-1] + dx / 2, @@ -1258,18 +1274,29 @@ class ManualAlignmentDialog(QDialog): # ------------------------------------------------------------------ def _on_auto_derotate(self): + """Seed every angle's rotation from the stage's reported angle. + + A starting point for nudging by eye, not an alignment: the stage's + sign convention relative to this module's is not knowable from the + file, so the sign that lines the scans up is whichever of the two looks + right in the overlay. Auto Cross-Correlate decides that from the images + instead, and is the button to reach for first. + """ + sign = -1.0 if self._derotate_sign_flipped else 1.0 + self._derotate_sign_flipped = not self._derotate_sign_flipped n_changed = 0 for a in range(self._sras.n_angles): if a == self._ref_angle_idx: continue - self._angle_params[a].rotation_deg = compute._theta_deg( + self._angle_params[a].rotation_deg = sign * compute._nominal_delta_deg( self._sras, a, self._ref_angle_idx) n_changed += 1 self._sync_active_spinboxes() self._rebuild_preview_canvas() self.lbl_status.setText( - f"Rotation set to the known scan angle for {n_changed} angle(s) " - "(translation left untouched).") + f"Rotation set to the stage angle ({'−' if sign < 0 else '+'}delta) " + f"for {n_changed} angle(s); translation untouched. Click again to " + "try the opposite sign.") def _on_auto_correlate(self): if not self._masks_ready: @@ -1277,13 +1304,14 @@ class ManualAlignmentDialog(QDialog): angles = [a for a in range(self._sras.n_angles) if a != self._ref_angle_idx] if not angles: return - use_mask = self.combo_correlate_source.currentIndex() == 1 worker = CrossCorrelateWorker( - self._sras, self._ref_angle_idx, angles, self._dc4_mv, self._pivot_mm, - use_mask=use_mask, dc_threshold_mv=self.spin_mask_threshold_mv.value(), - margin_frac=self.spin_correlate_margin.value()) + self._sras, self._ref_angle_idx, angles, self._dc4_mv, + sources=self.combo_correlate_source.currentData(), + dc_threshold_mv=self.spin_mask_threshold_mv.value(), + search_deg=self.spin_correlate_search_deg.value()) self._correlate_done_count = 0 self._correlate_total = len(angles) + self._fit_notes = {} self._set_controls_enabled(False) self.lbl_status.setText(f"Cross-correlating: 0/{self._correlate_total} angle(s)…") started = self._parent._run_worker( @@ -1298,8 +1326,10 @@ class ManualAlignmentDialog(QDialog): self.lbl_status.setText("Could not start cross-correlation (busy) — try again.") def _on_correlate_angle_done(self, angle_idx: int, rotation_deg: float, - shift_x_mm: float, shift_y_mm: float): + shift_x_mm: float, shift_y_mm: float, + score: float, source: str): self._angle_params[angle_idx] = ManualAngleParams(rotation_deg, (shift_x_mm, shift_y_mm)) + self._fit_notes[angle_idx] = (score, source) self._correlate_done_count += 1 self.lbl_status.setText( f"Cross-correlating: {self._correlate_done_count}/{self._correlate_total} angle(s)…") @@ -1311,20 +1341,47 @@ class ManualAlignmentDialog(QDialog): self._sync_active_spinboxes() self._rebuild_preview_canvas() self._set_controls_enabled(True) - source = "thresholded mask" if self.combo_correlate_source.currentIndex() == 1 \ - else "raw signal" self.lbl_status.setText( f"Cross-correlated {self._correlate_done_count} angle(s) against " - f"Angle {self._ref_angle_idx} using the {source}. Nudge from here " - "for any remaining fine correction.") + f"Angle {self._ref_angle_idx}.\n" + self._fit_report()) + + def _fit_report(self) -> str: + """Per-angle registration quality, worst first. + + Surfaced rather than buried because a single bad acquisition (stage + glitch, laser dropout) registers poorly and would otherwise be fused in + silently — seeing which angle it is, is what makes dropping it with + sras_edit_scans.py actionable. The deviation from the stage's own + reported angle is shown alongside: a large one means the search and the + stage disagree, which is either a genuine mechanical error or a sign + that this angle's fit is not to be trusted. + """ + if not self._fit_notes: + return "" + rows = sorted(self._fit_notes.items(), key=lambda kv: kv[1][0]) + worst = rows[0] + lines = [f"Worst fit: angle {worst[0]} (score {worst[1][0]:.3f}, " + f"{worst[1][1]})."] + drifted = [] + for a, _note in rows: + nominal = compute._nominal_delta_deg(self._sras, a, self._ref_angle_idx) + got = self._angle_params[a].rotation_deg + dev = min(abs(got - nominal), abs(got + nominal)) + if dev > 1.0: + drifted.append(f"{a} ({dev:.2f}°)") + if drifted: + lines.append("Rotation differs from the stage angle by >1° for " + "angle(s) " + ", ".join(drifted) + ".") + lines.append("Nudge from here for any remaining fine correction.") + return " ".join(lines) def _on_save(self): threshold = self.spin_mask_threshold_mv.value() resolved = dict(self._angle_params) # already concrete floats try: path = save_manual_alignment(self._sras, self._ref_angle_idx, threshold, resolved) - result = build_manual_alignment(self._sras, self._ref_angle_idx, threshold, - resolved, self._pivot_mm) + result = build_manual_alignment(self._sras, self._ref_angle_idx, + threshold, resolved) except OSError as exc: QMessageBox.warning(self, "Save Alignment Failed", str(exc)) return @@ -1348,6 +1405,7 @@ class ManualAlignmentDialog(QDialog): f"Could not delete the saved alignment file: {exc}") return self._angle_params = {a: ManualAngleParams() for a in range(self._sras.n_angles)} + self._fit_notes = {} self._sync_active_spinboxes() self._rebuild_preview_canvas() self.lbl_status.setText( diff --git a/sras_viewer_requirements.txt b/sras_viewer_requirements.txt index 05adf9a..1448342 100644 --- a/sras_viewer_requirements.txt +++ b/sras_viewer_requirements.txt @@ -1,3 +1,11 @@ PyQt6==6.10.2 numpy==2.4.1 matplotlib==3.10.8 +scipy==1.18.0 +# Angle alignment only: masked FFT phase correlation, which registers scans +# whose valid (scanned) regions differ in shape — see sras_compute's +# _masked_shift. +scikit-image==0.26.0 +# Optional: a faster rfft backend for the RF/FFT images (FFT Options -> pyFFTW). +# The viewer falls back to scipy.fft when it is not installed. +pyFFTW==0.15.1 diff --git a/sras_workers.py b/sras_workers.py index 5696b0e..334e655 100644 --- a/sras_workers.py +++ b/sras_workers.py @@ -282,10 +282,10 @@ class BatchCacheWorker(QObject): class AngleAlignmentWorker(QObject): - """Computes rotation+translation alignment for every angle in *sras*, - referenced to *ref_angle_idx*, from each angle's binarized CH4 mask. - Rotation is analytic (from sras.angles_deg); only translation is found by - phase correlation. + """Computes the rigid (rotation + translation, never scale) alignment for + every angle in *sras* against *ref_angle_idx*, by cross-correlating each + angle's CH4 image against the reference's. Both the rotation and the + translation are found from image content — see compute_angle_alignment. """ progress = pyqtSignal(int) # 0–100 finished = pyqtSignal(object, str) # AlignmentResult|None, error ("" = success) @@ -349,52 +349,53 @@ class Ch4MaskWorker(QObject): class CrossCorrelateWorker(QObject): - """FFT phase-correlation translation for each of *angle_indices* against - *ref_angle_idx*, for ManualAlignmentDialog's Auto Cross-Correlate button. + """Rigid registration (rotation + translation, never scale) of each of + *angle_indices* against *ref_angle_idx*, for ManualAlignmentDialog's Auto + Cross-Correlate button. - Runs on a background thread — a real many-angle, high-resolution scan's - correlation (even at its downsampled working resolution) can take long - enough that doing all of them on the GUI thread would visibly freeze the - dialog. Rotation is set to the same analytic scan-angle delta Auto - De-rotate uses alongside the correlated shift, since a translation - search is only meaningful once both angles' content is already oriented - the same way. dc4_mv/pivot_mm are the dialog's own already-in-memory - per-angle images/pivots — this worker does no fetching of its own. + Runs on a background thread — registering a real many-angle, + high-resolution scan takes long enough that doing it on the GUI thread + would visibly freeze the dialog. Rotation is *searched*, not taken from the + stage's reported angle: see compute.register_angle_to_reference, which + seeds from that angle but scores both of its signs and refines from there. + dc4_mv is the dialog's own already-in-memory per-angle CH4 image — this + worker does no fetching of its own. """ - angle_done = pyqtSignal(int, float, float, float) # angle_idx, rotation_deg, shift_x_mm, shift_y_mm + # angle_idx, rotation_deg, shift_x_mm, shift_y_mm, score, source + angle_done = pyqtSignal(int, float, float, float, float, str) finished = pyqtSignal() error = pyqtSignal(str) def __init__(self, sras: SrasFile, ref_angle_idx: int, angle_indices: list[int], - dc4_mv: dict[int, np.ndarray], pivot_mm: dict[int, tuple[float, float]], - *, use_mask: bool, dc_threshold_mv: float, margin_frac: float): + dc4_mv: dict[int, np.ndarray], *, + sources: tuple[str, ...], dc_threshold_mv: float, + search_deg: float): super().__init__() self._sras = sras self._ref = ref_angle_idx self._angles = angle_indices self._dc4_mv = dc4_mv - self._pivot_mm = pivot_mm - self._use_mask = use_mask + self._sources = sources self._threshold = dc_threshold_mv - self._margin = margin_frac + self._search_deg = search_deg - def _one(self, a: int) -> tuple[int, float, float, float]: - theta = compute._theta_deg(self._sras, a, self._ref) - dx, dy = compute.correlate_translation_mm( - self._sras, a, self._ref, self._dc4_mv, self._pivot_mm, - use_mask=self._use_mask, dc_threshold_mv=self._threshold, - margin_frac=self._margin) - return a, theta, dx, dy + def _one(self, a: int) -> tuple[int, compute.RigidFit]: + return a, compute.register_angle_to_reference( + self._sras, a, self._ref, self._dc4_mv, + dc_threshold_mv=self._threshold, sources=self._sources, + search_deg=self._search_deg) def run(self): try: - n_workers, _budget = compute.plan_angle_level(self._sras) + n_workers = compute._registration_workers( + self._sras, compute._DEFAULT_FINE_DIM) pool = ThreadPoolExecutor(max_workers=max(1, n_workers)) try: futures = [pool.submit(self._one, a) for a in self._angles] for fut in as_completed(futures): - a, theta, dx, dy = fut.result() - self.angle_done.emit(a, theta, dx, dy) + a, fit = fut.result() + self.angle_done.emit(a, fit.rotation_deg, fit.shift_mm[0], + fit.shift_mm[1], fit.score, fit.source) finally: pool.shutdown(wait=True) self.finished.emit() diff --git a/tools/make_test_sras.py b/tools/make_test_sras.py index 07386c3..71e93ef 100644 --- a/tools/make_test_sras.py +++ b/tools/make_test_sras.py @@ -125,6 +125,138 @@ def write(path: Path, n_angles: int = 3, seed: int = 0, return meta +# --------------------------------------------------------------------------- +# Rotating-sample scan: one shape, imaged at several known rotations +# --------------------------------------------------------------------------- +# +# The scan the angle-alignment path actually has to solve: every angle images +# the *same* sample at a different known rotation and offset, and a correct +# alignment stacks them all back into one shape. Two properties are +# deliberately hostile: +# +# * every angle gets a different window size and a different, meaningless +# stage x_start / y0 — alignment must ignore per-angle stage coordinates +# entirely, so any code that reads them will visibly fail here; +# * the pixel grid is strongly anisotropic (5 µm along x, 50 µm along y), +# like the real instrument, so any registration that rotates raw indices +# instead of millimetres shears the image and cannot converge. + +_ROT_DX_MM = 0.005 # x pitch, from velocity/laser_freq below +_ROT_DY_MM = 0.05 # row spacing +_ROT_BG_MV = 4.0 +_ROT_FG_MV = 160.0 + + +# How far the sample sits from the rotation axis. Non-zero on purpose: on the +# real instrument every angle's scan window is centred on the rotation axis +# while the sample is not, so each scan sees the sample somewhere else along a +# circle. That offset is exactly what a wrong rotation pivot turns into a ring +# of scans instead of a stack, so a centred test sample would hide the bug. +_ROT_SAMPLE_OFFSET_MM = (0.55, 0.40) + + +def _sample_shape_mv(u: np.ndarray, v: np.ndarray) -> np.ndarray: + """An asymmetric test sample in its own mm frame, chirally distinct at + every rotation (no 180° ambiguity) and with structure at several radii so + rotation is well determined.""" + u = u - _ROT_SAMPLE_OFFSET_MM[0] + v = v - _ROT_SAMPLE_OFFSET_MM[1] + img = np.full(u.shape, _ROT_BG_MV, dtype=np.float32) + img[((u / 0.85) ** 2 + (v / 0.40) ** 2) <= 1.0] = _ROT_FG_MV # bar + img[(np.abs(u - 0.55) <= 0.22) & (np.abs(v - 0.62) <= 0.22)] = _ROT_FG_MV # nub + img[((u + 0.75) ** 2 + (v + 0.30) ** 2) <= 0.20 ** 2] = _ROT_FG_MV # dot + return img + + +def _rot(theta_deg: float) -> np.ndarray: + t = np.radians(theta_deg) + c, s = np.cos(t), np.sin(t) + return np.array([[c, -s], [s, c]]) + + +def write_rotating(path: Path, n_angles: int = 5, samples_per_frame: int = 4, + seed: int = 0) -> dict: + """Write a v6 file whose CH4 DC image is one sample seen at n_angles known + rotations, and return the ground truth each angle should register to. + + ``truth[a] = (rotation_deg, (shift_x_mm, shift_y_mm))`` is the rigid map + from angle *a*'s local mm (origin at its own array center) to angle 0's — + exactly what ``register_angle_to_reference`` is supposed to recover. + """ + rng = np.random.default_rng(seed) + n_ch, bps = 3, 1 + cal = [(1.5625e-3, -87.04, 0.0), (2.0e-3, -60.0, 1.0e-3), (2.5e-3, -40.0, -2.0e-3)] + ymult_mv, yoff, yzero_mv = cal[2][0] * 1000, cal[2][1], cal[2][2] * 1000 + + stage_angles, geom, x_starts, y_starts, thetas, offsets = [], [], [], [], [], [] + for a in range(n_angles): + stage = -37.0 * a # what the rotation stage reports + stage_angles.append(stage) + # The true image rotation is the negative of the stage's reported + # angle: the stage's positive sense is the opposite of math-positive + # (x toward y) in scan mm. Nothing may depend on knowing that — the + # registration search tries both signs. + thetas.append(-stage) + offsets.append((0.0, 0.0) if a == 0 + else (float(rng.uniform(-0.3, 0.3)), float(rng.uniform(-0.3, 0.3)))) + # A different window per angle, all centred on the same array center — + # the real instrument grows each angle's axis-aligned bounding box to + # cover the rotated ROI. Sized so the off-axis sample stays inside every + # window at every angle, keeping the expected result unambiguous. + geom.append((88 + 8 * a, 780 + 60 * a)) + # Meaningless per-angle stage positions: correct alignment never reads + # them, so scattering them proves it. + x_starts.append(float(20.0 + rng.uniform(-6.0, 6.0))) + y_starts.append(float(30.0 + rng.uniform(-6.0, 6.0))) + + out = bytearray() + out += struct.pack( + HDR_FMT_V6, b"SRAS", 6, n_angles, + x_starts[0], y_starts[0], 1.0, 1.0, _ROT_DY_MM, + _VELOCITY_MM_S, _VELOCITY_MM_S / _ROT_DX_MM, # velocity/freq -> 5 µm pitch + samples_per_frame, _SAMPLE_RATE_HZ, bps, n_ch, + ) + out += np.array(stage_angles, dtype=">f4").tobytes() + for a, (n_rows, n_frames) in enumerate(geom): + out += struct.pack(GEO_FMT_V6, x_starts[a], 1.0, n_frames, n_rows) + for a, (n_rows, _) in enumerate(geom): + out += (y_starts[a] + np.arange(n_rows) * _ROT_DY_MM).astype(">f4").tobytes() + for ymult_v, yoff_a, yzero_v in cal: + p = _preamble(ymult_v, yoff_a, yzero_v) + out += struct.pack(">H", len(p)) + p + background = rng.integers(-8, 9, size=samples_per_frame, dtype=np.int8) + out += struct.pack(">I", samples_per_frame) + background.tobytes() + + truth, dc4_images = {}, [] + for a, (n_rows, n_frames) in enumerate(geom): + # Local mm of every pixel, measured from this angle's own array center. + lx = (np.arange(n_frames) - (n_frames - 1) / 2.0) * _ROT_DX_MM + ly = (np.arange(n_rows) - (n_rows - 1) / 2.0) * _ROT_DY_MM + gx, gy = np.meshgrid(lx, ly) + # local = R(theta) @ sample + offset, so sample = R(theta)^T @ (local - offset) + rel = np.stack([gx - offsets[a][0], gy - offsets[a][1]], axis=-1) + s = rel @ _rot(thetas[a]) # == rel @ R^T.T == R^T @ rel + dc4 = _sample_shape_mv(s[..., 0], s[..., 1]) + dc4_images.append(dc4) + + inv = _rot(-thetas[a]) + truth[a] = (-thetas[a], + tuple(float(v) for v in -(inv @ np.array(offsets[a])))) + + adc4 = np.clip(np.round((dc4 - yzero_mv) / ymult_mv + yoff), -128, 127).astype(np.int8) + block = np.zeros((n_rows, n_ch, n_frames, samples_per_frame), dtype=np.int8) + block[:, 2] = adc4[:, :, None] # CH4 carries the sample + block[:, 1] = 10 # CH3 flat + block[:, 0] = rng.integers(-40, 41, size=(n_rows, n_frames, samples_per_frame), + dtype=np.int8) # CH1 noise + out += block.tobytes() + + path.write_bytes(bytes(out)) + return {"n_angles": n_angles, "geometry": geom, "stage_angles_deg": stage_angles, + "truth": truth, "dc4_mv": dc4_images, "x_starts": x_starts, + "y_starts": y_starts, "dx_mm": _ROT_DX_MM, "dy_mm": _ROT_DY_MM} + + HDR_FMT_LEGACY = ">4sBHHffffIIdBB" diff --git a/tools/test_alignment.py b/tools/test_alignment.py new file mode 100644 index 0000000..f78f9fa --- /dev/null +++ b/tools/test_alignment.py @@ -0,0 +1,223 @@ +#!/usr/bin/env python3 +"""Angle-alignment tests: does registration actually stack the scans? + +Builds a synthetic scan in which one sample is imaged at several *known* +rotations and offsets (tools/make_test_sras.write_rotating) and checks that the +alignment path recovers them, that the shared canvas is angle 0's own pixel +grid extended, and that nothing in the result depends on any other angle's +stage coordinates. + +No Qt — this exercises sras_compute directly. See tools/test_gui.py for the +dialog and Aligned-View plumbing. + +Usage: python tools/test_alignment.py +""" + +import sys +import tempfile +from pathlib import Path + +import numpy as np + +sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) + +import sras_compute as compute # noqa: E402 +from sras_format import CH4_IDX, SrasFile, adc_to_mv # noqa: E402 +import tools.make_test_sras as gen # noqa: E402 + +# Registration is limited by how far a feature moves per degree: with this +# sample's ~1 mm radius and a ~16 µm registration pitch, a quarter degree is +# already sub-pixel, so it is the floor of what any metric can resolve here. +_ROT_TOL_DEG = 0.5 +_SHIFT_TOL_MM = 0.02 +_STACK_IOU_MIN = 0.90 +_THRESHOLD_MV = 80.0 + +_failures: list[str] = [] + + +def check(name: str, ok: bool, detail: str = ""): + print(f" {'PASS' if ok else 'FAIL'} {name}" + (f" — {detail}" if detail else "")) + if not ok: + _failures.append(name) + + +def dc4_images(sras: SrasFile) -> dict[int, np.ndarray]: + return {a: adc_to_mv(compute.compute_dc_image(sras, a, CH4_IDX), *sras.cal(CH4_IDX)) + for a in range(sras.n_angles)} + + +def mm_transform(sras: SrasFile, result, angle_idx: int) -> np.ndarray: + """Recover the pure mm-space rotation from a canvas->raw affine. + + matrix == D @ R^T @ A_out, where A_out and D only carry the canvas and + per-angle pixel pitches; undoing both must leave something orthonormal, or + the transform is smuggling in a scale or a shear. + """ + dx_a, dy_a = compute._pixel_pitch_mm(sras, angle_idx) + A_out = np.array([[0.0, result.canvas_dx_mm], [result.canvas_dy_mm, 0.0]]) + D = np.array([[0.0, 1.0 / dy_a], [1.0 / dx_a, 0.0]]) + return np.linalg.inv(D) @ result.per_angle[angle_idx].matrix @ np.linalg.inv(A_out) + + +def main() -> int: + tmpdir = Path(tempfile.mkdtemp(prefix="sras_align_")) + path = tmpdir / "rotating.sras" + meta = gen.write_rotating(path, n_angles=5) + sras = SrasFile(str(path)) + truth = meta["truth"] + + print(f"\nrotating-sample scan: {sras.n_angles} angles, " + f"shapes {[sras.image_shape(a) for a in range(sras.n_angles)]}") + + print("\nper-angle rigid registration (rotation + translation, no scale)") + dc4 = dc4_images(sras) + fits = {a: compute.register_angle_to_reference( + sras, a, 0, dc4, dc_threshold_mv=_THRESHOLD_MV) + for a in range(sras.n_angles)} + for a, fit in fits.items(): + t_rot, t_shift = truth[a] + rot_err = abs(fit.rotation_deg - t_rot) + shift_err = float(np.hypot(fit.shift_mm[0] - t_shift[0], + fit.shift_mm[1] - t_shift[1])) + check(f"angle {a} rotation within {_ROT_TOL_DEG}° of truth", + rot_err <= _ROT_TOL_DEG, + f"got {fit.rotation_deg:.3f}°, truth {t_rot:.3f}° (err {rot_err:.3f}°)") + check(f"angle {a} translation within {_SHIFT_TOL_MM} mm of truth", + shift_err <= _SHIFT_TOL_MM, f"err {shift_err:.4f} mm") + check("reference angle registers as exact identity", + fits[0] == compute.RigidFit(0.0, (0.0, 0.0), 1.0, "reference")) + + # The stage's rotational sense relative to this module's math-positive + # convention is not knowable from the file, and the old code hardcoded a + # guess. Flipping every reported angle must therefore change nothing: the + # search scores both signs and the images decide. + flipped = SrasFile(str(path)) + flipped.angles_deg = -flipped.angles_deg + flipped_fits = {a: compute.register_angle_to_reference( + flipped, a, 0, dc4, dc_threshold_mv=_THRESHOLD_MV) + for a in range(1, flipped.n_angles)} + check("negating every reported stage angle changes no fit", + all(flipped_fits[a] == fits[a] for a in flipped_fits), + str({a: (flipped_fits[a].rotation_deg, fits[a].rotation_deg) + for a in flipped_fits if flipped_fits[a] != fits[a]})) + + print("\nper-angle stage coordinates are not consulted") + # Move every non-reference angle's scan window somewhere else entirely. + # Only angle 0's coordinates may matter, so every fit must be untouched. + moved = SrasFile(str(path)) + for a in range(1, moved.n_angles): + moved.x_start_mm[a] += 13.5 * a + moved._y_pos_per_angle[a] = moved._y_pos_per_angle[a] - 9.25 * a + moved_dc4 = dc4_images(moved) + moved_fits = {a: compute.register_angle_to_reference( + moved, a, 0, moved_dc4, dc_threshold_mv=_THRESHOLD_MV) + for a in range(1, moved.n_angles)} + check("relocating every other angle's scan window changes no fit", + all(moved_fits[a] == fits[a] for a in moved_fits), + str({a: (round(moved_fits[a].rotation_deg, 4), fits[a].rotation_deg) + for a in moved_fits if moved_fits[a] != fits[a]})) + + print("\nshared canvas is angle 0's own pixel grid, extended") + result = compute.compute_angle_alignment(sras, 0, _THRESHOLD_MV) + t0 = result.per_angle[0] + check("angle 0's transform has no rotation, scale or shear", + np.allclose(t0.matrix, np.eye(2)), str(t0.matrix)) + check("angle 0 lands on whole canvas pixels (no resampling of the reference)", + np.allclose(t0.offset, np.round(t0.offset)), str(t0.offset)) + check("canvas pitch is angle 0's own pitch", + (result.canvas_dx_mm, result.canvas_dy_mm) + == compute._pixel_pitch_mm(sras, 0)) + + n_rows, n_cols = result.canvas_shape + x_axis = result.canvas_origin_mm[0] + np.arange(n_cols) * result.canvas_dx_mm + y_axis = result.canvas_origin_mm[1] + np.arange(n_rows) * result.canvas_dy_mm + row0, col0 = int(round(-t0.offset[0])), int(round(-t0.offset[1])) + a0_rows, a0_cols = sras.image_shape(0) + check("canvas X axis reproduces angle 0's own X coordinates", + np.allclose(x_axis[col0:col0 + a0_cols], sras.x_axis_mm(0))) + check("canvas Y axis reproduces angle 0's own Y coordinates", + np.allclose(y_axis[row0:row0 + a0_rows], sras.y_positions_mm(0))) + check("canvas covers every angle's footprint", + n_rows >= max(int(sras.n_rows[a]) for a in range(sras.n_angles)) + and n_cols >= max(int(sras.n_frames[a]) for a in range(sras.n_angles)), + str(result.canvas_shape)) + + print("\nno scaling anywhere in the per-angle transforms") + for a in range(sras.n_angles): + R = mm_transform(sras, result, a) + check(f"angle {a}'s mm-space transform is a pure rotation", + np.allclose(R @ R.T, np.eye(2), atol=1e-9) + and abs(abs(np.linalg.det(R)) - 1.0) < 1e-9, + f"det={np.linalg.det(R):.6f}") + + print("\nall angles stack into one shape") + aligned = {a: compute.apply_alignment(result, a, dc4[a]) + for a in range(sras.n_angles)} + base = aligned[0] >= _THRESHOLD_MV + for a in range(1, sras.n_angles): + other = aligned[a] >= _THRESHOLD_MV + iou = float((base & other).sum()) / max(1, int((base | other).sum())) + check(f"angle {a}'s aligned sample overlaps angle 0's (IoU >= {_STACK_IOU_MIN})", + iou >= _STACK_IOU_MIN, f"IoU {iou:.4f}") + + print("\ndownsampled preview lands where the full-resolution image does") + # ManualAlignmentDialog reprojects block-mean-downsampled masks, so the + # affine has to account for the factor. When it did not, every preview + # layer came out magnified by that factor and offset — the overlay showed a + # blown-up crop of each mask, which is not something you can align by eye. + pitch = (result.canvas_dx_mm, result.canvas_dy_mm) + a = sras.n_angles - 1 + p = result.per_angle[a] + full_mask = (dc4[a] >= _THRESHOLD_MV).astype(np.float32) + full = compute.reproject_mask( + sras, a, 0, full_mask, p.rotation_deg, p.shift_mm, pitch, + result.canvas_origin_mm, result.canvas_shape) + fy, fx = 4, 16 + small = compute.reproject_mask( + sras, a, 0, compute._block_mean_2d(full_mask, fy, fx), + p.rotation_deg, p.shift_mm, (pitch[0] * fx, pitch[1] * fy), + result.canvas_origin_mm, + (result.canvas_shape[0] // fy, result.canvas_shape[1] // fx), + src_downsample=(fy, fx)) + # Compare in mm, via each layer's own center of mass. + def com_mm(layer, px, py): + rows, cols = np.nonzero(layer > 0.5) + return np.array([cols.mean() * px, rows.mean() * py]) + d = com_mm(small, pitch[0] * fx, pitch[1] * fy) - com_mm(full, *pitch) + check("a downsampled preview layer lands within a preview pixel of the " + "full-resolution one", + abs(d[0]) <= abs(pitch[0] * fx) and abs(d[1]) <= abs(pitch[1] * fy), + f"offset {d[0]:+.4f}, {d[1]:+.4f} mm") + + print("\nmanual path reproduces the same geometry") + params = {a: compute.ManualAngleParams(t.rotation_deg, t.shift_mm) + for a, t in result.per_angle.items()} + manual = compute.build_manual_alignment(sras, 0, _THRESHOLD_MV, params) + check("build_manual_alignment matches compute_angle_alignment for the same params", + manual.canvas_shape == result.canvas_shape + and np.allclose(manual.canvas_origin_mm, result.canvas_origin_mm) + and all(np.allclose(manual.per_angle[a].matrix, result.per_angle[a].matrix) + and np.allclose(manual.per_angle[a].offset, result.per_angle[a].offset) + for a in range(sras.n_angles))) + + print("\nsidecar round-trip") + compute.save_manual_alignment(sras, 0, _THRESHOLD_MV, params) + loaded = compute.load_manual_alignment(sras) + check("sidecar reloads every angle's params", + loaded is not None + and all(np.isclose(loaded.per_angle[a].rotation_deg, params[a].rotation_deg) + and np.allclose(loaded.per_angle[a].shift_mm, params[a].shift_mm) + for a in range(sras.n_angles))) + check("sidecar deletes cleanly", compute.delete_manual_alignment(sras)) + + print() + if _failures: + print(f"{len(_failures)} FAILURE(S): " + ", ".join(_failures)) + return 1 + print("All alignment checks passed.") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tools/test_gui.py b/tools/test_gui.py index 482e11a..36bca72 100644 --- a/tools/test_gui.py +++ b/tools/test_gui.py @@ -27,7 +27,7 @@ from PyQt6.QtWidgets import QApplication, QMessageBox # noqa: E402 sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) import sras_compute as compute # noqa: E402 -from sras_format import CH1_IDX, CH3_IDX, CH4_IDX # noqa: E402 +from sras_format import CH1_IDX, CH3_IDX, CH4_IDX, SrasFile # noqa: E402 from sras_viewer import RoiQuad, SrasViewerWindow, VELOCITY_MODE_IDX # noqa: E402 import tools.make_test_sras as gen # noqa: E402 @@ -240,60 +240,49 @@ def main(): print("\nmanual alignment (Fusion)") check("manual alignment action enabled", win._manual_align_act.isEnabled()) - # --- Alignment pivot is a signal-weighted centroid, not the raw bbox -- - # center, and is independent of any DC threshold (so a threshold that - # happens to leave a real angle's binary mask empty can't silently - # degrade the pivot back to the bbox center). - corner_signal = np.zeros(s.image_shape(0), dtype=np.float32) - corner_signal[0, 0] = 1.0 # single spike -> weighted centroid is exact - expected_corner = (float(s.x_axis_mm(0)[0]), float(s.y_positions_mm(0)[0])) - centroid = compute._signal_centroid_mm(s, 0, corner_signal) - check("signal-weighted centroid of a single spike pixel is that pixel exactly", - np.allclose(centroid, expected_corner), f"{centroid} vs {expected_corner}") - bbox_center = compute._bbox_center_mm(s, 0) - check("signal centroid differs from the raw scan-window bbox center", - not np.allclose(centroid, bbox_center), - f"centroid {centroid} vs bbox center {bbox_center}") + # --- Local mm is anchored on each angle's array center, not its stage -- + # position: that is what makes a scan's placement independent of where its + # window happened to sit. (Registration accuracy itself is covered by + # tools/test_alignment.py, which has a synthetic sample to register.) + n_rows, n_frames = s.image_shape(0) + check("array center is the geometric center of the pixel grid", + np.allclose(compute._center_idx(s, 0), + [(n_rows - 1) / 2, (n_frames - 1) / 2])) + dx0, dy0 = compute._pixel_pitch_mm(s, 0) + check("local half-extent is derived from shape and pitch alone", + np.allclose(compute._local_half_extent_mm(s, 0), + [(n_frames - 1) / 2 * abs(dx0), (n_rows - 1) / 2 * abs(dy0)])) + identity = {a: compute.ManualAngleParams() for a in range(s.n_angles)} + origin_a, shape_a = compute.canvas_for_params(s, 0, (dx0, dy0), identity) + moved = SrasFile(str(path)) + for a in range(1, moved.n_angles): + moved.x_start_mm[a] += 7.5 + moved._y_pos_per_angle[a] = moved._y_pos_per_angle[a] + 3.25 + origin_b, shape_b = compute.canvas_for_params(moved, 0, (dx0, dy0), identity) + check("moving every non-reference angle's scan window leaves the canvas " + "unchanged (only angle 0's coordinates are used)", + shape_a == shape_b and np.allclose(origin_a, origin_b), + f"{origin_a} {shape_a} vs {origin_b} {shape_b}") - # compute_pivot_points_mm should reuse a pre-computed dc4_mv dict rather - # than recomputing from the real DC4 image (which has no such spike and - # would give a different answer if silently recomputed). - reused_pivot = compute.compute_pivot_points_mm(s, dc4_mv={0: corner_signal})[0] - check("compute_pivot_points_mm reuses a pre-computed dc4_mv dict", - np.allclose(reused_pivot, expected_corner)) + # --- Both signs of the stage's reported angle are searched -------------- + cands = compute._rotation_candidates(30.0, 6.0, 2.0) + check("rotation candidates bracket both signs of the stage angle", + min(cands) < -29.0 and max(cands) > 29.0, f"{min(cands)}..{max(cands)}") - # A perfectly flat signal carries no information to weight by, so it - # falls back to the bbox center rather than producing a NaN/degenerate - # centroid. - flat_signal = np.full(s.image_shape(0), 5.0, dtype=np.float32) - flat_centroid = compute._signal_centroid_mm(s, 0, flat_signal) - check("a perfectly flat signal falls back to the bbox center", - np.allclose(flat_centroid, bbox_center)) - - # --- Rotation sign convention: negative of the raw angles_deg delta ---- - check("_theta_deg negates the raw angles_deg delta (GR stage's positive " - "angle is the opposite rotational sense from this module's CCW " - "math convention)", - all(np.isclose(compute._theta_deg(s, a, 0), - -(float(s.angles_deg[a]) - float(s.angles_deg[0]))) - for a in range(s.n_angles))) - - # --- FFT phase correlation recovers a known synthetic pixel shift ------ - rng = np.random.default_rng(0) - corr_ref = np.zeros((40, 50), dtype=np.float32) - corr_ref[10:25, 15:35] = 1.0 - corr_ref += 0.05 * rng.standard_normal(corr_ref.shape).astype(np.float32) - corr_mov = np.roll(corr_ref, shift=(4, -7), axis=(0, 1)) - dr, dc = compute._phase_correlate_shift(corr_ref, corr_mov) - check("phase correlation recovers the shift that aligns mov onto ref", - (dr, dc) == (-4, 7), f"got (dr, dc)={(dr, dc)}") + # --- Whole-pixel translation must not wrap content around the edge ------ + arr = np.zeros((6, 6), dtype=np.float32) + arr[0, 0] = 1.0 + check("_shift_into zero-fills rather than wrapping", + compute._shift_into(arr, -1, -1).sum() == 0.0) + check("_shift_into moves content by exactly the requested offset", + compute._shift_into(arr, 2, 3)[2, 3] == 1.0) # --- Open: must NOT seed from the still-live automatic AlignmentResult -- - # The automatic result's translation comes from FFT phase correlation -- - # the very thing manual mode exists to work around -- so manual mode - # must start from identity (centroids coincide, zero shift) regardless - # of whatever the automatic run last computed. Only a previously *saved - # manual* alignment (sidecar) should ever seed this dialog. + # Manual mode exists to fix up whatever the automatic registration got + # wrong, so it must start from identity (every angle centered on the + # reference, no rotation) regardless of whatever the automatic run last + # computed. Only a previously *saved manual* alignment (sidecar) should + # ever seed this dialog. win._on_manual_alignment() check("dialog opened", win._manual_align_dialog is not None) dlg = win._manual_align_dialog @@ -343,39 +332,41 @@ def main(): check("a real Right-arrow key event nudged shift_x", dlg._angle_params[active].shift_mm[0] > before[0]) - # --- Auto De-rotate: rotation only, translation untouched --------------- + # --- Auto De-rotate: seeds rotation from the stage angle, no translation - shift_before_derotate = dlg._angle_params[active].shift_mm dlg._on_auto_derotate() - expected_theta = compute._theta_deg(s, active, dlg._ref_angle_idx) - check("auto de-rotate set the known analytic angle", - abs(dlg._angle_params[active].rotation_deg - expected_theta) < 1e-6) + nominal = compute._nominal_delta_deg(s, active, dlg._ref_angle_idx) + check("auto de-rotate seeded rotation from the stage's reported angle", + abs(dlg._angle_params[active].rotation_deg - nominal) < 1e-6) check("auto de-rotate left translation untouched", dlg._angle_params[active].shift_mm == shift_before_derotate) check("reference angle stays identity after auto de-rotate", dlg._angle_params[dlg._ref_angle_idx].rotation_deg == 0.0) + # Clicking again offers the other sign, since which one lines the scans up + # is not knowable from the file. + dlg._on_auto_derotate() + check("auto de-rotate offers the opposite sign on a second click", + abs(dlg._angle_params[active].rotation_deg + nominal) < 1e-6) - # --- Auto Cross-Correlate: rotation + FFT-correlated shift, backgrounded - + # --- Auto Cross-Correlate: searches rotation *and* translation ---------- check("cross-correlate action enabled once masks are ready", dlg.btn_auto_correlate.isEnabled()) - dlg._on_auto_correlate() - check("auto cross-correlate completed", wait_until( - lambda: not win._job_running("manual_align_correlate"), timeout_ms=30000)) - check("auto cross-correlate set the known analytic angle for every angle", - all(abs(dlg._angle_params[a].rotation_deg - - compute._theta_deg(s, a, dlg._ref_angle_idx)) < 1e-6 - for a in range(s.n_angles) if a != dlg._ref_angle_idx)) + for label_idx, (label, _sources) in enumerate(dlg._CORRELATE_SOURCES): + dlg.combo_correlate_source.setCurrentIndex(label_idx) + dlg._on_auto_correlate() + check(f"auto cross-correlate completed ({label})", wait_until( + lambda: not win._job_running("manual_align_correlate"), timeout_ms=60000)) + check(f"every non-reference angle got a fit ({label})", + all(a in dlg._fit_notes for a in range(s.n_angles) + if a != dlg._ref_angle_idx)) check("auto cross-correlate reference angle stays identity", dlg._angle_params[dlg._ref_angle_idx] == compute.ManualAngleParams()) check("auto cross-correlate re-enabled controls when done", dlg.grp_correlate.isEnabled() and dlg.btn_save.isEnabled()) check("preview canvas rebuilt after cross-correlate", len(dlg._preview_layers) == s.n_angles) - - # The thresholded-mask option should also work end to end. - dlg.combo_correlate_source.setCurrentIndex(1) # thresholded mask - dlg._on_auto_correlate() - check("auto cross-correlate (thresholded-mask option) completed", wait_until( - lambda: not win._job_running("manual_align_correlate"), timeout_ms=30000)) + check("fit quality is reported per angle", bool(dlg._fit_report()), + dlg._fit_report()) # --- Save ----------------------------------------------------------------- dlg._on_save() From dc513e3fd090f8999bac760c1a45318754fd6004 Mon Sep 17 00:00:00 2001 From: Thomas Ales Date: Thu, 6 Aug 2026 09:36:02 -0500 Subject: [PATCH 02/10] Housekeeping: extend .gitignore, drop stale v2-v5 format docs SRAS_FORMAT.md/.html documented only v2-v5 and were superseded by scan_format.md in July. Ignore .DS_Store, screenshots, pytest cache, and test scratch outputs. Co-Authored-By: Claude Fable 5 --- .gitignore | 6 + SRAS_FORMAT.html | 761 ----------------------------------------------- SRAS_FORMAT.md | 379 ----------------------- 3 files changed, 6 insertions(+), 1140 deletions(-) delete mode 100644 SRAS_FORMAT.html delete mode 100644 SRAS_FORMAT.md diff --git a/.gitignore b/.gitignore index 670a936..2694b86 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1,8 @@ __pycache__/ .venv/ +.DS_Store +.pytest_cache/ +*.png +*.sras +baseline*.txt +after*.txt diff --git a/SRAS_FORMAT.html b/SRAS_FORMAT.html deleted file mode 100644 index 204a4e5..0000000 --- a/SRAS_FORMAT.html +++ /dev/null @@ -1,761 +0,0 @@ - - - - - SRAS File Format Specification - - - - - - - - - - - - -

SRAS File Format Specification

-

Format family: .sras
-Byte order: Big-endian (network byte order) throughout, unless noted.
-Version history: v2 (baseline), v3 (scope calibration), v4 (background waveform), v5 (precomputed images + guaranteed frame count).

-
-

Table of Contents

-
    -
  1. Overview
  2. -
  3. Type notation
  4. -
  5. Version history
  6. -
  7. File structure - -
  8. -
  9. Derived quantities
  10. -
  11. ADC calibration
  12. -
  13. Waveform data layout detail
  14. -
  15. Size reference
  16. -
  17. Compatibility notes
  18. -
-
-

Overview

-

An SRAS file stores the raw RF waveforms captured during a Surface-acoustic-wave Resonance And Spectroscopy (SRAS) scan, along with the scan geometry and scope calibration metadata needed to interpret them.

-

A scan consists of one or more angles (rotation positions of the sample), each containing a 2-D raster of rows × frames. At every pixel, n_channels waveforms of samples_per_frame ADC counts are stored. Channel order is fixed:

- - - - - - - - - - - - - - - - - - - - - - - - - -
IndexHardware channelSignal
0CH1RF acoustic packet (AC-coupled)
1CH3Bias A — DC mean used for masking
2CH4Bias B — DC mean used for masking
-
-

Type notation

- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
SymbolC typeSizeNotes
u8uint8_t1 byteunsigned
u16uint16_t2 bytesbig-endian
u32uint32_t4 bytesbig-endian
i8int8_t1 bytesigned, used for ADC samples when bytes_per_sample == 1
i16int16_t2 bytesbig-endian signed, used when bytes_per_sample == 2
f32float4 bytesbig-endian IEEE 754 single
f64double8 bytesbig-endian IEEE 754 double
char[N]—N bytesraw bytes, no null terminator unless noted
utf8[N]—N bytesUTF-8 string, length-prefixed (see preamble section)
-
-

Version history

- - - - - - - - - - - - - - - - - - - - - - - - - -
VersionAdded
2Baseline: fixed header, angle table, row table, raw waveform data. No scope calibration (fallback constants used by readers).
3Per-channel Tektronix WFMOutpre preamble strings carrying YMULT / YOFF / YZERO calibration.
4Background waveform section: one CH1 reference shot subtracted from each CH1 frame before FFT.
5(this document) Version byte incremented to 5. n_frames_hdr is now the actual acquired frame count (authoritative). PREC section appended after waveform data with precomputed FFT-peak and DC images for instant re-display.
-
-

v2 note: Version 1 is not defined; version 2 is the lowest observed in the field.

-
-
-

File structure

-
┌─────────────────────────────────────────────┐
-│  1. Fixed header                  (43 bytes) │  all versions
-├─────────────────────────────────────────────┤
-│  2. Angle table          (n_angles × 4 bytes)│  all versions
-├─────────────────────────────────────────────┤
-│  3. Row position table     (n_rows × 4 bytes)│  all versions
-├─────────────────────────────────────────────┤
-│  4. Channel preambles           (variable)   │  v3+
-├─────────────────────────────────────────────┤
-│  5. Background waveform         (variable)   │  v4+
-├─────────────────────────────────────────────┤
-│  6. Waveform data               (variable)   │  all versions
-├─────────────────────────────────────────────┤
-│  7. PREC section                (variable)   │  v5 only
-└─────────────────────────────────────────────┘
-
-
-

1. Fixed header (43 bytes, all versions)

- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
OffsetSizeTypeFieldDescription
04char[4]magicSRAS (ASCII, no null terminator). Reject file if this does not match.
41u8versionFormat version. This document describes version 5.
52u16n_anglesNumber of scan angles (rotation positions). ≥ 1.
72u16n_rowsNumber of scan rows (Y positions). ≥ 1.
94f32x_start_mmX position of the first frame in the first row, in millimetres.
134f32x_delta_mmPre-computed pixel pitch in mm (velocity_mm_s / laser_freq_hz). Provided for convenience; readers should prefer the derived value.
174f32velocity_mm_sScanner stage velocity, mm s⁻¹. Used together with laser_freq_hz to compute pixel pitch.
214f32laser_freq_hzLaser repetition rate, Hz.
254u32n_frames_hdrv2–v4: the configured frame count written before acquisition; may exceed actual frames acquired (use file-size arithmetic to obtain the true count). v5: the actual acquired frame count — authoritative; readers must not re-derive it from file size.
294u32samples_per_frameADC samples per waveform (spf).
338f64sample_rate_hzOscilloscope sample rate, Hz (e.g. 5 × 10⁹ for 5 GS/s).
411u8bytes_per_sampleADC word size: 1 → i8, 2 → i16 (big-endian).
421u8n_channelsNumber of channels per frame. Currently always 3.
-
-

2. Angle table (all versions)

-

Immediately follows the fixed header.

-
n_angles × f32   — scan angle in degrees
-
-

Each entry is a big-endian f32 giving the sample rotation angle in degrees at which that angle index was acquired.

-
-

3. Row position table (all versions)

-

Immediately follows the angle table.

-
n_rows × f32   — Y position of each row, in millimetres
-
-
-

4. Channel preambles (v3+)

-

One entry per channel, in channel-index order (CH1 first).

-
for each channel:
-    u16          preamble_length   — byte count of the UTF-8 string that follows
-    utf8[N]      preamble          — Tektronix WFMOutpre string
-
-

The preamble is the oscilloscope's WFMOutpre response string. Readers extract the following keys (case-insensitive, space-separated value):

- - - - - - - - - - - - - - - - - - - - - - - - - -
KeyStored unitConversion to mV
YMULTV count⁻¹multiply by 1000
YOFFADC countsused directly
YZEROVmultiply by 1000
-

v2 fallback: when preambles are absent, readers use:

-
    -
  • YMULT = 1.5625 mV count⁻¹ (50 mV/div, 8 div, 8-bit ADC)
  • -
  • YOFF = −87.04 ADC counts (scope position = −2.72 div)
  • -
  • YZERO = 0 mV
  • -
-
-

5. Background waveform (v4+)

-
u32          n_bg_samples   — number of i8 ADC samples that follow
-i8[n_bg]     background     — one representative CH1 background shot
-
-

The background waveform has the same samples_per_frame length as a normal CH1 waveform. It is subtracted from each CH1 waveform before FFT processing when background subtraction is enabled. When n_bg_samples == 0 the section is present but empty.

-
-

6. Waveform data (all versions)

-

Begins immediately after the fixed header (v2), preambles (v3), or background waveform (v4+). The waveform data is a flat, contiguous array with the following logical shape, stored in row-major (C) order:

-
waveform_data[n_angles][n_rows][n_channels][n_frames][samples_per_frame]
-
-

Each element is a signed ADC count of size bytes_per_sample:

-
    -
  • bytes_per_sample == 1 → i8
  • -
  • bytes_per_sample == 2 → i16 big-endian
  • -
-

Total byte count:

-
waveform_bytes = n_angles × n_rows × n_channels × n_frames × samples_per_frame × bytes_per_sample
-
-

Index semantics

- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
DimensionRangeMeaning
[a]0 … n_angles−1Scan angle (rotation position)
[r]0 … n_rows−1Row (Y position); row 0 is the first acquired
[c]0 … n_channels−1Channel (0=CH1 RF, 1=CH3 Bias A, 2=CH4 Bias B)
[f]0 … n_frames−1Frame (X position) within the row
[s]0 … spf−1Sample index within the waveform
-

Frame-count determination

-
    -
  • v5: use n_frames_hdr directly; do not use file-size arithmetic.
  • -
  • v2–v4: n_frames = floor((file_bytes_after_header_sections) / (bytes_per_sample × n_angles × n_rows × n_channels × samples_per_frame)). Any remainder bytes are a partial trailing row and are discarded.
  • -
-
-

7. PREC section (v5)

-

The PREC section is appended immediately after the waveform data and is present if and only if version == 5 and the file size exceeds waveform_end_offset.

-
waveform_end_offset = data_offset + waveform_bytes
-
-

where data_offset is the file offset of the first waveform byte (the byte immediately after the background waveform, or after the angle/row tables for v2 files).

-

PREC header (8 bytes)

- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
Offset (relative)SizeTypeFieldDescription
04char[4]prec_magicPREC (ASCII). Absent or wrong magic → ignore section.
41u8prec_versionPREC format version. Currently 1.
51u8flagsBitmask (see below).
62u16n_storedNumber of angle entries that follow. 0 ≤ n_stored ≤ n_angles.
-
Flags byte
- - - - - - - - - - - - - - - - - - - - -
BitMaskMeaning when set
00x01bg_sub_applied — background waveform was subtracted from CH1 before the FFT when these images were computed.
1–7—Reserved, must be zero on write; readers must ignore.
-

PREC angle entries

-

Repeated n_stored times, in arbitrary angle-index order:

-
for each stored angle:
-    u16                   angle_idx      — index into the angle table (0-based)
-    f32[n_rows×n_frames]  peak_freq_mhz  — CH1 FFT peak frequency, MHz, row-major
-    f32[n_rows×n_frames]  dc4_mv         — CH4 waveform mean, mV, row-major
-    f32[n_rows×n_frames]  dc3_mv         — CH3 waveform mean, mV, row-major
-
-

All image arrays are f32 big-endian, stored in row-major order: element [r][f] is at offset (r × n_frames + f) × 4 bytes within the array.

-

peak_freq_mhz is computed without any DC-threshold masking (i.e. the FFT is run on every pixel unconditionally). Readers apply the dc4_mv threshold at display time:

-
pixel is valid  ⟺  dc4_mv[r][f]  ≥  threshold_mv
-display_value   =   peak_freq_mhz[r][f]  if valid, else 0
-
-

dc4_mv / dc3_mv are the mean of all ADC samples in the respective channel waveform, converted to millivolts using the channel calibration:

-
dc_mv = (adc_mean − YOFF) × YMULT + YZERO
-
-

When readers must bypass the PREC fast path

-

Readers must fall back to real-time FFT computation (ignoring stored peak_freq_mhz) when any of the following are true:

-
    -
  • Time-domain gating is active (zeroing samples outside a time window changes peak frequency).
  • -
  • Zero-padding (n_fft ≠ samples_per_frame) is requested (changes bin spacing).
  • -
  • The reader's background-subtraction setting does not match flags.bg_sub_applied.
  • -
-
-

Derived quantities

-
pixel_pitch_mm  = velocity_mm_s / laser_freq_hz
-
-x_axis_mm[f]    = x_start_mm + f × pixel_pitch_mm          (f = 0 … n_frames−1)
-
-time_axis_ns[s] = s / sample_rate_hz × 1e9                 (s = 0 … spf−1)
-
-freq_axis_mhz[k] = k × sample_rate_hz / (n_fft × 1e6)      (k = 0 … n_fft/2)
-                   where n_fft = samples_per_frame unless zero-padding is active
-
-velocity_ms[r][f] = peak_freq_mhz[r][f] × grating_um       (grating_um user-supplied)
-
-
-

ADC calibration

-

Convert raw ADC counts to millivolts:

-
voltage_mv = (adc_count − YOFF) × YMULT_mv + YZERO_mv
-
-

Invert (mV → ADC count):

-
adc_count = (voltage_mv − YZERO_mv) / YMULT_mv + YOFF
-
-

where YMULT_mv is YMULT in mV count⁻¹ (= scope YMULT in V count⁻¹ × 1000).

-
-

Waveform data layout detail

-

For a scan with n_angles=2, n_rows=3, n_channels=3, n_frames=4, spf=5 the layout is:

-
angle 0
-  row 0
-    CH1: [s0 s1 s2 s3 s4] [s0 s1 s2 s3 s4] [s0 s1 s2 s3 s4] [s0 s1 s2 s3 s4]
-          frame 0           frame 1           frame 2           frame 3
-    CH3: …(same layout)…
-    CH4: …(same layout)…
-  row 1
-    …
-  row 2
-    …
-angle 1
-  …
-
-

The flat byte offset of sample s of frame f, channel c, row r, angle a is:

-
offset = data_offset
-       + (a × n_rows × n_channels × n_frames × spf
-         + r × n_channels × n_frames × spf
-         + c × n_frames × spf
-         + f × spf
-         + s)
-       × bytes_per_sample
-
-
-

Size reference

-

Approximate sizes for representative scans (bytes_per_sample = 1, n_channels = 3).

- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
n_anglesn_rowsn_framesspfWaveform dataPREC section
1500500400300 MB12 MB
45005004001.2 GB48 MB
1200020004004.8 GB48 MB
42000200040019.2 GB192 MB
82000200040038.4 GB384 MB
162000200040076.8 GB768 MB
-

PREC section size formula:

-
prec_bytes = 8  +  n_stored × (2 + 3 × n_rows × n_frames × 4)
-
-
-

Compatibility notes

-

Reading v5 files with a v4 reader

-

A v4 reader that only accepts versions {2, 3, 4} will reject a v5 file with an "unsupported version" error. This is intentional: a v4 reader would derive n_frames from the file size, incorrectly including the PREC bytes in the sample count, producing a silently wrong reshape.

-

Producing v5 files

-

v5 files are produced by the SRAS viewer's "Pre-process and Save as v5" action. The procedure is:

-
    -
  1. Copy the source file (any version) verbatim.
  2. -
  3. Set version = 5 at byte offset 4.
  4. -
  5. Set n_frames_hdr at byte offset 25 to the actual acquired frame count.
  6. -
  7. Truncate the copy to data_offset + waveform_bytes (removes any pre-existing stale PREC tail).
  8. -
  9. Compute peak_freq_mhz, dc4_mv, and dc3_mv for every angle using chunked FFT.
  10. -
  11. Append the PREC section.
  12. -
-

Partially-written PREC sections

-

If n_stored < n_angles (e.g. pre-processing was interrupted), the file is still valid. Readers use stored images for the angles present in the PREC section and fall back to real-time FFT for the remainder. Readers must check angle_idx bounds on each entry and stop parsing on an out-of-range value.

- - - - - \ No newline at end of file diff --git a/SRAS_FORMAT.md b/SRAS_FORMAT.md deleted file mode 100644 index 4bc0346..0000000 --- a/SRAS_FORMAT.md +++ /dev/null @@ -1,379 +0,0 @@ -# SRAS File Format Specification - -**Format family:** `.sras` -**Byte order:** Big-endian (network byte order) throughout, unless noted. -**Version history:** v2 (baseline), v3 (scope calibration), v4 (background waveform), v5 (precomputed images + guaranteed frame count). - ---- - -## Table of Contents - -1. [Overview](#overview) -2. [Type notation](#type-notation) -3. [Version history](#version-history) -4. [File structure](#file-structure) - - [Fixed header (all versions)](#1-fixed-header-43-bytes-all-versions) - - [Angle table (all versions)](#2-angle-table-all-versions) - - [Row position table (all versions)](#3-row-position-table-all-versions) - - [Channel preambles (v3+)](#4-channel-preambles-v3) - - [Background waveform (v4+)](#5-background-waveform-v4) - - [Waveform data (all versions)](#6-waveform-data-all-versions) - - [PREC section (v5)](#7-prec-section-v5) -5. [Derived quantities](#derived-quantities) -6. [ADC calibration](#adc-calibration) -7. [Waveform data layout detail](#waveform-data-layout-detail) -8. [Size reference](#size-reference) -9. [Compatibility notes](#compatibility-notes) - ---- - -## Overview - -An SRAS file stores the raw RF waveforms captured during a Surface-acoustic-wave Resonance And Spectroscopy (SRAS) scan, along with the scan geometry and scope calibration metadata needed to interpret them. - -A scan consists of one or more **angles** (rotation positions of the sample), each containing a 2-D raster of **rows** × **frames**. At every pixel, `n_channels` waveforms of `samples_per_frame` ADC counts are stored. Channel order is fixed: - -| Index | Hardware channel | Signal | -|-------|-----------------|--------| -| 0 | CH1 | RF acoustic packet (AC-coupled) | -| 1 | CH3 | Bias A — DC mean used for masking | -| 2 | CH4 | Bias B — DC mean used for masking | - ---- - -## Type notation - -| Symbol | C type | Size | Notes | -|--------|--------|------|-------| -| `u8` | `uint8_t` | 1 byte | unsigned | -| `u16` | `uint16_t` | 2 bytes | big-endian | -| `u32` | `uint32_t` | 4 bytes | big-endian | -| `i8` | `int8_t` | 1 byte | signed, used for ADC samples when `bytes_per_sample == 1` | -| `i16` | `int16_t` | 2 bytes | big-endian signed, used when `bytes_per_sample == 2` | -| `f32` | `float` | 4 bytes | big-endian IEEE 754 single | -| `f64` | `double` | 8 bytes | big-endian IEEE 754 double | -| `char[N]` | — | N bytes | raw bytes, no null terminator unless noted | -| `utf8[N]` | — | N bytes | UTF-8 string, length-prefixed (see preamble section) | - ---- - -## Version history - -| Version | Added | -|---------|-------| -| 2 | Baseline: fixed header, angle table, row table, raw waveform data. No scope calibration (fallback constants used by readers). | -| 3 | Per-channel Tektronix WFMOutpre preamble strings carrying YMULT / YOFF / YZERO calibration. | -| 4 | Background waveform section: one CH1 reference shot subtracted from each CH1 frame before FFT. | -| 5 | **(this document)** Version byte incremented to 5. `n_frames_hdr` is now the *actual* acquired frame count (authoritative). PREC section appended after waveform data with precomputed FFT-peak and DC images for instant re-display. | - -> **v2 note:** Version 1 is not defined; version 2 is the lowest observed in the field. - ---- - -## File structure - -``` -┌─────────────────────────────────────────────┐ -│ 1. Fixed header (43 bytes) │ all versions -├─────────────────────────────────────────────┤ -│ 2. Angle table (n_angles × 4 bytes)│ all versions -├─────────────────────────────────────────────┤ -│ 3. Row position table (n_rows × 4 bytes)│ all versions -├─────────────────────────────────────────────┤ -│ 4. Channel preambles (variable) │ v3+ -├─────────────────────────────────────────────┤ -│ 5. Background waveform (variable) │ v4+ -├─────────────────────────────────────────────┤ -│ 6. Waveform data (variable) │ all versions -├─────────────────────────────────────────────┤ -│ 7. PREC section (variable) │ v5 only -└─────────────────────────────────────────────┘ -``` - ---- - -### 1. Fixed header (43 bytes, all versions) - -| Offset | Size | Type | Field | Description | -|--------|------|------|-------|-------------| -| 0 | 4 | `char[4]` | `magic` | `SRAS` (ASCII, no null terminator). Reject file if this does not match. | -| 4 | 1 | `u8` | `version` | Format version. This document describes version **5**. | -| 5 | 2 | `u16` | `n_angles` | Number of scan angles (rotation positions). ≥ 1. | -| 7 | 2 | `u16` | `n_rows` | Number of scan rows (Y positions). ≥ 1. | -| 9 | 4 | `f32` | `x_start_mm` | X position of the first frame in the first row, in millimetres. | -| 13 | 4 | `f32` | `x_delta_mm` | Pre-computed pixel pitch in mm (`velocity_mm_s / laser_freq_hz`). Provided for convenience; readers should prefer the derived value. | -| 17 | 4 | `f32` | `velocity_mm_s` | Scanner stage velocity, mm s⁻¹. Used together with `laser_freq_hz` to compute pixel pitch. | -| 21 | 4 | `f32` | `laser_freq_hz` | Laser repetition rate, Hz. | -| 25 | 4 | `u32` | `n_frames_hdr` | **v2–v4:** the *configured* frame count written before acquisition; may exceed actual frames acquired (use file-size arithmetic to obtain the true count). **v5:** the *actual* acquired frame count — authoritative; readers must not re-derive it from file size. | -| 29 | 4 | `u32` | `samples_per_frame` | ADC samples per waveform (`spf`). | -| 33 | 8 | `f64` | `sample_rate_hz` | Oscilloscope sample rate, Hz (e.g. 5 × 10⁹ for 5 GS/s). | -| 41 | 1 | `u8` | `bytes_per_sample` | ADC word size: `1` → `i8`, `2` → `i16` (big-endian). | -| 42 | 1 | `u8` | `n_channels` | Number of channels per frame. Currently always `3`. | - ---- - -### 2. Angle table (all versions) - -Immediately follows the fixed header. - -``` -n_angles × f32 — scan angle in degrees -``` - -Each entry is a big-endian `f32` giving the sample rotation angle in degrees at which that angle index was acquired. - ---- - -### 3. Row position table (all versions) - -Immediately follows the angle table. - -``` -n_rows × f32 — Y position of each row, in millimetres -``` - ---- - -### 4. Channel preambles (v3+) - -One entry per channel, in channel-index order (CH1 first). - -``` -for each channel: - u16 preamble_length — byte count of the UTF-8 string that follows - utf8[N] preamble — Tektronix WFMOutpre string -``` - -The preamble is the oscilloscope's `WFMOutpre` response string. Readers extract the following keys (case-insensitive, space-separated value): - -| Key | Stored unit | Conversion to mV | -|-----|-------------|-----------------| -| `YMULT` | V count⁻¹ | multiply by 1000 | -| `YOFF` | ADC counts | used directly | -| `YZERO` | V | multiply by 1000 | - -**v2 fallback:** when preambles are absent, readers use: -- `YMULT` = 1.5625 mV count⁻¹ (50 mV/div, 8 div, 8-bit ADC) -- `YOFF` = −87.04 ADC counts (scope position = −2.72 div) -- `YZERO` = 0 mV - ---- - -### 5. Background waveform (v4+) - -``` -u32 n_bg_samples — number of i8 ADC samples that follow -i8[n_bg] background — one representative CH1 background shot -``` - -The background waveform has the same `samples_per_frame` length as a normal CH1 waveform. It is subtracted from each CH1 waveform before FFT processing when background subtraction is enabled. When `n_bg_samples == 0` the section is present but empty. - ---- - -### 6. Waveform data (all versions) - -Begins immediately after the fixed header (v2), preambles (v3), or background waveform (v4+). The waveform data is a flat, contiguous array with the following logical shape, stored in row-major (C) order: - -``` -waveform_data[n_angles][n_rows][n_channels][n_frames][samples_per_frame] -``` - -Each element is a signed ADC count of size `bytes_per_sample`: -- `bytes_per_sample == 1` → `i8` -- `bytes_per_sample == 2` → `i16` big-endian - -**Total byte count:** - -``` -waveform_bytes = n_angles × n_rows × n_channels × n_frames × samples_per_frame × bytes_per_sample -``` - -#### Index semantics - -| Dimension | Range | Meaning | -|-----------|-------|---------| -| `[a]` | 0 … n_angles−1 | Scan angle (rotation position) | -| `[r]` | 0 … n_rows−1 | Row (Y position); row 0 is the first acquired | -| `[c]` | 0 … n_channels−1 | Channel (0=CH1 RF, 1=CH3 Bias A, 2=CH4 Bias B) | -| `[f]` | 0 … n_frames−1 | Frame (X position) within the row | -| `[s]` | 0 … spf−1 | Sample index within the waveform | - -#### Frame-count determination - -- **v5:** use `n_frames_hdr` directly; do not use file-size arithmetic. -- **v2–v4:** `n_frames = floor((file_bytes_after_header_sections) / (bytes_per_sample × n_angles × n_rows × n_channels × samples_per_frame))`. Any remainder bytes are a partial trailing row and are discarded. - ---- - -### 7. PREC section (v5) - -The PREC section is appended immediately after the waveform data and is present if and only if `version == 5` and the file size exceeds `waveform_end_offset`. - -``` -waveform_end_offset = data_offset + waveform_bytes -``` - -where `data_offset` is the file offset of the first waveform byte (the byte immediately after the background waveform, or after the angle/row tables for v2 files). - -#### PREC header (8 bytes) - -| Offset (relative) | Size | Type | Field | Description | -|-------------------|------|------|-------|-------------| -| 0 | 4 | `char[4]` | `prec_magic` | `PREC` (ASCII). Absent or wrong magic → ignore section. | -| 4 | 1 | `u8` | `prec_version` | PREC format version. Currently `1`. | -| 5 | 1 | `u8` | `flags` | Bitmask (see below). | -| 6 | 2 | `u16` | `n_stored` | Number of angle entries that follow. 0 ≤ `n_stored` ≤ `n_angles`. | - -##### Flags byte - -| Bit | Mask | Meaning when set | -|-----|------|-----------------| -| 0 | `0x01` | `bg_sub_applied` — background waveform was subtracted from CH1 before the FFT when these images were computed. | -| 1–7 | — | Reserved, must be zero on write; readers must ignore. | - -#### PREC angle entries - -Repeated `n_stored` times, in arbitrary angle-index order: - -``` -for each stored angle: - u16 angle_idx — index into the angle table (0-based) - f32[n_rows×n_frames] peak_freq_mhz — CH1 FFT peak frequency, MHz, row-major - f32[n_rows×n_frames] dc4_mv — CH4 waveform mean, mV, row-major - f32[n_rows×n_frames] dc3_mv — CH3 waveform mean, mV, row-major -``` - -All image arrays are `f32` big-endian, stored in row-major order: element `[r][f]` is at offset `(r × n_frames + f) × 4` bytes within the array. - -**`peak_freq_mhz`** is computed without any DC-threshold masking (i.e. the FFT is run on every pixel unconditionally). Readers apply the `dc4_mv` threshold at display time: - -``` -pixel is valid ⟺ dc4_mv[r][f] ≥ threshold_mv -display_value = peak_freq_mhz[r][f] if valid, else 0 -``` - -**`dc4_mv` / `dc3_mv`** are the mean of all ADC samples in the respective channel waveform, converted to millivolts using the channel calibration: - -``` -dc_mv = (adc_mean − YOFF) × YMULT + YZERO -``` - -#### When readers must bypass the PREC fast path - -Readers must fall back to real-time FFT computation (ignoring stored `peak_freq_mhz`) when any of the following are true: - -- Time-domain gating is active (zeroing samples outside a time window changes peak frequency). -- Zero-padding (`n_fft ≠ samples_per_frame`) is requested (changes bin spacing). -- The reader's background-subtraction setting does not match `flags.bg_sub_applied`. - ---- - -## Derived quantities - -``` -pixel_pitch_mm = velocity_mm_s / laser_freq_hz - -x_axis_mm[f] = x_start_mm + f × pixel_pitch_mm (f = 0 … n_frames−1) - -time_axis_ns[s] = s / sample_rate_hz × 1e9 (s = 0 … spf−1) - -freq_axis_mhz[k] = k × sample_rate_hz / (n_fft × 1e6) (k = 0 … n_fft/2) - where n_fft = samples_per_frame unless zero-padding is active - -velocity_ms[r][f] = peak_freq_mhz[r][f] × grating_um (grating_um user-supplied) -``` - ---- - -## ADC calibration - -Convert raw ADC counts to millivolts: - -``` -voltage_mv = (adc_count − YOFF) × YMULT_mv + YZERO_mv -``` - -Invert (mV → ADC count): - -``` -adc_count = (voltage_mv − YZERO_mv) / YMULT_mv + YOFF -``` - -where `YMULT_mv` is YMULT in mV count⁻¹ (= scope YMULT in V count⁻¹ × 1000). - ---- - -## Waveform data layout detail - -For a scan with `n_angles=2`, `n_rows=3`, `n_channels=3`, `n_frames=4`, `spf=5` the layout is: - -``` -angle 0 - row 0 - CH1: [s0 s1 s2 s3 s4] [s0 s1 s2 s3 s4] [s0 s1 s2 s3 s4] [s0 s1 s2 s3 s4] - frame 0 frame 1 frame 2 frame 3 - CH3: …(same layout)… - CH4: …(same layout)… - row 1 - … - row 2 - … -angle 1 - … -``` - -The flat byte offset of sample `s` of frame `f`, channel `c`, row `r`, angle `a` is: - -``` -offset = data_offset - + (a × n_rows × n_channels × n_frames × spf - + r × n_channels × n_frames × spf - + c × n_frames × spf - + f × spf - + s) - × bytes_per_sample -``` - ---- - -## Size reference - -Approximate sizes for representative scans (`bytes_per_sample = 1`, `n_channels = 3`). - -| n_angles | n_rows | n_frames | spf | Waveform data | PREC section | -|----------|--------|----------|-----|---------------|-------------| -| 1 | 500 | 500 | 400 | 300 MB | 12 MB | -| 4 | 500 | 500 | 400 | 1.2 GB | 48 MB | -| 1 | 2000 | 2000 | 400 | 4.8 GB | 48 MB | -| 4 | 2000 | 2000 | 400 | 19.2 GB | 192 MB | -| 8 | 2000 | 2000 | 400 | 38.4 GB | 384 MB | -| 16 | 2000 | 2000 | 400 | 76.8 GB | 768 MB | - -**PREC section size formula:** - -``` -prec_bytes = 8 + n_stored × (2 + 3 × n_rows × n_frames × 4) -``` - ---- - -## Compatibility notes - -### Reading v5 files with a v4 reader - -A v4 reader that only accepts versions `{2, 3, 4}` will reject a v5 file with an "unsupported version" error. This is intentional: a v4 reader would derive `n_frames` from the file size, incorrectly including the PREC bytes in the sample count, producing a silently wrong reshape. - -### Producing v5 files - -v5 files are produced by the SRAS viewer's **"Pre-process and Save as v5"** action. The procedure is: - -1. Copy the source file (any version) verbatim. -2. Set `version = 5` at byte offset 4. -3. Set `n_frames_hdr` at byte offset 25 to the actual acquired frame count. -4. Truncate the copy to `data_offset + waveform_bytes` (removes any pre-existing stale PREC tail). -5. Compute `peak_freq_mhz`, `dc4_mv`, and `dc3_mv` for every angle using chunked FFT. -6. Append the PREC section. - -### Partially-written PREC sections - -If `n_stored < n_angles` (e.g. pre-processing was interrupted), the file is still valid. Readers use stored images for the angles present in the PREC section and fall back to real-time FFT for the remainder. Readers must check `angle_idx` bounds on each entry and stop parsing on an out-of-range value. From 007089dd48ce33b5ac8a747f7cef7296704ae3dc Mon Sep 17 00:00:00 2001 From: Thomas Ales Date: Thu, 6 Aug 2026 09:45:35 -0500 Subject: [PATCH 03/10] Convert test scripts to pytest; extend the equivalence harness - pyproject.toml replaces sras_viewer_requirements.txt (same pins) and adds a dev extra with pytest. - tools/test_refactor.py, test_alignment.py, test_gui.py become tests/test_compute.py, tests/test_alignment.py, tests/test_gui.py with assertions preserved verbatim. test_gui.py stays one ordered integration sequence over a shared module-scoped window. - check_equivalence.py: drop the dead pre-refactor monolith shim (and the _compute_angle_alignment alias it consumed), extend the pad sweep to (1, 2, 4, 8, 40), add legacy-v4 and big-endian int16 legs (new bps=2 option in make_test_sras) so the padded FFT path and the >i2 memmap path are in the baseline before the FFT rewrite. Co-Authored-By: Claude Fable 5 --- pyproject.toml | 35 +++ sras_compute.py | 4 - sras_viewer_requirements.txt | 11 - tests/conftest.py | 8 + tests/test_alignment.py | 212 +++++++++++++++ tests/test_compute.py | 277 +++++++++++++++++++ tests/test_gui.py | 501 +++++++++++++++++++++++++++++++++++ tools/check_equivalence.py | 51 ++-- tools/make_test_sras.py | 13 +- tools/test_alignment.py | 223 ---------------- tools/test_gui.py | 476 --------------------------------- tools/test_refactor.py | 358 ------------------------- 12 files changed, 1065 insertions(+), 1104 deletions(-) create mode 100644 pyproject.toml delete mode 100644 sras_viewer_requirements.txt create mode 100644 tests/conftest.py create mode 100644 tests/test_alignment.py create mode 100644 tests/test_compute.py create mode 100644 tests/test_gui.py delete mode 100644 tools/test_alignment.py delete mode 100644 tools/test_gui.py delete mode 100644 tools/test_refactor.py diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..ca708ca --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,35 @@ +[build-system] +requires = ["setuptools>=68"] +build-backend = "setuptools.build_meta" + +[project] +name = "sras-viewer" +version = "0.1.0" +description = "Viewer and processing tools for SRAS .sras scan files" +requires-python = ">=3.12" +dependencies = [ + "PyQt6==6.10.2", + "numpy==2.4.1", + "matplotlib==3.10.8", + "scipy==1.18.0", + # Angle alignment only: masked FFT phase correlation (skimage.registration). + "scikit-image==0.26.0", + # Faster rfft backend; the viewer falls back to scipy.fft without it. + "pyFFTW==0.15.1", +] + +[project.optional-dependencies] +dev = ["pytest"] + +[tool.setuptools] +py-modules = [ + "sras_format", + "sras_compute", + "sras_workers", + "sras_viewer", + "sras_average", + "sras_edit_scans", +] + +[tool.pytest.ini_options] +testpaths = ["tests"] diff --git a/sras_compute.py b/sras_compute.py index 28c8b3d..bf48c6b 100644 --- a/sras_compute.py +++ b/sras_compute.py @@ -1188,10 +1188,6 @@ def compute_angle_alignment(sras: SrasFile, ref_angle_idx: int, return result -# Back-compat alias for the pre-split private name (used by tooling). -_compute_angle_alignment = compute_angle_alignment - - # --------------------------------------------------------------------------- # Manual alignment (Fusion menu -> Manual Alignment... dialog) # diff --git a/sras_viewer_requirements.txt b/sras_viewer_requirements.txt deleted file mode 100644 index 1448342..0000000 --- a/sras_viewer_requirements.txt +++ /dev/null @@ -1,11 +0,0 @@ -PyQt6==6.10.2 -numpy==2.4.1 -matplotlib==3.10.8 -scipy==1.18.0 -# Angle alignment only: masked FFT phase correlation, which registers scans -# whose valid (scanned) regions differ in shape — see sras_compute's -# _masked_shift. -scikit-image==0.26.0 -# Optional: a faster rfft backend for the RF/FFT images (FFT Options -> pyFFTW). -# The viewer falls back to scipy.fft when it is not installed. -pyFFTW==0.15.1 diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..c9df201 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,8 @@ +"""Shared test setup: repo-root imports and the offscreen Qt platform.""" + +import os +import sys +from pathlib import Path + +os.environ.setdefault("QT_QPA_PLATFORM", "offscreen") +sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) diff --git a/tests/test_alignment.py b/tests/test_alignment.py new file mode 100644 index 0000000..18be785 --- /dev/null +++ b/tests/test_alignment.py @@ -0,0 +1,212 @@ +"""Angle-alignment tests: does registration actually stack the scans? + +Builds a synthetic scan in which one sample is imaged at several *known* +rotations and offsets (tools/make_test_sras.write_rotating) and checks that the +alignment path recovers them, that the shared canvas is angle 0's own pixel +grid extended, and that nothing in the result depends on any other angle's +stage coordinates. + +No Qt — this exercises sras_compute directly. See tests/test_gui.py for the +dialog and Aligned-View plumbing. +""" + +from types import SimpleNamespace + +import numpy as np +import pytest + +import sras_compute as compute +from sras_format import CH4_IDX, SrasFile, adc_to_mv +import tools.make_test_sras as gen + +# Registration is limited by how far a feature moves per degree: with this +# sample's ~1 mm radius and a ~16 µm registration pitch, a quarter degree is +# already sub-pixel, so it is the floor of what any metric can resolve here. +_ROT_TOL_DEG = 0.5 +_SHIFT_TOL_MM = 0.02 +_STACK_IOU_MIN = 0.90 +_THRESHOLD_MV = 80.0 + + +def dc4_images(sras: SrasFile) -> dict[int, np.ndarray]: + return {a: adc_to_mv(compute.compute_dc_image(sras, a, CH4_IDX), *sras.cal(CH4_IDX)) + for a in range(sras.n_angles)} + + +def mm_transform(sras: SrasFile, result, angle_idx: int) -> np.ndarray: + """Recover the pure mm-space rotation from a canvas->raw affine. + + matrix == D @ R^T @ A_out, where A_out and D only carry the canvas and + per-angle pixel pitches; undoing both must leave something orthonormal, or + the transform is smuggling in a scale or a shear. + """ + dx_a, dy_a = compute._pixel_pitch_mm(sras, angle_idx) + A_out = np.array([[0.0, result.canvas_dx_mm], [result.canvas_dy_mm, 0.0]]) + D = np.array([[0.0, 1.0 / dy_a], [1.0 / dx_a, 0.0]]) + return np.linalg.inv(D) @ result.per_angle[angle_idx].matrix @ np.linalg.inv(A_out) + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory): + """The rotating-sample scan plus everything computed from it once.""" + tmpdir = tmp_path_factory.mktemp("sras_align") + path = tmpdir / "rotating.sras" + meta = gen.write_rotating(path, n_angles=5) + sras = SrasFile(str(path)) + dc4 = dc4_images(sras) + fits = {a: compute.register_angle_to_reference( + sras, a, 0, dc4, dc_threshold_mv=_THRESHOLD_MV) + for a in range(sras.n_angles)} + result = compute.compute_angle_alignment(sras, 0, _THRESHOLD_MV) + return SimpleNamespace(path=path, sras=sras, truth=meta["truth"], + dc4=dc4, fits=fits, result=result) + + +def test_registration_recovers_truth(rig): + """Per-angle rigid registration (rotation + translation, no scale).""" + for a, fit in rig.fits.items(): + t_rot, t_shift = rig.truth[a] + rot_err = abs(fit.rotation_deg - t_rot) + shift_err = float(np.hypot(fit.shift_mm[0] - t_shift[0], + fit.shift_mm[1] - t_shift[1])) + assert rot_err <= _ROT_TOL_DEG, \ + (f"angle {a}: got {fit.rotation_deg:.3f}°, truth {t_rot:.3f}° " + f"(err {rot_err:.3f}°)") + assert shift_err <= _SHIFT_TOL_MM, f"angle {a}: err {shift_err:.4f} mm" + assert rig.fits[0] == compute.RigidFit(0.0, (0.0, 0.0), 1.0, "reference"), \ + "reference angle registers as exact identity" + + +def test_stage_angle_sign_is_not_trusted(rig): + # The stage's rotational sense relative to this module's math-positive + # convention is not knowable from the file, and the old code hardcoded a + # guess. Flipping every reported angle must therefore change nothing: the + # search scores both signs and the images decide. + flipped = SrasFile(str(rig.path)) + flipped.angles_deg = -flipped.angles_deg + flipped_fits = {a: compute.register_angle_to_reference( + flipped, a, 0, rig.dc4, dc_threshold_mv=_THRESHOLD_MV) + for a in range(1, flipped.n_angles)} + mismatches = {a: (flipped_fits[a].rotation_deg, rig.fits[a].rotation_deg) + for a in flipped_fits if flipped_fits[a] != rig.fits[a]} + assert not mismatches, \ + f"negating every reported stage angle changed fits: {mismatches}" + + +def test_stage_coordinates_are_not_consulted(rig): + # Move every non-reference angle's scan window somewhere else entirely. + # Only angle 0's coordinates may matter, so every fit must be untouched. + moved = SrasFile(str(rig.path)) + for a in range(1, moved.n_angles): + moved.x_start_mm[a] += 13.5 * a + moved._y_pos_per_angle[a] = moved._y_pos_per_angle[a] - 9.25 * a + moved_dc4 = dc4_images(moved) + moved_fits = {a: compute.register_angle_to_reference( + moved, a, 0, moved_dc4, dc_threshold_mv=_THRESHOLD_MV) + for a in range(1, moved.n_angles)} + mismatches = {a: (round(moved_fits[a].rotation_deg, 4), rig.fits[a].rotation_deg) + for a in moved_fits if moved_fits[a] != rig.fits[a]} + assert not mismatches, \ + f"relocating every other angle's scan window changed fits: {mismatches}" + + +def test_canvas_is_reference_grid_extended(rig): + sras, result = rig.sras, rig.result + t0 = result.per_angle[0] + assert np.allclose(t0.matrix, np.eye(2)), \ + f"angle 0's transform has rotation/scale/shear: {t0.matrix}" + assert np.allclose(t0.offset, np.round(t0.offset)), \ + f"angle 0 does not land on whole canvas pixels: {t0.offset}" + assert ((result.canvas_dx_mm, result.canvas_dy_mm) + == compute._pixel_pitch_mm(sras, 0)), \ + "canvas pitch is angle 0's own pitch" + + n_rows, n_cols = result.canvas_shape + x_axis = result.canvas_origin_mm[0] + np.arange(n_cols) * result.canvas_dx_mm + y_axis = result.canvas_origin_mm[1] + np.arange(n_rows) * result.canvas_dy_mm + row0, col0 = int(round(-t0.offset[0])), int(round(-t0.offset[1])) + a0_rows, a0_cols = sras.image_shape(0) + assert np.allclose(x_axis[col0:col0 + a0_cols], sras.x_axis_mm(0)), \ + "canvas X axis reproduces angle 0's own X coordinates" + assert np.allclose(y_axis[row0:row0 + a0_rows], sras.y_positions_mm(0)), \ + "canvas Y axis reproduces angle 0's own Y coordinates" + assert (n_rows >= max(int(sras.n_rows[a]) for a in range(sras.n_angles)) + and n_cols >= max(int(sras.n_frames[a]) for a in range(sras.n_angles))), \ + f"canvas does not cover every angle's footprint: {result.canvas_shape}" + + +def test_transforms_are_pure_rotations(rig): + """No scaling anywhere in the per-angle transforms.""" + for a in range(rig.sras.n_angles): + R = mm_transform(rig.sras, rig.result, a) + assert (np.allclose(R @ R.T, np.eye(2), atol=1e-9) + and abs(abs(np.linalg.det(R)) - 1.0) < 1e-9), \ + f"angle {a}: det={np.linalg.det(R):.6f}" + + +def test_all_angles_stack(rig): + aligned = {a: compute.apply_alignment(rig.result, a, rig.dc4[a]) + for a in range(rig.sras.n_angles)} + base = aligned[0] >= _THRESHOLD_MV + for a in range(1, rig.sras.n_angles): + other = aligned[a] >= _THRESHOLD_MV + iou = float((base & other).sum()) / max(1, int((base | other).sum())) + assert iou >= _STACK_IOU_MIN, f"angle {a}: IoU {iou:.4f}" + + +def test_downsampled_preview_lands_with_full_res(rig): + # ManualAlignmentDialog reprojects block-mean-downsampled masks, so the + # affine has to account for the factor. When it did not, every preview + # layer came out magnified by that factor and offset — the overlay showed a + # blown-up crop of each mask, which is not something you can align by eye. + sras, result = rig.sras, rig.result + pitch = (result.canvas_dx_mm, result.canvas_dy_mm) + a = sras.n_angles - 1 + p = result.per_angle[a] + full_mask = (rig.dc4[a] >= _THRESHOLD_MV).astype(np.float32) + full = compute.reproject_mask( + sras, a, 0, full_mask, p.rotation_deg, p.shift_mm, pitch, + result.canvas_origin_mm, result.canvas_shape) + fy, fx = 4, 16 + small = compute.reproject_mask( + sras, a, 0, compute._block_mean_2d(full_mask, fy, fx), + p.rotation_deg, p.shift_mm, (pitch[0] * fx, pitch[1] * fy), + result.canvas_origin_mm, + (result.canvas_shape[0] // fy, result.canvas_shape[1] // fx), + src_downsample=(fy, fx)) + + # Compare in mm, via each layer's own center of mass. + def com_mm(layer, px, py): + rows, cols = np.nonzero(layer > 0.5) + return np.array([cols.mean() * px, rows.mean() * py]) + + d = com_mm(small, pitch[0] * fx, pitch[1] * fy) - com_mm(full, *pitch) + assert (abs(d[0]) <= abs(pitch[0] * fx) and abs(d[1]) <= abs(pitch[1] * fy)), \ + f"downsampled preview offset {d[0]:+.4f}, {d[1]:+.4f} mm" + + +def test_manual_path_reproduces_geometry(rig): + sras, result = rig.sras, rig.result + params = {a: compute.ManualAngleParams(t.rotation_deg, t.shift_mm) + for a, t in result.per_angle.items()} + manual = compute.build_manual_alignment(sras, 0, _THRESHOLD_MV, params) + assert (manual.canvas_shape == result.canvas_shape + and np.allclose(manual.canvas_origin_mm, result.canvas_origin_mm) + and all(np.allclose(manual.per_angle[a].matrix, result.per_angle[a].matrix) + and np.allclose(manual.per_angle[a].offset, result.per_angle[a].offset) + for a in range(sras.n_angles))), \ + "build_manual_alignment matches compute_angle_alignment for the same params" + + +def test_sidecar_roundtrip(rig): + sras, result = rig.sras, rig.result + params = {a: compute.ManualAngleParams(t.rotation_deg, t.shift_mm) + for a, t in result.per_angle.items()} + compute.save_manual_alignment(sras, 0, _THRESHOLD_MV, params) + loaded = compute.load_manual_alignment(sras) + assert (loaded is not None + and all(np.isclose(loaded.per_angle[a].rotation_deg, params[a].rotation_deg) + and np.allclose(loaded.per_angle[a].shift_mm, params[a].shift_mm) + for a in range(sras.n_angles))), \ + "sidecar reloads every angle's params" + assert compute.delete_manual_alignment(sras), "sidecar deletes cleanly" diff --git a/tests/test_compute.py b/tests/test_compute.py new file mode 100644 index 0000000..e08af39 --- /dev/null +++ b/tests/test_compute.py @@ -0,0 +1,277 @@ +"""Behavioural tests for the compute/format layer. + +Covers what the golden-hash harness can't: the v6->v7 cache round-trip +(including block carry-forward), parallel-vs-serial identity, the no-mask +fast path, and the ROI bounding-box mask optimisation. +""" + +import subprocess +import sys +from pathlib import Path + +import numpy as np + +import sras_compute as compute +from sras_compute import ( + cache_file, compute_dc_image, compute_rf_image, dc_image_mv, +) +from sras_format import CH3_IDX, CH4_IDX, SrasFile, adc_to_mv +import tools.make_test_sras as gen + +REPO = Path(__file__).resolve().parent.parent + + +def test_cache_roundtrip(tmp_path): + """v6 -> v7 for DC, then FFT, asserting the first block survives the + second write (the carry-forward path in write_v7_cache).""" + path = tmp_path / "roundtrip.sras" + gen.write(path, n_angles=3, seed=1, samples_per_frame=64) + + src = SrasFile(str(path)) + assert src.version == 6, f"got v{src.version}" + expect_dc3 = [dc_image_mv(src, a, CH3_IDX) for a in range(src.n_angles)] + expect_dc4 = [dc_image_mv(src, a, CH4_IDX) for a in range(src.n_angles)] + expect_fft = [compute_rf_image(src, a, dc_threshold_mv=None, apply_bg_sub=True) + for a in range(src.n_angles)] + + err = cache_file(str(path), "dc", True) + assert err == "", err + + after_dc = SrasFile(str(path)) + assert after_dc.version == 7, f"got v{after_dc.version}" + assert all(x is not None for x in after_dc.precomputed_dc3_mv) + assert all(np.allclose(after_dc.precomputed_dc3_mv[a], expect_dc3[a], atol=1e-4) + for a in range(after_dc.n_angles)) + assert all(np.allclose(after_dc.precomputed_dc4_mv[a], expect_dc4[a], atol=1e-4) + for a in range(after_dc.n_angles)) + assert all(x is None for x in after_dc.precomputed_freq_mhz), "no fft block yet" + assert (after_dc.precomputed_dc3_mv[0].dtype == np.float32 + and after_dc.precomputed_dc3_mv[0].dtype.byteorder in ("=", "|")), \ + "cached images are native float32" + assert after_dc.precomputed_dc3_mv[0].flags.writeable + + err = cache_file(str(path), "fft", True) + assert err == "", err + + both = SrasFile(str(path)) + assert all(x is not None for x in both.precomputed_freq_mhz) + assert all(np.allclose(both.precomputed_freq_mhz[a], expect_fft[a], atol=1e-3) + for a in range(both.n_angles)) + assert all(np.allclose(both.precomputed_dc3_mv[a], expect_dc3[a], atol=1e-4) + for a in range(both.n_angles)), \ + "DC block carried forward through the FFT write" + assert both.precomputed_bg_sub is True + + # The fast path must reproduce a fresh compute, and masking must still + # apply on top of a cached (unmasked) image. + fresh = SrasFile(str(path)) + fresh.precomputed_freq_mhz = [None] * fresh.n_angles + dc4 = dc_image_mv(both, 0, CH4_IDX) + thr = float(np.median(dc4)) + assert np.allclose( + compute_rf_image(both, 0, dc_threshold_mv=None, apply_bg_sub=True), + compute_rf_image(fresh, 0, dc_threshold_mv=None, apply_bg_sub=True), + atol=1e-3), "cached fast path == fresh compute (unmasked)" + assert np.allclose( + compute_rf_image(both, 0, dc_threshold_mv=thr, apply_bg_sub=True), + compute_rf_image(fresh, 0, dc_threshold_mv=thr, apply_bg_sub=True), + atol=1e-3), "cached fast path == fresh compute (masked)" + + # Waveform data must be byte-identical to the pre-cache file. + orig = tmp_path / "roundtrip_orig.sras" + gen.write(orig, n_angles=3, seed=1, samples_per_frame=64) + o, n = SrasFile(str(orig)), SrasFile(str(path)) + assert all(np.array_equal(np.asarray(o.data[a]), np.asarray(n.data[a])) + for a in range(o.n_angles)), \ + "waveform data untouched by the cache write" + + +def test_partial_v7_cache(tmp_path): + """Only some angles cached: uncached angles must compute, not read zeros. + This is the v5 bug the ragged normalisation fixed, checked via v7.""" + path = tmp_path / "partial.sras" + gen.write(path, n_angles=3, seed=2, samples_per_frame=64) + + src = SrasFile(str(path)) + expected = [compute_rf_image(src, a, dc_threshold_mv=None, apply_bg_sub=True) + for a in range(src.n_angles)] + partial = [expected[0], None, expected[2]] # angle 1 deliberately absent + src.write_v7_cache(new_freq_mhz=partial, new_bg_sub=True) + + reread = SrasFile(str(path)) + assert reread.precomputed_freq_mhz[1] is None + assert (reread.precomputed_freq_mhz[0] is not None + and reread.precomputed_freq_mhz[2] is not None) + img1 = compute_rf_image(reread, 1, dc_threshold_mv=None, apply_bg_sub=True) + assert np.any(img1 != 0) and np.allclose(img1, expected[1], atol=1e-3), \ + "uncached angle computes rather than returning zeros" + + +def test_parallel_identity(tmp_path, monkeypatch): + """Forcing 1 worker vs many must give identical output — catches + chunk-boundary and race bugs.""" + path = tmp_path / "parallel.sras" + # Many rows, so the row loop actually splits into several chunks. + n_rows, n_frames, spf = 48, 9, 256 + gen.write(path, n_angles=1, seed=3, samples_per_frame=spf, + geometry=[(n_rows, n_frames)]) + sras = SrasFile(str(path)) + + # Shrink the budget so chunk_rows collapses to 1 and every row is + # its own chunk — the worst case for boundary bugs. + monkeypatch.setattr(compute, "_TOTAL_BYTES_BUDGET", 8 * n_frames * spf * 4) + + monkeypatch.setattr(compute, "_MAX_WORKERS", 1) + dc_serial = compute_dc_image(sras, 0, CH4_IDX) + rf_serial = compute_rf_image(sras, 0, dc_threshold_mv=None, apply_bg_sub=True) + dc4 = adc_to_mv(dc_serial, *sras.cal(CH4_IDX)) + thr = float(np.median(dc4)) + rf_masked_serial = compute_rf_image(sras, 0, dc_threshold_mv=thr, + apply_bg_sub=True) + + chunk_rows, n_workers = compute._plan_chunks( + n_rows, n_frames, spf, live_multiplier=compute._FFT_LIVE_MULTIPLIER) + assert n_workers == 1, f"serial plan uses 1 worker (chunk_rows={chunk_rows})" + assert chunk_rows < n_rows, \ + f"work actually splits into multiple chunks ({chunk_rows} of {n_rows} rows)" + + monkeypatch.setattr(compute, "_MAX_WORKERS", 8) + chunk_rows, n_workers = compute._plan_chunks( + n_rows, n_frames, spf, live_multiplier=compute._FFT_LIVE_MULTIPLIER) + assert n_workers > 1, \ + f"parallel plan uses >1 worker (chunk_rows={chunk_rows} workers={n_workers})" + + dc_par = compute_dc_image(sras, 0, CH4_IDX) + rf_par = compute_rf_image(sras, 0, dc_threshold_mv=None, apply_bg_sub=True) + rf_masked_par = compute_rf_image(sras, 0, dc_threshold_mv=thr, apply_bg_sub=True) + + assert np.array_equal(dc_serial, dc_par), "dc image identical" + assert np.array_equal(rf_serial, rf_par), "rf image identical (unmasked)" + assert np.array_equal(rf_masked_serial, rf_masked_par), \ + "rf image identical (masked)" + + +def test_nomask_equals_low_threshold(tmp_path): + """dc_threshold_mv=None must equal a threshold below every pixel, while + skipping the CH4 read.""" + path = tmp_path / "nomask.sras" + gen.write(path, n_angles=2, seed=4, samples_per_frame=128) + sras = SrasFile(str(path)) + for a in range(sras.n_angles): + none_img = compute_rf_image(sras, a, dc_threshold_mv=None, apply_bg_sub=True) + low_img = compute_rf_image(sras, a, dc_threshold_mv=-1e9, apply_bg_sub=True) + assert np.array_equal(none_img, low_img), \ + f"angle {a}: None == -1e9 threshold" + assert len(np.unique(none_img)) > 1, \ + f"angle {a}: image is degenerate ({len(np.unique(none_img))} unique)" + + +def test_roi_mask(): + """The bbox-restricted mask must equal a full-grid point-in-polygon test.""" + from matplotlib.path import Path as MplPath + from sras_viewer import RoiQuad + + rng = np.random.default_rng(0) + x = np.linspace(-2.0, 3.0, 137) + y = np.linspace(1.0, 4.0, 91) + + cases = { + "axis-aligned rect": np.array([[0.0, 1.5], [1.0, 1.5], [1.0, 3.0], [0.0, 3.0]]), + "skewed quad": np.array([[-0.5, 1.2], [1.7, 1.9], [1.2, 3.4], [-1.0, 2.6]]), + "entirely outside": np.array([[8.0, 8.0], [9.0, 8.0], [9.0, 9.0], [8.0, 9.0]]), + "covers whole grid": np.array([[-9.0, -9.0], [9.0, -9.0], [9.0, 9.0], [-9.0, 9.0]]), + "straddles left edge": np.array([[-4.0, 2.0], [0.5, 2.0], [0.5, 3.0], [-4.0, 3.0]]), + } + for _ in range(5): + cases[f"random {_}"] = rng.uniform([-2.5, 0.5], [3.5, 4.5], size=(4, 2)) + + for name, pts in cases.items(): + roi = RoiQuad(pts) + fast = roi.mask_for_grid(x, y) + X, Y = np.meshgrid(x.astype(np.float64), y.astype(np.float64)) + slow = MplPath(pts).contains_points( + np.column_stack([X.ravel(), Y.ravel()])).reshape(X.shape) + assert np.array_equal(fast, slow), f"{name} ({int(slow.sum())} px inside)" + + # Descending y axis (images are stored top-down in some scans). + roi = RoiQuad(cases["skewed quad"]) + y_desc = y[::-1] + fast = roi.mask_for_grid(x, y_desc) + X, Y = np.meshgrid(x.astype(np.float64), y_desc.astype(np.float64)) + slow = MplPath(cases["skewed quad"]).contains_points( + np.column_stack([X.ravel(), Y.ravel()])).reshape(X.shape) + assert np.array_equal(fast, slow), "descending y axis" + + +def test_legacy_parse(tmp_path): + """v2-v4 parsing against known written data.""" + for version in (2, 3, 4): + path = tmp_path / f"legacy_v{version}.sras" + meta = gen.write_legacy(path, version=version, n_angles=2, n_rows=4, + n_frames=12, samples_per_frame=32, seed=version) + s = SrasFile(str(path)) + assert s.version == version, f"got v{s.version}" + assert list(s.n_rows) == [4, 4] and list(s.n_frames) == [12, 12], \ + f"rows={list(s.n_rows)} frames={list(s.n_frames)}" + assert all(np.array_equal(np.asarray(s.data[a]), meta["data"][a]) + for a in range(s.n_angles)), \ + f"v{version} waveform data matches what was written" + assert (s.background is not None) == (version >= 4), \ + f"v{version} background {'present' if version >= 4 else 'absent'}" + assert (isinstance(s.precomputed_freq_mhz, list) + and len(s.precomputed_freq_mhz) == s.n_angles), \ + f"v{version} precomputed stores are ragged lists" + # DC image must equal a direct mean of the known input. + expect = meta["data"][0][:, CH3_IDX, :, :].astype(np.float64).mean(axis=-1) + assert np.allclose(compute_dc_image(s, 0, CH3_IDX), expect, atol=1e-3), \ + f"v{version} DC image equals a direct mean" + + +def test_sras_average(tmp_path): + """The sras_average.py CLI: frame averaging with remainder handling.""" + src = tmp_path / "legacy_v4.sras" + meta = gen.write_legacy(src, version=4, n_angles=2, n_rows=4, n_frames=12, + samples_per_frame=32, seed=4) + dst = tmp_path / "legacy_v4_avg.sras" + proc = subprocess.run( + [sys.executable, str(REPO / "sras_average.py"), str(src), str(dst), "--n", "4"], + capture_output=True, text=True, cwd=REPO) + assert proc.returncode == 0, (proc.stderr or proc.stdout).strip()[-200:] + + avg = SrasFile(str(dst)) + assert avg.version == 4 + assert list(avg.n_frames) == [3, 3], f"{list(avg.n_frames)}" + assert (avg.n_angles == 2 and list(avg.n_rows) == [4, 4] + and avg.n_channels == meta["n_channels"]) + assert np.allclose(avg.ch_ymult_mv, SrasFile(str(src)).ch_ymult_mv), \ + "calibration preserved" + assert np.array_equal(avg.background, SrasFile(str(src)).background), \ + "background preserved" + src_data = meta["data"] + expect0 = src_data[0][:, :, 0:4, :].astype(np.float32).mean(axis=2).astype(np.int16) + assert np.array_equal(np.asarray(avg.data[0])[:, :, 0, :], expect0), \ + "first averaged group equals the mean of its 4 source frames" + + # Remainder handling: 12 frames / 5 -> 2 full groups + 1 partial. + dst2 = tmp_path / "legacy_v4_avg5.sras" + subprocess.run([sys.executable, str(REPO / "sras_average.py"), + str(src), str(dst2), "--n", "5"], + capture_output=True, text=True, cwd=REPO) + assert list(SrasFile(str(dst2)).n_frames) == [3, 3], \ + "partial trailing group kept by default" + dst3 = tmp_path / "legacy_v4_avg5d.sras" + subprocess.run([sys.executable, str(REPO / "sras_average.py"), + str(src), str(dst3), "--n", "5", "--discard-remainder"], + capture_output=True, text=True, cwd=REPO) + assert list(SrasFile(str(dst3)).n_frames) == [2, 2], \ + "--discard-remainder drops the partial group" + + +def test_unsupported_version_reported(tmp_path): + """cache_file must report, not raise, for a file it can't handle.""" + bogus = tmp_path / "bogus.sras" + bogus.write_bytes(b"SRAS" + bytes([99]) + b"\x00" * 200) + err = cache_file(str(bogus), "dc", True) + assert err, "bad version returns an error string" + missing = cache_file(str(tmp_path / "does_not_exist.sras"), "dc", True) + assert missing, "missing file returns an error string" diff --git a/tests/test_gui.py b/tests/test_gui.py new file mode 100644 index 0000000..7b1ebec --- /dev/null +++ b/tests/test_gui.py @@ -0,0 +1,501 @@ +"""Headless GUI test: drives SrasViewerWindow through the real Qt widgets, +signals and worker threads under the offscreen platform plugin. + +Covers the interactions a manual smoke test would: load, switch angles and +channels, background DC precompute, lazy FFT compute, threshold and bg-sub +changes, angle alignment, manual angle alignment, aligned view, ROI +draw/move, and CSV export. + +NOTE: this module is one ordered integration sequence over a single shared +window — the tests build on each other's state and must run in definition +order (pytest's default within a module). Run the whole module, not single +tests. +""" + +import json +from types import SimpleNamespace +from unittest.mock import patch + +import numpy as np +import pytest +from PyQt6.QtCore import QEventLoop, Qt, QTimer +from PyQt6.QtTest import QTest +from PyQt6.QtWidgets import QApplication, QMessageBox + +import sras_compute as compute +from sras_format import CH1_IDX, CH3_IDX, CH4_IDX, SrasFile +from sras_viewer import RoiQuad, SrasViewerWindow, VELOCITY_MODE_IDX +import tools.make_test_sras as gen + + +def pump(ms: int = 250): + """Run the event loop for a while so queued signals and worker threads + make progress.""" + loop = QEventLoop() + QTimer.singleShot(ms, loop.quit) + loop.exec() + + +def wait_until(pred, timeout_ms: int = 20000, step: int = 100) -> bool: + waited = 0 + while waited < timeout_ms: + if pred(): + return True + pump(step) + waited += step + return pred() + + +@pytest.fixture(scope="module") +def ctx(tmp_path_factory): + """The shared window, test file, and cross-test state for the sequence.""" + app = QApplication.instance() or QApplication([]) + tmpdir = tmp_path_factory.mktemp("sras_gui") + path = tmpdir / "gui.sras" + gen.write(path, n_angles=4, seed=11, samples_per_frame=256) + + win = SrasViewerWindow() + win.show() + errors: list[str] = [] + # Capture anything the app reports as an error via the status bar. + win.statusBar().messageChanged.connect( + lambda m: errors.append(m) if m and "error" in m.lower() else None) + + c = SimpleNamespace(app=app, win=win, path=path, tmpdir=tmpdir, + errors=errors, s=None) + yield c + if win.isVisible(): + win.close() + pump(400) + + +def test_load(ctx): + win = ctx.win + win._load_file(str(ctx.path)) + assert wait_until(lambda: win._sras is not None), "file loaded" + ctx.s = s = win._sras + assert s.version == 6, f"v{s.version}" + assert win.combo_channel.currentIndex() == CH4_IDX, "defaults to CH4" + assert win._current_image is not None, "image displayed" + assert win.spin_angle.maximum() == s.n_angles - 1, \ + "angle spinbox ranges over all angles" + assert win._info["Angles"].text() == f"Angles: {s.n_angles}", \ + win._info["Angles"].text() + + +def test_dc_precompute_all_angles(ctx): + win, s = ctx.win, ctx.s + ok = wait_until(lambda: all((a, CH4_IDX) in win._dc_cache + and (a, CH3_IDX) in win._dc_cache + for a in range(s.n_angles))) + assert ok, f"every angle cached for CH3 and CH4 ({len(win._dc_cache)} entries)" + assert "ready for all angles" in win.lbl_dc_precompute.text(), \ + win.lbl_dc_precompute.text() + + +def test_angle_switching_from_cache(ctx): + win, s = ctx.win, ctx.s + for a in range(s.n_angles): + win.spin_angle.setValue(a) + win._on_view_changed() + pump(60) + expected = win._sras.image_shape(a) + assert win._current_image.shape == expected, \ + f"angle {a} shows its own geometry {expected}, got {win._current_image.shape}" + assert not win._job_running("compute"), \ + "no compute job needed for cached DC angles" + + +def test_channel_switching(ctx): + win = ctx.win + win.spin_angle.setValue(0) + win._on_view_changed() + pump(60) + win.combo_channel.setCurrentIndex(CH3_IDX) + assert wait_until(lambda: win._current_ch == CH3_IDX), "CH3 displayed" + + win.combo_channel.setCurrentIndex(CH1_IDX) + assert wait_until( + lambda: win._current_ch == CH1_IDX and not win._job_running("compute")), \ + "CH1 (FFT) computed" + assert len(win._fft_cache) > 0, "FFT result cached" + ctx.rf_img = win._current_image + assert len(np.unique(ctx.rf_img)) > 1, \ + f"FFT image is degenerate ({len(np.unique(ctx.rf_img))} unique values)" + + +def test_velocity_mode(ctx): + """Velocity mode is a pure post-multiply, no recompute.""" + win = ctx.win + ctx.n_fft_before = len(win._fft_cache) + win.combo_channel.setCurrentIndex(VELOCITY_MODE_IDX) + assert wait_until( + lambda: win._current_ch == VELOCITY_MODE_IDX + and not win._job_running("compute")), "velocity displayed" + grating = win.spin_grating_um.value() + assert np.allclose(win._current_image, ctx.rf_img * grating, atol=1e-3), \ + "velocity == freq x grating" + assert len(win._fft_cache) == ctx.n_fft_before, \ + f"velocity reused the cached FFT ({ctx.n_fft_before} -> {len(win._fft_cache)})" + assert win.grp_velocity.isVisible(), "grating spinbox visible in velocity mode" + + +def test_threshold_change_recomputes(ctx): + """A threshold change is a genuine cache-key change.""" + win = ctx.win + win.combo_channel.setCurrentIndex(CH1_IDX) + wait_until(lambda: not win._job_running("compute")) + dc4 = win._dc_cache[(0, CH4_IDX)] + win.spin_threshold_mv.setValue(float(np.median(dc4))) + win._on_threshold_changed() + assert wait_until( + lambda: not win._job_running("compute") + and len(win._fft_cache) > ctx.n_fft_before), "recomputed at new threshold" + n_zero = int((win._current_image == 0).sum()) + assert n_zero > 0, \ + f"masking zeroed some pixels ({n_zero} of {win._current_image.size})" + + +def test_bg_sub_toggle(ctx): + win = ctx.win + n_before = len(win._fft_cache) + win.chk_bg_sub.setChecked(False) + assert wait_until( + lambda: not win._job_running("compute") and len(win._fft_cache) > n_before), \ + "recomputed without bg-sub" + win.chk_bg_sub.setChecked(True) + pump(200) + assert not win._job_running("compute"), \ + "returning to bg-sub was a cache hit (no recompute)" + + +def test_roi_and_csv_export(ctx): + win, s = ctx.win, ctx.s + x = s.x_axis_mm(0) + y = s.y_positions_mm(0) + roi = RoiQuad.from_bbox(float(x[1]), float(y[1]), + float(x[-2]), float(y[-2])) + win.image_canvas.set_roi(roi) + pump(120) + assert win.image_canvas.get_roi() is not None, "ROI registered" + assert ("pixels inside" in win.lbl_roi_npix.text() + and win.lbl_roi_npix.text() != "pixels inside: —"), \ + win.lbl_roi_npix.text() + npix = int(win.lbl_roi_npix.text().split(":")[1]) + assert 0 < npix <= win._current_image.size, f"{npix}" + assert win.btn_export_roi.isEnabled(), "Export ROI enabled" + + csv_path = ctx.tmpdir / "roi.csv" + with patch("sras_viewer.QFileDialog.getSaveFileName", + return_value=(str(csv_path), "")): + win._on_export_roi_csv() + assert csv_path.exists(), "ROI CSV written" + body = [l for l in csv_path.read_text().splitlines() if not l.startswith("#")] + assert len(body) == npix + 1, \ + f"ROI CSV has {len(body)} lines for {npix} pixels (want header + one per pixel)" + + img_csv = ctx.tmpdir / "img.csv" + with patch("sras_viewer.QFileDialog.getSaveFileName", + return_value=(str(img_csv), "")): + win._on_export_csv() + assert img_csv.exists(), "image CSV written" + arr = np.loadtxt(img_csv, delimiter=",") + assert (arr.shape == win._current_image.shape + and np.allclose(arr, win._current_image, rtol=1e-5, atol=1e-4)), \ + "image CSV round-trips the displayed image" + + +def test_roi_survives_switches(ctx): + win = ctx.win + win.spin_angle.setValue(1) + win._on_view_changed() + wait_until(lambda: not win._job_running("compute")) + assert win.image_canvas.get_roi() is not None, \ + "ROI still present after angle switch" + win.combo_channel.setCurrentIndex(CH4_IDX) + wait_until(lambda: win._current_ch == CH4_IDX) + assert win.image_canvas.get_roi() is not None, \ + "ROI still present after channel switch" + + +def test_angle_alignment(ctx): + win, s = ctx.win, ctx.s + win.spin_angle.setValue(0) + win._on_view_changed() + wait_until(lambda: not win._job_running("compute")) + assert win._alignment_act.isEnabled(), "alignment action enabled" + win._on_angle_alignment() + assert wait_until( + lambda: win._alignment_result is not None and not win._job_running("align"), + timeout_ms=60000), "alignment completed" + + r = win._alignment_result + assert len(r.per_angle) == s.n_angles, "transform for every angle" + assert all(r.canvas_shape[0] >= int(s.n_rows[a]) + and r.canvas_shape[1] >= int(s.n_frames[a]) + for a in range(s.n_angles)), \ + f"canvas is at least as large as any single angle: {r.canvas_shape}" + assert r.per_angle[r.ref_angle_idx].shift_mm == (0.0, 0.0), \ + "reference angle has zero shift" + assert win.chk_aligned_view.isEnabled() and win.chk_aligned_view.isChecked(), \ + "Aligned View auto-enabled and checked" + pump(200) + assert win.image_canvas._img_shape == r.canvas_shape, \ + f"{win.image_canvas._img_shape} vs {r.canvas_shape}" + + win.chk_aligned_view.setChecked(False) + pump(200) + assert win.image_canvas._img_shape == s.image_shape(0), \ + f"unchecking returns to the raw per-angle grid: {win.image_canvas._img_shape}" + + +def test_manual_alignment_geometry(ctx): + """Local mm is anchored on each angle's array center, not its stage + position: that is what makes a scan's placement independent of where its + window happened to sit. (Registration accuracy itself is covered by + tests/test_alignment.py, which has a synthetic sample to register.)""" + win, s = ctx.win, ctx.s + assert win._manual_align_act.isEnabled(), "manual alignment action enabled" + + n_rows, n_frames = s.image_shape(0) + assert np.allclose(compute._center_idx(s, 0), + [(n_rows - 1) / 2, (n_frames - 1) / 2]), \ + "array center is the geometric center of the pixel grid" + dx0, dy0 = compute._pixel_pitch_mm(s, 0) + assert np.allclose(compute._local_half_extent_mm(s, 0), + [(n_frames - 1) / 2 * abs(dx0), (n_rows - 1) / 2 * abs(dy0)]), \ + "local half-extent is derived from shape and pitch alone" + identity = {a: compute.ManualAngleParams() for a in range(s.n_angles)} + origin_a, shape_a = compute.canvas_for_params(s, 0, (dx0, dy0), identity) + moved = SrasFile(str(ctx.path)) + for a in range(1, moved.n_angles): + moved.x_start_mm[a] += 7.5 + moved._y_pos_per_angle[a] = moved._y_pos_per_angle[a] + 3.25 + origin_b, shape_b = compute.canvas_for_params(moved, 0, (dx0, dy0), identity) + assert shape_a == shape_b and np.allclose(origin_a, origin_b), \ + ("moving every non-reference angle's scan window must leave the canvas " + f"unchanged: {origin_a} {shape_a} vs {origin_b} {shape_b}") + + # Both signs of the stage's reported angle are searched. + cands = compute._rotation_candidates(30.0, 6.0, 2.0) + assert min(cands) < -29.0 and max(cands) > 29.0, f"{min(cands)}..{max(cands)}" + + # Whole-pixel translation must not wrap content around the edge. + arr = np.zeros((6, 6), dtype=np.float32) + arr[0, 0] = 1.0 + assert compute._shift_into(arr, -1, -1).sum() == 0.0, \ + "_shift_into zero-fills rather than wrapping" + assert compute._shift_into(arr, 2, 3)[2, 3] == 1.0, \ + "_shift_into moves content by exactly the requested offset" + + +def test_manual_dialog_opens_at_identity(ctx): + """Open must NOT seed from the still-live automatic AlignmentResult. + Manual mode exists to fix up whatever the automatic registration got + wrong, so it must start from identity (every angle centered on the + reference, no rotation) regardless of whatever the automatic run last + computed. Only a previously *saved manual* alignment (sidecar) should + ever seed this dialog.""" + win, s = ctx.win, ctx.s + win._on_manual_alignment() + assert win._manual_align_dialog is not None, "dialog opened" + ctx.dlg = dlg = win._manual_align_dialog + assert not win._job_running("manual_align_masks"), \ + "mask prep needed no background worker (already DC-cached)" + assert all(dlg._angle_params[a] == compute.ManualAngleParams() + for a in range(s.n_angles)), \ + "no manual sidecar yet -> dialog starts at identity, not the automatic result" + + +def test_reference_angle_is_locked(ctx): + dlg = ctx.dlg + dlg.combo_active_angle.setCurrentIndex(dlg._ref_angle_idx) + pump(30) + before_ref = dlg._angle_params[dlg._ref_angle_idx] + dlg._on_nudge_translate(1, 0, False) + dlg._on_nudge_rotate(1, False) + assert not dlg.grp_manual_adjust.isEnabled(), "reference angle group disabled" + assert dlg._angle_params[dlg._ref_angle_idx] == before_ref, \ + "reference angle untouched by nudge attempts" + + +def test_nudges(ctx): + """Nudging a real angle (fine + coarse, translate + rotate).""" + dlg, s = ctx.dlg, ctx.s + ctx.active = active = 1 if s.n_angles > 1 else 0 + dlg.combo_active_angle.setCurrentIndex(active) + pump(30) + before = dlg._angle_params[active].shift_mm + dlg._on_nudge_translate(1, 0, False) # fine +X + fine_step = dlg.spin_step_translate_mm.value() + assert abs(dlg._angle_params[active].shift_mm[0] - (before[0] + fine_step)) < 1e-9, \ + "fine translate nudge moved shift_x by exactly one fine step" + + before = dlg._angle_params[active].shift_mm + dlg._on_nudge_translate(0, -1, True) # coarse -Y + coarse_step = fine_step * dlg.spin_step_multiplier.value() + assert abs(dlg._angle_params[active].shift_mm[1] - (before[1] - coarse_step)) < 1e-9, \ + "coarse translate nudge uses the multiplier" + + before_rot = dlg._angle_params[active].rotation_deg + dlg._on_nudge_rotate(1, False) + assert dlg._angle_params[active].rotation_deg != before_rot, \ + "rotate nudge changed rotation_deg" + assert len(dlg._preview_layers) == s.n_angles, \ + "preview canvas rebuilt for every angle after a rotation nudge" + + # Real key-event wiring (proves keyPressEvent -> signal -> slot). + before = dlg._angle_params[active].shift_mm + QTest.keyClick(dlg.canvas, Qt.Key.Key_Right) + assert dlg._angle_params[active].shift_mm[0] > before[0], \ + "a real Right-arrow key event nudged shift_x" + + +def test_auto_derotate(ctx): + """Auto De-rotate: seeds rotation from the stage angle, no translation.""" + dlg, s, active = ctx.dlg, ctx.s, ctx.active + shift_before_derotate = dlg._angle_params[active].shift_mm + dlg._on_auto_derotate() + nominal = compute._nominal_delta_deg(s, active, dlg._ref_angle_idx) + assert abs(dlg._angle_params[active].rotation_deg - nominal) < 1e-6, \ + "auto de-rotate seeded rotation from the stage's reported angle" + assert dlg._angle_params[active].shift_mm == shift_before_derotate, \ + "auto de-rotate left translation untouched" + assert dlg._angle_params[dlg._ref_angle_idx].rotation_deg == 0.0, \ + "reference angle stays identity after auto de-rotate" + # Clicking again offers the other sign, since which one lines the scans up + # is not knowable from the file. + dlg._on_auto_derotate() + assert abs(dlg._angle_params[active].rotation_deg + nominal) < 1e-6, \ + "auto de-rotate offers the opposite sign on a second click" + + +def test_auto_cross_correlate(ctx): + """Auto Cross-Correlate: searches rotation *and* translation.""" + win, dlg, s = ctx.win, ctx.dlg, ctx.s + assert dlg.btn_auto_correlate.isEnabled(), \ + "cross-correlate action enabled once masks are ready" + for label_idx, (label, _sources) in enumerate(dlg._CORRELATE_SOURCES): + dlg.combo_correlate_source.setCurrentIndex(label_idx) + dlg._on_auto_correlate() + assert wait_until( + lambda: not win._job_running("manual_align_correlate"), + timeout_ms=60000), f"auto cross-correlate completed ({label})" + assert all(a in dlg._fit_notes for a in range(s.n_angles) + if a != dlg._ref_angle_idx), \ + f"every non-reference angle got a fit ({label})" + assert dlg._angle_params[dlg._ref_angle_idx] == compute.ManualAngleParams(), \ + "auto cross-correlate reference angle stays identity" + assert dlg.grp_correlate.isEnabled() and dlg.btn_save.isEnabled(), \ + "auto cross-correlate re-enabled controls when done" + assert len(dlg._preview_layers) == s.n_angles, \ + "preview canvas rebuilt after cross-correlate" + assert dlg._fit_report(), "fit quality is reported per angle" + + +def test_save_sidecar(ctx): + win, dlg, s, active = ctx.win, ctx.dlg, ctx.s, ctx.active + dlg._on_save() + sidecar = compute.sidecar_path(s.path) + assert sidecar.exists(), "sidecar file written" + ctx.sidecar = sidecar + ctx.sidecar_raw = raw = json.loads(sidecar.read_text()) + assert raw.get("schema_version") == compute._SIDECAR_SCHEMA_VERSION, \ + "sidecar schema_version is current" + assert all(raw.get("per_angle", {}).get(str(a), {}).get("rotation_deg") + == dlg._angle_params[a].rotation_deg for a in range(s.n_angles)), \ + "sidecar per_angle round-trips the dialog's resolved params" + assert (win._alignment_result is not None + and win._alignment_result.per_angle[active].rotation_deg + == dlg._angle_params[active].rotation_deg), \ + "main window's alignment_result replaced by the manual build" + assert win.chk_aligned_view.isEnabled() and win.chk_aligned_view.isChecked(), \ + "Aligned View auto-enabled after Save" + + +def test_stale_schema_sidecar_ignored(ctx): + """An old-schema sidecar (pre-pivot/sign fix) is treated as absent.""" + s, raw, sidecar = ctx.s, ctx.sidecar_raw, ctx.sidecar + stale = dict(raw) + stale["schema_version"] = compute._SIDECAR_SCHEMA_VERSION - 1 + sidecar.write_text(json.dumps(stale)) + assert compute.load_manual_alignment(s) is None, \ + "a sidecar with an old schema_version is not loaded" + sidecar.write_text(json.dumps(raw)) # restore for the rest of the sequence + + +def test_clear_with_confirmation(ctx): + win, dlg, s = ctx.win, ctx.dlg, ctx.s + with patch("sras_viewer.QMessageBox.question", + return_value=QMessageBox.StandardButton.Yes): + dlg._on_clear() + assert not ctx.sidecar.exists(), "sidecar file deleted" + assert all(dlg._angle_params[a] == compute.ManualAngleParams() + for a in range(s.n_angles)), "dialog params reset to identity" + assert win._alignment_result is None, "main window alignment_result cleared" + assert (not win.chk_aligned_view.isEnabled() + and not win.chk_aligned_view.isChecked()), \ + "Aligned View disabled after Clear" + + dlg.close() + pump(150) + assert win._manual_align_dialog is None, "dialog reference released on close" + + +def test_sidecar_restored_on_reload(ctx): + win, active = ctx.win, ctx.active + win._on_manual_alignment() + dlg = win._manual_align_dialog + dlg.combo_active_angle.setCurrentIndex(active) + pump(30) + dlg._on_auto_derotate() + dlg._on_nudge_translate(1, 1, True) + saved_rotation = dlg._angle_params[active].rotation_deg + saved_shift = dlg._angle_params[active].shift_mm + dlg._on_save() + dlg.close() + pump(150) + + old_sras_id = id(win._sras) + win._load_file(str(ctx.path)) # reload the same file fresh + assert wait_until( + lambda: win._sras is not None and id(win._sras) != old_sras_id), \ + "file reloaded" + ctx.s = win._sras + assert win._manual_align_dialog is None, \ + "manual dialog force-closed by a reload" + assert win._alignment_result is not None, \ + "reload restores the saved manual alignment automatically" + assert abs(win._alignment_result.per_angle[active].rotation_deg + - saved_rotation) < 1e-9, "restored rotation matches what was saved" + assert win._alignment_result.per_angle[active].shift_mm == saved_shift, \ + "restored shift matches what was saved" + assert win.chk_aligned_view.isChecked(), \ + "Aligned View auto-checked after restoring a saved alignment" + + +def test_pixel_inspector(ctx): + win = ctx.win + win.chk_aligned_view.setChecked(False) + pump(100) + win._on_pixel_clicked(0, 0) + pump(150) + assert win.lbl_wave_hint.isHidden(), "waveform hint hidden after a click" + win.combo_channel.setCurrentIndex(CH1_IDX) + wait_until(lambda: not win._job_running("compute")) + win._on_pixel_clicked(1, 1) + pump(150) + assert len(win.wave_canvas.ax_wave.lines) > 0, \ + f"RF waveform panel rendered ({len(win.wave_canvas.ax_wave.lines)} lines)" + + +def test_shutdown(ctx): + win = ctx.win + win.close() + pump(400) + assert len(win._jobs) == 0, f"all background jobs released: {list(win._jobs)}" + + +def test_no_status_bar_errors(ctx): + unexpected = [e for e in ctx.errors if e] + assert not unexpected, f"status-bar errors seen: {unexpected}" diff --git a/tools/check_equivalence.py b/tools/check_equivalence.py index fa11968..37934dc 100644 --- a/tools/check_equivalence.py +++ b/tools/check_equivalence.py @@ -1,13 +1,9 @@ #!/usr/bin/env python3 -"""Golden-output equivalence harness for the sras-viewer refactor. +"""Golden-output equivalence harness for compute-path refactors. Computes a battery of DC / FFT / alignment outputs and prints a stable hash -for each. Run it on the pre-refactor commit to capture a baseline, then again -after the refactor and diff the two reports — every line must match. - -Imports work against both the pre-refactor monolith (`sras_viewer`) and the -post-refactor split (`sras_format` + `sras_compute`), so the *same* script -produces both sides of the comparison. +for each. Run it before a refactor to capture a baseline, then again after +and diff the two reports — every line must match. Hashes canonicalise to native little-endian float64 before hashing, so a deliberate dtype/byte-order change that preserves values does not show up as @@ -29,23 +25,11 @@ import numpy as np sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) -# --- Import shim: split modules if present, else the monolith -------------- -try: - from sras_format import SrasFile, CH1_IDX, CH3_IDX, CH4_IDX, adc_to_mv - import sras_compute as C - _LAYOUT = "split" -except ImportError: - import sras_viewer as _V - from sras_viewer import SrasFile, CH1_IDX, CH3_IDX, CH4_IDX, adc_to_mv - C = _V - _LAYOUT = "monolith" - -compute_dc_image = C.compute_dc_image -compute_rf_image = C.compute_rf_image -compute_alignment = C._compute_angle_alignment -apply_alignment = C.apply_alignment - -import tools.make_test_sras as gen # noqa: E402 +from sras_format import SrasFile, CH3_IDX, CH4_IDX, adc_to_mv # noqa: E402 +from sras_compute import ( # noqa: E402 + apply_alignment, compute_angle_alignment, compute_dc_image, compute_rf_image, +) +import tools.make_test_sras as gen # noqa: E402 def h(arr) -> str: @@ -109,7 +93,7 @@ def check_file(path: Path, lines: list[str], tag: str, for bg in (False, True): if bg and s.background is None: continue - for pad in (1, 2): + for pad in (1, 2, 4, 8, 40): n_fft = s.samples_per_frame * pad if pad > 1 else None for ti, thr in enumerate(thresholds): img = compute_rf_image(s, a, dc_threshold_mv=thr, @@ -135,7 +119,7 @@ def check_alignment(path: Path, lines: list[str], tag: str): sras.ch_ymult_mv[CH4_IDX], sras.ch_yoff_adc[CH4_IDX], sras.ch_yzero_mv[CH4_IDX]) thr = float(np.median(dc4)) - res = compute_alignment(sras, 0, thr) + res = compute_angle_alignment(sras, 0, thr) report(lines, f"[{tag}] align canvas_shape", str(res.canvas_shape)) report(lines, f"[{tag}] align canvas_origin", f"{res.canvas_origin_mm[0]:.9g},{res.canvas_origin_mm[1]:.9g}") @@ -165,7 +149,7 @@ def main(): help="directory for generated synthetic files") args = p.parse_args() - lines = [f"# layout: {_LAYOUT}", f"# numpy: {np.__version__}"] + lines = [f"# numpy: {np.__version__}"] scratch = Path(args.scratch) synth = scratch / "equiv_synth.sras" @@ -179,6 +163,19 @@ def main(): gen.write(synth_odd, n_angles=2, seed=7, samples_per_frame=37) check_file(synth_odd, lines, "odd", angles=[0, 1], n_rows=None) + # A legacy v4 file exercises the uniform-geometry legacy layout through + # the same DC/FFT battery. + synth_v4 = scratch / "equiv_synth_v4.sras" + gen.write_legacy(synth_v4, version=4, n_angles=2, n_rows=6, + n_frames=14, samples_per_frame=48, seed=5) + check_file(synth_v4, lines, "v4", angles=[0, 1], n_rows=None) + + # A big-endian int16 v6 file (real acquisitions are >i2; the other + # synthetics are int8). + synth_i16 = scratch / "equiv_synth_i16.sras" + gen.write(synth_i16, n_angles=2, seed=9, samples_per_frame=64, bps=2) + check_file(synth_i16, lines, "int16", angles=[0, 1], n_rows=None) + if args.real: real = Path(args.real) if real.exists(): diff --git a/tools/make_test_sras.py b/tools/make_test_sras.py index 71e93ef..92b0b54 100644 --- a/tools/make_test_sras.py +++ b/tools/make_test_sras.py @@ -42,12 +42,12 @@ def _preamble(ymult_v: float, yoff_adc: float, yzero_v: float) -> bytes: def build(n_angles: int, seed: int, samples_per_frame: int, - geometry: list[tuple[int, int]] | None = None) -> tuple[bytes, dict]: + geometry: list[tuple[int, int]] | None = None, + bps: int = 1) -> tuple[bytes, dict]: rng = np.random.default_rng(seed) src_geom = geometry or _GEOMETRY geom = [src_geom[a % len(src_geom)] for a in range(n_angles)] n_ch = 3 - bps = 1 angles_deg = np.linspace(0.0, 60.0, n_angles, dtype=np.float32) # Distinct calibration per channel so a swapped-channel bug is visible. @@ -100,7 +100,9 @@ def build(n_angles: int, seed: int, samples_per_frame: int, block[r, 1, f] = np.int8((a * 7 + r * 3 + f) % 100 - 50) block[r, 2, f] = np.int8((a * 5 + r * 11 + f * 2) % 120 - 60) waveforms.append(block) - out += block.tobytes() + # bps=2 stores the same values big-endian int16, exercising the + # reader's >i2 memmap path. + out += (block.astype(">i2") if bps == 2 else block).tobytes() meta = { "n_angles": n_angles, @@ -119,8 +121,9 @@ def build(n_angles: int, seed: int, samples_per_frame: int, def write(path: Path, n_angles: int = 3, seed: int = 0, samples_per_frame: int = 64, - geometry: list[tuple[int, int]] | None = None) -> dict: - payload, meta = build(n_angles, seed, samples_per_frame, geometry) + geometry: list[tuple[int, int]] | None = None, + bps: int = 1) -> dict: + payload, meta = build(n_angles, seed, samples_per_frame, geometry, bps=bps) path.write_bytes(payload) return meta diff --git a/tools/test_alignment.py b/tools/test_alignment.py deleted file mode 100644 index f78f9fa..0000000 --- a/tools/test_alignment.py +++ /dev/null @@ -1,223 +0,0 @@ -#!/usr/bin/env python3 -"""Angle-alignment tests: does registration actually stack the scans? - -Builds a synthetic scan in which one sample is imaged at several *known* -rotations and offsets (tools/make_test_sras.write_rotating) and checks that the -alignment path recovers them, that the shared canvas is angle 0's own pixel -grid extended, and that nothing in the result depends on any other angle's -stage coordinates. - -No Qt — this exercises sras_compute directly. See tools/test_gui.py for the -dialog and Aligned-View plumbing. - -Usage: python tools/test_alignment.py -""" - -import sys -import tempfile -from pathlib import Path - -import numpy as np - -sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) - -import sras_compute as compute # noqa: E402 -from sras_format import CH4_IDX, SrasFile, adc_to_mv # noqa: E402 -import tools.make_test_sras as gen # noqa: E402 - -# Registration is limited by how far a feature moves per degree: with this -# sample's ~1 mm radius and a ~16 µm registration pitch, a quarter degree is -# already sub-pixel, so it is the floor of what any metric can resolve here. -_ROT_TOL_DEG = 0.5 -_SHIFT_TOL_MM = 0.02 -_STACK_IOU_MIN = 0.90 -_THRESHOLD_MV = 80.0 - -_failures: list[str] = [] - - -def check(name: str, ok: bool, detail: str = ""): - print(f" {'PASS' if ok else 'FAIL'} {name}" + (f" — {detail}" if detail else "")) - if not ok: - _failures.append(name) - - -def dc4_images(sras: SrasFile) -> dict[int, np.ndarray]: - return {a: adc_to_mv(compute.compute_dc_image(sras, a, CH4_IDX), *sras.cal(CH4_IDX)) - for a in range(sras.n_angles)} - - -def mm_transform(sras: SrasFile, result, angle_idx: int) -> np.ndarray: - """Recover the pure mm-space rotation from a canvas->raw affine. - - matrix == D @ R^T @ A_out, where A_out and D only carry the canvas and - per-angle pixel pitches; undoing both must leave something orthonormal, or - the transform is smuggling in a scale or a shear. - """ - dx_a, dy_a = compute._pixel_pitch_mm(sras, angle_idx) - A_out = np.array([[0.0, result.canvas_dx_mm], [result.canvas_dy_mm, 0.0]]) - D = np.array([[0.0, 1.0 / dy_a], [1.0 / dx_a, 0.0]]) - return np.linalg.inv(D) @ result.per_angle[angle_idx].matrix @ np.linalg.inv(A_out) - - -def main() -> int: - tmpdir = Path(tempfile.mkdtemp(prefix="sras_align_")) - path = tmpdir / "rotating.sras" - meta = gen.write_rotating(path, n_angles=5) - sras = SrasFile(str(path)) - truth = meta["truth"] - - print(f"\nrotating-sample scan: {sras.n_angles} angles, " - f"shapes {[sras.image_shape(a) for a in range(sras.n_angles)]}") - - print("\nper-angle rigid registration (rotation + translation, no scale)") - dc4 = dc4_images(sras) - fits = {a: compute.register_angle_to_reference( - sras, a, 0, dc4, dc_threshold_mv=_THRESHOLD_MV) - for a in range(sras.n_angles)} - for a, fit in fits.items(): - t_rot, t_shift = truth[a] - rot_err = abs(fit.rotation_deg - t_rot) - shift_err = float(np.hypot(fit.shift_mm[0] - t_shift[0], - fit.shift_mm[1] - t_shift[1])) - check(f"angle {a} rotation within {_ROT_TOL_DEG}° of truth", - rot_err <= _ROT_TOL_DEG, - f"got {fit.rotation_deg:.3f}°, truth {t_rot:.3f}° (err {rot_err:.3f}°)") - check(f"angle {a} translation within {_SHIFT_TOL_MM} mm of truth", - shift_err <= _SHIFT_TOL_MM, f"err {shift_err:.4f} mm") - check("reference angle registers as exact identity", - fits[0] == compute.RigidFit(0.0, (0.0, 0.0), 1.0, "reference")) - - # The stage's rotational sense relative to this module's math-positive - # convention is not knowable from the file, and the old code hardcoded a - # guess. Flipping every reported angle must therefore change nothing: the - # search scores both signs and the images decide. - flipped = SrasFile(str(path)) - flipped.angles_deg = -flipped.angles_deg - flipped_fits = {a: compute.register_angle_to_reference( - flipped, a, 0, dc4, dc_threshold_mv=_THRESHOLD_MV) - for a in range(1, flipped.n_angles)} - check("negating every reported stage angle changes no fit", - all(flipped_fits[a] == fits[a] for a in flipped_fits), - str({a: (flipped_fits[a].rotation_deg, fits[a].rotation_deg) - for a in flipped_fits if flipped_fits[a] != fits[a]})) - - print("\nper-angle stage coordinates are not consulted") - # Move every non-reference angle's scan window somewhere else entirely. - # Only angle 0's coordinates may matter, so every fit must be untouched. - moved = SrasFile(str(path)) - for a in range(1, moved.n_angles): - moved.x_start_mm[a] += 13.5 * a - moved._y_pos_per_angle[a] = moved._y_pos_per_angle[a] - 9.25 * a - moved_dc4 = dc4_images(moved) - moved_fits = {a: compute.register_angle_to_reference( - moved, a, 0, moved_dc4, dc_threshold_mv=_THRESHOLD_MV) - for a in range(1, moved.n_angles)} - check("relocating every other angle's scan window changes no fit", - all(moved_fits[a] == fits[a] for a in moved_fits), - str({a: (round(moved_fits[a].rotation_deg, 4), fits[a].rotation_deg) - for a in moved_fits if moved_fits[a] != fits[a]})) - - print("\nshared canvas is angle 0's own pixel grid, extended") - result = compute.compute_angle_alignment(sras, 0, _THRESHOLD_MV) - t0 = result.per_angle[0] - check("angle 0's transform has no rotation, scale or shear", - np.allclose(t0.matrix, np.eye(2)), str(t0.matrix)) - check("angle 0 lands on whole canvas pixels (no resampling of the reference)", - np.allclose(t0.offset, np.round(t0.offset)), str(t0.offset)) - check("canvas pitch is angle 0's own pitch", - (result.canvas_dx_mm, result.canvas_dy_mm) - == compute._pixel_pitch_mm(sras, 0)) - - n_rows, n_cols = result.canvas_shape - x_axis = result.canvas_origin_mm[0] + np.arange(n_cols) * result.canvas_dx_mm - y_axis = result.canvas_origin_mm[1] + np.arange(n_rows) * result.canvas_dy_mm - row0, col0 = int(round(-t0.offset[0])), int(round(-t0.offset[1])) - a0_rows, a0_cols = sras.image_shape(0) - check("canvas X axis reproduces angle 0's own X coordinates", - np.allclose(x_axis[col0:col0 + a0_cols], sras.x_axis_mm(0))) - check("canvas Y axis reproduces angle 0's own Y coordinates", - np.allclose(y_axis[row0:row0 + a0_rows], sras.y_positions_mm(0))) - check("canvas covers every angle's footprint", - n_rows >= max(int(sras.n_rows[a]) for a in range(sras.n_angles)) - and n_cols >= max(int(sras.n_frames[a]) for a in range(sras.n_angles)), - str(result.canvas_shape)) - - print("\nno scaling anywhere in the per-angle transforms") - for a in range(sras.n_angles): - R = mm_transform(sras, result, a) - check(f"angle {a}'s mm-space transform is a pure rotation", - np.allclose(R @ R.T, np.eye(2), atol=1e-9) - and abs(abs(np.linalg.det(R)) - 1.0) < 1e-9, - f"det={np.linalg.det(R):.6f}") - - print("\nall angles stack into one shape") - aligned = {a: compute.apply_alignment(result, a, dc4[a]) - for a in range(sras.n_angles)} - base = aligned[0] >= _THRESHOLD_MV - for a in range(1, sras.n_angles): - other = aligned[a] >= _THRESHOLD_MV - iou = float((base & other).sum()) / max(1, int((base | other).sum())) - check(f"angle {a}'s aligned sample overlaps angle 0's (IoU >= {_STACK_IOU_MIN})", - iou >= _STACK_IOU_MIN, f"IoU {iou:.4f}") - - print("\ndownsampled preview lands where the full-resolution image does") - # ManualAlignmentDialog reprojects block-mean-downsampled masks, so the - # affine has to account for the factor. When it did not, every preview - # layer came out magnified by that factor and offset — the overlay showed a - # blown-up crop of each mask, which is not something you can align by eye. - pitch = (result.canvas_dx_mm, result.canvas_dy_mm) - a = sras.n_angles - 1 - p = result.per_angle[a] - full_mask = (dc4[a] >= _THRESHOLD_MV).astype(np.float32) - full = compute.reproject_mask( - sras, a, 0, full_mask, p.rotation_deg, p.shift_mm, pitch, - result.canvas_origin_mm, result.canvas_shape) - fy, fx = 4, 16 - small = compute.reproject_mask( - sras, a, 0, compute._block_mean_2d(full_mask, fy, fx), - p.rotation_deg, p.shift_mm, (pitch[0] * fx, pitch[1] * fy), - result.canvas_origin_mm, - (result.canvas_shape[0] // fy, result.canvas_shape[1] // fx), - src_downsample=(fy, fx)) - # Compare in mm, via each layer's own center of mass. - def com_mm(layer, px, py): - rows, cols = np.nonzero(layer > 0.5) - return np.array([cols.mean() * px, rows.mean() * py]) - d = com_mm(small, pitch[0] * fx, pitch[1] * fy) - com_mm(full, *pitch) - check("a downsampled preview layer lands within a preview pixel of the " - "full-resolution one", - abs(d[0]) <= abs(pitch[0] * fx) and abs(d[1]) <= abs(pitch[1] * fy), - f"offset {d[0]:+.4f}, {d[1]:+.4f} mm") - - print("\nmanual path reproduces the same geometry") - params = {a: compute.ManualAngleParams(t.rotation_deg, t.shift_mm) - for a, t in result.per_angle.items()} - manual = compute.build_manual_alignment(sras, 0, _THRESHOLD_MV, params) - check("build_manual_alignment matches compute_angle_alignment for the same params", - manual.canvas_shape == result.canvas_shape - and np.allclose(manual.canvas_origin_mm, result.canvas_origin_mm) - and all(np.allclose(manual.per_angle[a].matrix, result.per_angle[a].matrix) - and np.allclose(manual.per_angle[a].offset, result.per_angle[a].offset) - for a in range(sras.n_angles))) - - print("\nsidecar round-trip") - compute.save_manual_alignment(sras, 0, _THRESHOLD_MV, params) - loaded = compute.load_manual_alignment(sras) - check("sidecar reloads every angle's params", - loaded is not None - and all(np.isclose(loaded.per_angle[a].rotation_deg, params[a].rotation_deg) - and np.allclose(loaded.per_angle[a].shift_mm, params[a].shift_mm) - for a in range(sras.n_angles))) - check("sidecar deletes cleanly", compute.delete_manual_alignment(sras)) - - print() - if _failures: - print(f"{len(_failures)} FAILURE(S): " + ", ".join(_failures)) - return 1 - print("All alignment checks passed.") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/tools/test_gui.py b/tools/test_gui.py deleted file mode 100644 index 36bca72..0000000 --- a/tools/test_gui.py +++ /dev/null @@ -1,476 +0,0 @@ -#!/usr/bin/env python3 -"""Headless GUI test: drives SrasViewerWindow through the real Qt widgets, -signals and worker threads under the offscreen platform plugin. - -Covers the interactions a manual smoke test would: load, switch angles and -channels, background DC precompute, lazy FFT compute, threshold and bg-sub -changes, angle alignment, manual angle alignment, aligned view, ROI -draw/move, and CSV export. - -Usage: QT_QPA_PLATFORM=offscreen python tools/test_gui.py [file.sras] -""" - -import json -import os -import sys -import tempfile -from pathlib import Path -from unittest.mock import patch - -os.environ.setdefault("QT_QPA_PLATFORM", "offscreen") - -import numpy as np # noqa: E402 -from PyQt6.QtCore import QEventLoop, Qt, QTimer # noqa: E402 -from PyQt6.QtTest import QTest # noqa: E402 -from PyQt6.QtWidgets import QApplication, QMessageBox # noqa: E402 - -sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) - -import sras_compute as compute # noqa: E402 -from sras_format import CH1_IDX, CH3_IDX, CH4_IDX, SrasFile # noqa: E402 -from sras_viewer import RoiQuad, SrasViewerWindow, VELOCITY_MODE_IDX # noqa: E402 -import tools.make_test_sras as gen # noqa: E402 - -_failures: list[str] = [] - - -def check(name: str, ok: bool, detail: str = ""): - print(f" {'PASS' if ok else 'FAIL'} {name}" + (f" — {detail}" if detail else "")) - if not ok: - _failures.append(name) - - -def pump(ms: int = 250): - """Run the event loop for a while so queued signals and worker threads - make progress.""" - loop = QEventLoop() - QTimer.singleShot(ms, loop.quit) - loop.exec() - - -def wait_until(pred, timeout_ms: int = 20000, step: int = 100) -> bool: - waited = 0 - while waited < timeout_ms: - if pred(): - return True - pump(step) - waited += step - return pred() - - -def main(): - app = QApplication(sys.argv) - errors: list[str] = [] - - tmpdir = Path(tempfile.mkdtemp(prefix="sras_gui_")) - path = Path(sys.argv[1]) if len(sys.argv) > 1 else tmpdir / "gui.sras" - if len(sys.argv) <= 1: - gen.write(path, n_angles=4, seed=11, samples_per_frame=256) - - print(f"\nloading {path.name}") - win = SrasViewerWindow() - win.show() - # Capture anything the app reports as an error via the status bar. - win.statusBar().messageChanged.connect( - lambda m: errors.append(m) if m and "error" in m.lower() else None) - - win._load_file(str(path)) - check("file loaded", wait_until(lambda: win._sras is not None)) - s = win._sras - check("parsed as v6", s.version == 6, f"v{s.version}") - check("defaults to CH4", win.combo_channel.currentIndex() == CH4_IDX) - check("image displayed", win._current_image is not None) - check("angle spinbox ranges over all angles", - win.spin_angle.maximum() == s.n_angles - 1) - check("scan info populated", - win._info["Angles"].text() == f"Angles: {s.n_angles}", - win._info["Angles"].text()) - - print("\nbackground DC precompute (all angles)") - ok = wait_until(lambda: all((a, CH4_IDX) in win._dc_cache - and (a, CH3_IDX) in win._dc_cache - for a in range(s.n_angles))) - check("every angle cached for CH3 and CH4", ok, - f"{len(win._dc_cache)} entries") - check("status label reports completion", - "ready for all angles" in win.lbl_dc_precompute.text(), - win.lbl_dc_precompute.text()) - - print("\nangle switching (DC, should be served from cache)") - for a in range(s.n_angles): - win.spin_angle.setValue(a) - win._on_view_changed() - pump(60) - expected = win._sras.image_shape(a) - check(f"angle {a} shows its own geometry {expected}", - win._current_image.shape == expected, - str(win._current_image.shape)) - check("no compute job needed for cached DC angles", - not win._job_running("compute")) - - print("\nchannel switching") - win.spin_angle.setValue(0) - win._on_view_changed() - pump(60) - win.combo_channel.setCurrentIndex(CH3_IDX) - check("CH3 displayed", wait_until(lambda: win._current_ch == CH3_IDX)) - - win.combo_channel.setCurrentIndex(CH1_IDX) - check("CH1 (FFT) computed", wait_until( - lambda: win._current_ch == CH1_IDX and not win._job_running("compute"))) - check("FFT result cached", len(win._fft_cache) > 0, f"{len(win._fft_cache)} keys") - rf_img = win._current_image - check("FFT image is non-degenerate", len(np.unique(rf_img)) > 1, - f"{len(np.unique(rf_img))} unique values") - - print("\nvelocity mode (pure post-multiply, no recompute)") - n_fft_before = len(win._fft_cache) - win.combo_channel.setCurrentIndex(VELOCITY_MODE_IDX) - check("velocity displayed", wait_until( - lambda: win._current_ch == VELOCITY_MODE_IDX and not win._job_running("compute"))) - grating = win.spin_grating_um.value() - check("velocity == freq x grating", - np.allclose(win._current_image, rf_img * grating, atol=1e-3)) - check("velocity reused the cached FFT", len(win._fft_cache) == n_fft_before, - f"{n_fft_before} -> {len(win._fft_cache)}") - check("grating spinbox visible in velocity mode", win.grp_velocity.isVisible()) - - print("\nthreshold change (genuine cache-key change)") - win.combo_channel.setCurrentIndex(CH1_IDX) - wait_until(lambda: not win._job_running("compute")) - dc4 = win._dc_cache[(0, CH4_IDX)] - win.spin_threshold_mv.setValue(float(np.median(dc4))) - win._on_threshold_changed() - check("recomputed at new threshold", wait_until( - lambda: not win._job_running("compute") and len(win._fft_cache) > n_fft_before)) - check("masking zeroed some pixels", - int((win._current_image == 0).sum()) > 0, - f"{int((win._current_image == 0).sum())} of {win._current_image.size}") - - print("\nbackground subtraction toggle") - n_before = len(win._fft_cache) - win.chk_bg_sub.setChecked(False) - check("recomputed without bg-sub", wait_until( - lambda: not win._job_running("compute") and len(win._fft_cache) > n_before)) - win.chk_bg_sub.setChecked(True) - pump(200) - check("returning to bg-sub was a cache hit (no recompute)", - not win._job_running("compute")) - - print("\nROI") - x = s.x_axis_mm(0) - y = s.y_positions_mm(0) - roi = RoiQuad.from_bbox(float(x[1]), float(y[1]), - float(x[-2]), float(y[-2])) - win.image_canvas.set_roi(roi) - pump(120) - check("ROI registered", win.image_canvas.get_roi() is not None) - check("pixel count reported", - "pixels inside" in win.lbl_roi_npix.text() - and win.lbl_roi_npix.text() != "pixels inside: —", - win.lbl_roi_npix.text()) - npix = int(win.lbl_roi_npix.text().split(":")[1]) - check("ROI pixel count is plausible", - 0 < npix <= win._current_image.size, f"{npix}") - check("Export ROI enabled", win.btn_export_roi.isEnabled()) - - csv_path = tmpdir / "roi.csv" - with patch("sras_viewer.QFileDialog.getSaveFileName", - return_value=(str(csv_path), "")): - win._on_export_roi_csv() - check("ROI CSV written", csv_path.exists()) - if csv_path.exists(): - body = [l for l in csv_path.read_text().splitlines() if not l.startswith("#")] - check("ROI CSV has header + one line per pixel", - len(body) == npix + 1, f"{len(body)} lines for {npix} pixels") - - img_csv = tmpdir / "img.csv" - with patch("sras_viewer.QFileDialog.getSaveFileName", - return_value=(str(img_csv), "")): - win._on_export_csv() - check("image CSV written", img_csv.exists()) - if img_csv.exists(): - arr = np.loadtxt(img_csv, delimiter=",") - check("image CSV round-trips the displayed image", - arr.shape == win._current_image.shape - and np.allclose(arr, win._current_image, rtol=1e-5, atol=1e-4)) - - print("\nROI survives angle and channel switches") - win.spin_angle.setValue(1) - win._on_view_changed() - wait_until(lambda: not win._job_running("compute")) - check("ROI still present after angle switch", - win.image_canvas.get_roi() is not None) - win.combo_channel.setCurrentIndex(CH4_IDX) - wait_until(lambda: win._current_ch == CH4_IDX) - check("ROI still present after channel switch", - win.image_canvas.get_roi() is not None) - - print("\nangle alignment (Fusion)") - win.spin_angle.setValue(0) - win._on_view_changed() - wait_until(lambda: not win._job_running("compute")) - check("alignment action enabled", win._alignment_act.isEnabled()) - win._on_angle_alignment() - check("alignment completed", wait_until( - lambda: win._alignment_result is not None and not win._job_running("align"), - timeout_ms=60000)) - if win._alignment_result is not None: - r = win._alignment_result - check("transform for every angle", len(r.per_angle) == s.n_angles) - check("canvas is at least as large as any single angle", - all(r.canvas_shape[0] >= int(s.n_rows[a]) - and r.canvas_shape[1] >= int(s.n_frames[a]) - for a in range(s.n_angles)), str(r.canvas_shape)) - check("reference angle has zero shift", - r.per_angle[r.ref_angle_idx].shift_mm == (0.0, 0.0)) - check("Aligned View auto-enabled and checked", - win.chk_aligned_view.isEnabled() and win.chk_aligned_view.isChecked()) - pump(200) - check("displayed image is on the alignment canvas", - win.image_canvas._img_shape == r.canvas_shape, - f"{win.image_canvas._img_shape} vs {r.canvas_shape}") - - win.chk_aligned_view.setChecked(False) - pump(200) - check("unchecking returns to the raw per-angle grid", - win.image_canvas._img_shape == s.image_shape(0), - str(win.image_canvas._img_shape)) - - print("\nmanual alignment (Fusion)") - check("manual alignment action enabled", win._manual_align_act.isEnabled()) - - # --- Local mm is anchored on each angle's array center, not its stage -- - # position: that is what makes a scan's placement independent of where its - # window happened to sit. (Registration accuracy itself is covered by - # tools/test_alignment.py, which has a synthetic sample to register.) - n_rows, n_frames = s.image_shape(0) - check("array center is the geometric center of the pixel grid", - np.allclose(compute._center_idx(s, 0), - [(n_rows - 1) / 2, (n_frames - 1) / 2])) - dx0, dy0 = compute._pixel_pitch_mm(s, 0) - check("local half-extent is derived from shape and pitch alone", - np.allclose(compute._local_half_extent_mm(s, 0), - [(n_frames - 1) / 2 * abs(dx0), (n_rows - 1) / 2 * abs(dy0)])) - identity = {a: compute.ManualAngleParams() for a in range(s.n_angles)} - origin_a, shape_a = compute.canvas_for_params(s, 0, (dx0, dy0), identity) - moved = SrasFile(str(path)) - for a in range(1, moved.n_angles): - moved.x_start_mm[a] += 7.5 - moved._y_pos_per_angle[a] = moved._y_pos_per_angle[a] + 3.25 - origin_b, shape_b = compute.canvas_for_params(moved, 0, (dx0, dy0), identity) - check("moving every non-reference angle's scan window leaves the canvas " - "unchanged (only angle 0's coordinates are used)", - shape_a == shape_b and np.allclose(origin_a, origin_b), - f"{origin_a} {shape_a} vs {origin_b} {shape_b}") - - # --- Both signs of the stage's reported angle are searched -------------- - cands = compute._rotation_candidates(30.0, 6.0, 2.0) - check("rotation candidates bracket both signs of the stage angle", - min(cands) < -29.0 and max(cands) > 29.0, f"{min(cands)}..{max(cands)}") - - # --- Whole-pixel translation must not wrap content around the edge ------ - arr = np.zeros((6, 6), dtype=np.float32) - arr[0, 0] = 1.0 - check("_shift_into zero-fills rather than wrapping", - compute._shift_into(arr, -1, -1).sum() == 0.0) - check("_shift_into moves content by exactly the requested offset", - compute._shift_into(arr, 2, 3)[2, 3] == 1.0) - - # --- Open: must NOT seed from the still-live automatic AlignmentResult -- - # Manual mode exists to fix up whatever the automatic registration got - # wrong, so it must start from identity (every angle centered on the - # reference, no rotation) regardless of whatever the automatic run last - # computed. Only a previously *saved manual* alignment (sidecar) should - # ever seed this dialog. - win._on_manual_alignment() - check("dialog opened", win._manual_align_dialog is not None) - dlg = win._manual_align_dialog - check("mask prep needed no background worker (already DC-cached)", - not win._job_running("manual_align_masks")) - check("no manual sidecar yet -> dialog starts at identity, not the " - "automatic result", - all(dlg._angle_params[a] == compute.ManualAngleParams() - for a in range(s.n_angles))) - - # --- Reference angle is locked ------------------------------------------- - dlg.combo_active_angle.setCurrentIndex(dlg._ref_angle_idx) - pump(30) - before_ref = dlg._angle_params[dlg._ref_angle_idx] - dlg._on_nudge_translate(1, 0, False) - dlg._on_nudge_rotate(1, False) - check("reference angle group disabled", not dlg.grp_manual_adjust.isEnabled()) - check("reference angle untouched by nudge attempts", - dlg._angle_params[dlg._ref_angle_idx] == before_ref) - - # --- Nudging a real angle (fine + coarse, translate + rotate) ----------- - active = 1 if s.n_angles > 1 else 0 - dlg.combo_active_angle.setCurrentIndex(active) - pump(30) - before = dlg._angle_params[active].shift_mm - dlg._on_nudge_translate(1, 0, False) # fine +X - fine_step = dlg.spin_step_translate_mm.value() - check("fine translate nudge moved shift_x by exactly one fine step", - abs(dlg._angle_params[active].shift_mm[0] - (before[0] + fine_step)) < 1e-9) - - before = dlg._angle_params[active].shift_mm - dlg._on_nudge_translate(0, -1, True) # coarse -Y - coarse_step = fine_step * dlg.spin_step_multiplier.value() - check("coarse translate nudge uses the multiplier", - abs(dlg._angle_params[active].shift_mm[1] - (before[1] - coarse_step)) < 1e-9) - - before_rot = dlg._angle_params[active].rotation_deg - dlg._on_nudge_rotate(1, False) - check("rotate nudge changed rotation_deg", - dlg._angle_params[active].rotation_deg != before_rot) - check("preview canvas rebuilt for every angle after a rotation nudge", - len(dlg._preview_layers) == s.n_angles) - - # --- Real key-event wiring (proves keyPressEvent -> signal -> slot) ----- - before = dlg._angle_params[active].shift_mm - QTest.keyClick(dlg.canvas, Qt.Key.Key_Right) - check("a real Right-arrow key event nudged shift_x", - dlg._angle_params[active].shift_mm[0] > before[0]) - - # --- Auto De-rotate: seeds rotation from the stage angle, no translation - - shift_before_derotate = dlg._angle_params[active].shift_mm - dlg._on_auto_derotate() - nominal = compute._nominal_delta_deg(s, active, dlg._ref_angle_idx) - check("auto de-rotate seeded rotation from the stage's reported angle", - abs(dlg._angle_params[active].rotation_deg - nominal) < 1e-6) - check("auto de-rotate left translation untouched", - dlg._angle_params[active].shift_mm == shift_before_derotate) - check("reference angle stays identity after auto de-rotate", - dlg._angle_params[dlg._ref_angle_idx].rotation_deg == 0.0) - # Clicking again offers the other sign, since which one lines the scans up - # is not knowable from the file. - dlg._on_auto_derotate() - check("auto de-rotate offers the opposite sign on a second click", - abs(dlg._angle_params[active].rotation_deg + nominal) < 1e-6) - - # --- Auto Cross-Correlate: searches rotation *and* translation ---------- - check("cross-correlate action enabled once masks are ready", - dlg.btn_auto_correlate.isEnabled()) - for label_idx, (label, _sources) in enumerate(dlg._CORRELATE_SOURCES): - dlg.combo_correlate_source.setCurrentIndex(label_idx) - dlg._on_auto_correlate() - check(f"auto cross-correlate completed ({label})", wait_until( - lambda: not win._job_running("manual_align_correlate"), timeout_ms=60000)) - check(f"every non-reference angle got a fit ({label})", - all(a in dlg._fit_notes for a in range(s.n_angles) - if a != dlg._ref_angle_idx)) - check("auto cross-correlate reference angle stays identity", - dlg._angle_params[dlg._ref_angle_idx] == compute.ManualAngleParams()) - check("auto cross-correlate re-enabled controls when done", - dlg.grp_correlate.isEnabled() and dlg.btn_save.isEnabled()) - check("preview canvas rebuilt after cross-correlate", - len(dlg._preview_layers) == s.n_angles) - check("fit quality is reported per angle", bool(dlg._fit_report()), - dlg._fit_report()) - - # --- Save ----------------------------------------------------------------- - dlg._on_save() - sidecar = compute.sidecar_path(s.path) - check("sidecar file written", sidecar.exists()) - raw = json.loads(sidecar.read_text()) if sidecar.exists() else {} - check("sidecar schema_version is current", - raw.get("schema_version") == compute._SIDECAR_SCHEMA_VERSION) - check("sidecar per_angle round-trips the dialog's resolved params", - all(raw.get("per_angle", {}).get(str(a), {}).get("rotation_deg") - == dlg._angle_params[a].rotation_deg for a in range(s.n_angles))) - check("main window's alignment_result replaced by the manual build", - win._alignment_result is not None - and win._alignment_result.per_angle[active].rotation_deg - == dlg._angle_params[active].rotation_deg) - check("Aligned View auto-enabled after Save", - win.chk_aligned_view.isEnabled() and win.chk_aligned_view.isChecked()) - - # --- An old-schema sidecar (pre-pivot/sign fix) is treated as absent ------ - stale = dict(raw) - stale["schema_version"] = compute._SIDECAR_SCHEMA_VERSION - 1 - sidecar.write_text(json.dumps(stale)) - check("a sidecar with an old schema_version is not loaded", - compute.load_manual_alignment(s) is None) - sidecar.write_text(json.dumps(raw)) # restore for the rest of this section - - # --- Clear (with confirmation) -------------------------------------------- - with patch("sras_viewer.QMessageBox.question", - return_value=QMessageBox.StandardButton.Yes): - dlg._on_clear() - check("sidecar file deleted", not sidecar.exists()) - check("dialog params reset to identity", - all(dlg._angle_params[a] == compute.ManualAngleParams() - for a in range(s.n_angles))) - check("main window alignment_result cleared", win._alignment_result is None) - check("Aligned View disabled after Clear", - not win.chk_aligned_view.isEnabled() and not win.chk_aligned_view.isChecked()) - - dlg.close() - pump(150) - check("dialog reference released on close", win._manual_align_dialog is None) - - # --- Sidecar auto-restore on next load ------------------------------------ - win._on_manual_alignment() - dlg = win._manual_align_dialog - dlg.combo_active_angle.setCurrentIndex(active) - pump(30) - dlg._on_auto_derotate() - dlg._on_nudge_translate(1, 1, True) - saved_rotation = dlg._angle_params[active].rotation_deg - saved_shift = dlg._angle_params[active].shift_mm - dlg._on_save() - dlg.close() - pump(150) - - old_sras_id = id(win._sras) - win._load_file(str(path)) # reload the same file fresh - check("file reloaded", wait_until( - lambda: win._sras is not None and id(win._sras) != old_sras_id)) - s = win._sras - check("manual dialog force-closed by a reload", win._manual_align_dialog is None) - check("reload restores the saved manual alignment automatically", - win._alignment_result is not None) - if win._alignment_result is not None: - check("restored rotation matches what was saved", - abs(win._alignment_result.per_angle[active].rotation_deg - - saved_rotation) < 1e-9) - check("restored shift matches what was saved", - win._alignment_result.per_angle[active].shift_mm == saved_shift) - check("Aligned View auto-checked after restoring a saved alignment", - win.chk_aligned_view.isChecked()) - - print("\npixel inspector") - win.chk_aligned_view.setChecked(False) - pump(100) - win._on_pixel_clicked(0, 0) - pump(150) - check("waveform hint hidden after a click", win.lbl_wave_hint.isHidden()) - win.combo_channel.setCurrentIndex(CH1_IDX) - wait_until(lambda: not win._job_running("compute")) - win._on_pixel_clicked(1, 1) - pump(150) - check("RF waveform panel rendered", - len(win.wave_canvas.ax_wave.lines) > 0, - f"{len(win.wave_canvas.ax_wave.lines)} lines") - - print("\nshutdown") - win.close() - pump(400) - check("all background jobs released", len(win._jobs) == 0, - f"{list(win._jobs)}") - - print() - unexpected = [e for e in errors if e] - if unexpected: - print(f"status-bar errors seen: {unexpected}") - _failures.append("status-bar errors") - - if _failures: - print(f"{len(_failures)} FAILURE(S): " + ", ".join(_failures)) - return 1 - print("All GUI checks passed.") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/tools/test_refactor.py b/tools/test_refactor.py deleted file mode 100644 index 0cc50f6..0000000 --- a/tools/test_refactor.py +++ /dev/null @@ -1,358 +0,0 @@ -#!/usr/bin/env python3 -"""Behavioural tests for the sras-viewer refactor. - -Covers what the golden-hash harness can't: the v6->v7 cache round-trip -(including block carry-forward), parallel-vs-serial identity, the no-mask -fast path, and the ROI bounding-box mask optimisation. - -Usage: python tools/test_refactor.py [--scratch DIR] -""" - -import argparse -import shutil -import sys -import tempfile -from pathlib import Path - -import numpy as np - -sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) - -import sras_compute as compute # noqa: E402 -from sras_compute import ( # noqa: E402 - cache_file, compute_dc_image, compute_rf_image, dc_image_mv, -) -from sras_format import CH3_IDX, CH4_IDX, SrasFile, adc_to_mv # noqa: E402 -import tools.make_test_sras as gen # noqa: E402 - -_failures: list[str] = [] - - -def check(name: str, ok: bool, detail: str = ""): - print(f" {'PASS' if ok else 'FAIL'} {name}" + (f" — {detail}" if detail else "")) - if not ok: - _failures.append(name) - - -def test_cache_roundtrip(scratch: Path): - """v6 -> v7 for DC, then FFT, asserting the first block survives the - second write (the carry-forward path in write_v7_cache).""" - print("\ncache round-trip (v6 -> v7, both blocks)") - path = scratch / "roundtrip.sras" - gen.write(path, n_angles=3, seed=1, samples_per_frame=64) - - src = SrasFile(str(path)) - check("source is v6", src.version == 6, f"got v{src.version}") - expect_dc3 = [dc_image_mv(src, a, CH3_IDX) for a in range(src.n_angles)] - expect_dc4 = [dc_image_mv(src, a, CH4_IDX) for a in range(src.n_angles)] - expect_fft = [compute_rf_image(src, a, dc_threshold_mv=None, apply_bg_sub=True) - for a in range(src.n_angles)] - - err = cache_file(str(path), "dc", True) - check("dc cache_file succeeded", err == "", err) - - after_dc = SrasFile(str(path)) - check("version flipped to 7", after_dc.version == 7, f"got v{after_dc.version}") - check("dc3 stored for every angle", - all(x is not None for x in after_dc.precomputed_dc3_mv)) - check("dc3 values round-trip", - all(np.allclose(after_dc.precomputed_dc3_mv[a], expect_dc3[a], atol=1e-4) - for a in range(after_dc.n_angles))) - check("dc4 values round-trip", - all(np.allclose(after_dc.precomputed_dc4_mv[a], expect_dc4[a], atol=1e-4) - for a in range(after_dc.n_angles))) - check("no fft block yet", - all(x is None for x in after_dc.precomputed_freq_mhz)) - check("cached images are native float32", - after_dc.precomputed_dc3_mv[0].dtype == np.float32 - and after_dc.precomputed_dc3_mv[0].dtype.byteorder in ("=", "|"), - str(after_dc.precomputed_dc3_mv[0].dtype.byteorder)) - check("cached images are writable", - after_dc.precomputed_dc3_mv[0].flags.writeable) - - err = cache_file(str(path), "fft", True) - check("fft cache_file succeeded", err == "", err) - - both = SrasFile(str(path)) - check("fft stored for every angle", - all(x is not None for x in both.precomputed_freq_mhz)) - check("fft values round-trip", - all(np.allclose(both.precomputed_freq_mhz[a], expect_fft[a], atol=1e-3) - for a in range(both.n_angles))) - check("DC block carried forward through the FFT write", - all(np.allclose(both.precomputed_dc3_mv[a], expect_dc3[a], atol=1e-4) - for a in range(both.n_angles))) - check("bg_sub flag persisted", both.precomputed_bg_sub is True) - - # The fast path must reproduce a fresh compute, and masking must still - # apply on top of a cached (unmasked) image. - fresh = SrasFile(str(path)) - fresh.precomputed_freq_mhz = [None] * fresh.n_angles - dc4 = dc_image_mv(both, 0, CH4_IDX) - thr = float(np.median(dc4)) - check("cached fast path == fresh compute (unmasked)", - np.allclose(compute_rf_image(both, 0, dc_threshold_mv=None, apply_bg_sub=True), - compute_rf_image(fresh, 0, dc_threshold_mv=None, apply_bg_sub=True), - atol=1e-3)) - check("cached fast path == fresh compute (masked)", - np.allclose(compute_rf_image(both, 0, dc_threshold_mv=thr, apply_bg_sub=True), - compute_rf_image(fresh, 0, dc_threshold_mv=thr, apply_bg_sub=True), - atol=1e-3)) - - # Waveform data must be byte-identical to the pre-cache file. - orig = scratch / "roundtrip_orig.sras" - gen.write(orig, n_angles=3, seed=1, samples_per_frame=64) - o, n = SrasFile(str(orig)), SrasFile(str(path)) - check("waveform data untouched by the cache write", - all(np.array_equal(np.asarray(o.data[a]), np.asarray(n.data[a])) - for a in range(o.n_angles))) - - -def test_partial_v7_cache(scratch: Path): - """Only some angles cached: uncached angles must compute, not read zeros. - This is the v5 bug the ragged normalisation fixed, checked via v7.""" - print("\npartial cache (only some angles stored)") - path = scratch / "partial.sras" - gen.write(path, n_angles=3, seed=2, samples_per_frame=64) - - src = SrasFile(str(path)) - expected = [compute_rf_image(src, a, dc_threshold_mv=None, apply_bg_sub=True) - for a in range(src.n_angles)] - partial = [expected[0], None, expected[2]] # angle 1 deliberately absent - src.write_v7_cache(new_freq_mhz=partial, new_bg_sub=True) - - reread = SrasFile(str(path)) - check("angle 1 is not cached", reread.precomputed_freq_mhz[1] is None) - check("angles 0 and 2 are cached", - reread.precomputed_freq_mhz[0] is not None - and reread.precomputed_freq_mhz[2] is not None) - img1 = compute_rf_image(reread, 1, dc_threshold_mv=None, apply_bg_sub=True) - check("uncached angle computes rather than returning zeros", - np.any(img1 != 0) and np.allclose(img1, expected[1], atol=1e-3)) - - -def test_parallel_identity(scratch: Path): - """Forcing 1 worker vs many must give identical output — catches - chunk-boundary and race bugs.""" - print("\nparallel vs serial identity") - path = scratch / "parallel.sras" - # Many rows, so the row loop actually splits into several chunks. - n_rows, n_frames, spf = 48, 9, 256 - gen.write(path, n_angles=1, seed=3, samples_per_frame=spf, - geometry=[(n_rows, n_frames)]) - sras = SrasFile(str(path)) - - saved_budget, saved_workers = compute._TOTAL_BYTES_BUDGET, compute._MAX_WORKERS - try: - # Shrink the budget so chunk_rows collapses to 1 and every row is - # its own chunk — the worst case for boundary bugs. - compute._TOTAL_BYTES_BUDGET = 8 * n_frames * spf * 4 - - compute._MAX_WORKERS = 1 - dc_serial = compute_dc_image(sras, 0, CH4_IDX) - rf_serial = compute_rf_image(sras, 0, dc_threshold_mv=None, apply_bg_sub=True) - dc4 = adc_to_mv(dc_serial, *sras.cal(CH4_IDX)) - thr = float(np.median(dc4)) - rf_masked_serial = compute_rf_image(sras, 0, dc_threshold_mv=thr, - apply_bg_sub=True) - - chunk_rows, n_workers = compute._plan_chunks( - n_rows, n_frames, spf, live_multiplier=compute._FFT_LIVE_MULTIPLIER) - check("serial plan uses 1 worker", n_workers == 1, f"chunk_rows={chunk_rows}") - check("work actually splits into multiple chunks", chunk_rows < n_rows, - f"chunk_rows={chunk_rows} of {n_rows} rows") - - compute._MAX_WORKERS = 8 - chunk_rows, n_workers = compute._plan_chunks( - n_rows, n_frames, spf, live_multiplier=compute._FFT_LIVE_MULTIPLIER) - check("parallel plan uses >1 worker", n_workers > 1, - f"chunk_rows={chunk_rows} workers={n_workers}") - - dc_par = compute_dc_image(sras, 0, CH4_IDX) - rf_par = compute_rf_image(sras, 0, dc_threshold_mv=None, apply_bg_sub=True) - rf_masked_par = compute_rf_image(sras, 0, dc_threshold_mv=thr, apply_bg_sub=True) - - check("dc image identical", np.array_equal(dc_serial, dc_par)) - check("rf image identical (unmasked)", np.array_equal(rf_serial, rf_par)) - check("rf image identical (masked)", - np.array_equal(rf_masked_serial, rf_masked_par)) - finally: - compute._TOTAL_BYTES_BUDGET, compute._MAX_WORKERS = saved_budget, saved_workers - - -def test_nomask_equals_low_threshold(scratch: Path): - """dc_threshold_mv=None must equal a threshold below every pixel, while - skipping the CH4 read.""" - print("\nno-mask path") - path = scratch / "nomask.sras" - gen.write(path, n_angles=2, seed=4, samples_per_frame=128) - sras = SrasFile(str(path)) - for a in range(sras.n_angles): - none_img = compute_rf_image(sras, a, dc_threshold_mv=None, apply_bg_sub=True) - low_img = compute_rf_image(sras, a, dc_threshold_mv=-1e9, apply_bg_sub=True) - check(f"angle {a}: None == -1e9 threshold", - np.array_equal(none_img, low_img)) - check(f"angle {a}: image is non-degenerate", - len(np.unique(none_img)) > 1, f"{len(np.unique(none_img))} unique") - - -def test_roi_mask(): - """The bbox-restricted mask must equal a full-grid point-in-polygon test.""" - print("\nROI mask (bbox fast path vs full grid)") - from matplotlib.path import Path as MplPath - from sras_viewer import RoiQuad - - rng = np.random.default_rng(0) - x = np.linspace(-2.0, 3.0, 137) - y = np.linspace(1.0, 4.0, 91) - - cases = { - "axis-aligned rect": np.array([[0.0, 1.5], [1.0, 1.5], [1.0, 3.0], [0.0, 3.0]]), - "skewed quad": np.array([[-0.5, 1.2], [1.7, 1.9], [1.2, 3.4], [-1.0, 2.6]]), - "entirely outside": np.array([[8.0, 8.0], [9.0, 8.0], [9.0, 9.0], [8.0, 9.0]]), - "covers whole grid": np.array([[-9.0, -9.0], [9.0, -9.0], [9.0, 9.0], [-9.0, 9.0]]), - "straddles left edge": np.array([[-4.0, 2.0], [0.5, 2.0], [0.5, 3.0], [-4.0, 3.0]]), - } - for _ in range(5): - cases[f"random {_}"] = rng.uniform([-2.5, 0.5], [3.5, 4.5], size=(4, 2)) - - for name, pts in cases.items(): - roi = RoiQuad(pts) - fast = roi.mask_for_grid(x, y) - X, Y = np.meshgrid(x.astype(np.float64), y.astype(np.float64)) - slow = MplPath(pts).contains_points( - np.column_stack([X.ravel(), Y.ravel()])).reshape(X.shape) - check(f"{name} ({int(slow.sum())} px inside)", np.array_equal(fast, slow)) - - # Descending y axis (images are stored top-down in some scans). - roi = RoiQuad(cases["skewed quad"]) - y_desc = y[::-1] - fast = roi.mask_for_grid(x, y_desc) - X, Y = np.meshgrid(x.astype(np.float64), y_desc.astype(np.float64)) - slow = MplPath(cases["skewed quad"]).contains_points( - np.column_stack([X.ravel(), Y.ravel()])).reshape(X.shape) - check("descending y axis", np.array_equal(fast, slow)) - - -def test_legacy_parse_and_average(scratch: Path): - """v2-v4 parsing plus the sras_average.py rewrite (which now streams via - SrasFile rather than slurping the whole file).""" - import subprocess - print("\nlegacy formats (v2-v4) and sras_average") - repo = Path(__file__).resolve().parent.parent - - for version in (2, 3, 4): - path = scratch / f"legacy_v{version}.sras" - meta = gen.write_legacy(path, version=version, n_angles=2, n_rows=4, - n_frames=12, samples_per_frame=32, seed=version) - s = SrasFile(str(path)) - check(f"v{version} parses", s.version == version, f"got v{s.version}") - check(f"v{version} geometry uniform across angles", - list(s.n_rows) == [4, 4] and list(s.n_frames) == [12, 12], - f"rows={list(s.n_rows)} frames={list(s.n_frames)}") - check(f"v{version} waveform data matches what was written", - all(np.array_equal(np.asarray(s.data[a]), meta["data"][a]) - for a in range(s.n_angles))) - check(f"v{version} background {'present' if version >= 4 else 'absent'}", - (s.background is not None) == (version >= 4)) - check(f"v{version} precomputed stores are ragged lists", - isinstance(s.precomputed_freq_mhz, list) - and len(s.precomputed_freq_mhz) == s.n_angles) - # DC image must equal a direct mean of the known input. - expect = meta["data"][0][:, CH3_IDX, :, :].astype(np.float64).mean(axis=-1) - check(f"v{version} DC image equals a direct mean", - np.allclose(compute_dc_image(s, 0, CH3_IDX), expect, atol=1e-3)) - - src = scratch / "legacy_v4.sras" - meta = gen.write_legacy(src, version=4, n_angles=2, n_rows=4, n_frames=12, - samples_per_frame=32, seed=4) - dst = scratch / "legacy_v4_avg.sras" - if dst.exists(): - dst.unlink() - proc = subprocess.run( - [sys.executable, str(repo / "sras_average.py"), str(src), str(dst), "--n", "4"], - capture_output=True, text=True, cwd=repo) - check("sras_average ran", proc.returncode == 0, - (proc.stderr or proc.stdout).strip()[-200:]) - - if dst.exists(): - avg = SrasFile(str(dst)) - check("averaged file parses", avg.version == 4) - check("frame count divided by 4", - list(avg.n_frames) == [3, 3], f"{list(avg.n_frames)}") - check("angles/rows/channels unchanged", - avg.n_angles == 2 and list(avg.n_rows) == [4, 4] - and avg.n_channels == meta["n_channels"]) - check("calibration preserved", - np.allclose(avg.ch_ymult_mv, SrasFile(str(src)).ch_ymult_mv)) - check("background preserved", - np.array_equal(avg.background, SrasFile(str(src)).background)) - src_data = meta["data"] - expect0 = src_data[0][:, :, 0:4, :].astype(np.float32).mean(axis=2).astype(np.int16) - check("first averaged group equals the mean of its 4 source frames", - np.array_equal(np.asarray(avg.data[0])[:, :, 0, :], expect0)) - - # Remainder handling: 12 frames / 5 -> 2 full groups + 1 partial. - dst2 = scratch / "legacy_v4_avg5.sras" - subprocess.run([sys.executable, str(repo / "sras_average.py"), - str(src), str(dst2), "--n", "5"], - capture_output=True, text=True, cwd=repo) - if dst2.exists(): - check("partial trailing group kept by default", - list(SrasFile(str(dst2)).n_frames) == [3, 3], - f"{list(SrasFile(str(dst2)).n_frames)}") - dst3 = scratch / "legacy_v4_avg5d.sras" - subprocess.run([sys.executable, str(repo / "sras_average.py"), - str(src), str(dst3), "--n", "5", "--discard-remainder"], - capture_output=True, text=True, cwd=repo) - if dst3.exists(): - check("--discard-remainder drops the partial group", - list(SrasFile(str(dst3)).n_frames) == [2, 2], - f"{list(SrasFile(str(dst3)).n_frames)}") - - -def test_unsupported_version_reported(scratch: Path): - """cache_file must report, not raise, for a file it can't handle.""" - print("\nerror reporting") - bogus = scratch / "bogus.sras" - bogus.write_bytes(b"SRAS" + bytes([99]) + b"\x00" * 200) - err = cache_file(str(bogus), "dc", True) - check("bad version returns an error string", bool(err), err) - missing = cache_file(str(scratch / "does_not_exist.sras"), "dc", True) - check("missing file returns an error string", bool(missing), missing) - - -def main(): - p = argparse.ArgumentParser(description=__doc__) - p.add_argument("--scratch") - args = p.parse_args() - - tmp = None - if args.scratch: - scratch = Path(args.scratch) - scratch.mkdir(parents=True, exist_ok=True) - else: - tmp = tempfile.mkdtemp(prefix="sras_test_") - scratch = Path(tmp) - - try: - test_cache_roundtrip(scratch) - test_partial_v7_cache(scratch) - test_parallel_identity(scratch) - test_nomask_equals_low_threshold(scratch) - test_roi_mask() - test_legacy_parse_and_average(scratch) - test_unsupported_version_reported(scratch) - finally: - if tmp: - shutil.rmtree(tmp, ignore_errors=True) - - print() - if _failures: - print(f"{len(_failures)} FAILURE(S): " + ", ".join(_failures)) - sys.exit(1) - print("All checks passed.") - - -if __name__ == "__main__": - main() From 11ff3b62e26506e6d0030167eb2ebf5280b5e74b Mon Sep 17 00:00:00 2001 From: Thomas Ales Date: Thu, 6 Aug 2026 10:20:14 -0500 Subject: [PATCH 04/10] Rewrite the FFT peak search: block-parallel zoom refinement MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit At pad 40 the old path materialised a ~9 GB padded spectrum per row, which collapsed the chunk planner to one worker and one rfft call with workers=1 — synthesis ran single-threaded, ~1 hour per angle on real files. The padded spectrum is never materialised now. Each block of 512 waveforms gets a coarse rfft at next_fast_len(2*spf); every coarse bin within 0.7 of its row's max (plus the DC-adjacent window, which coarse DC suppression would otherwise blind) is refined onto the exact n_fft grid by a small complex gemm. The selected bin is bit-identical to the full padded argmax — enforced by test_zoom_identity, a 25-seed fuzz test over adversarial spectra, and a clean golden-hash diff against the pre-rewrite baseline across pads {1,2,4,8,40}, masked/unmasked, bg on/off, int8/int16, and both backends. Blocks fan out over a persistent thread pool; pyFFTW runs through per-thread FFTW_MEASURE builder plans with wisdom persisted to ~/.cache/sras-viewer, and threadpoolctl clamps BLAS under the pool. compute_rf_image(exact=True) (or SRAS_FFT_EXACT=1) keeps the reference padded path for audits. tools/bench_fft.py measures: pad 40, 16 cores, 8192x2500 synthetic — exact serial 717 wf/s -> zoom pool 25100 wf/s (35x, pyFFTW backend; 19x scipy), every variant verified equal to the reference. Co-Authored-By: Claude Fable 5 --- pyproject.toml | 2 + sras_compute.py | 363 +++++++++++++++++++++++++++++++++++++----- tests/test_compute.py | 105 ++++++++++-- tools/bench_fft.py | 106 ++++++++++++ 4 files changed, 522 insertions(+), 54 deletions(-) create mode 100644 tools/bench_fft.py diff --git a/pyproject.toml b/pyproject.toml index ca708ca..ae73ea4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -16,6 +16,8 @@ dependencies = [ "scikit-image==0.26.0", # Faster rfft backend; the viewer falls back to scipy.fft without it. "pyFFTW==0.15.1", + # Clamps BLAS threading under the FFT worker pool. + "threadpoolctl==3.6.0", ] [project.optional-dependencies] diff --git a/sras_compute.py b/sras_compute.py index bf48c6b..5700b02 100644 --- a/sras_compute.py +++ b/sras_compute.py @@ -6,8 +6,10 @@ multiprocessing child can import it without loading Qt or matplotlib — which matters because Python 3.14 on macOS spawns rather than forks. """ +import atexit import json import os +import threading from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from pathlib import Path @@ -24,12 +26,16 @@ from sras_format import CH1_IDX, CH3_IDX, CH4_IDX, SrasFile, adc_to_mv try: import pyfftw - pyfftw.interfaces.cache.enable() PYFFTW_AVAILABLE = True except ImportError: PYFFTW_AVAILABLE = False -_fft_backend = "numpy" # "numpy" or "pyfftw"; set via set_fft_backend() +try: + from threadpoolctl import threadpool_limits +except ImportError: + threadpool_limits = None + +_fft_backend = "numpy" # "numpy" (scipy.fft) or "pyfftw"; set via set_fft_backend() def set_fft_backend(name: str): @@ -52,6 +58,218 @@ def _do_rfft(x: np.ndarray, n: int | None = None, axis: int = -1, return scipy_fft.rfft(x, n=n, axis=axis, workers=workers) +# --------------------------------------------------------------------------- +# FFT worker pool and per-thread pyFFTW plans +# --------------------------------------------------------------------------- + +_FFT_BLOCK = 512 # waveforms per FFT task. The knee on a 16-core machine: + # smaller blocks serialise on GIL-held numpy dispatch, + # larger ones lose cache residency and task granularity. +_ZOOM_MIN_PAD = 4 # zoom refinement engages at n_fft >= _ZOOM_MIN_PAD * spf +_FFT_EXACT_ENV = bool(int(os.environ.get("SRAS_FFT_EXACT", "0") or 0)) + +_pool_lock = threading.Lock() +_pool: ThreadPoolExecutor | None = None + + +def _fft_pool() -> ThreadPoolExecutor: + """The persistent process-wide pool for FFT block tasks.""" + global _pool + with _pool_lock: + if _pool is None: + _pool = ThreadPoolExecutor(max_workers=_MAX_WORKERS, + thread_name_prefix="sras-fft") + atexit.register(_pool.shutdown, wait=False, cancel_futures=True) + return _pool + + +_WISDOM_PATH = Path.home() / ".cache" / "sras-viewer" / "fftw_wisdom" +_wisdom_lock = threading.Lock() +_wisdom_loaded = False +_fftw_local = threading.local() + + +def _load_wisdom_once(): + """Import saved FFTW wisdom so FFTW_MEASURE planning is a one-time cost + per machine. Purely an optimisation: failures are ignored.""" + global _wisdom_loaded + with _wisdom_lock: + if _wisdom_loaded: + return + _wisdom_loaded = True + try: + pyfftw.import_wisdom(_WISDOM_PATH.read_bytes().split(b"\x00\n")) + except Exception: + pass + + +def _save_wisdom(): + with _wisdom_lock: + try: + _WISDOM_PATH.parent.mkdir(parents=True, exist_ok=True) + _WISDOM_PATH.write_bytes(b"\x00\n".join(pyfftw.export_wisdom())) + except Exception: + pass + + +def _fftw_block_rfft(waves: np.ndarray, n: int) -> np.ndarray: + """rfft of a (B, spf) float32 block via a cached per-thread FFTW plan. + + Plans have a fixed (_FFT_BLOCK, spf) input shape so each worker thread + plans once per transform length; a remainder block runs through the same + plan with its tail rows ignored. The returned array is the plan's output + buffer — consume it before the next call on the same thread. + """ + n_wf, spf = waves.shape + plans = getattr(_fftw_local, "plans", None) + if plans is None: + plans = _fftw_local.plans = {} + key = (_FFT_BLOCK, spf, n) + plan = plans.get(key) + if plan is None: + _load_wisdom_once() + buf = pyfftw.empty_aligned((_FFT_BLOCK, spf), dtype="float32") + # No overwrite_input: FFTW must not scribble on input_array, whose + # zero-padded tail (columns spf..n) is zeroed exactly once here. + plan = pyfftw.builders.rfft(buf, n=n, axis=-1, threads=1, + planner_effort="FFTW_MEASURE") + plan.input_array[:] = 0.0 + plans[key] = plan + _save_wisdom() + inp = plan.input_array + inp[:n_wf, :spf] = waves + return plan()[:n_wf] + + +def _block_rfft(waves: np.ndarray, n: int) -> np.ndarray: + """Single-threaded rfft of one block; outer parallelism comes from the + pool, so the transform itself must not spin up threads.""" + if _fft_backend == "pyfftw" and PYFFTW_AVAILABLE: + return _fftw_block_rfft(waves, n) + return scipy_fft.rfft(waves, n=n, axis=-1, workers=1) + + +# --------------------------------------------------------------------------- +# Zoom peak search: coarse rfft + local fine DFT around the winning bin +# --------------------------------------------------------------------------- + +_ZOOM_HALFWIDTH = 0.75 # refinement window half-width, in coarse spacings. + # Every fine bin lies within 0.5 spacings of its + # nearest coarse bin, and that bin is guaranteed to + # be a candidate (see _ZOOM_CAND_RATIO), so 0.5 + # suffices; 0.75 adds rounding margin. +_ZOOM_CAND_RATIO = 0.7 # refine every coarse bin within this power ratio of + # its row's coarse maximum. Quarter-natural-bin + # scalloping at the 2x-oversampled coarse grid can + # understate a peak by at most ~19% in power, so 0.7 + # keeps a wide margin — near-contenders are resolved + # on the fine grid, never ranked from coarse samples. + + +@dataclass +class _ZoomPlan: + """Constants for the coarse+refine peak search, built once per + compute_rf_image call and shared across worker threads. *phases* is a + lazily-filled window-start -> phase-vector cache; a benign duplicate + compute under concurrency is cheaper than locking.""" + n_fft: int + n_coarse: int + m: int # fine bins per refinement window + n_bins_fine: int + E: np.ndarray # (spf, m) complex64 fine-DFT matrix, relative bins + j: np.ndarray # (spf,) float64 sample indices + phases: dict + + +def _zoom_plan(spf: int, n_fft: int) -> _ZoomPlan: + # 2x-oversampled coarse grid: the padded power spectrum is a trig + # polynomial of degree spf-1, so at 2x sampling its global max cannot + # hide between coarse bins. + n_coarse = scipy_fft.next_fast_len(2 * spf, real=True) + n_bins_fine = n_fft // 2 + 1 + m = int(np.ceil(2 * _ZOOM_HALFWIDTH * n_fft / n_coarse)) + 1 + m = min(m, n_bins_fine - 1) + j = np.arange(spf, dtype=np.float64) + E = np.exp((-2j * np.pi / n_fft) * np.outer(j, np.arange(m))).astype(np.complex64) + return _ZoomPlan(n_fft, n_coarse, m, n_bins_fine, E, j, {}) + + +def _window_start(k_c: np.ndarray, zp: _ZoomPlan) -> np.ndarray: + """First fine bin of the refinement window around each coarse bin. + Clipped to [1, ...] so fine bin 0 stays excluded (DC suppression).""" + k0 = np.floor((k_c - _ZOOM_HALFWIDTH) * (zp.n_fft / zp.n_coarse)).astype(np.int64) + return np.clip(k0, 1, max(1, zp.n_bins_fine - zp.m)) + + +def _refine_window(waves: np.ndarray, rows: np.ndarray, k0: int, zp: _ZoomPlan, + best_pow: np.ndarray, best_bin: np.ndarray): + """Evaluate fine bins k0..k0+m-1 for the given rows and fold the result + into the per-row best (power, bin), preserving np.argmax's lowest-bin + tie-break.""" + ph = zp.phases.get(k0) + if ph is None: + ph = np.exp((-2j * np.pi * k0 / zp.n_fft) * zp.j).astype(np.complex64) + zp.phases[k0] = ph + F = (waves[rows] * ph) @ zp.E + q = F.real ** 2 + q += F.imag ** 2 + i = np.argmax(q, axis=1) + p = q[np.arange(len(rows)), i] + b = k0 + i + upd = (p > best_pow[rows]) | ((p == best_pow[rows]) & (b < best_bin[rows])) + ridx = rows[upd] + best_pow[ridx] = p[upd] + best_bin[ridx] = b[upd] + + +def _peak_bins_zoom(waves: np.ndarray, zp: _ZoomPlan) -> np.ndarray: + """Peak fine-bin per waveform without materialising the padded spectrum: + coarse rfft, then a small fine DFT (one gemm per shared window) on the + exact n_fft bin grid. Identity with the full padded argmax is enforced + by tests/test_compute.py::test_zoom_identity and the golden-hash sweep.""" + n_wf = waves.shape[0] + S = _block_rfft(waves, zp.n_coarse) + P = S.real ** 2 + P += S.imag ** 2 + P[:, 0] = 0.0 + p_c = np.max(P, axis=1) + + best_pow = np.full(n_wf, -1.0, dtype=np.float32) + best_bin = np.full(n_wf, np.iinfo(np.int64).max, dtype=np.int64) + + # For a clean signal this yields one or two windows; a noise spectrum + # (many near-equal peaks) yields a dozen or so — still a tiny fraction + # of the padded grid. + thr = np.where(p_c > 0, np.float32(_ZOOM_CAND_RATIO) * p_c, + np.float32(np.inf)) + rows_c, bins_c = np.nonzero(P >= thr[:, None]) + k0_c = _window_start(bins_c, zp) + # Zeroing the coarse DC bin (suppression) blinds the candidate scan to + # fine bins closer to DC than the first coarse sample — where the + # DC-leakage skirt of an un-subtracted offset peaks. Always refine the + # DC-adjacent window too. + rows_c = np.concatenate([rows_c, np.arange(n_wf)]) + k0_c = np.concatenate([k0_c, np.ones(n_wf, dtype=np.int64)]) + pair = np.unique(np.stack([rows_c, k0_c], axis=1), axis=0) + for k0 in np.unique(pair[:, 1]): + _refine_window(waves, pair[pair[:, 1] == k0, 0], int(k0), zp, + best_pow, best_bin) + + # An all-zero spectrum must reproduce argmax-of-zeros = bin 0. + best_bin[p_c == 0.0] = 0 + return best_bin + + +def _peak_bins_direct(waves: np.ndarray, n_len: int) -> np.ndarray: + """Peak bin per waveform via the full transform — for pad factors too + small for the zoom search to pay.""" + S = _block_rfft(waves, n_len) + power = S.real ** 2 + power += S.imag ** 2 + power[:, 0] = 0.0 + return np.argmax(power, axis=1) + + # --------------------------------------------------------------------------- # Chunking / parallel budget # --------------------------------------------------------------------------- @@ -76,17 +294,21 @@ _TOTAL_BYTES_BUDGET = int(os.environ.get("SRAS_MEM_BUDGET_MB", 1024)) * 1024 * 1 _CHUNK_ROWS_MAX = 32 # cap for small scans (original behavior) _MAX_WORKERS = int(os.environ.get("SRAS_MAX_WORKERS", 0)) or (os.cpu_count() or 4) -# An rfft chunk holds, live at once: the float32 input, the complex64 -# transform, and the float32 power spectrum — roughly 3x the input buffer. -_FFT_LIVE_MULTIPLIER = 3 - - def _chunk_rows_for(n_frames: int, samples_per_frame: int, budget: int = _TOTAL_BYTES_BUDGET) -> int: bytes_per_row = max(1, n_frames * samples_per_frame * 4) # float32 return int(max(1, min(_CHUNK_ROWS_MAX, budget // bytes_per_row))) +def _plan_fft_rows(n_frames: int, samples_per_frame: int, budget: int) -> int: + """Rows per outer chunk for the block-FFT path. Only the float32 + waveform buffer scales with the chunk (in-flight block spectra total a + few MB across the whole pool), so budget it with 2x slack and let the + block fan-out saturate the pool regardless of pad factor.""" + bytes_per_row = max(1, n_frames * samples_per_frame * 4 * 2) + return int(max(1, min(_CHUNK_ROWS_MAX, budget // bytes_per_row))) + + def _plan_chunks(n_rows: int, n_frames: int, samples_per_frame: int, live_multiplier: int = 1, max_workers: int | None = None, @@ -215,7 +437,8 @@ def compute_rf_image(sras: SrasFile, angle_idx: int, dc4_mv: np.ndarray | None = None, max_workers: int | None = None, budget: int | None = None, - should_stop=None) -> np.ndarray: + should_stop=None, + exact: bool = False) -> np.ndarray: """FFT of each CH1 waveform; pixel = peak frequency in MHz. Pixels where CH4_dc < dc_threshold_mv are set to 0 — and the FFT is @@ -233,6 +456,13 @@ def compute_rf_image(sras: SrasFile, angle_idx: int, DC-channel precompute cache), pass it as *dc4_mv* (mV, shape (n_rows, n_frames)) to reuse it instead of re-reading CH4 here. + At pad factors >= _ZOOM_MIN_PAD the padded spectrum is never + materialised: a coarse rfft finds each peak and a local fine DFT + resolves it on the exact n_fft bin grid (_peak_bins_zoom). *exact* + forces the reference full-padded transform instead — it exists for + tests, audits, and the SRAS_FFT_EXACT=1 escape hatch, and is slow and + memory-hungry at high pad. + Fast path: if the file has a precomputed peak-frequency image for this angle (v5 PREC or v7 CACH), zero-padding is off, and the bg-sub flag matches, the stored image is used directly — no FFT is run. @@ -257,17 +487,31 @@ def compute_rf_image(sras: SrasFile, angle_idx: int, return freq_img # ---- Chunked FFT path -------------------------------------------------- - freq_axis = sras.freq_axis_mhz(n_fft) + exact = exact or _FFT_EXACT_ENV + spf = sras.samples_per_frame + n_len = n_fft if n_fft is not None else spf + freq32 = sras.freq_axis_mhz(n_fft).astype(np.float32) img = np.zeros((n_rows, n_frames), dtype=np.float32) - n_fft_bins = n_fft if n_fft is not None else sras.samples_per_frame - chunk_rows, n_workers = _plan_chunks( - n_rows, n_frames, max(sras.samples_per_frame, n_fft_bins), - live_multiplier=_FFT_LIVE_MULTIPLIER, max_workers=max_workers, - budget=budget) background = sras.background if (apply_bg_sub and sras.background is not None) else None cal4 = sras.cal(CH4_IDX) - def chunk(r0: int, r1: int): + zp = (_zoom_plan(spf, n_fft) + if not exact and n_fft is not None and n_fft >= _ZOOM_MIN_PAD * spf + else None) + total = _TOTAL_BYTES_BUDGET if budget is None else max(1, budget) + cap = max_workers if max_workers is not None else _MAX_WORKERS + if exact: + # The reference path materialises the full padded spectrum, so rows + # are budgeted against it (complex64 + power + temp per bin) and the + # chunk runs as one serial transform. + bytes_per_row = max(1, n_frames * (4 * spf + 16 * (n_len // 2 + 1))) + chunk_rows = int(max(1, min(_CHUNK_ROWS_MAX, total // bytes_per_row))) + cap = 1 + else: + chunk_rows = _plan_fft_rows(n_frames, spf, total) + pool = _fft_pool() if cap > 1 else None + + def process(r0: int, r1: int): if dc_threshold_mv is None: valid = None else: @@ -280,36 +524,73 @@ def compute_rf_image(sras: SrasFile, angle_idx: int, if not valid.any(): return - # Index the raw memmap slice with the boolean mask *before* - # converting dtype — this is a lazy view until touched, so only the - # selected elements are actually read from disk; masked-out pixels' - # pages are never paged in at all. - raw = data[r0:r1, CH1_IDX, :, :] - waves = (raw[valid] if valid is not None else raw).astype(np.float32) + counts = (valid.sum(axis=1) if valid is not None + else np.full(r1 - r0, n_frames, dtype=np.int64)) + offs = np.concatenate(([0], np.cumsum(counts))) + n_wf = int(offs[-1]) + waves = np.empty((n_wf, spf), dtype=np.float32) - if background is not None: - waves -= background # background is 1-D (spf,) + def read_row(i: int): + if should_stop is not None and should_stop(): + return + # Index the raw memmap slice with the boolean mask *before* + # converting dtype — this is a lazy view until touched, so only + # the selected elements are actually read from disk; masked-out + # pixels' pages are never paged in at all. + raw = data[r0 + i, CH1_IDX] + dst = waves[offs[i]:offs[i + 1]] + dst[:] = raw[valid[i]] if valid is not None else raw + if background is not None: + dst -= background # background is 1-D (spf,) - # Parallelism comes from the outer chunk loop, so keep the inner - # transform single-threaded to avoid oversubscribing the machine. - spectrum = _do_rfft(waves, n=n_fft, axis=-1, workers=1) - del waves - # |z|^2 without np.abs()'s extra full-size temporary. - power = spectrum.real ** 2 - power += spectrum.imag ** 2 - del spectrum - # [..., 0] not [:, 0]: the unmasked path keeps the (rows, frames, - # bins) shape, where [:, 0] would blank a whole frame. - power[..., 0] = 0.0 # suppress DC bin - peak_bins = np.argmax(power, axis=-1) - del power - - if valid is not None: - img[r0:r1][valid] = freq_axis[peak_bins] + if pool is None: + for i in range(r1 - r0): + read_row(i) else: - img[r0:r1] = freq_axis[peak_bins] + list(pool.map(read_row, range(r1 - r0))) - _map_row_chunks(n_rows, chunk_rows, n_workers, chunk, should_stop=should_stop) + out = np.empty(n_wf, dtype=np.float32) + + def fft_block(b0: int): + if should_stop is not None and should_stop(): + return + b1 = min(b0 + _FFT_BLOCK, n_wf) + w = waves[b0:b1] + bins = (_peak_bins_zoom(w, zp) if zp is not None + else _peak_bins_direct(w, n_len)) + out[b0:b1] = freq32[bins] + + if exact: + spectrum = _do_rfft(waves, n=n_fft, axis=-1, workers=1) + power = spectrum.real ** 2 + power += spectrum.imag ** 2 + power[:, 0] = 0.0 # suppress DC bin + out[:] = freq32[np.argmax(power, axis=1)] + elif pool is None: + for b0 in range(0, n_wf, _FFT_BLOCK): + fft_block(b0) + else: + list(pool.map(fft_block, range(0, n_wf, _FFT_BLOCK))) + + # Boolean scatter is row-major, matching the read pass's + # concatenation order. + if valid is not None: + img[r0:r1][valid] = out + else: + img[r0:r1] = out.reshape(r1 - r0, n_frames) + + # BLAS must not thread under the pool (the fine-DFT gemm would multiply + # against the pool's own workers). + limiter = (threadpool_limits(limits=1) + if pool is not None and threadpool_limits is not None else None) + try: + for r0 in range(0, n_rows, chunk_rows): + if should_stop is not None and should_stop(): + break + process(r0, min(r0 + chunk_rows, n_rows)) + finally: + if limiter is not None: + limiter.unregister() return img diff --git a/tests/test_compute.py b/tests/test_compute.py index e08af39..408d79f 100644 --- a/tests/test_compute.py +++ b/tests/test_compute.py @@ -10,6 +10,7 @@ import sys from pathlib import Path import numpy as np +import pytest import sras_compute as compute from sras_compute import ( @@ -117,9 +118,17 @@ def test_parallel_identity(tmp_path, monkeypatch): geometry=[(n_rows, n_frames)]) sras = SrasFile(str(path)) - # Shrink the budget so chunk_rows collapses to 1 and every row is - # its own chunk — the worst case for boundary bugs. + # Shrink the budget so the outer row loop splits into many chunks, and + # the block size so every chunk splits into many FFT tasks — the worst + # case for boundary bugs. monkeypatch.setattr(compute, "_TOTAL_BYTES_BUDGET", 8 * n_frames * spf * 4) + monkeypatch.setattr(compute, "_FFT_BLOCK", 4) + fft_rows = compute._plan_fft_rows(n_frames, spf, compute._TOTAL_BYTES_BUDGET) + assert fft_rows < n_rows, \ + f"FFT work actually splits into multiple chunks ({fft_rows} of {n_rows})" + dc_rows = compute._chunk_rows_for(n_frames, spf, compute._TOTAL_BYTES_BUDGET) + assert dc_rows < n_rows, \ + f"DC work actually splits into multiple chunks ({dc_rows} of {n_rows})" monkeypatch.setattr(compute, "_MAX_WORKERS", 1) dc_serial = compute_dc_image(sras, 0, CH4_IDX) @@ -128,27 +137,97 @@ def test_parallel_identity(tmp_path, monkeypatch): thr = float(np.median(dc4)) rf_masked_serial = compute_rf_image(sras, 0, dc_threshold_mv=thr, apply_bg_sub=True) - - chunk_rows, n_workers = compute._plan_chunks( - n_rows, n_frames, spf, live_multiplier=compute._FFT_LIVE_MULTIPLIER) - assert n_workers == 1, f"serial plan uses 1 worker (chunk_rows={chunk_rows})" - assert chunk_rows < n_rows, \ - f"work actually splits into multiple chunks ({chunk_rows} of {n_rows} rows)" + rf_pad_serial = compute_rf_image(sras, 0, dc_threshold_mv=thr, + apply_bg_sub=True, n_fft=spf * 8) monkeypatch.setattr(compute, "_MAX_WORKERS", 8) - chunk_rows, n_workers = compute._plan_chunks( - n_rows, n_frames, spf, live_multiplier=compute._FFT_LIVE_MULTIPLIER) - assert n_workers > 1, \ - f"parallel plan uses >1 worker (chunk_rows={chunk_rows} workers={n_workers})" - dc_par = compute_dc_image(sras, 0, CH4_IDX) rf_par = compute_rf_image(sras, 0, dc_threshold_mv=None, apply_bg_sub=True) rf_masked_par = compute_rf_image(sras, 0, dc_threshold_mv=thr, apply_bg_sub=True) + rf_pad_par = compute_rf_image(sras, 0, dc_threshold_mv=thr, + apply_bg_sub=True, n_fft=spf * 8) assert np.array_equal(dc_serial, dc_par), "dc image identical" assert np.array_equal(rf_serial, rf_par), "rf image identical (unmasked)" assert np.array_equal(rf_masked_serial, rf_masked_par), \ "rf image identical (masked)" + assert np.array_equal(rf_pad_serial, rf_pad_par), \ + "rf image identical (masked, padded/zoom)" + + +@pytest.mark.parametrize("spf,bps", [(64, 2), (37, 1)]) +def test_zoom_identity(tmp_path, monkeypatch, spf, bps): + """The zoom peak search must reproduce the full padded-rfft argmax + bit-for-bit, across pad factors, masking, bg-sub, dtype, and backend.""" + path = tmp_path / f"zoom_{spf}.sras" + gen.write(path, n_angles=2, seed=6, samples_per_frame=spf, bps=bps) + sras = SrasFile(str(path)) + dc4 = dc_image_mv(sras, 0, CH4_IDX) + thr = float(np.median(dc4)) + + backends = ["numpy"] + (["pyfftw"] if compute.PYFFTW_AVAILABLE else []) + for backend in backends: + monkeypatch.setattr(compute, "_fft_backend", backend) + for pad in (4, 8, 40): + n_fft = spf * pad + for thr_v in (None, thr): + for bg in (False, True): + ref = compute_rf_image(sras, 0, dc_threshold_mv=thr_v, + apply_bg_sub=bg, n_fft=n_fft, + exact=True) + zoom = compute_rf_image(sras, 0, dc_threshold_mv=thr_v, + apply_bg_sub=bg, n_fft=n_fft) + diff = int((ref != zoom).sum()) + assert diff == 0, \ + (f"{diff} px differ: backend={backend} pad={pad} " + f"thr={thr_v} bg={bg} spf={spf}") + + # A threshold above every pixel masks everything: both paths must agree + # on an all-zero image. + all_masked = compute_rf_image(sras, 0, dc_threshold_mv=1e9, n_fft=spf * 8) + assert not all_masked.any() + + +def test_zoom_identity_fuzz(): + """Hammer _peak_bins_zoom directly with adversarial spectra: noise, + un-subtracted DC offsets, on-bin and off-bin tones, near-tie tone pairs, + and all-zero rows.""" + import scipy.fft as scipy_fft + + rng = np.random.default_rng(42) + for _ in range(25): + spf = int(rng.integers(16, 220)) + pad = int(rng.choice([4, 5, 8, 16, 40])) + n_fft = spf * pad + n_wf = 24 + w = rng.normal(scale=20.0, size=(n_wf, spf)) + t = np.arange(spf) + # rows 0-5: pure/noisy tones (some off-bin), row 6-7: near-tie pair, + # row 8: big DC offset, row 9: all zeros, rest: plain noise. + for r in range(6): + f = rng.uniform(1.0, spf / 2 - 1) + w[r] = 60 * np.sin(2 * np.pi * f * t / spf) + w[r] * (r % 2) + f1, f2 = rng.uniform(2.0, spf / 2 - 2, size=2) + w[6] = 50 * np.sin(2 * np.pi * f1 * t / spf) \ + + 49.9 * np.sin(2 * np.pi * f2 * t / spf) + w[7] = 50 * np.sin(2 * np.pi * f1 * t / spf) \ + + 50 * np.cos(2 * np.pi * f2 * t / spf) + w[8] = 90 + rng.normal(scale=5.0, size=spf) + w[9] = 0.0 + w = w.astype(np.float32) + + S = scipy_fft.rfft(w, n=n_fft, axis=-1, workers=1) + P = S.real ** 2 + P += S.imag ** 2 + P[:, 0] = 0.0 + ref = np.argmax(P, axis=1) + + zp = compute._zoom_plan(spf, n_fft) + got = compute._peak_bins_zoom(w, zp) + bad = np.nonzero(ref != got)[0] + assert not len(bad), \ + (f"spf={spf} pad={pad}: rows {bad.tolist()} picked " + f"{got[bad].tolist()} instead of {ref[bad].tolist()}") def test_nomask_equals_low_threshold(tmp_path): diff --git a/tools/bench_fft.py b/tools/bench_fft.py new file mode 100644 index 0000000..d71285f --- /dev/null +++ b/tools/bench_fft.py @@ -0,0 +1,106 @@ +#!/usr/bin/env python3 +"""Benchmark the FFT peak-search path: exact vs zoom, serial vs pooled. + +Reports wall time, waveforms/s, CPU utilization (utime+stime over wall, in +cores), and verifies every variant against the exact reference image. + +Usage: + python tools/bench_fft.py # synthetic, pads 1/8/40 + python tools/bench_fft.py --pads 40 --spf 2500 --rows 8 --frames 1024 + python tools/bench_fft.py --real /path/big.sras --real-rows 32 --pads 40 +""" + +import argparse +import resource +import sys +import tempfile +import time +from pathlib import Path + +import numpy as np + +sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) + +import sras_compute as compute # noqa: E402 +from sras_compute import compute_rf_image, set_fft_backend # noqa: E402 +from sras_format import SrasFile # noqa: E402 +import tools.make_test_sras as gen # noqa: E402 +from tools.check_equivalence import row_slice # noqa: E402 + + +def _timed(fn): + r0 = resource.getrusage(resource.RUSAGE_SELF) + t0 = time.perf_counter() + out = fn() + wall = time.perf_counter() - t0 + r1 = resource.getrusage(resource.RUSAGE_SELF) + cpu = (r1.ru_utime - r0.ru_utime) + (r1.ru_stime - r0.ru_stime) + return out, wall, cpu / max(wall, 1e-9) + + +def bench(sras, pads, backends): + n_wf = sum(int(sras.n_rows[a]) * int(sras.n_frames[a]) + for a in range(sras.n_angles)) + spf = sras.samples_per_frame + print(f"{n_wf} waveforms x {spf} samples, {sras.n_angles} angle(s)") + print(f"{'pad':>4} {'backend':>8} {'variant':>16} {'wall':>9} " + f"{'wf/s':>10} {'util':>6} match") + + for pad in pads: + n_fft = spf * pad if pad > 1 else None + for backend in backends: + set_fft_backend(backend) + + def run(**kw): + imgs = [compute_rf_image(sras, a, dc_threshold_mv=None, + apply_bg_sub=True, n_fft=n_fft, **kw) + for a in range(sras.n_angles)] + return np.concatenate([i.ravel() for i in imgs]) + + ref, wall, util = _timed(lambda: run(exact=True)) + rows = [("exact(serial)", ref, wall, util, True)] + for label, kw in (("zoom(serial)", dict(max_workers=1)), + ("zoom(pool)", {})): + img, wall, util = _timed(lambda: run(**kw)) + rows.append((label, img, wall, util, bool(np.array_equal(img, ref)))) + for label, img, wall, util, ok in rows: + print(f"{pad:>4} {backend:>8} {label:>16} {wall:>8.2f}s " + f"{n_wf / wall:>10.0f} {util:>5.1f}x " + f"{'OK' if ok else 'MISMATCH'}") + + +def main(): + p = argparse.ArgumentParser(description=__doc__) + p.add_argument("--pads", default="1,8,40", + help="comma-separated pad factors (default 1,8,40)") + p.add_argument("--spf", type=int, default=2500) + p.add_argument("--rows", type=int, default=8) + p.add_argument("--frames", type=int, default=1024) + p.add_argument("--backends", default=None, + help="comma-separated (default: numpy,pyfftw if available)") + p.add_argument("--real", help="path to a real .sras file") + p.add_argument("--real-rows", type=int, default=32, + help="rows of angle 0 to use from the real file") + args = p.parse_args() + + pads = [int(x) for x in args.pads.split(",")] + if args.backends: + backends = args.backends.split(",") + else: + backends = ["numpy"] + (["pyfftw"] if compute.PYFFTW_AVAILABLE else []) + + if args.real: + sras = row_slice(SrasFile(args.real), 0, args.real_rows) + sras.data = [sras.data[0]] + sras.n_angles = 1 + bench(sras, pads, backends) + else: + with tempfile.TemporaryDirectory(prefix="sras_bench_") as tmp: + path = Path(tmp) / "bench.sras" + gen.write(path, n_angles=1, seed=0, samples_per_frame=args.spf, + geometry=[(args.rows, args.frames)]) + bench(SrasFile(str(path)), pads, backends) + + +if __name__ == "__main__": + main() From 9a9a2557d6d47b0823592ea74bfcd220c15e0e61 Mon Sep 17 00:00:00 2001 From: Thomas Ales Date: Thu, 6 Aug 2026 10:26:50 -0500 Subject: [PATCH 05/10] Persist FFT settings, fix backend naming, parallelise batch DC caching MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - FFT backend and pad factor persist across sessions via QSettings (IniFormat; tests redirect the settings path for hermeticity). - The default backend was labelled "NumPy FFT" but always dispatched to scipy.fft — rename the canonical value to "scipy" ("numpy" stays as a legacy alias) and fix the dialog label. - cache_file: DC caching fans out over angles via _parallel_map with per-angle budgets (the DcPrecomputeWorker pattern); FFT caching stays serial per angle because compute_rf_image now parallelises internally over blocks. Documented that the v7 FFT cache is natural-resolution (pad 1) by design. Co-Authored-By: Claude Fable 5 --- sras_compute.py | 33 ++++++++++++++++++++++++--------- sras_viewer.py | 31 ++++++++++++++++++++++--------- tests/conftest.py | 10 +++++++++- tests/test_compute.py | 2 +- tools/bench_fft.py | 4 ++-- 5 files changed, 58 insertions(+), 22 deletions(-) diff --git a/sras_compute.py b/sras_compute.py index 5700b02..fc7b30e 100644 --- a/sras_compute.py +++ b/sras_compute.py @@ -35,15 +35,17 @@ try: except ImportError: threadpool_limits = None -_fft_backend = "numpy" # "numpy" (scipy.fft) or "pyfftw"; set via set_fft_backend() +_fft_backend = "scipy" # "scipy" or "pyfftw"; set via set_fft_backend() def set_fft_backend(name: str): """Select the rfft implementation. Module-level state, so it must be set explicitly inside each multiprocessing child — it does not survive a - spawn.""" + spawn. "numpy" is accepted as a legacy alias for "scipy".""" global _fft_backend - _fft_backend = name if (name != "pyfftw" or PYFFTW_AVAILABLE) else "numpy" + if name == "numpy": + name = "scipy" + _fft_backend = name if (name == "pyfftw" and PYFFTW_AVAILABLE) else "scipy" def get_fft_backend() -> str: @@ -599,13 +601,17 @@ def compute_rf_image(sras: SrasFile, angle_idx: int, # --------------------------------------------------------------------------- def cache_file(path: str, mode: str, apply_bg_sub: bool, - fft_backend: str = "numpy", max_workers: int = 0) -> str: + fft_backend: str = "scipy", max_workers: int = 0) -> str: """Compute and store DC or FFT images for every angle of one file, converting v6 → v7 in place. Returns "" on success or an error message. Module-level and picklable so it can run in a ProcessPoolExecutor. The FFT backend and worker cap are passed explicitly because module globals do not survive a spawn. + + The stored FFT cache is always natural-resolution (pad 1): the v7 SFFT + block records no pad factor, and padded views compute live fast enough + (see _peak_bins_zoom) that caching them is not worth a format change. """ global _MAX_WORKERS try: @@ -621,17 +627,26 @@ def cache_file(path: str, mode: str, apply_bg_sub: bool, "can be batch-cached") n = sras.n_angles + n_workers, angle_budget = plan_angle_level(sras) if mode == "dc": - dc3 = [adc_to_mv(compute_dc_image(sras, a, CH3_IDX), *sras.cal(CH3_IDX)) - for a in range(n)] - dc4 = [adc_to_mv(compute_dc_image(sras, a, CH4_IDX), *sras.cal(CH4_IDX)) - for a in range(n)] + dc3 = _parallel_map( + lambda a: adc_to_mv( + compute_dc_image(sras, a, CH3_IDX, max_workers=1, + budget=angle_budget), *sras.cal(CH3_IDX)), + range(n), n_workers) + dc4 = _parallel_map( + lambda a: adc_to_mv( + compute_dc_image(sras, a, CH4_IDX, max_workers=1, + budget=angle_budget), *sras.cal(CH4_IDX)), + range(n), n_workers) sras.write_v7_cache(new_dc3_mv=dc3, new_dc4_mv=dc4) else: effective_bg = apply_bg_sub and sras.background is not None # dc_threshold_mv=None: store unmasked images and mask at display # time (same convention as v5's PREC block). Skipping the mask - # also skips reading CH4 entirely. + # also skips reading CH4 entirely. The FFT path parallelises + # internally over blocks, so angles run one at a time with the + # full budget. freq = [compute_rf_image(sras, a, dc_threshold_mv=None, apply_bg_sub=effective_bg) for a in range(n)] diff --git a/sras_viewer.py b/sras_viewer.py index 54272c0..e28fa73 100644 --- a/sras_viewer.py +++ b/sras_viewer.py @@ -25,7 +25,7 @@ from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg, NavigationToolb from matplotlib.figure import Figure from matplotlib.patches import Polygon from matplotlib.path import Path as MplPath -from PyQt6.QtCore import QObject, Qt, QThread, pyqtSignal +from PyQt6.QtCore import QObject, QSettings, Qt, QThread, pyqtSignal from PyQt6.QtGui import QAction, QKeyEvent from PyQt6.QtWidgets import ( QApplication, QButtonGroup, QCheckBox, QComboBox, QDialog, QDialogButtonBox, @@ -655,22 +655,22 @@ class FftOptionsDialog(QDialog): grp_backend = QGroupBox("FFT Backend") bl = QVBoxLayout(grp_backend) - self._btn_numpy = QRadioButton("NumPy FFT (always available)") + self._btn_scipy = QRadioButton("SciPy FFT (pocketfft) (always available)") self._btn_pyfftw = QRadioButton( "pyFFTW (faster for large arrays)" if PYFFTW_AVAILABLE else "pyFFTW (not installed — run: pip install pyfftw)") self._btn_pyfftw.setEnabled(PYFFTW_AVAILABLE) self._backend_group = QButtonGroup(self) - self._backend_group.addButton(self._btn_numpy, id=0) + self._backend_group.addButton(self._btn_scipy, id=0) self._backend_group.addButton(self._btn_pyfftw, id=1) if current_backend == "pyfftw" and PYFFTW_AVAILABLE: self._btn_pyfftw.setChecked(True) else: - self._btn_numpy.setChecked(True) + self._btn_scipy.setChecked(True) - bl.addWidget(self._btn_numpy) + bl.addWidget(self._btn_scipy) bl.addWidget(self._btn_pyfftw) layout.addWidget(grp_backend) @@ -737,7 +737,7 @@ class FftOptionsDialog(QDialog): f"(at grating = {self._grating_um:.2f} µm)") def get_backend(self) -> str: - return "pyfftw" if self._btn_pyfftw.isChecked() and PYFFTW_AVAILABLE else "numpy" + return "pyfftw" if self._btn_pyfftw.isChecked() and PYFFTW_AVAILABLE else "scipy" def get_pad_factor(self) -> int: return max(1, self._spin_pad.value()) @@ -1451,8 +1451,18 @@ class SrasViewerWindow(QMainWindow): self._jobs: dict[str, tuple] = {} self._progress_dlgs: dict[str, QProgressDialog] = {} - # FFT settings (configured via FFT Options dialog) - self._fft_pad_factor: int = 1 # 1 = no padding + # FFT settings (configured via FFT Options dialog, persisted across + # sessions). IniFormat: predictable cross-platform and redirectable + # in tests. + self._settings = QSettings(QSettings.Format.IniFormat, + QSettings.Scope.UserScope, + "sras-viewer", "sras-viewer") + compute.set_fft_backend(str(self._settings.value("fft/backend", "scipy"))) + try: + pad = int(self._settings.value("fft/pad_factor", 1)) + except (TypeError, ValueError): + pad = 1 + self._fft_pad_factor: int = max(1, min(256, pad)) # 1 = no padding # Convert menu: batch DC/FFT compute-and-store (v6 -> v7) self._batch_errors: list[str] = [] @@ -1846,7 +1856,8 @@ class SrasViewerWindow(QMainWindow): self._batch_fft_act = QAction("Batch Compute FFT and Sto&re…", self) self._batch_fft_act.setStatusTip( "Select .sras files and compute+store FFT peak-frequency images " - "for every angle, converting v6 files to v7 in place.") + "for every angle, converting v6 files to v7 in place. Stored " + "images are natural-resolution (pad 1); padded views compute live.") self._batch_fft_act.triggered.connect(lambda: self._on_batch_compute("fft")) convert_menu.addAction(self._batch_fft_act) @@ -2748,6 +2759,8 @@ class SrasViewerWindow(QMainWindow): return compute.set_fft_backend(dlg.get_backend()) self._fft_pad_factor = dlg.get_pad_factor() + self._settings.setValue("fft/backend", compute.get_fft_backend()) + self._settings.setValue("fft/pad_factor", self._fft_pad_factor) # Pad factor changes the FFT bin count, so it genuinely invalidates # the cached raw FFT (part of the cache key) — _refresh_display() # recomputes only on a cache miss. diff --git a/tests/conftest.py b/tests/conftest.py index c9df201..5486080 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,8 +1,16 @@ -"""Shared test setup: repo-root imports and the offscreen Qt platform.""" +"""Shared test setup: repo-root imports, the offscreen Qt platform, and +hermetic QSettings (tests must not read or write the user's real viewer +settings).""" import os import sys +import tempfile from pathlib import Path os.environ.setdefault("QT_QPA_PLATFORM", "offscreen") sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) + +from PyQt6.QtCore import QSettings # noqa: E402 + +QSettings.setPath(QSettings.Format.IniFormat, QSettings.Scope.UserScope, + tempfile.mkdtemp(prefix="sras_qsettings_")) diff --git a/tests/test_compute.py b/tests/test_compute.py index 408d79f..3c6e400 100644 --- a/tests/test_compute.py +++ b/tests/test_compute.py @@ -165,7 +165,7 @@ def test_zoom_identity(tmp_path, monkeypatch, spf, bps): dc4 = dc_image_mv(sras, 0, CH4_IDX) thr = float(np.median(dc4)) - backends = ["numpy"] + (["pyfftw"] if compute.PYFFTW_AVAILABLE else []) + backends = ["scipy"] + (["pyfftw"] if compute.PYFFTW_AVAILABLE else []) for backend in backends: monkeypatch.setattr(compute, "_fft_backend", backend) for pad in (4, 8, 40): diff --git a/tools/bench_fft.py b/tools/bench_fft.py index d71285f..68a1cc3 100644 --- a/tools/bench_fft.py +++ b/tools/bench_fft.py @@ -77,7 +77,7 @@ def main(): p.add_argument("--rows", type=int, default=8) p.add_argument("--frames", type=int, default=1024) p.add_argument("--backends", default=None, - help="comma-separated (default: numpy,pyfftw if available)") + help="comma-separated (default: scipy,pyfftw if available)") p.add_argument("--real", help="path to a real .sras file") p.add_argument("--real-rows", type=int, default=32, help="rows of angle 0 to use from the real file") @@ -87,7 +87,7 @@ def main(): if args.backends: backends = args.backends.split(",") else: - backends = ["numpy"] + (["pyfftw"] if compute.PYFFTW_AVAILABLE else []) + backends = ["scipy"] + (["pyfftw"] if compute.PYFFTW_AVAILABLE else []) if args.real: sras = row_slice(SrasFile(args.real), 0, args.real_rows) From e80717d5c570949fb9e96d68c7c0f65f00d72f3f Mon Sep 17 00:00:00 2001 From: Thomas Ales Date: Thu, 6 Aug 2026 10:32:06 -0500 Subject: [PATCH 06/10] Dedup workers: one _PooledWorker base for the three pool fan-outs DcPrecomputeWorker, Ch4MaskWorker and CrossCorrelateWorker shared the same submit/as_completed/emit/shutdown skeleton; they are now constructors plus _plan/_items/_one/_emit hooks. All three inherit the cancellation-aware shutdown (wait=not stopped, cancel_futures) that only DcPrecomputeWorker had before, and the per-run budget attributes are initialised in __init__ instead of appearing mid-run. Also removes the viewer's write-only _dc_precompute_worker plumbing (job tracking already owns worker lifetime via _jobs). Co-Authored-By: Claude Fable 5 --- sras_viewer.py | 7 +-- sras_workers.py | 136 +++++++++++++++++++++++------------------------- 2 files changed, 67 insertions(+), 76 deletions(-) diff --git a/sras_viewer.py b/sras_viewer.py index e28fa73..c1c232f 100644 --- a/sras_viewer.py +++ b/sras_viewer.py @@ -1478,7 +1478,6 @@ class SrasViewerWindow(QMainWindow): # same combination is free. self._dc_cache: dict[tuple[int, int], np.ndarray] = {} self._fft_cache: dict[tuple[int, bool, int | None, float], np.ndarray] = {} - self._dc_precompute_worker: DcPrecomputeWorker | None = None self._dc_generation: int = 0 # Angle alignment ("Fusion" menu) @@ -2463,8 +2462,7 @@ class SrasViewerWindow(QMainWindow): n_angles = self._sras.n_angles worker = DcPrecomputeWorker(self._sras) - self._dc_precompute_worker = worker - started = self._run_worker( + self._run_worker( "dc_precompute", worker, connect=( ("angle_done", lambda a, dc3, dc4, g=generation: @@ -2473,10 +2471,7 @@ class SrasViewerWindow(QMainWindow): f"DC precompute error: {msg}", 5000)), ), quit_on=("finished", "error"), - on_done=lambda: setattr(self, "_dc_precompute_worker", None), ) - if not started: - self._dc_precompute_worker = None def _on_dc_precompute_angle_done(self, generation: int, angle_idx: int, dc3_mv: np.ndarray, dc4_mv: np.ndarray, diff --git a/sras_workers.py b/sras_workers.py index 334e655..5919b79 100644 --- a/sras_workers.py +++ b/sras_workers.py @@ -54,6 +54,34 @@ class CancellableWorker(QObject): return self._stop +class _PooledWorker(CancellableWorker): + """Fans a per-item computation across a thread pool, emitting each result + from this worker's own thread as it lands (never from a pool thread). + + Subclasses provide _plan() -> n_workers (stashing whatever per-run + context they need), _items(), _one(item) -> result, and _emit(result). + On stop(): queued items are dropped, in-flight ones are not waited for — + that is what keeps closing the window responsive on a large scan. + """ + finished = pyqtSignal() + error = pyqtSignal(str) + + def run(self): + try: + pool = ThreadPoolExecutor(max_workers=max(1, self._plan())) + try: + futures = [pool.submit(self._one, it) for it in self._items()] + for fut in as_completed(futures): + if self._stop: + break + self._emit(fut.result()) + finally: + pool.shutdown(wait=not self._stop, cancel_futures=True) + self.finished.emit() + except Exception as exc: + self.error.emit(str(exc)) + + class LoadWorker(QObject): finished = pyqtSignal(object) # SrasFile | None error = pyqtSignal(str) @@ -117,29 +145,29 @@ class ComputeWorker(CancellableWorker): self.error.emit(str(exc)) -class DcPrecomputeWorker(CancellableWorker): +class DcPrecomputeWorker(_PooledWorker): """Computes CH3/CH4 DC images for every angle in the background. DC images are cheap (a per-waveform mean, no FFT) compared to the CH1/Velocity FFT, so precomputing them for the whole file right after load makes switching angles instant while on a DC channel, and also means the FFT masking step (which needs a DC4 image) rarely has to wait on anything. - - Angles are computed on a thread pool — the work is a pure mean over the - waveform block, so it is I/O- and bandwidth-bound and embarrassingly - parallel. Results are emitted one at a time as they land (out of angle - order), and always from this worker's own thread: nothing emits a Qt - signal from a pool thread. """ angle_done = pyqtSignal(int, np.ndarray, np.ndarray) # angle_idx, dc3_mv, dc4_mv - finished = pyqtSignal() - error = pyqtSignal(str) def __init__(self, sras: SrasFile): super().__init__() self._sras = sras + self._angle_budget = 0 - def _one_angle(self, a: int) -> tuple[int, np.ndarray, np.ndarray]: + def _plan(self) -> int: + n_workers, self._angle_budget = compute.plan_angle_level(self._sras) + return n_workers + + def _items(self): + return range(self._sras.n_angles) + + def _one(self, a: int) -> tuple[int, np.ndarray, np.ndarray]: # max_workers=1 *and* a budget share: this call is one of several # concurrent angles, and both the thread count and the buffer size # have to be divided (see compute.plan_angle_level). @@ -149,26 +177,8 @@ class DcPrecomputeWorker(CancellableWorker): dc_image_mv(self._sras, a, CH3_IDX, **kw), dc_image_mv(self._sras, a, CH4_IDX, **kw)) - def run(self): - try: - n = self._sras.n_angles - n_workers, self._angle_budget = compute.plan_angle_level(self._sras) - pool = ThreadPoolExecutor(max_workers=n_workers) - try: - futures = {pool.submit(self._one_angle, a): a for a in range(n)} - for fut in as_completed(futures): - if self._stop: - break - a, dc3, dc4 = fut.result() - self.angle_done.emit(a, dc3, dc4) - finally: - # cancel_futures drops the queued angles; should_stop lets the - # in-flight ones bail within a chunk. Not waiting here is what - # keeps closing the window responsive on a large scan. - pool.shutdown(wait=not self._stop, cancel_futures=True) - self.finished.emit() - except Exception as exc: - self.error.emit(str(exc)) + def _emit(self, result): + self.angle_done.emit(*result) class BatchCacheWorker(QObject): @@ -306,7 +316,7 @@ class AngleAlignmentWorker(QObject): self.finished.emit(None, str(exc)) -class Ch4MaskWorker(QObject): +class Ch4MaskWorker(_PooledWorker): """Fetches each requested angle's CH4 (Bias B) DC image in mV, for ManualAlignmentDialog's initial threshold-mask overlay. @@ -320,35 +330,29 @@ class Ch4MaskWorker(QObject): ManualAlignmentDialog._start_mask_prep). """ angle_done = pyqtSignal(int, np.ndarray) # angle_idx, dc4_mv - finished = pyqtSignal() - error = pyqtSignal(str) def __init__(self, sras: SrasFile, angle_indices: list[int]): super().__init__() self._sras = sras self._angles = angle_indices + self._budget = 0 - def run(self): - try: - n_workers, budget = compute.plan_angle_level(self._sras) - pool = ThreadPoolExecutor(max_workers=n_workers) - try: - futures = { - pool.submit(dc_image_mv, self._sras, a, CH4_IDX, - max_workers=1, budget=budget): a - for a in self._angles - } - for fut in as_completed(futures): - a = futures[fut] - self.angle_done.emit(a, fut.result()) - finally: - pool.shutdown(wait=True) - self.finished.emit() - except Exception as exc: - self.error.emit(str(exc)) + def _plan(self) -> int: + n_workers, self._budget = compute.plan_angle_level(self._sras) + return n_workers + + def _items(self): + return self._angles + + def _one(self, a: int) -> tuple[int, np.ndarray]: + return a, dc_image_mv(self._sras, a, CH4_IDX, + max_workers=1, budget=self._budget) + + def _emit(self, result): + self.angle_done.emit(*result) -class CrossCorrelateWorker(QObject): +class CrossCorrelateWorker(_PooledWorker): """Rigid registration (rotation + translation, never scale) of each of *angle_indices* against *ref_angle_idx*, for ManualAlignmentDialog's Auto Cross-Correlate button. @@ -363,8 +367,6 @@ class CrossCorrelateWorker(QObject): """ # angle_idx, rotation_deg, shift_x_mm, shift_y_mm, score, source angle_done = pyqtSignal(int, float, float, float, float, str) - finished = pyqtSignal() - error = pyqtSignal(str) def __init__(self, sras: SrasFile, ref_angle_idx: int, angle_indices: list[int], dc4_mv: dict[int, np.ndarray], *, @@ -379,25 +381,19 @@ class CrossCorrelateWorker(QObject): self._threshold = dc_threshold_mv self._search_deg = search_deg + def _plan(self) -> int: + return compute._registration_workers(self._sras, compute._DEFAULT_FINE_DIM) + + def _items(self): + return self._angles + def _one(self, a: int) -> tuple[int, compute.RigidFit]: return a, compute.register_angle_to_reference( self._sras, a, self._ref, self._dc4_mv, dc_threshold_mv=self._threshold, sources=self._sources, search_deg=self._search_deg) - def run(self): - try: - n_workers = compute._registration_workers( - self._sras, compute._DEFAULT_FINE_DIM) - pool = ThreadPoolExecutor(max_workers=max(1, n_workers)) - try: - futures = [pool.submit(self._one, a) for a in self._angles] - for fut in as_completed(futures): - a, fit = fut.result() - self.angle_done.emit(a, fit.rotation_deg, fit.shift_mm[0], - fit.shift_mm[1], fit.score, fit.source) - finally: - pool.shutdown(wait=True) - self.finished.emit() - except Exception as exc: - self.error.emit(str(exc)) + def _emit(self, result): + a, fit = result + self.angle_done.emit(a, fit.rotation_deg, fit.shift_mm[0], + fit.shift_mm[1], fit.score, fit.source) From 00a7afade0e1d525bc088f4bb3008c5ca06e0543 Mon Sep 17 00:00:00 2001 From: Thomas Ales Date: Thu, 6 Aug 2026 10:40:42 -0500 Subject: [PATCH 07/10] Dedup format layer; public accessors replace private reach-throughs MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - SrasFile: _parse_v6 now retains per-angle x_delta and the verbatim preamble/background byte spans, and gains public data_offset, y_pos_per_angle, and iter_angle_blocks() (which now owns the ragged block-offset walk used three separate places before). - sras_edit_scans: the 70-line re-parse of the v6 header sections (_read_v6_sections/_reread_span/_v6_angle_offsets) collapses into a _write_v6 that consumes SrasFile directly — verified byte-identical round-trip on v6 int8/int16 and legacy files. Eight print-and-exit pairs become _die(). - tools/make_test_sras imports the struct layouts from sras_format and the rotation matrix from sras_compute instead of re-declaring them (byte assembly stays independent of the reader). - sras_compute: block_mean_2d, pixel_pitch_mm, nominal_delta_deg made public (they were GUI-facing); registration_workers() and default_max_workers() wrap the remaining private reach-throughs from sras_workers. Co-Authored-By: Claude Fable 5 --- sras_average.py | 2 +- sras_compute.py | 30 +++++--- sras_edit_scans.py | 148 +++++++++---------------------------- sras_format.py | 61 ++++++++++++--- sras_viewer.py | 8 +- sras_workers.py | 5 +- tests/test_alignment.py | 8 +- tests/test_gui.py | 6 +- tools/check_equivalence.py | 4 +- tools/make_test_sras.py | 18 ++--- 10 files changed, 134 insertions(+), 156 deletions(-) diff --git a/sras_average.py b/sras_average.py index d997a15..42aeaca 100644 --- a/sras_average.py +++ b/sras_average.py @@ -47,7 +47,7 @@ def read_header_sections(sras: SrasFile) -> bytes: lost in a re-encode.""" with open(sras.path, "rb") as f: f.seek(HDR_SIZE) - return f.read(sras._data_offset - HDR_SIZE) + return f.read(sras.data_offset - HDR_SIZE) def average_rows(block: np.ndarray, n: int, discard_remainder: bool) -> np.ndarray: diff --git a/sras_compute.py b/sras_compute.py index fc7b30e..0305c54 100644 --- a/sras_compute.py +++ b/sras_compute.py @@ -768,7 +768,7 @@ def _skimage_phase_cross_correlation(): # ---- Geometry: local mm, ref mm, and the one affine builder --------------- -def _pixel_pitch_mm(sras: SrasFile, angle_idx: int) -> tuple[float, float]: +def pixel_pitch_mm(sras: SrasFile, angle_idx: int) -> tuple[float, float]: """(dx, dy) mm/pixel for one angle: dx is the file-wide constant pixel_x_mm; dy is this angle's own row spacing (assumed uniform, the same assumption _redraw_image makes when it builds the display extent). dy keeps @@ -803,7 +803,7 @@ def _local_half_extent_mm(sras: SrasFile, angle_idx: int) -> tuple[float, float] """(half width, half height) in mm from this angle's array center to the center of its outermost pixel.""" n_rows, n_frames = sras.image_shape(angle_idx) - dx, dy = _pixel_pitch_mm(sras, angle_idx) + dx, dy = pixel_pitch_mm(sras, angle_idx) return (n_frames - 1) / 2.0 * abs(dx), (n_rows - 1) / 2.0 * abs(dy) @@ -813,7 +813,7 @@ def _rotation_matrix(theta_deg: float) -> np.ndarray: return np.array([[c, -s], [s, c]]) # CCW rotation acting on (x, y) -def _nominal_delta_deg(sras: SrasFile, angle_idx: int, ref_idx: int) -> float: +def nominal_delta_deg(sras: SrasFile, angle_idx: int, ref_idx: int) -> float: """The rotation-stage's own reported angle change between two angles. Used only to *seed* the rotation search, never as the answer: the stage's @@ -888,7 +888,7 @@ def apply_alignment(result: AlignmentResult, angle_idx: int, img: np.ndarray, mode="constant", cval=0.0) -def _block_mean_2d(img: np.ndarray, fy: int, fx: int) -> np.ndarray: +def block_mean_2d(img: np.ndarray, fy: int, fx: int) -> np.ndarray: """Block-mean by independent row/column factors. Independent factors matter because the raw grid is strongly anisotropic (5 µm along x, 50 µm along y): a single square factor would either alias along x or throw away rows.""" @@ -944,10 +944,10 @@ def _prepare_reg_image(sras: SrasFile, angle_idx: int, img: np.ndarray, resampled onto the registration grid. Pre-averaging matters: the raw grid is 10x finer along x than along y, so sampling it directly at the (much coarser) isotropic registration pitch would alias badly along x.""" - dx, dy = _pixel_pitch_mm(sras, angle_idx) + dx, dy = pixel_pitch_mm(sras, angle_idx) fx = max(1, int(pitch_mm / abs(dx))) fy = max(1, int(pitch_mm / abs(dy))) - small = _block_mean_2d(np.asarray(img, dtype=np.float32), fy, fx) + small = block_mean_2d(np.asarray(img, dtype=np.float32), fy, fx) n_rows, n_frames = img.shape return _RegImage( small, dx * fx, dy * fy, @@ -1087,7 +1087,7 @@ def _reg_pitch_and_size(sras: SrasFile, max_dim: int, span = diag * margin native = float(np.sqrt( abs(sras.pixel_x_mm) - * max(abs(_pixel_pitch_mm(sras, a)[1]) for a in range(sras.n_angles)))) + * max(abs(pixel_pitch_mm(sras, a)[1]) for a in range(sras.n_angles)))) pitch = max(span / max_dim, native) n = int(scipy_fft.next_fast_len(max(16, int(np.ceil(span / pitch))))) return pitch, n @@ -1111,6 +1111,16 @@ def _registration_workers(sras: SrasFile, fine_dim: int) -> int: _TOTAL_BYTES_BUDGET // max(1, per_worker)))) +def registration_workers(sras: SrasFile) -> int: + """Public: concurrent-registration cap at the default fine grid.""" + return _registration_workers(sras, _DEFAULT_FINE_DIM) + + +def default_max_workers() -> int: + """Public: the module-wide worker cap (SRAS_MAX_WORKERS or cpu count).""" + return _MAX_WORKERS + + def _rotation_candidates(nominal_deg: float, search_deg: float, step_deg: float) -> list[float]: """Coarse rotation candidates: a window around *both* signs of the stage's @@ -1254,7 +1264,7 @@ def register_angle_to_reference( if not candidates: return RigidFit(0.0, (0.0, 0.0), -1.0, "none") - nominal = _nominal_delta_deg(sras, angle_idx, ref_angle_idx) + nominal = nominal_delta_deg(sras, angle_idx, ref_angle_idx) thetas = _rotation_candidates(nominal, search_deg, coarse_step_deg) # ---- Stage 1: coarse sweep, every source ------------------------------ @@ -1363,7 +1373,7 @@ def build_canvas_affine(sras: SrasFile, angle_idx: int, ref_angle_idx: int, (rows, cols) block-mean factor already applied to the image the caller will resample — 1:1 for the raw image, coarser for the manual-alignment preview's downsampled masks.""" - dx_a, dy_a = _pixel_pitch_mm(sras, angle_idx) + dx_a, dy_a = pixel_pitch_mm(sras, angle_idx) n_rows, n_frames = sras.image_shape(angle_idx) fy, fx = (max(1, int(v)) for v in src_downsample) if (fy, fx) != (1, 1): @@ -1398,7 +1408,7 @@ def _result_from_params(sras: SrasFile, ref_angle_idx: int, the shared canvas, then build each angle's canvas->raw affine. Pure matrix and bbox math, so it is cheap enough to call synchronously on the GUI thread on every manual edit.""" - pitch = _pixel_pitch_mm(sras, ref_angle_idx) + pitch = pixel_pitch_mm(sras, ref_angle_idx) canvas_origin_mm, canvas_shape = canvas_for_params( sras, ref_angle_idx, pitch, params, snap=True) diff --git a/sras_edit_scans.py b/sras_edit_scans.py index 8774daf..040e0d5 100644 --- a/sras_edit_scans.py +++ b/sras_edit_scans.py @@ -25,15 +25,17 @@ import struct import sys from pathlib import Path -from sras_format import ( - GEO_FMT_V6, GEO_SIZE_V6, HDR_FMT, HDR_FMT_V6, HDR_SIZE, HDR_SIZE_V6, - SrasFile, -) +from sras_format import GEO_FMT_V6, HDR_FMT, HDR_FMT_V6, HDR_SIZE, SrasFile _LEGACY_VERSIONS = (2, 3, 4, 5) _V6_VERSIONS = (6, 7) +def _die(msg: str): + print(f"Error: {msg}", file=sys.stderr) + sys.exit(1) + + def parse_args(): p = argparse.ArgumentParser( description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) @@ -106,7 +108,7 @@ def _write_legacy(sras: SrasFile, keep: list[int], out_path: Path): angle_table_size = sras.n_angles * 4 with open(sras.path, "rb") as f: f.seek(HDR_SIZE + angle_table_size) - shared_mid = f.read(sras._data_offset - (HDR_SIZE + angle_table_size)) + shared_mid = f.read(sras.data_offset - (HDR_SIZE + angle_table_size)) angle_bytes = n_rows * n_ch * n_frames * spf * bps @@ -115,128 +117,55 @@ def _write_legacy(sras: SrasFile, keep: list[int], out_path: Path): fout.write(sras.angles_deg[keep].astype(">f4").tobytes()) fout.write(shared_mid) for a in keep: - _copy_range(fin, fout, sras._data_offset + a * angle_bytes, angle_bytes) + _copy_range(fin, fout, sras.data_offset + a * angle_bytes, angle_bytes) # --------------------------------------------------------------------------- # v6/v7: per-angle geometry, ragged waveform blocks # --------------------------------------------------------------------------- -def _read_v6_sections(path: Path) -> dict: - """Raw, low-level read of everything before the waveform data. - - SrasFile._parse_v6 reads and discards each angle's x_delta (it's not - part of the display-facing geometry it exposes), so a round-trip through - SrasFile would silently drop that field. Re-parsing here keeps every - byte of the Per-Angle Geometry Table intact. - """ - with open(path, "rb") as f: - hdr_raw = f.read(HDR_SIZE_V6) - (magic, ver, n_angles_declared, x_start_nom, y_start_nom, x_delta_nom, - y_delta_nom, row_spacing, vel, freq, spf, sr, bps, - n_ch) = struct.unpack(HDR_FMT_V6, hdr_raw) - - angles_deg = list(struct.unpack(f">{n_angles_declared}f", - f.read(n_angles_declared * 4))) - - geo = [struct.unpack(GEO_FMT_V6, f.read(GEO_SIZE_V6)) - for _ in range(n_angles_declared)] # (x_start, x_delta, n_frames, n_rows) - - row_table = [f.read(geo[a][3] * 4) for a in range(n_angles_declared)] - - preamble_start = f.tell() - for _ in range(n_ch): - (length,) = struct.unpack(">H", f.read(2)) - f.read(length) - preambles_raw = _reread_span(f, preamble_start) - - bg_start = f.tell() - (n_bg,) = struct.unpack(">I", f.read(4)) - f.read(n_bg) - background_raw = _reread_span(f, bg_start) - - data_offset = f.tell() - - return { - "version": ver, "n_angles_declared": n_angles_declared, - "x_start_nom": x_start_nom, "y_start_nom": y_start_nom, - "x_delta_nom": x_delta_nom, "y_delta_nom": y_delta_nom, - "row_spacing": row_spacing, "vel": vel, "freq": freq, - "spf": spf, "sr": sr, "bps": bps, "n_ch": n_ch, - "angles_deg": angles_deg, "geo": geo, "row_table": row_table, - "preambles_raw": preambles_raw, "background_raw": background_raw, - "data_offset": data_offset, - } - - -def _reread_span(f, start: int) -> bytes: - end = f.tell() - f.seek(start) - span = f.read(end - start) - f.seek(end) - return span - - -def _v6_angle_offsets(sections: dict, file_size: int) -> list[tuple[int, int]]: - """(offset, nbytes) of each declared angle's waveform block, stopping at - the first angle whose data isn't fully on disk (aborted scan).""" - n_ch, spf, bps = sections["n_ch"], sections["spf"], sections["bps"] - offsets = [] - offset = sections["data_offset"] - for xs, xd, nf, nr in sections["geo"]: - nbytes = nr * n_ch * nf * spf * bps - if offset + nbytes > file_size: - break - offsets.append((offset, nbytes)) - offset += nbytes - return offsets - - -def _write_v6(in_path: Path, sections: dict, keep: list[int], out_path: Path): - file_size = in_path.stat().st_size - offsets = _v6_angle_offsets(sections, file_size) - geo = sections["geo"] - +def _write_v6(sras: SrasFile, keep: list[int], out_path: Path): header = struct.pack( - HDR_FMT_V6, b"SRAS", sections["version"], len(keep), - sections["x_start_nom"], sections["y_start_nom"], - sections["x_delta_nom"], sections["y_delta_nom"], - sections["row_spacing"], sections["vel"], sections["freq"], - sections["spf"], sections["sr"], sections["bps"], sections["n_ch"], + HDR_FMT_V6, b"SRAS", sras.version, len(keep), + sras.x_start_nominal_mm, sras.y_start_nominal_mm, + sras.x_delta_nominal_mm, sras.y_delta_nominal_mm, + sras.row_spacing_mm, sras.velocity_mm_s, sras.laser_freq_hz, + sras.samples_per_frame, sras.sample_rate_hz, + sras.bytes_per_sample, sras.n_channels, ) - with open(in_path, "rb") as fin, open(out_path, "wb") as fout: + blocks = {a: (offset, nbytes) for a, offset, nbytes in sras.iter_angle_blocks()} + with open(sras.path, "rb") as fin, open(out_path, "wb") as fout: fout.write(header) - fout.write(struct.pack(f">{len(keep)}f", - *[sections["angles_deg"][i] for i in keep])) + fout.write(sras.angles_deg[keep].astype(">f4").tobytes()) for i in keep: - fout.write(struct.pack(GEO_FMT_V6, *geo[i])) + fout.write(struct.pack( + GEO_FMT_V6, float(sras.x_start_mm[i]), + float(sras.x_delta_mm_per_angle[i]), + int(sras.n_frames[i]), int(sras.n_rows[i]))) for i in keep: - fout.write(sections["row_table"][i]) - fout.write(sections["preambles_raw"]) - fout.write(sections["background_raw"]) + fout.write(sras.y_pos_per_angle[i].astype(">f4").tobytes()) + fout.write(sras.preambles_raw) + fout.write(sras.background_raw) for i in keep: - off, nbytes = offsets[i] - _copy_range(fin, fout, off, nbytes) + offset, nbytes = blocks[i] + _copy_range(fin, fout, offset, nbytes) def main(): args = parse_args() in_path = Path(args.input) if not in_path.exists(): - print(f"Error: input file not found: {in_path}", file=sys.stderr) - sys.exit(1) + _die(f"input file not found: {in_path}") print(f"Reading {in_path} ...", flush=True) try: sras = SrasFile(str(in_path)) except ValueError as e: - print(f"Error: {e}", file=sys.stderr) - sys.exit(1) + _die(str(e)) if sras.version not in (*_LEGACY_VERSIONS, *_V6_VERSIONS): - print(f"Error: unsupported .sras version: {sras.version}", file=sys.stderr) - sys.exit(1) + _die(f"unsupported .sras version: {sras.version}") aborted_note = " (scan aborted; trailing angle(s) already excluded)" if sras.scan_aborted else "" print(f" Version : v{sras.version}", flush=True) @@ -247,16 +176,13 @@ def main(): return if not args.output: - print("Error: output path required unless --list is given.", file=sys.stderr) - sys.exit(1) + _die("output path required unless --list is given.") if not (args.drop or args.keep): - print("Error: specify --drop or --keep (see --list for indices).", file=sys.stderr) - sys.exit(1) + _die("specify --drop or --keep (see --list for indices).") out_path = Path(args.output) if out_path.resolve() == in_path.resolve(): - print("Error: output path must differ from input path.", file=sys.stderr) - sys.exit(1) + _die("output path must differ from input path.") try: if args.drop: @@ -265,12 +191,10 @@ def main(): else: keep = sorted(_parse_index_list(args.keep, sras.n_angles)) except ValueError as e: - print(f"Error: {e}", file=sys.stderr) - sys.exit(1) + _die(str(e)) if not keep: - print("Error: at least one angle must remain.", file=sys.stderr) - sys.exit(1) + _die("at least one angle must remain.") dropped = [a for a in range(sras.n_angles) if a not in keep] print(f"\nDropping angle(s): {dropped}") @@ -280,7 +204,7 @@ def main(): if sras.version in _LEGACY_VERSIONS: _write_legacy(sras, keep, out_path) else: - _write_v6(in_path, _read_v6_sections(in_path), keep, out_path) + _write_v6(sras, keep, out_path) in_mb = in_path.stat().st_size / 1024**2 out_mb = out_path.stat().st_size / 1024**2 diff --git a/sras_format.py b/sras_format.py index 0b2715d..46d8003 100644 --- a/sras_format.py +++ b/sras_format.py @@ -375,21 +375,34 @@ class SrasFile: angles = np.frombuffer(f.read(n_angles * 4), dtype=">f4").astype(np.float32) x_start = np.empty(n_angles, dtype=np.float64) + x_delta = np.empty(n_angles, dtype=np.float64) n_frames = np.empty(n_angles, dtype=np.int64) n_rows = np.empty(n_angles, dtype=np.int64) for a in range(n_angles): xs, xd, nf, nr = _read_struct(f, GEO_FMT_V6) - x_start[a], n_frames[a], n_rows[a] = xs, nf, nr + x_start[a], x_delta[a], n_frames[a], n_rows[a] = xs, xd, nf, nr y_pos_per_angle = [ np.frombuffer(f.read(int(n_rows[a]) * 4), dtype=">f4").astype(np.float32) for a in range(n_angles) ] + # Verbatim on-disk spans of the preamble and background sections, + # kept so file-rewriting tools (sras_edit_scans) can carry them + # over byte-for-byte without re-parsing. + span_start = f.tell() self._set_calibration(_read_preambles(f, n_ch), n_ch) - self.background = _read_background(f) + span_end = f.tell() + f.seek(span_start) + self.preambles_raw = f.read(span_end - span_start) - data_offset = f.tell() + span_start = span_end + self.background = _read_background(f) + span_end = f.tell() + f.seek(span_start) + self.background_raw = f.read(span_end - span_start) + + data_offset = span_end self._data_offset = data_offset @@ -428,6 +441,7 @@ class SrasFile: self.scan_aborted = n_complete < n_angles_declared self.angles_deg = angles[:n_complete] self.x_start_mm = x_start[:n_complete] + self.x_delta_mm_per_angle = x_delta[:n_complete] self.n_frames = n_frames[:n_complete] self.n_rows = n_rows[:n_complete] self._y_pos_per_angle = y_pos_per_angle[:n_complete] @@ -436,6 +450,37 @@ class SrasFile: if self.version == 7 and offset < file_size: self._parse_cach_section(offset) + # ------------------------------------------------------------------ + # Byte-layout accessors (public: used by file-rewriting tools) + # ------------------------------------------------------------------ + + @property + def data_offset(self) -> int: + """File offset where the waveform data begins (headers end).""" + return self._data_offset + + @property + def y_pos_per_angle(self) -> list[np.ndarray]: + """Per-angle Y row positions (mm). The list and its arrays are the + live parsed state — tools that reproject may replace entries.""" + return self._y_pos_per_angle + + @y_pos_per_angle.setter + def y_pos_per_angle(self, value: list[np.ndarray]): + self._y_pos_per_angle = value + + def iter_angle_blocks(self): + """Yields (angle_idx, byte_offset, byte_count) for each complete + angle's waveform block. Works for every version: legacy files have + uniform per-angle geometry, so the same walk applies.""" + offset = self._data_offset + for a in range(self.n_angles): + nbytes = (int(self.n_rows[a]) * self.n_channels + * int(self.n_frames[a]) * self.samples_per_frame + * self.bytes_per_sample) + yield a, offset, nbytes + offset += nbytes + # ------------------------------------------------------------------ # v7 cache tail (CACH section: precomputed DC / FFT images) # ------------------------------------------------------------------ @@ -445,12 +490,10 @@ class SrasFile: start), derived purely from the header + Per-Angle Geometry Table — independent of whether a cache tail is actually present. Used by both the parser and the in-place writer.""" - waveform_bytes = sum( - int(self.n_rows[a]) * self.n_channels * int(self.n_frames[a]) - * self.samples_per_frame * self.bytes_per_sample - for a in range(self.n_angles) - ) - return self._data_offset + int(waveform_bytes) + end = self._data_offset + for _, offset, nbytes in self.iter_angle_blocks(): + end = offset + nbytes + return end def _read_cache_block(self, f, hdr_fmt: str, magic: bytes, stores: list[list]) -> int | None: diff --git a/sras_viewer.py b/sras_viewer.py index c1c232f..23f9344 100644 --- a/sras_viewer.py +++ b/sras_viewer.py @@ -1127,7 +1127,7 @@ class ManualAlignmentDialog(QDialog): threshold = self.spin_mask_threshold_mv.value() fy, fx = self._downsample self._masks_small = { - a: compute._block_mean_2d((img >= threshold).astype(np.float32), fy, fx) + a: compute.block_mean_2d((img >= threshold).astype(np.float32), fy, fx) for a, img in self._dc4_mv.items() } @@ -1143,7 +1143,7 @@ class ManualAlignmentDialog(QDialog): mask-threshold change, Auto De-rotate, a rotation nudge/edit of the active angle. NOT triggered by a translation-only nudge — see _refresh_active_preview_layer.""" - dx_ref, dy_ref = compute._pixel_pitch_mm(self._sras, self._ref_angle_idx) + dx_ref, dy_ref = compute.pixel_pitch_mm(self._sras, self._ref_angle_idx) fy, fx = self._downsample pitch = (dx_ref * fx, dy_ref * fy) origin, shape = compute.canvas_for_params( @@ -1288,7 +1288,7 @@ class ManualAlignmentDialog(QDialog): for a in range(self._sras.n_angles): if a == self._ref_angle_idx: continue - self._angle_params[a].rotation_deg = sign * compute._nominal_delta_deg( + self._angle_params[a].rotation_deg = sign * compute.nominal_delta_deg( self._sras, a, self._ref_angle_idx) n_changed += 1 self._sync_active_spinboxes() @@ -1364,7 +1364,7 @@ class ManualAlignmentDialog(QDialog): f"{worst[1][1]})."] drifted = [] for a, _note in rows: - nominal = compute._nominal_delta_deg(self._sras, a, self._ref_angle_idx) + nominal = compute.nominal_delta_deg(self._sras, a, self._ref_angle_idx) got = self._angle_params[a].rotation_deg dev = min(abs(got - nominal), abs(got + nominal)) if dev > 1.0: diff --git a/sras_workers.py b/sras_workers.py index 5919b79..4a75dd8 100644 --- a/sras_workers.py +++ b/sras_workers.py @@ -253,7 +253,8 @@ class BatchCacheWorker(QObject): for path in paths: try: err = cache_file(path, self._mode, self._apply_bg_sub, - compute.get_fft_backend(), compute._MAX_WORKERS) + compute.get_fft_backend(), + compute.default_max_workers()) except Exception as exc: err = str(exc) done += 1 @@ -382,7 +383,7 @@ class CrossCorrelateWorker(_PooledWorker): self._search_deg = search_deg def _plan(self) -> int: - return compute._registration_workers(self._sras, compute._DEFAULT_FINE_DIM) + return compute.registration_workers(self._sras) def _items(self): return self._angles diff --git a/tests/test_alignment.py b/tests/test_alignment.py index 18be785..04bb85f 100644 --- a/tests/test_alignment.py +++ b/tests/test_alignment.py @@ -40,7 +40,7 @@ def mm_transform(sras: SrasFile, result, angle_idx: int) -> np.ndarray: per-angle pixel pitches; undoing both must leave something orthonormal, or the transform is smuggling in a scale or a shear. """ - dx_a, dy_a = compute._pixel_pitch_mm(sras, angle_idx) + dx_a, dy_a = compute.pixel_pitch_mm(sras, angle_idx) A_out = np.array([[0.0, result.canvas_dx_mm], [result.canvas_dy_mm, 0.0]]) D = np.array([[0.0, 1.0 / dy_a], [1.0 / dx_a, 0.0]]) return np.linalg.inv(D) @ result.per_angle[angle_idx].matrix @ np.linalg.inv(A_out) @@ -99,7 +99,7 @@ def test_stage_coordinates_are_not_consulted(rig): moved = SrasFile(str(rig.path)) for a in range(1, moved.n_angles): moved.x_start_mm[a] += 13.5 * a - moved._y_pos_per_angle[a] = moved._y_pos_per_angle[a] - 9.25 * a + moved.y_pos_per_angle[a] = moved.y_pos_per_angle[a] - 9.25 * a moved_dc4 = dc4_images(moved) moved_fits = {a: compute.register_angle_to_reference( moved, a, 0, moved_dc4, dc_threshold_mv=_THRESHOLD_MV) @@ -118,7 +118,7 @@ def test_canvas_is_reference_grid_extended(rig): assert np.allclose(t0.offset, np.round(t0.offset)), \ f"angle 0 does not land on whole canvas pixels: {t0.offset}" assert ((result.canvas_dx_mm, result.canvas_dy_mm) - == compute._pixel_pitch_mm(sras, 0)), \ + == compute.pixel_pitch_mm(sras, 0)), \ "canvas pitch is angle 0's own pitch" n_rows, n_cols = result.canvas_shape @@ -169,7 +169,7 @@ def test_downsampled_preview_lands_with_full_res(rig): result.canvas_origin_mm, result.canvas_shape) fy, fx = 4, 16 small = compute.reproject_mask( - sras, a, 0, compute._block_mean_2d(full_mask, fy, fx), + sras, a, 0, compute.block_mean_2d(full_mask, fy, fx), p.rotation_deg, p.shift_mm, (pitch[0] * fx, pitch[1] * fy), result.canvas_origin_mm, (result.canvas_shape[0] // fy, result.canvas_shape[1] // fx), diff --git a/tests/test_gui.py b/tests/test_gui.py index 7b1ebec..cb0bc28 100644 --- a/tests/test_gui.py +++ b/tests/test_gui.py @@ -261,7 +261,7 @@ def test_manual_alignment_geometry(ctx): assert np.allclose(compute._center_idx(s, 0), [(n_rows - 1) / 2, (n_frames - 1) / 2]), \ "array center is the geometric center of the pixel grid" - dx0, dy0 = compute._pixel_pitch_mm(s, 0) + dx0, dy0 = compute.pixel_pitch_mm(s, 0) assert np.allclose(compute._local_half_extent_mm(s, 0), [(n_frames - 1) / 2 * abs(dx0), (n_rows - 1) / 2 * abs(dy0)]), \ "local half-extent is derived from shape and pitch alone" @@ -270,7 +270,7 @@ def test_manual_alignment_geometry(ctx): moved = SrasFile(str(ctx.path)) for a in range(1, moved.n_angles): moved.x_start_mm[a] += 7.5 - moved._y_pos_per_angle[a] = moved._y_pos_per_angle[a] + 3.25 + moved.y_pos_per_angle[a] = moved.y_pos_per_angle[a] + 3.25 origin_b, shape_b = compute.canvas_for_params(moved, 0, (dx0, dy0), identity) assert shape_a == shape_b and np.allclose(origin_a, origin_b), \ ("moving every non-reference angle's scan window must leave the canvas " @@ -356,7 +356,7 @@ def test_auto_derotate(ctx): dlg, s, active = ctx.dlg, ctx.s, ctx.active shift_before_derotate = dlg._angle_params[active].shift_mm dlg._on_auto_derotate() - nominal = compute._nominal_delta_deg(s, active, dlg._ref_angle_idx) + nominal = compute.nominal_delta_deg(s, active, dlg._ref_angle_idx) assert abs(dlg._angle_params[active].rotation_deg - nominal) < 1e-6, \ "auto de-rotate seeded rotation from the stage's reported angle" assert dlg._angle_params[active].shift_mm == shift_before_derotate, \ diff --git a/tools/check_equivalence.py b/tools/check_equivalence.py index 37934dc..9488040 100644 --- a/tools/check_equivalence.py +++ b/tools/check_equivalence.py @@ -47,8 +47,8 @@ def row_slice(sras: SrasFile, angle_idx: int, n_rows: int) -> SrasFile: view.n_rows[angle_idx] = n view.data = list(sras.data) view.data[angle_idx] = sras.data[angle_idx][:n] - view._y_pos_per_angle = list(sras._y_pos_per_angle) - view._y_pos_per_angle[angle_idx] = sras._y_pos_per_angle[angle_idx][:n] + view.y_pos_per_angle = list(sras.y_pos_per_angle) + view.y_pos_per_angle[angle_idx] = sras.y_pos_per_angle[angle_idx][:n] return view diff --git a/tools/make_test_sras.py b/tools/make_test_sras.py index 92b0b54..0a3651b 100644 --- a/tools/make_test_sras.py +++ b/tools/make_test_sras.py @@ -11,12 +11,19 @@ Usage: import argparse import struct +import sys from pathlib import Path import numpy as np -HDR_FMT_V6 = ">4sBHfffffffIdBB" -GEO_FMT_V6 = ">ffIH" +sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) + +# Single source of truth for the byte layout: the reader's own constants. +# The byte *assembly* below stays independent, so a writer bug can't be +# masked by a matching reader bug. +from sras_format import HDR_FMT as HDR_FMT_LEGACY # noqa: E402 +from sras_format import GEO_FMT_V6, HDR_FMT_V6 # noqa: E402 +from sras_compute import _rotation_matrix as _rot # noqa: E402 # Per-angle (n_rows, n_frames) — deliberately different per angle so ragged # geometry handling is actually exercised. @@ -171,12 +178,6 @@ def _sample_shape_mv(u: np.ndarray, v: np.ndarray) -> np.ndarray: return img -def _rot(theta_deg: float) -> np.ndarray: - t = np.radians(theta_deg) - c, s = np.cos(t), np.sin(t) - return np.array([[c, -s], [s, c]]) - - def write_rotating(path: Path, n_angles: int = 5, samples_per_frame: int = 4, seed: int = 0) -> dict: """Write a v6 file whose CH4 DC image is one sample seen at n_angles known @@ -260,7 +261,6 @@ def write_rotating(path: Path, n_angles: int = 5, samples_per_frame: int = 4, "y_starts": y_starts, "dx_mm": _ROT_DX_MM, "dy_mm": _ROT_DY_MM} -HDR_FMT_LEGACY = ">4sBHHffffIIdBB" def write_legacy(path: Path, version: int = 4, n_angles: int = 2, From 26a34f7436fc9c4ebaec9eb09915aa4753a9b0fc Mon Sep 17 00:00:00 2001 From: Thomas Ales Date: Thu, 6 Aug 2026 10:47:34 -0500 Subject: [PATCH 08/10] Split sras_viewer.py into a package (pure move) The 2790-line module becomes sras_viewer/: common.py (constants + layout helpers), canvases.py (RoiQuad, ImageCanvas, WaveformCanvas, ManualAlignOverlayCanvas), dialogs.py (FftOptionsDialog, ManualAlignmentDialog), main_window.py (SrasViewerWindow + main), with __init__ re-exporting the public names and __main__ keeping `python -m sras_viewer` working. pyproject gains a `sras-viewer` console script. Code moved verbatim; only import headers are new (pyflakes-clean). tests/test_gui.py patch targets follow the classes to their new modules. Co-Authored-By: Claude Fable 5 --- pyproject.toml | 5 +- sras_viewer.py | 2796 ------------------------------------ sras_viewer/__init__.py | 24 + sras_viewer/__main__.py | 4 + sras_viewer/canvases.py | 542 +++++++ sras_viewer/common.py | 115 ++ sras_viewer/dialogs.py | 767 ++++++++++ sras_viewer/main_window.py | 1403 ++++++++++++++++++ tests/test_gui.py | 6 +- 9 files changed, 2862 insertions(+), 2800 deletions(-) delete mode 100644 sras_viewer.py create mode 100644 sras_viewer/__init__.py create mode 100644 sras_viewer/__main__.py create mode 100644 sras_viewer/canvases.py create mode 100644 sras_viewer/common.py create mode 100644 sras_viewer/dialogs.py create mode 100644 sras_viewer/main_window.py diff --git a/pyproject.toml b/pyproject.toml index ae73ea4..edad96f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -23,15 +23,18 @@ dependencies = [ [project.optional-dependencies] dev = ["pytest"] +[project.scripts] +sras-viewer = "sras_viewer.main_window:main" + [tool.setuptools] py-modules = [ "sras_format", "sras_compute", "sras_workers", - "sras_viewer", "sras_average", "sras_edit_scans", ] +packages = ["sras_viewer"] [tool.pytest.ini_options] testpaths = ["tests"] diff --git a/sras_viewer.py b/sras_viewer.py deleted file mode 100644 index 23f9344..0000000 --- a/sras_viewer.py +++ /dev/null @@ -1,2796 +0,0 @@ -#!/usr/bin/env python3 -""" -SRAS Scan File Viewer -PyQt6 application for visualizing channel data from .sras binary scan files. - -Channel semantics (fixed by sc3_aui_app.py acquisition settings): - CH1 — RF Acoustic Packet (AC-coupled, 100 mV/div): FFT → peak frequency - CH3 — Bias A (DC-coupled, 50 mV/div): waveform mean - CH4 — Bias B (DC-coupled, 50 mV/div): waveform mean - -RF images are masked: pixels where CH4_dc < dc_threshold show 0. - -File parsing lives in sras_format, image/alignment math in sras_compute, and -background workers in sras_workers — the first two import neither Qt nor -matplotlib so multiprocessing children can load them cheaply. -""" - -import faulthandler -import sys -from pathlib import Path - -import matplotlib as mpl -import numpy as np -from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg, NavigationToolbar2QT -from matplotlib.figure import Figure -from matplotlib.patches import Polygon -from matplotlib.path import Path as MplPath -from PyQt6.QtCore import QObject, QSettings, Qt, QThread, pyqtSignal -from PyQt6.QtGui import QAction, QKeyEvent -from PyQt6.QtWidgets import ( - QApplication, QButtonGroup, QCheckBox, QComboBox, QDialog, QDialogButtonBox, - QDoubleSpinBox, QFileDialog, QFormLayout, QFrame, QGroupBox, QHBoxLayout, - QLabel, QMainWindow, QMessageBox, QProgressDialog, QPushButton, QRadioButton, - QScrollArea, QSizePolicy, QSpinBox, QSplitter, QVBoxLayout, QWidget, -) - -import sras_compute as compute -from sras_compute import ( - PYFFTW_AVAILABLE, ManualAngleParams, apply_alignment, build_manual_alignment, - delete_manual_alignment, load_manual_alignment, save_manual_alignment, - sidecar_path, -) -from sras_format import ( - CH1_IDX, CH3_IDX, CH4_IDX, CH_NAMES, SrasFile, adc_to_mv, mv_to_adc, - _FALLBACK_YMULT_MV, _FALLBACK_YOFF_ADC, -) -from sras_workers import ( - AngleAlignmentWorker, BatchCacheWorker, Ch4MaskWorker, ComputeWorker, - CrossCorrelateWorker, DcPrecomputeWorker, LoadWorker, -) - -faulthandler.enable() # print a native stack trace on SIGSEGV/SIGABRT/etc. - -# --------------------------------------------------------------------------- -# Display constants -# --------------------------------------------------------------------------- - -CH_LABELS = [ - "CH1 — RF (FFT peak freq)", - "CH3 — Bias A (DC mean)", - "CH4 — Bias B (DC mean)", - "CH1 — Velocity (SRAS)", -] - -# Combo index for the derived velocity mode (uses CH1_IDX data) -VELOCITY_MODE_IDX = 3 -# All modes that operate on CH1 waveforms -CH1_DERIVED_MODES = (CH1_IDX, VELOCITY_MODE_IDX) - -CMAPS = ["gray", "viridis", "plasma", "inferno", "hot", "jet", "RdBu_r", "seismic"] - -# (mode_str, status-bar unit, colorbar label) per channel index -_CHANNEL_DISPLAY = { - CH1_IDX: ("RF", "Peak frequency (MHz)", "MHz"), - CH3_IDX: ("DC", "DC mean (mV)", "mV"), - CH4_IDX: ("DC", "DC mean (mV)", "mV"), - VELOCITY_MODE_IDX: ("Velocity", "Velocity (m/s)", "m/s"), -} - -_CSS_HINT = "font-size: 11px; color: #aaa;" -_CSS_INFO = "font-size: 11px;" -_CSS_MUTED = "color: #888; font-size: 11px;" -_CSS_WARN = "color: #e07000; font-size: 11px;" -_CSS_BUSY = "color: #4a90d9; font-size: 11px;" - -# Side-panel column widths (the scroll areas that hold the controls). -_LEFT_PANEL_W = 288 -_RIGHT_PANEL_W = 272 - -# Minimum width for a spin box so its value + suffix are never clipped. -_SPIN_MIN_W = 96 - - -# --------------------------------------------------------------------------- -# Small layout helpers -# --------------------------------------------------------------------------- - -def _wrap_label(text: str = "", css: str | None = None) -> QLabel: - """A word-wrapped QLabel that reports its *wrapped* height to the layout. - - A plain word-wrapped QLabel advertises a single-line minimum height, so in a - fixed-width column the layout happily shrinks it and the extra lines get - clipped. Enabling height-for-width makes the box layout ask for the real - height at the column's width instead. - """ - lbl = QLabel(text) - lbl.setWordWrap(True) - sp = lbl.sizePolicy() - sp.setVerticalPolicy(QSizePolicy.Policy.Minimum) - sp.setHeightForWidth(True) - lbl.setSizePolicy(sp) - if css: - lbl.setStyleSheet(css) - return lbl - - -def _group(title: str) -> tuple[QGroupBox, QVBoxLayout]: - """A group box with consistent, non-cramped internal margins.""" - grp = QGroupBox(title) - lay = QVBoxLayout(grp) - lay.setContentsMargins(10, 8, 10, 10) - lay.setSpacing(6) - return grp, lay - - -def _form() -> QFormLayout: - """A label/field form layout for a narrow side panel.""" - form = QFormLayout() - form.setContentsMargins(0, 0, 0, 0) - form.setHorizontalSpacing(8) - form.setVerticalSpacing(6) - form.setLabelAlignment(Qt.AlignmentFlag.AlignRight - | Qt.AlignmentFlag.AlignVCenter) - form.setFormAlignment(Qt.AlignmentFlag.AlignLeft | Qt.AlignmentFlag.AlignTop) - form.setFieldGrowthPolicy( - QFormLayout.FieldGrowthPolicy.AllNonFixedFieldsGrow) - form.setRowWrapPolicy(QFormLayout.RowWrapPolicy.DontWrapRows) - return form - - -def _scroll_panel(inner: QWidget, width: int) -> QScrollArea: - """Put a side panel in a fixed-width scroll area. - - Without this the panels are sized by the window: a short window squeezes the - controls past their minimum heights, which is what makes text overlap the - widget below it. Scrolling keeps every control at its natural size. - """ - area = QScrollArea() - area.setWidget(inner) - area.setWidgetResizable(True) - area.setFrameShape(QFrame.Shape.NoFrame) - area.setHorizontalScrollBarPolicy(Qt.ScrollBarPolicy.ScrollBarAlwaysOff) - area.setVerticalScrollBarPolicy(Qt.ScrollBarPolicy.ScrollBarAsNeeded) - area.setFixedWidth(width) - area.viewport().setAutoFillBackground(False) - inner.setAutoFillBackground(False) - return area - - -# --------------------------------------------------------------------------- -# 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): - 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 - - # 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: str, - vmin: float, vmax: float, xlabel: str, ylabel: str, title: str, - colorbar_label: str = ""): - 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 - - im = self.ax.imshow( - img, aspect="auto", origin="upper", - extent=extent, cmap=cmap, vmin=vmin, vmax=vmax, - interpolation="nearest", - ) - cb = self.figure.colorbar(im, ax=self.ax, fraction=0.046, pad=0.04) - 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] - - self._draw_roi() - self.draw_idle() - - 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() - - -# --------------------------------------------------------------------------- -# FFT Options dialog -# --------------------------------------------------------------------------- - -class FftOptionsDialog(QDialog): - """Configure FFT backend and zero-padding. - - Changes take effect only when the user clicks Apply. Cancel discards - all pending edits. The live 'frequency resolution' label updates as - the user adjusts the pad factor so they can see the trade-off before - committing. - """ - - def __init__(self, parent=None, *, - current_backend: str, - current_pad_factor: int, - samples_per_frame: int | None, - sample_rate_hz: float | None, - grating_um: float): - super().__init__(parent) - self.setWindowTitle("FFT Options") - self.setModal(True) - self.setMinimumWidth(380) - - self._samples_per_frame = samples_per_frame - self._sample_rate_hz = sample_rate_hz - self._grating_um = grating_um - - layout = QVBoxLayout(self) - - # ---- Backend --------------------------------------------------- - grp_backend = QGroupBox("FFT Backend") - bl = QVBoxLayout(grp_backend) - - self._btn_scipy = QRadioButton("SciPy FFT (pocketfft) (always available)") - self._btn_pyfftw = QRadioButton( - "pyFFTW (faster for large arrays)" if PYFFTW_AVAILABLE - else "pyFFTW (not installed — run: pip install pyfftw)") - self._btn_pyfftw.setEnabled(PYFFTW_AVAILABLE) - - self._backend_group = QButtonGroup(self) - self._backend_group.addButton(self._btn_scipy, id=0) - self._backend_group.addButton(self._btn_pyfftw, id=1) - - if current_backend == "pyfftw" and PYFFTW_AVAILABLE: - self._btn_pyfftw.setChecked(True) - else: - self._btn_scipy.setChecked(True) - - bl.addWidget(self._btn_scipy) - bl.addWidget(self._btn_pyfftw) - layout.addWidget(grp_backend) - - # ---- Zero-padding ---------------------------------------------- - grp_zp = QGroupBox("Zero-Padding") - zl = QVBoxLayout(grp_zp) - - pad_row = QHBoxLayout() - pad_row.addWidget(QLabel("Pad factor:")) - self._spin_pad = QSpinBox() - self._spin_pad.setRange(1, 256) - self._spin_pad.setValue(max(1, current_pad_factor)) - self._spin_pad.setToolTip( - "Multiply the waveform length by this factor via zero-padding\n" - "before computing the FFT.\n" - "1 = no padding (natural length).\n" - "Powers of 2 (2, 4, 8 …) give the best performance." - ) - self._spin_pad.valueChanged.connect(self._update_info) - pad_row.addWidget(self._spin_pad) - zl.addLayout(pad_row) - - self._lbl_nfft = QLabel() - self._lbl_freq_res = QLabel() - self._lbl_vel_res = QLabel() - for lbl in (self._lbl_nfft, self._lbl_freq_res, self._lbl_vel_res): - lbl.setStyleSheet(_CSS_HINT) - zl.addWidget(lbl) - - layout.addWidget(grp_zp) - - # ---- Buttons --------------------------------------------------- - buttons = QDialogButtonBox() - buttons.addButton("Apply", QDialogButtonBox.ButtonRole.AcceptRole - ).clicked.connect(self.accept) - buttons.addButton("Cancel", QDialogButtonBox.ButtonRole.RejectRole - ).clicked.connect(self.reject) - layout.addWidget(buttons) - - self._update_info() - - def _update_info(self): - spf = self._samples_per_frame - sr = self._sample_rate_hz - pad = self._spin_pad.value() - - if spf is None or sr is None: - self._lbl_nfft.setText("Load a file to preview FFT parameters.") - self._lbl_freq_res.setText("") - self._lbl_vel_res.setText("") - return - - n_fft = spf * pad - freq_res_hz = sr / n_fft - freq_res_mhz = freq_res_hz / 1e6 - # v (m/s) = freq (MHz) × grating (µm) - vel_res_ms = freq_res_mhz * self._grating_um - - self._lbl_nfft.setText(f"FFT points: {spf} × {pad} = {n_fft:,}") - self._lbl_freq_res.setText( - f"Frequency bin: {freq_res_mhz:.4f} MHz ({freq_res_hz / 1e3:.2f} kHz)") - self._lbl_vel_res.setText( - f"Velocity bin: {vel_res_ms:.3f} m/s " - f"(at grating = {self._grating_um:.2f} µm)") - - def get_backend(self) -> str: - return "pyfftw" if self._btn_pyfftw.isChecked() and PYFFTW_AVAILABLE else "scipy" - - def get_pad_factor(self) -> int: - return max(1, self._spin_pad.value()) - - -# --------------------------------------------------------------------------- -# Manual alignment dialog (Fusion -> Manual Alignment...) -# --------------------------------------------------------------------------- - -class ManualAlignOverlayCanvas(FigureCanvasQTAgg): - """Renders ManualAlignmentDialog's multi-angle mask overlay and turns - keyboard input into translate/rotate nudge requests for whichever angle - the dialog currently has active. - - A pure input+render widget — it holds no alignment state and never - touches SrasFile itself; ManualAlignmentDialog 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.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) - - -class ManualAlignmentDialog(QDialog): - """Non-modal manual angle-alignment editor (Fusion -> Manual Alignment...). - - Shows every angle's binarized CH4 (Bias B) mask overlaid in a distinct - color at partial opacity on one shared canvas, so translation/rotation - misalignment is visible by eye. Reference angle (always index 0) is - ground truth and never moves; every other angle is aligned to it. The - user picks an "active" angle and nudges its rotation+translation with - the keyboard; Auto Cross-Correlate finds every non-reference angle's - rotation *and* translation by registering its image against the - reference's (see compute.register_angle_to_reference) — meant to get every - angle stacked on top of each other so keyboard nudging only has to make - small corrections, not find an alignment from scratch; Auto De-rotate is - the weaker fallback that just seeds rotation from the stage's reported - angle, leaving translation alone. Save writes a JSON sidecar next to the - .sras file and hands a freshly-built, full-resolution AlignmentResult back - to the main window — the exact same object shape compute_angle_alignment - produces, so every existing Aligned-View code path (apply_alignment, - _aligned_canvas_axes, the pixel-inspector inverse-transform) works - completely unmodified. - - Non-modal by design (shown via .show(), never .exec() or setModal(True)) - so the user can still interact with the main window. Talks back to - SrasViewerWindow two ways: it reuses parent._run_worker/_jobs directly - for its background mask-fetch and cross-correlate steps, so the main - window's existing shutdown/lifecycle plumbing covers both for free, and - it emits alignment_saved / alignment_cleared signals for the two moments - that should actually mutate the main window's persistent state — - everything else (nudging, Auto De-rotate, Auto Cross-Correlate, threshold - edits) stays purely local to this dialog until Save. - """ - - alignment_saved = pyqtSignal(object, str) # AlignmentResult, sidecar path (str) - alignment_cleared = pyqtSignal() - - _PREVIEW_MARGIN_FRAC = 0.15 - _BASE_ALPHA = 0.42 - _ACTIVE_ALPHA = 0.75 - _MAX_PREVIEW_DIM = 1024 - - # (label, sources passed to compute.register_angle_to_reference). "Both" - # registers on each and keeps whichever scores higher per angle, which - # costs roughly double but removes the failure mode where the single - # chosen source is the one that happens to be uninformative for one angle. - _CORRELATE_SOURCES = ( - ("Both, keep best (recommended)", ("signal", "mask")), - ("Raw signal", ("signal",)), - ("Thresholded mask", ("mask",)), - ) - - def __init__(self, parent: "SrasViewerWindow", sras: SrasFile, *, - ref_angle_idx: int, dc_threshold_mv: float, - seed_per_angle: dict[int, ManualAngleParams] | None, - cached_dc4_mv: dict[int, np.ndarray]): - super().__init__(parent) - self._parent = parent - self._sras = sras - self._ref_angle_idx = ref_angle_idx - self._downsample = (1, 1) # (rows, cols) block-mean factors - self._dc4_mv: dict[int, np.ndarray] = {} - self._masks_small: dict[int, np.ndarray] = {} - self._preview_layers: dict[int, np.ndarray] = {} - self._preview_origin_mm = (0.0, 0.0) - self._preview_shape = (1, 1) - self._preview_pitch_mm = (1.0, 1.0) - self._masks_ready = False - self._fit_notes: dict[int, tuple[float, str]] = {} - self._derotate_sign_flipped = False - - self.setWindowTitle(f"Manual Alignment — {sras.path.name}") - self.resize(1150, 760) - - self._seed_initial_params(seed_per_angle) - n = sras.n_angles - cmap = mpl.colormaps["tab10"] if n <= 10 else mpl.colormaps["tab20"] - self._angle_colors = {a: cmap(a % cmap.N)[:3] for a in range(n)} - self._active_angle = 1 if ref_angle_idx == 0 and n > 1 else 0 - - self._build_ui(dc_threshold_mv) - self._set_controls_enabled(False) # re-enabled once masks are ready - self._start_mask_prep(cached_dc4_mv) - - def showEvent(self, event): - super().showEvent(event) - self.canvas.setFocus() - - # ------------------------------------------------------------------ - # Construction - # ------------------------------------------------------------------ - - def _seed_initial_params(self, seed_per_angle: dict[int, ManualAngleParams] | None): - seed = seed_per_angle or {} - self._angle_params: dict[int, ManualAngleParams] = { - a: (ManualAngleParams(seed[a].rotation_deg, seed[a].shift_mm) - if a in seed else ManualAngleParams()) - for a in range(self._sras.n_angles) - } - self._angle_params[self._ref_angle_idx] = ManualAngleParams() - - def _build_ui(self, dc_threshold_mv: float): - root = QHBoxLayout(self) - - self.canvas = ManualAlignOverlayCanvas() - left = QWidget() - left_l = QVBoxLayout(left) - left_l.setContentsMargins(0, 0, 0, 0) - left_l.setSpacing(4) - left_l.addWidget(NavigationToolbar2QT(self.canvas, left)) - left_l.addWidget(self.canvas) - root.addWidget(left, stretch=1) - - panel = QWidget() - panel_l = QVBoxLayout(panel) - panel_l.setContentsMargins(0, 0, 0, 0) - panel_l.setSpacing(8) - - # ---- Active Angle ------------------------------------------------- - grp_angle, al = _group("Active Angle") - self.combo_active_angle = QComboBox() - for a in range(self._sras.n_angles): - label = f"Angle {a} ({self._sras.angles_deg[a]:.1f}°)" - if a == self._ref_angle_idx: - label += " [reference]" - self.combo_active_angle.addItem(label) - al.addWidget(self.combo_active_angle) - self.lbl_active_note = _wrap_label("", _CSS_WARN) - al.addWidget(self.lbl_active_note) - panel_l.addWidget(grp_angle) - - # ---- Manual Adjustment --------------------------------------------- - self.grp_manual_adjust, mform_box = _group("Manual Adjustment") - mform = _form() - self.spin_active_rotation_deg = QDoubleSpinBox() - self.spin_active_rotation_deg.setRange(-3600.0, 3600.0) - self.spin_active_rotation_deg.setDecimals(3) - self.spin_active_rotation_deg.setSuffix(" °") - self.spin_active_rotation_deg.setMinimumWidth(_SPIN_MIN_W) - mform.addRow("Rotation:", self.spin_active_rotation_deg) - - self.spin_active_shift_x_mm = QDoubleSpinBox() - self.spin_active_shift_x_mm.setRange(-1e5, 1e5) - self.spin_active_shift_x_mm.setDecimals(4) - self.spin_active_shift_x_mm.setSuffix(" mm") - self.spin_active_shift_x_mm.setMinimumWidth(_SPIN_MIN_W) - mform.addRow("Shift X:", self.spin_active_shift_x_mm) - - self.spin_active_shift_y_mm = QDoubleSpinBox() - self.spin_active_shift_y_mm.setRange(-1e5, 1e5) - self.spin_active_shift_y_mm.setDecimals(4) - self.spin_active_shift_y_mm.setSuffix(" mm") - self.spin_active_shift_y_mm.setMinimumWidth(_SPIN_MIN_W) - mform.addRow("Shift Y:", self.spin_active_shift_y_mm) - mform_box.addLayout(mform) - panel_l.addWidget(self.grp_manual_adjust) - - # ---- Nudge Step Sizes ------------------------------------------------ - self.grp_step_sizes, sl = _group("Nudge Step Sizes") - sform = _form() - self.spin_step_translate_mm = QDoubleSpinBox() - self.spin_step_translate_mm.setRange(0.0001, 1000.0) - self.spin_step_translate_mm.setDecimals(4) - self.spin_step_translate_mm.setSuffix(" mm") - self.spin_step_translate_mm.setValue(0.01) - self.spin_step_translate_mm.setMinimumWidth(_SPIN_MIN_W) - sform.addRow("Translate step:", self.spin_step_translate_mm) - - self.spin_step_rotate_deg = QDoubleSpinBox() - self.spin_step_rotate_deg.setRange(0.001, 90.0) - self.spin_step_rotate_deg.setDecimals(3) - self.spin_step_rotate_deg.setSuffix(" °") - self.spin_step_rotate_deg.setValue(0.1) - self.spin_step_rotate_deg.setMinimumWidth(_SPIN_MIN_W) - sform.addRow("Rotate step:", self.spin_step_rotate_deg) - - self.spin_step_multiplier = QDoubleSpinBox() - self.spin_step_multiplier.setRange(1.0, 1000.0) - self.spin_step_multiplier.setDecimals(1) - self.spin_step_multiplier.setValue(10.0) - self.spin_step_multiplier.setMinimumWidth(_SPIN_MIN_W) - sform.addRow("Coarse × (Shift):", self.spin_step_multiplier) - sl.addLayout(sform) - sl.addWidget(_wrap_label( - "Arrow keys nudge X/Y translation; Q/E nudge rotation (CCW/CW). " - "Hold Shift for the coarse step. Click the image once so it has " - "keyboard focus.", _CSS_HINT)) - panel_l.addWidget(self.grp_step_sizes) - - # ---- Mask Threshold --------------------------------------------------- - self.grp_mask_threshold, tl = _group("Mask Threshold") - tform = _form() - self.spin_mask_threshold_mv = QDoubleSpinBox() - self.spin_mask_threshold_mv.setRange(-500.0, 500.0) - self.spin_mask_threshold_mv.setDecimals(3) - self.spin_mask_threshold_mv.setSuffix(" mV") - self.spin_mask_threshold_mv.setValue(dc_threshold_mv) - self.spin_mask_threshold_mv.setMinimumWidth(_SPIN_MIN_W) - tform.addRow("DC threshold:", self.spin_mask_threshold_mv) - tl.addLayout(tform) - panel_l.addWidget(self.grp_mask_threshold) - - # ---- Cross-Correlate (FFT) ----------------------------------------- - self.grp_correlate, cl = _group("Cross-Correlate (FFT)") - cform = _form() - self.combo_correlate_source = QComboBox() - for label, sources in self._CORRELATE_SOURCES: - self.combo_correlate_source.addItem(label, sources) - cform.addRow("Correlate on:", self.combo_correlate_source) - - self.spin_correlate_search_deg = QDoubleSpinBox() - self.spin_correlate_search_deg.setRange(0.0, 180.0) - self.spin_correlate_search_deg.setSingleStep(1.0) - self.spin_correlate_search_deg.setDecimals(1) - self.spin_correlate_search_deg.setSuffix(" °") - self.spin_correlate_search_deg.setValue(6.0) - self.spin_correlate_search_deg.setMinimumWidth(_SPIN_MIN_W) - cform.addRow("Rotation search (±):", self.spin_correlate_search_deg) - cl.addLayout(cform) - self.btn_auto_correlate = QPushButton("Auto Cross-Correlate (vs Reference)") - cl.addWidget(self.btn_auto_correlate) - cl.addWidget(_wrap_label( - "Finds each non-reference angle's rotation *and* translation by " - "cross-correlating its image against the reference's — the stage's " - "reported angle is only the starting point of the search, and both " - "of its signs are tried. Run this first, then nudge only for small " - "corrections.", _CSS_HINT)) - panel_l.addWidget(self.grp_correlate) - - # ---- Actions ------------------------------------------------------ - grp_actions, acl = _group("Actions") - self.btn_auto_derotate = QPushButton("Auto De-rotate (use known angles)") - self.btn_save = QPushButton("Save Alignment") - self.btn_clear = QPushButton("Clear Alignment…") - self.btn_close = QPushButton("Close") - for btn in (self.btn_auto_derotate, self.btn_save, self.btn_clear, self.btn_close): - acl.addWidget(btn) - panel_l.addWidget(grp_actions) - - self.lbl_status = _wrap_label("", _CSS_MUTED) - panel_l.addWidget(self.lbl_status) - panel_l.addStretch() - - root.addWidget(_scroll_panel(panel, 320)) - - self.combo_active_angle.currentIndexChanged.connect(self._on_active_angle_changed) - self.spin_active_rotation_deg.editingFinished.connect(self._on_rotation_spin_edited) - self.spin_active_shift_x_mm.editingFinished.connect(self._on_shift_spin_edited) - self.spin_active_shift_y_mm.editingFinished.connect(self._on_shift_spin_edited) - self.spin_mask_threshold_mv.editingFinished.connect(self._on_mask_threshold_edited) - self.btn_auto_derotate.clicked.connect(self._on_auto_derotate) - self.btn_auto_correlate.clicked.connect(self._on_auto_correlate) - self.btn_save.clicked.connect(self._on_save) - self.btn_clear.clicked.connect(self._on_clear) - self.btn_close.clicked.connect(self.close) - self.canvas.nudge_translate.connect(self._on_nudge_translate) - self.canvas.nudge_rotate.connect(self._on_nudge_rotate) - - self.combo_active_angle.blockSignals(True) - self.combo_active_angle.setCurrentIndex(self._active_angle) - self.combo_active_angle.blockSignals(False) - self._on_active_angle_changed(self._active_angle) - - # ------------------------------------------------------------------ - # Mask preparation (initial CH4 fetch + threshold + downsample) - # ------------------------------------------------------------------ - - def _start_mask_prep(self, cached_dc4_mv: dict[int, np.ndarray]): - self._dc4_mv = dict(cached_dc4_mv) - missing = [a for a in range(self._sras.n_angles) if a not in self._dc4_mv] - if not missing: - self._finish_mask_prep() - return - self.lbl_status.setText(f"Preparing masks: 0/{len(missing)} angle(s) needed…") - started = self._parent._run_worker( - "manual_align_masks", Ch4MaskWorker(self._sras, missing), - connect=( - ("angle_done", self._on_mask_angle_done), - ("error", lambda msg: self.lbl_status.setText(f"Mask prep error: {msg}")), - ), - on_done=self._finish_mask_prep) - if not started: - self.lbl_status.setText( - "Could not start mask preparation (busy) — close and reopen.") - - def _on_mask_angle_done(self, angle_idx: int, dc4_mv: np.ndarray): - self._dc4_mv[angle_idx] = dc4_mv - self.lbl_status.setText( - f"Preparing masks: {len(self._dc4_mv)}/{self._sras.n_angles} ready…") - - def _finish_mask_prep(self): - if len(self._dc4_mv) < self._sras.n_angles: - return # a mask-worker error left some angles unfetched - # Rows and columns get their own factor. A real scan is ~7500 frames - # wide but only ~750 rows tall, so one shared factor sized for the - # frames would throw away 8x more row detail than the preview needs and - # leave the overlay too coarse in y to judge alignment by eye. - max_rows = max(img.shape[0] for img in self._dc4_mv.values()) - max_cols = max(img.shape[1] for img in self._dc4_mv.values()) - self._downsample = ( - max(1, int(np.ceil(max_rows / self._MAX_PREVIEW_DIM))), - max(1, int(np.ceil(max_cols / self._MAX_PREVIEW_DIM)))) - self._recompute_masks_small() - self._rebuild_preview_canvas() - self._set_controls_enabled(True) - self.lbl_status.setText("Ready.") - - def _recompute_masks_small(self): - """Threshold + downsample every angle's already-in-memory full-res - CH4 mV image. Cheap (a compare + block-mean), so this re-runs in - full whenever the mask-threshold spin box changes — no re-fetch. - Purely for the overlay's visuals: no alignment geometry depends on this - threshold, only which pixels the overlay paints.""" - threshold = self.spin_mask_threshold_mv.value() - fy, fx = self._downsample - self._masks_small = { - a: compute.block_mean_2d((img >= threshold).astype(np.float32), fy, fx) - for a, img in self._dc4_mv.items() - } - - # ------------------------------------------------------------------ - # Preview canvas: full rebuild vs. incremental single-layer refresh - # ------------------------------------------------------------------ - - def _rebuild_preview_canvas(self): - """Full geometry rebuild: recomputes the shared preview canvas's - origin/shape (rotation can grow the union bbox — translation alone - cannot, per the padding baked in via _PREVIEW_MARGIN_FRAC) and every - angle's reprojected mask layer. Triggered by: dialog open, - mask-threshold change, Auto De-rotate, a rotation nudge/edit of the - active angle. NOT triggered by a translation-only nudge — see - _refresh_active_preview_layer.""" - dx_ref, dy_ref = compute.pixel_pitch_mm(self._sras, self._ref_angle_idx) - fy, fx = self._downsample - pitch = (dx_ref * fx, dy_ref * fy) - origin, shape = compute.canvas_for_params( - self._sras, self._ref_angle_idx, pitch, self._angle_params, - margin_frac=self._PREVIEW_MARGIN_FRAC, snap=False) - self._preview_origin_mm, self._preview_shape = origin, shape - self._preview_pitch_mm = pitch - self._preview_layers = { - a: self._reproject(a) for a in range(self._sras.n_angles) - } - self._redraw_overlay() - - def _reproject(self, angle_idx: int) -> np.ndarray: - """One angle's downsampled mask on the current preview canvas. - src_downsample must match _masks_small's block-mean factors, or the - layer lands magnified and offset instead of where the alignment - actually puts it.""" - p = self._angle_params[angle_idx] - return compute.reproject_mask( - self._sras, angle_idx, self._ref_angle_idx, - self._masks_small[angle_idx], p.rotation_deg, p.shift_mm, - self._preview_pitch_mm, self._preview_origin_mm, self._preview_shape, - src_downsample=self._downsample) - - def _refresh_active_preview_layer(self): - """Cheap path for a translation-only nudge/edit of the active angle: - reproject just that one angle's downsampled mask onto the *existing* - preview canvas — every other angle's cached layer is untouched.""" - self._preview_layers[self._active_angle] = self._reproject(self._active_angle) - self._redraw_overlay() - - def _redraw_overlay(self): - """Alpha-composite every angle's colored mask layer into one RGBA - image ("all thresholds overlaid with varying opacity"). Each angle - keeps a fixed, distinct color regardless of which is active; the - active angle is drawn last (on top) at a visibly higher alpha so - it's easy to track while nudging.""" - if not self._preview_layers: - return # mask prep hasn't finished yet — nothing to draw - n_rows, n_cols = self._preview_shape - rgba = np.zeros((n_rows, n_cols, 4), dtype=np.float32) - order = sorted(range(self._sras.n_angles), key=lambda a: a == self._active_angle) - for a in order: - layer = self._preview_layers.get(a) - if layer is None: - continue - alpha = self._ACTIVE_ALPHA if a == self._active_angle else self._BASE_ALPHA - color = self._angle_colors[a] - fg_a = layer * alpha - for c in range(3): - rgba[..., c] = color[c] * fg_a + rgba[..., c] * rgba[..., 3] * (1 - fg_a) - rgba[..., 3] = fg_a + rgba[..., 3] * (1 - fg_a) - - x0, y0 = self._preview_origin_mm - dx, dy = self._preview_pitch_mm - x_axis = x0 + np.arange(n_cols) * dx - y_axis = y0 + np.arange(n_rows) * dy - extent = [x_axis[0] - dx / 2, x_axis[-1] + dx / 2, - y_axis[-1] + dy / 2, y_axis[0] - dy / 2] - title = (f"Angle {self._active_angle} active " - f"({self._sras.angles_deg[self._active_angle]:.1f}°)") - self.canvas.show_overlay(rgba, extent, title) - - # ------------------------------------------------------------------ - # Angle selection / nudge / edit handlers - # ------------------------------------------------------------------ - - def _on_active_angle_changed(self, angle_idx: int): - self._active_angle = angle_idx - is_ref = angle_idx == self._ref_angle_idx - self.grp_manual_adjust.setEnabled(self._masks_ready and not is_ref) - self.lbl_active_note.setText( - "Reference angle — defines the shared origin, not adjustable." if is_ref else "") - self._sync_active_spinboxes() - self._redraw_overlay() - - def _sync_active_spinboxes(self): - p = self._angle_params[self._active_angle] - for spin, val in ((self.spin_active_rotation_deg, p.rotation_deg), - (self.spin_active_shift_x_mm, p.shift_mm[0]), - (self.spin_active_shift_y_mm, p.shift_mm[1])): - spin.blockSignals(True) - spin.setValue(val) - spin.blockSignals(False) - - def _on_nudge_translate(self, dir_x: int, dir_y: int, coarse: bool): - if not self._masks_ready or self._active_angle == self._ref_angle_idx: - return - step = self.spin_step_translate_mm.value() - if coarse: - step *= self.spin_step_multiplier.value() - p = self._angle_params[self._active_angle] - p.shift_mm = (p.shift_mm[0] + dir_x * step, p.shift_mm[1] + dir_y * step) - self._sync_active_spinboxes() - self._refresh_active_preview_layer() - - def _on_nudge_rotate(self, direction: int, coarse: bool): - if not self._masks_ready or self._active_angle == self._ref_angle_idx: - return - step = self.spin_step_rotate_deg.value() - if coarse: - step *= self.spin_step_multiplier.value() - self._angle_params[self._active_angle].rotation_deg += direction * step - self._sync_active_spinboxes() - self._rebuild_preview_canvas() - - def _on_rotation_spin_edited(self): - if self._active_angle == self._ref_angle_idx: - return - self._angle_params[self._active_angle].rotation_deg = self.spin_active_rotation_deg.value() - self._rebuild_preview_canvas() - - def _on_shift_spin_edited(self): - if self._active_angle == self._ref_angle_idx: - return - p = self._angle_params[self._active_angle] - p.shift_mm = (self.spin_active_shift_x_mm.value(), self.spin_active_shift_y_mm.value()) - self._refresh_active_preview_layer() - - def _on_mask_threshold_edited(self): - if not self._masks_ready: - return - self._recompute_masks_small() - self._rebuild_preview_canvas() - - # ------------------------------------------------------------------ - # Actions - # ------------------------------------------------------------------ - - def _on_auto_derotate(self): - """Seed every angle's rotation from the stage's reported angle. - - A starting point for nudging by eye, not an alignment: the stage's - sign convention relative to this module's is not knowable from the - file, so the sign that lines the scans up is whichever of the two looks - right in the overlay. Auto Cross-Correlate decides that from the images - instead, and is the button to reach for first. - """ - sign = -1.0 if self._derotate_sign_flipped else 1.0 - self._derotate_sign_flipped = not self._derotate_sign_flipped - n_changed = 0 - for a in range(self._sras.n_angles): - if a == self._ref_angle_idx: - continue - self._angle_params[a].rotation_deg = sign * compute.nominal_delta_deg( - self._sras, a, self._ref_angle_idx) - n_changed += 1 - self._sync_active_spinboxes() - self._rebuild_preview_canvas() - self.lbl_status.setText( - f"Rotation set to the stage angle ({'−' if sign < 0 else '+'}delta) " - f"for {n_changed} angle(s); translation untouched. Click again to " - "try the opposite sign.") - - def _on_auto_correlate(self): - if not self._masks_ready: - return - angles = [a for a in range(self._sras.n_angles) if a != self._ref_angle_idx] - if not angles: - return - worker = CrossCorrelateWorker( - self._sras, self._ref_angle_idx, angles, self._dc4_mv, - sources=self.combo_correlate_source.currentData(), - dc_threshold_mv=self.spin_mask_threshold_mv.value(), - search_deg=self.spin_correlate_search_deg.value()) - self._correlate_done_count = 0 - self._correlate_total = len(angles) - self._fit_notes = {} - self._set_controls_enabled(False) - self.lbl_status.setText(f"Cross-correlating: 0/{self._correlate_total} angle(s)…") - started = self._parent._run_worker( - "manual_align_correlate", worker, - connect=( - ("angle_done", self._on_correlate_angle_done), - ("error", self._on_correlate_error), - ), - on_done=self._finish_auto_correlate) - if not started: - self._set_controls_enabled(True) - self.lbl_status.setText("Could not start cross-correlation (busy) — try again.") - - def _on_correlate_angle_done(self, angle_idx: int, rotation_deg: float, - shift_x_mm: float, shift_y_mm: float, - score: float, source: str): - self._angle_params[angle_idx] = ManualAngleParams(rotation_deg, (shift_x_mm, shift_y_mm)) - self._fit_notes[angle_idx] = (score, source) - self._correlate_done_count += 1 - self.lbl_status.setText( - f"Cross-correlating: {self._correlate_done_count}/{self._correlate_total} angle(s)…") - - def _on_correlate_error(self, msg: str): - self.lbl_status.setText(f"Cross-correlation error: {msg}") - - def _finish_auto_correlate(self): - self._sync_active_spinboxes() - self._rebuild_preview_canvas() - self._set_controls_enabled(True) - self.lbl_status.setText( - f"Cross-correlated {self._correlate_done_count} angle(s) against " - f"Angle {self._ref_angle_idx}.\n" + self._fit_report()) - - def _fit_report(self) -> str: - """Per-angle registration quality, worst first. - - Surfaced rather than buried because a single bad acquisition (stage - glitch, laser dropout) registers poorly and would otherwise be fused in - silently — seeing which angle it is, is what makes dropping it with - sras_edit_scans.py actionable. The deviation from the stage's own - reported angle is shown alongside: a large one means the search and the - stage disagree, which is either a genuine mechanical error or a sign - that this angle's fit is not to be trusted. - """ - if not self._fit_notes: - return "" - rows = sorted(self._fit_notes.items(), key=lambda kv: kv[1][0]) - worst = rows[0] - lines = [f"Worst fit: angle {worst[0]} (score {worst[1][0]:.3f}, " - f"{worst[1][1]})."] - drifted = [] - for a, _note in rows: - nominal = compute.nominal_delta_deg(self._sras, a, self._ref_angle_idx) - got = self._angle_params[a].rotation_deg - dev = min(abs(got - nominal), abs(got + nominal)) - if dev > 1.0: - drifted.append(f"{a} ({dev:.2f}°)") - if drifted: - lines.append("Rotation differs from the stage angle by >1° for " - "angle(s) " + ", ".join(drifted) + ".") - lines.append("Nudge from here for any remaining fine correction.") - return " ".join(lines) - - def _on_save(self): - threshold = self.spin_mask_threshold_mv.value() - resolved = dict(self._angle_params) # already concrete floats - try: - path = save_manual_alignment(self._sras, self._ref_angle_idx, threshold, resolved) - result = build_manual_alignment(self._sras, self._ref_angle_idx, - threshold, resolved) - except OSError as exc: - QMessageBox.warning(self, "Save Alignment Failed", str(exc)) - return - self.lbl_status.setText(f"Saved to {path.name}.") - self.alignment_saved.emit(result, str(path)) - - def _on_clear(self): - reply = QMessageBox.question( - self, "Clear Alignment", - "This resets every angle back to raw/unaligned (0° rotation, no " - "shift) and deletes the saved alignment file for this scan, if " - "any. This cannot be undone. Continue?", - QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No, - QMessageBox.StandardButton.No) - if reply != QMessageBox.StandardButton.Yes: - return - try: - existed = delete_manual_alignment(self._sras) - except OSError as exc: - QMessageBox.warning(self, "Clear Alignment Failed", - f"Could not delete the saved alignment file: {exc}") - return - self._angle_params = {a: ManualAngleParams() for a in range(self._sras.n_angles)} - self._fit_notes = {} - self._sync_active_spinboxes() - self._rebuild_preview_canvas() - self.lbl_status.setText( - "Alignment cleared; saved file removed." if existed - else "Alignment cleared (there was no saved file).") - self.alignment_cleared.emit() - - def _set_controls_enabled(self, enabled: bool): - self._masks_ready = enabled - self.combo_active_angle.setEnabled(enabled) - self.grp_manual_adjust.setEnabled(enabled and self._active_angle != self._ref_angle_idx) - self.grp_step_sizes.setEnabled(enabled) - self.grp_mask_threshold.setEnabled(enabled) - self.grp_correlate.setEnabled(enabled) - self.btn_auto_derotate.setEnabled(enabled) - self.btn_save.setEnabled(enabled) - self.btn_clear.setEnabled(enabled) - - -# --------------------------------------------------------------------------- -# Main window -# --------------------------------------------------------------------------- - -class SrasViewerWindow(QMainWindow): - def __init__(self, initial_path: str | None = None): - super().__init__() - self.setWindowTitle("SRAS Scan Viewer") - self.resize(1560, 840) - self.setMinimumSize(960, 560) - self.setAcceptDrops(True) - - self._sras: SrasFile | None = None - self._current_image: np.ndarray | None = None - self._current_angle: int = 0 - self._current_ch: int = 0 - self._pending_angle: int = 0 - self._pending_ch: int = 0 - self._pending_bg_sub: bool = True - self._pending_threshold: float = 50.0 # mV - self._pending_fft_pad_factor: int = 1 - - # Live background jobs, keyed by role — see _run_worker. - self._jobs: dict[str, tuple] = {} - self._progress_dlgs: dict[str, QProgressDialog] = {} - - # FFT settings (configured via FFT Options dialog, persisted across - # sessions). IniFormat: predictable cross-platform and redirectable - # in tests. - self._settings = QSettings(QSettings.Format.IniFormat, - QSettings.Scope.UserScope, - "sras-viewer", "sras-viewer") - compute.set_fft_backend(str(self._settings.value("fft/backend", "scipy"))) - try: - pad = int(self._settings.value("fft/pad_factor", 1)) - except (TypeError, ValueError): - pad = 1 - self._fft_pad_factor: int = max(1, min(256, pad)) # 1 = no padding - - # Convert menu: batch DC/FFT compute-and-store (v6 -> v7) - self._batch_errors: list[str] = [] - - # Display-only settings (colormap, grating) never trigger a - # recompute — they're applied to cached data on redraw. DC images - # (CH3/CH4) are cheap and precomputed for every angle in the - # background right after load. CH1/Velocity FFT images are - # computed lazily (with a progress popup) the first time an - # angle/threshold combination is viewed — using the cached DC4 - # image to skip the FFT entirely for masked-out pixels — and - # cached per (angle, bg_sub, n_fft, threshold) so revisiting the - # same combination is free. - self._dc_cache: dict[tuple[int, int], np.ndarray] = {} - self._fft_cache: dict[tuple[int, bool, int | None, float], np.ndarray] = {} - self._dc_generation: int = 0 - - # Angle alignment ("Fusion" menu) - self._alignment_result = None - self._alignment_generation: int = 0 - self._aligned_cache: dict[tuple, np.ndarray] = {} - self._manual_align_dialog: ManualAlignmentDialog | None = None - - self._build_ui() - - if initial_path: - self._load_file(initial_path) - - # ------------------------------------------------------------------ - # Background job plumbing - # ------------------------------------------------------------------ - - def _run_worker(self, key: str, worker: QObject, *, - connect: tuple = (), quit_on: tuple = ("finished",), - on_done=None) -> bool: - """Move *worker* onto its own QThread and start it. Returns False if - a job under *key* is already running. - - Centralises two lifetime hazards that each cost a process abort: - - 1. The job is claimed in self._jobs *before* start() and before - anything below that can pump the Qt event loop (a - QProgressDialog.show() does on first display). If it weren't, a - re-entrant editingFinished could slip past the busy check, start a - second thread, and then have the first call's own assignment - clobber — and destroy while still running — that second QThread. - - 2. thread.finished fires as the thread winds down but does not - guarantee the OS thread has joined. Dropping the last reference to - a QThread whose thread is still running logs "QThread: Destroyed - while thread is still running" and aborts, so wait() first. - """ - if key in self._jobs: - return False - - thread = QThread() - self._jobs[key] = (thread, worker, on_done) # claim before anything pumps - worker.moveToThread(thread) - thread.started.connect(worker.run) - for signal_name, slot in connect: - getattr(worker, signal_name).connect(slot) - for signal_name in quit_on: - getattr(worker, signal_name).connect(thread.quit) - thread.finished.connect(lambda k=key: self._on_job_finished(k)) - thread.start() - return True - - def _on_job_finished(self, key: str): - job = self._jobs.pop(key, None) - if job is None: - return - thread, _worker, on_done = job - thread.wait() # join before releasing our last reference - if on_done is not None: - on_done() - - def _job_running(self, key: str) -> bool: - return key in self._jobs - - # ------------------------------------------------------------------ - # UI construction - # ------------------------------------------------------------------ - - def _build_ui(self): - central = QWidget() - self.setCentralWidget(central) - root = QHBoxLayout(central) - root.setContentsMargins(8, 8, 8, 8) - root.setSpacing(8) - - root.addWidget(self._build_left_panel()) - root.addWidget(self._build_canvases(), stretch=1) - root.addWidget(self._build_right_panel()) - - self.statusBar().showMessage("Open an .sras file to begin.") - self._build_menus() - - def _build_left_panel(self) -> QWidget: - panel = QWidget() - panel_layout = QVBoxLayout(panel) - panel_layout.setContentsMargins(0, 0, 0, 0) - panel_layout.setSpacing(8) - - # ---- File ------------------------------------------------------- - grp_file, fl = _group("File") - self.btn_open = QPushButton("Open .sras…") - self.btn_open.clicked.connect(self._on_open) - self.lbl_filename = _wrap_label("No file loaded", _CSS_MUTED) - fl.addWidget(self.btn_open) - fl.addWidget(self.lbl_filename) - panel_layout.addWidget(grp_file) - - # ---- Scan info -------------------------------------------------- - grp_info, il = _group("Scan Info") - il.setSpacing(3) - self._info = {} - for key in ("Angles", "Rows", "Frames / row", "Samples / frame", - "Sample rate", "X start", "Pixel Δx", "Laser freq"): - lbl = _wrap_label(f"{key}: —", _CSS_INFO) - il.addWidget(lbl) - self._info[key] = lbl - - # frame-count / format notes - self.lbl_frame_warn = _wrap_label("", _CSS_WARN) - il.addWidget(self.lbl_frame_warn) - - # background DC-precompute progress - self.lbl_dc_precompute = _wrap_label("", _CSS_BUSY) - il.addWidget(self.lbl_dc_precompute) - panel_layout.addWidget(grp_info) - - # ---- View settings ---------------------------------------------- - grp_view, vl = _group("View Settings") - - view_form = _form() - - self.spin_angle = QSpinBox() - self.spin_angle.setRange(0, 0) - self.spin_angle.setEnabled(False) - self.spin_angle.setMinimumWidth(64) - self.spin_angle.editingFinished.connect(self._on_view_changed) - self.lbl_angle_deg = QLabel("—") - angle_field = QWidget() - ar = QHBoxLayout(angle_field) - ar.setContentsMargins(0, 0, 0, 0) - ar.setSpacing(6) - ar.addWidget(self.spin_angle) - ar.addWidget(self.lbl_angle_deg) - ar.addStretch() - view_form.addRow("Angle:", angle_field) - - self.combo_channel = QComboBox() - self.combo_channel.addItems(CH_LABELS) - self.combo_channel.setEnabled(False) - self.combo_channel.setSizePolicy(QSizePolicy.Policy.Expanding, - QSizePolicy.Policy.Fixed) - self.combo_channel.setSizeAdjustPolicy( - QComboBox.SizeAdjustPolicy.AdjustToMinimumContentsLengthWithIcon) - self.combo_channel.setMinimumContentsLength(12) - self.combo_channel.currentIndexChanged.connect(self._on_channel_changed) - view_form.addRow("Channel:", self.combo_channel) - vl.addLayout(view_form) - - sep = QFrame() - sep.setFrameShape(QFrame.Shape.HLine) - sep.setStyleSheet("color: #555;") - vl.addWidget(sep) - - # DC threshold (for RF / CH1 masking) - self.grp_threshold, tl = _group("RF Mask Threshold (CH1 only)") - thr_form = _form() - self.spin_threshold_mv = QDoubleSpinBox() - self.spin_threshold_mv.setRange(-500.0, 500.0) - self.spin_threshold_mv.setDecimals(3) - self.spin_threshold_mv.setSingleStep(0.025) - self.spin_threshold_mv.setSuffix(" mV") - self.spin_threshold_mv.setValue(50.0) - self.spin_threshold_mv.setEnabled(False) - self.spin_threshold_mv.setMinimumWidth(_SPIN_MIN_W) - self.spin_threshold_mv.editingFinished.connect(self._on_threshold_changed) - thr_form.addRow("DC threshold:", self.spin_threshold_mv) - tl.addLayout(thr_form) - self.lbl_threshold_adc = _wrap_label( - f"≈ {mv_to_adc(50.0):.1f} ADC counts", _CSS_MUTED) - tl.addWidget(self.lbl_threshold_adc) - vl.addWidget(self.grp_threshold) - - # Background subtraction (v4+ files only) - self.chk_bg_sub = QCheckBox("Background subtraction (CH1 only)") - self.chk_bg_sub.setChecked(True) - self.chk_bg_sub.setEnabled(False) - self.chk_bg_sub.setToolTip( - "Subtract the stored background waveform from each CH1 frame\n" - "before computing the FFT (v4+ files only)." - ) - self.chk_bg_sub.toggled.connect(self._on_bg_sub_toggled) - vl.addWidget(self.chk_bg_sub) - - # Aligned View (Fusion → Angle Alignment result) - self.chk_aligned_view = QCheckBox("Aligned View (Fusion)") - self.chk_aligned_view.setChecked(False) - self.chk_aligned_view.setEnabled(False) - self.chk_aligned_view.setToolTip( - "Show the current angle/channel resampled onto the shared,\n" - "rotation+translation-aligned canvas from Fusion → Angle\n" - "Alignment. Uncheck to see the raw per-angle scan grid." - ) - self.chk_aligned_view.toggled.connect(self._on_aligned_view_toggled) - vl.addWidget(self.chk_aligned_view) - - self.btn_export_csv = QPushButton("Export Image as CSV…") - self.btn_export_csv.setEnabled(False) - self.btn_export_csv.setToolTip( - "Save the current CH1 image (one scan row per CSV line).") - self.btn_export_csv.clicked.connect(self._on_export_csv) - vl.addWidget(self.btn_export_csv) - - panel_layout.addWidget(grp_view) - - # ---- ROI --------------------------------------------------------- - grp_roi, rl = _group("ROI (Region of Interest)") - - self.btn_draw_roi = QPushButton("Draw ROI") - self.btn_draw_roi.setCheckable(True) - self.btn_draw_roi.setEnabled(False) - self.btn_draw_roi.setToolTip( - "Arm next click+drag on the image to draw a new ROI\n" - "(replaces any existing one). Click again to cancel.\n" - "After drawing, drag inside to move, or grab corners to reshape.\n" - "The ROI is persistent across channels / modes / angles." - ) - self.btn_draw_roi.toggled.connect(self._on_draw_roi_toggled) - rl.addWidget(self.btn_draw_roi) - - self.btn_clear_roi = QPushButton("Clear ROI") - self.btn_clear_roi.setEnabled(False) - self.btn_clear_roi.clicked.connect(self._on_clear_roi) - rl.addWidget(self.btn_clear_roi) - - self.btn_export_roi = QPushButton("Export ROI as CSV…") - self.btn_export_roi.setEnabled(False) - self.btn_export_roi.setToolTip( - "Save every pixel whose centre lies inside the ROI as CSV.\n" - "Columns: row, frame, x_mm, y_mm, value.\n" - "Corner coordinates of the quad are written in the file header." - ) - self.btn_export_roi.clicked.connect(self._on_export_roi_csv) - rl.addWidget(self.btn_export_roi) - - self.lbl_roi_center = _wrap_label("centroid: —", _CSS_HINT) - self.lbl_roi_size = _wrap_label("bbox: —", _CSS_HINT) - self.lbl_roi_npix = _wrap_label("pixels inside: —", _CSS_HINT) - for lbl in (self.lbl_roi_center, self.lbl_roi_size, self.lbl_roi_npix): - rl.addWidget(lbl) - - panel_layout.addWidget(grp_roi) - panel_layout.addStretch() - return _scroll_panel(panel, _LEFT_PANEL_W) - - def _build_canvases(self) -> QWidget: - splitter = QSplitter(Qt.Orientation.Vertical) - splitter.setChildrenCollapsible(False) - - img_widget = QWidget() - img_vl = QVBoxLayout(img_widget) - img_vl.setContentsMargins(0, 0, 0, 0) - img_vl.setSpacing(4) - self.image_canvas = ImageCanvas() - self.image_canvas.setMinimumHeight(220) - self.image_canvas.pixel_clicked.connect(self._on_pixel_clicked) - self.image_canvas.roi_changed.connect(self._update_roi_ui) - self.image_canvas.draw_mode_changed.connect(self._on_draw_mode_changed) - img_vl.addWidget(NavigationToolbar2QT(self.image_canvas, img_widget)) - img_vl.addWidget(self.image_canvas) - splitter.addWidget(img_widget) - - wave_widget = QWidget() - wave_vl = QVBoxLayout(wave_widget) - wave_vl.setContentsMargins(0, 0, 0, 0) - wave_vl.setSpacing(4) - self.lbl_wave_hint = QLabel( - "Click a pixel in the image above to inspect its waveform.") - self.lbl_wave_hint.setAlignment(Qt.AlignmentFlag.AlignCenter) - self.lbl_wave_hint.setStyleSheet(_CSS_MUTED) - self.wave_canvas = WaveformCanvas() - self.wave_canvas.setMinimumHeight(150) - wave_vl.addWidget(self.lbl_wave_hint) - wave_vl.addWidget(self.wave_canvas) - splitter.addWidget(wave_widget) - - splitter.setStretchFactor(0, 3) - splitter.setStretchFactor(1, 1) - splitter.setSizes([580, 250]) - return splitter - - def _build_right_panel(self) -> QWidget: - # Velocity settings (visible only in velocity mode) - self.grp_velocity, vel_l = _group("Velocity Settings (CH1 only)") - vel_form = _form() - self.spin_grating_um = QDoubleSpinBox() - self.spin_grating_um.setRange(0.1, 1000.0) - self.spin_grating_um.setDecimals(2) - self.spin_grating_um.setSingleStep(0.5) - self.spin_grating_um.setSuffix(" µm") - self.spin_grating_um.setValue(25) - self.spin_grating_um.setEnabled(False) - self.spin_grating_um.setMinimumWidth(_SPIN_MIN_W) - self.spin_grating_um.editingFinished.connect(self._on_grating_changed) - vel_form.addRow("Grating size:", self.spin_grating_um) - vel_l.addLayout(vel_form) - vel_l.addWidget(_wrap_label("v (m/s) = freq (MHz) × grating (µm)", - "font-size: 10px; color: #888;")) - self.grp_velocity.setVisible(False) - - grp_display, dl = _group("Display Options") - - cmap_form = _form() - self.combo_cmap = QComboBox() - self.combo_cmap.addItems(CMAPS) - self.combo_cmap.setCurrentText("gray") - self.combo_cmap.setEnabled(False) - self.combo_cmap.setSizePolicy(QSizePolicy.Policy.Expanding, - QSizePolicy.Policy.Fixed) - self.combo_cmap.currentIndexChanged.connect(self._on_cmap_changed) - cmap_form.addRow("Colormap:", self.combo_cmap) - dl.addLayout(cmap_form) - - self.chk_auto = QCheckBox("Auto-scale colormap") - self.chk_auto.setChecked(True) - self.chk_auto.toggled.connect(self._on_autoscale_toggled) - dl.addWidget(self.chk_auto) - - range_form = _form() - for label, attr in (("min:", "spin_vmin"), ("max:", "spin_vmax")): - spin = QDoubleSpinBox() - spin.setRange(-1e9, 1e9) - spin.setDecimals(4) - spin.setEnabled(False) - spin.setMinimumWidth(_SPIN_MIN_W) - spin.editingFinished.connect(self._on_manual_range_changed) - setattr(self, attr, spin) - range_form.addRow(label, spin) - dl.addLayout(range_form) - - right_panel = QWidget() - layout = QVBoxLayout(right_panel) - layout.setContentsMargins(0, 0, 0, 0) - layout.setSpacing(8) - layout.addWidget(self.grp_velocity) - layout.addWidget(grp_display) - layout.addStretch() - return _scroll_panel(right_panel, _RIGHT_PANEL_W) - - def _build_menus(self): - menubar = self.menuBar() - - fft_menu = menubar.addMenu("&FFT") - fft_act = QAction("FFT &Options…", self) - fft_act.setStatusTip("Configure FFT backend and zero-padding") - fft_act.triggered.connect(self._on_fft_options) - fft_menu.addAction(fft_act) - - fusion_menu = menubar.addMenu("&Fusion") - self._alignment_act = QAction("Angle &Alignment", self) - self._alignment_act.setStatusTip( - "Compute a rotation+translation alignment across all angles " - "(from CH4 masks) and enable Aligned View. Requires >1 angle.") - self._alignment_act.setEnabled(False) - self._alignment_act.triggered.connect(self._on_angle_alignment) - fusion_menu.addAction(self._alignment_act) - - self._manual_align_act = QAction("&Manual Alignment…", self) - self._manual_align_act.setStatusTip( - "Open an interactive dialog to align angles by eye: overlaid CH4 " - "threshold masks, keyboard nudge (translate + rotate), auto " - "de-rotate to the known scan angles, and save/clear a persistent " - "alignment.") - self._manual_align_act.setEnabled(False) - self._manual_align_act.triggered.connect(self._on_manual_alignment) - fusion_menu.addAction(self._manual_align_act) - - convert_menu = menubar.addMenu("&Convert") - self._batch_dc_act = QAction("Batch Compute DC and &Store…", self) - self._batch_dc_act.setStatusTip( - "Select .sras files and compute+store DC images (CH3/CH4 mean) " - "for every angle, converting v6 files to v7 in place.") - self._batch_dc_act.triggered.connect(lambda: self._on_batch_compute("dc")) - convert_menu.addAction(self._batch_dc_act) - - self._batch_fft_act = QAction("Batch Compute FFT and Sto&re…", self) - self._batch_fft_act.setStatusTip( - "Select .sras files and compute+store FFT peak-frequency images " - "for every angle, converting v6 files to v7 in place. Stored " - "images are natural-resolution (pad 1); padded views compute live.") - self._batch_fft_act.triggered.connect(lambda: self._on_batch_compute("fft")) - convert_menu.addAction(self._batch_fft_act) - - # ------------------------------------------------------------------ - # Drag-and-drop - # ------------------------------------------------------------------ - - def dragEnterEvent(self, event): - urls = event.mimeData().urls() - if urls and urls[0].toLocalFile().lower().endswith(".sras"): - event.acceptProposedAction() - - def dropEvent(self, event): - self._load_file(event.mimeData().urls()[0].toLocalFile()) - - # ------------------------------------------------------------------ - # File loading - # ------------------------------------------------------------------ - - def _on_open(self): - path, _ = QFileDialog.getOpenFileName( - self, "Open SRAS File", "", "SRAS Files (*.sras);;All Files (*)") - if path: - self._load_file(path) - - def _load_file(self, path: str): - started = self._run_worker( - "load", LoadWorker(path), - connect=( - ("finished", self._on_load_done), - ("error", lambda msg: self.statusBar().showMessage(f"Error: {msg}")), - ), - ) - if not started: - return - self.btn_open.setEnabled(False) - self.statusBar().showMessage(f"Loading {Path(path).name}…") - self._show_progress("main", f"Loading {Path(path).name}…") - - def _on_load_done(self, sras): - self._close_progress("main") - self.btn_open.setEnabled(True) - if sras is None: - return - self._sras = sras - self._current_image = None - - # A manual-alignment dialog bound to the previous file must not - # survive a reload — its per-angle state (and the sras it was - # constructed against) no longer matches the new file's geometry. - if self._manual_align_dialog is not None: - self._manual_align_dialog.close() - self._manual_align_dialog = None - - # Caches (and any in-flight DC precompute) belong to the previous - # file's geometry — discard and start fresh. Bumping the generation - # counters makes any still-running worker's result get dropped when - # it lands. - self._dc_cache = {} - self._fft_cache = {} - self._dc_generation += 1 - self.lbl_dc_precompute.setText("") - - self._alignment_result = None - self._aligned_cache = {} - self._alignment_generation += 1 - self.chk_aligned_view.blockSignals(True) - self.chk_aligned_view.setChecked(False) - self.chk_aligned_view.setEnabled(False) - self.chk_aligned_view.blockSignals(False) - - # Silently restore a previously-saved manual alignment, if any, so - # the work survives closing and reopening the file. - sidecar = load_manual_alignment(sras) - if sidecar is not None: - try: - self._alignment_result = build_manual_alignment( - sras, sidecar.ref_angle_idx, sidecar.dc_threshold_mv, - sidecar.per_angle) - self.chk_aligned_view.blockSignals(True) - self.chk_aligned_view.setChecked(True) - self.chk_aligned_view.blockSignals(False) - self.statusBar().showMessage( - f"Restored saved manual alignment from " - f"{sidecar_path(sras.path).name}") - except Exception as exc: - # A corrupt/foreign sidecar or a rescan that shrank n_angles - # below ref_angle_idx must not block opening the .sras file. - self.statusBar().showMessage( - f"Could not restore saved alignment: {exc}") - - # A ROI from the previous file no longer matches the new scan's - # geometry, so discard it on every load. - self.image_canvas.clear_roi() - - self.lbl_filename.setText(sras.path.name) - - self.spin_angle.blockSignals(True) - self.spin_angle.setRange(0, max(0, sras.n_angles - 1)) - self.spin_angle.setValue(0) - self.spin_angle.blockSignals(False) - - # DC channels are cheap and give an instant, fluid overview of a - # scan; CH1/Velocity require an FFT per pixel that can take minutes - # on a large scan, so don't default to it. - self.combo_channel.blockSignals(True) - self.combo_channel.setCurrentIndex(CH4_IDX) - self.combo_channel.blockSignals(False) - - self._update_controls_enabled(True) - self._on_threshold_changed() # refresh ADC label with file calibration - self._on_view_changed() - self._start_dc_precompute() - - # ------------------------------------------------------------------ - # Scan info panel - # ------------------------------------------------------------------ - - def _update_scan_info_labels(self): - s = self._sras - if s is None: - return - a = self.spin_angle.value() - for key, text in ( - ("Angles", f"{s.n_angles}"), - ("Rows", f"{s.n_rows[a]}"), - ("Frames / row", f"{s.n_frames[a]}"), - ("Samples / frame", f"{s.samples_per_frame}"), - ("Sample rate", f"{s.sample_rate_hz / 1e9:.4g} GS/s"), - ("X start", f"{s.x_start_mm[a]:.4g} mm"), - ("Pixel Δx", f"{s.pixel_x_mm * 1e3:.3g} µm"), - ("Laser freq", f"{s.laser_freq_hz / 1e3:.4g} kHz"), - ): - self._info[key].setText(f"{key}: {text}") - - notes = [] - if s.frame_count_mismatch: - notes.append(f"! Header n_frames={s.n_frames_header}, " - f"actual={s.n_frames[a]} (scanner bug — corrected)") - if s.scan_aborted: - notes.append(f"! Scan aborted: {s.n_angles}/{s.n_angles_declared} " - "angles complete") - if s.background is not None: - notes.append(f"Background waveform: {len(s.background)} samples") - if s.version in (6, 7): - notes.append("v6/v7 format: rows / frames / x_start are per-angle") - - n_dc = sum(1 for x in s.precomputed_dc4_mv if x is not None) - n_fft = sum(1 for x in s.precomputed_freq_mhz if x is not None) - if n_dc or n_fft: - bg_note = " (bg-sub)" if s.precomputed_bg_sub else " (no bg-sub)" - notes.append( - f"Cached images: DC {n_dc}/{s.n_angles} angles, " - f"FFT {n_fft}/{s.n_angles} angles{bg_note if n_fft else ''} " - "— display is instant for cached angles") - elif s.version == 7: - notes.append("v7 format: no cache blocks stored yet") - self.lbl_frame_warn.setText("\n".join(notes)) - - # ------------------------------------------------------------------ - # Controls - # ------------------------------------------------------------------ - - def _update_controls_enabled(self, enabled: bool): - s = self._sras - has_file = enabled and s is not None - ch_idx = self.combo_channel.currentIndex() - is_ch1 = enabled and ch_idx in CH1_DERIVED_MODES - is_vel = enabled and ch_idx == VELOCITY_MODE_IDX - - self.spin_angle.setEnabled(has_file and s.n_angles > 1) - self.combo_channel.setEnabled(enabled) - self.combo_cmap.setEnabled(enabled) - self.chk_auto.setEnabled(enabled) - manual = enabled and not self.chk_auto.isChecked() - self.spin_vmin.setEnabled(manual) - self.spin_vmax.setEnabled(manual) - - # Threshold and bg-sub apply to all CH1 modes - self.spin_threshold_mv.setEnabled(is_ch1) - self.chk_bg_sub.setEnabled(has_file and s.background is not None and is_ch1) - self.spin_grating_um.setEnabled(is_vel) - self.grp_velocity.setVisible(is_vel) - - self.btn_export_csv.setEnabled(is_ch1 and self._current_image is not None) - # ROI: always usable once a file is loaded (independent of channel) - self.btn_draw_roi.setEnabled(has_file) - - # Batch Convert actions pick their own files, independent of - # whatever's currently open — only gated on no batch already running. - can_batch = not self._job_running("batch") - self._batch_dc_act.setEnabled(can_batch) - self._batch_fft_act.setEnabled(can_batch) - - self._alignment_act.setEnabled( - has_file and s.n_angles > 1 and not self._job_running("align")) - self._manual_align_act.setEnabled( - has_file and s.n_angles > 1 and not self._job_running("align")) - self.chk_aligned_view.setEnabled(enabled and self._alignment_result is not None) - self._update_roi_ui() - - def _on_channel_changed(self): - self._update_controls_enabled(self._sras is not None) - self._on_view_changed() - - def _on_bg_sub_toggled(self): - # Background subtraction changes the FFT input, so it genuinely - # invalidates the cached raw FFT (the cache key includes it) — - # _refresh_display() recomputes only on a miss for the new state. - if self._sras is not None and self.combo_channel.currentIndex() in CH1_DERIVED_MODES: - self._refresh_display() - - def _on_grating_changed(self): - # Grating is a pure post-multiply on the cached frequency image — - # never needs a recompute. - if self._sras is not None and self.combo_channel.currentIndex() == VELOCITY_MODE_IDX: - self._refresh_display() - - def _on_threshold_changed(self): - mv = self.spin_threshold_mv.value() - cal = (self._sras.cal(CH4_IDX) if self._sras is not None - else (_FALLBACK_YMULT_MV, _FALLBACK_YOFF_ADC, 0.0)) - self.lbl_threshold_adc.setText(f"≈ {mv_to_adc(mv, *cal):.1f} ADC counts") - # Threshold decides which pixels get an FFT at all, so changing it is - # a genuine cache-key change — but the recompute reuses the cached DC4 - # image to skip masked-out pixels. - if self._sras is not None and self.combo_channel.currentIndex() in CH1_DERIVED_MODES: - self._refresh_display() - - def _on_autoscale_toggled(self, checked: bool): - manual = not checked - self.spin_vmin.setEnabled(manual and self._sras is not None) - self.spin_vmax.setEnabled(manual and self._sras is not None) - if self._sras is not None and self._current_image is not None: - self._redraw_image(self._current_image) - - def _on_manual_range_changed(self): - if not self.chk_auto.isChecked() and self._current_image is not None: - self._redraw_image(self._current_image) - - def _on_cmap_changed(self): - # Colormap is purely how the existing image is rendered. - if self._current_image is not None: - self._redraw_image(self._current_image) - - def _on_view_changed(self): - if self._sras is None: - return - idx = self.spin_angle.value() - self.lbl_angle_deg.setText(f"({self._sras.angles_deg[idx]:.1f}°)") - self._update_scan_info_labels() - self._refresh_display() - - def _on_aligned_view_toggled(self, checked: bool): - if self._current_image is not None: - self._redraw_image(self._current_image) - - # ------------------------------------------------------------------ - # CSV export - # ------------------------------------------------------------------ - - def _on_export_csv(self): - if self._current_image is None or self._sras is None: - return - default_name = (f"{self._sras.path.stem}_angle{self._current_angle}" - f"_{CH_NAMES[self._current_ch]}.csv") - path, _ = QFileDialog.getSaveFileName( - self, "Export Image as CSV", - str(self._sras.path.parent / default_name), - "CSV files (*.csv);;All files (*)") - if not path: - return - np.savetxt(path, self._current_image, delimiter=",", fmt="%.6g") - self.statusBar().showMessage(f"Exported {Path(path).name}") - - def _on_export_roi_csv(self): - if self._current_image is None or self._sras is None: - return - roi = self.image_canvas.get_roi() - if roi is None: - self.statusBar().showMessage("No ROI — draw one first") - return - s = self._sras - x_axis = s.x_axis_mm(self._current_angle) - y_axis = s.y_positions_mm(self._current_angle) - mask = roi.mask_for_grid(x_axis, y_axis) - if not mask.any(): - self.statusBar().showMessage("ROI does not overlap any pixel") - return - img = self._current_image - if img.shape != mask.shape: - self.statusBar().showMessage( - f"ROI shape {mask.shape} does not match image {img.shape}") - return - - X, Y = np.meshgrid(np.asarray(x_axis, dtype=np.float64), - np.asarray(y_axis, dtype=np.float64)) - rows_idx, frames_idx = np.where(mask) - n_pix = int(mask.sum()) - - ch_name = CH_NAMES[self._current_ch] - angle = self._current_angle - default_name = f"{s.path.stem}_angle{angle}_{ch_name}_ROI.csv" - path, _ = QFileDialog.getSaveFileName( - self, "Export ROI as CSV", - str(s.path.parent / default_name), - "CSV files (*.csv);;All files (*)") - if not path: - return - - corners_str = " ".join(f"({p[0]:.6g},{p[1]:.6g})" for p in roi.corners()) - header = ( - f"# ROI quad corners (BL BR TR TL) mm: {corners_str}\n" - f"# source: {s.path.name}, channel={ch_name}, " - f"angle_idx={angle}, angle_deg={s.angles_deg[angle]:.4g}\n" - f"# n_pixels={n_pix}\n" - "row,frame,x_mm,y_mm,value" - ) - data = np.column_stack([ - rows_idx.astype(np.int64), frames_idx.astype(np.int64), - X[mask], Y[mask], img[mask].astype(np.float64), - ]) - # integer columns first, floats after — use a per-column format list - np.savetxt(path, data, delimiter=",", - fmt=["%d", "%d", "%.6g", "%.6g", "%.6g"], - header=header, comments="") - self.statusBar().showMessage( - f"Exported ROI ({n_pix} pixels) to {Path(path).name}") - - # ------------------------------------------------------------------ - # ROI - # ------------------------------------------------------------------ - - def _on_draw_roi_toggled(self, checked: bool): - if checked: - self.image_canvas.start_drawing() - self.statusBar().showMessage( - "Click and drag on the image to draw a new rectangle.") - else: - self.image_canvas.cancel_drawing() - - def _on_draw_mode_changed(self, active: bool): - # Keep the toggle button's visual state in sync with the canvas. - self.btn_draw_roi.blockSignals(True) - self.btn_draw_roi.setChecked(active) - self.btn_draw_roi.blockSignals(False) - - def _on_clear_roi(self): - self.image_canvas.clear_roi() - self.statusBar().showMessage("ROI cleared") - - def _update_roi_ui(self): - roi = self.image_canvas.get_roi() - if roi is None: - self.lbl_roi_center.setText("centroid: —") - self.lbl_roi_size.setText("bbox: —") - self.lbl_roi_npix.setText("pixels inside: —") - self.btn_clear_roi.setEnabled(False) - self.btn_export_roi.setEnabled(False) - return - - cen = roi.centroid() - bbox = roi.bbox_size() - self.lbl_roi_center.setText(f"centroid: ({cen[0]:.3f}, {cen[1]:.3f}) mm") - self.lbl_roi_size.setText(f"bbox: {bbox[0]:.3f} × {bbox[1]:.3f} mm") - - npix = 0 - if self._sras is not None: - try: - # Deliberately always the raw per-angle grid, even when - # Aligned View is on: _on_export_roi_csv also exports on the - # raw grid (never synthetically-resampled pixels), so this - # readout must match what Export ROI actually writes. - mask = roi.mask_for_grid( - self._sras.x_axis_mm(self._current_angle), - self._sras.y_positions_mm(self._current_angle)) - npix = int(mask.sum()) - except Exception: - npix = 0 - self.lbl_roi_npix.setText(f"pixels inside: {npix}") - self.btn_clear_roi.setEnabled(True) - self.btn_export_roi.setEnabled(self._current_image is not None and npix > 0) - - # ------------------------------------------------------------------ - # Display - # ------------------------------------------------------------------ - - def _current_n_fft(self) -> int | None: - if self._fft_pad_factor <= 1 or self._sras is None: - return None - return self._sras.samples_per_frame * self._fft_pad_factor - - def _scale_for_display(self, freq_mhz: np.ndarray, ch_idx: int) -> np.ndarray: - """Velocity is a pure post-multiply of the (already DC-masked) - cached frequency image — never worth a recompute on its own.""" - if ch_idx == VELOCITY_MODE_IDX: - return freq_mhz * self.spin_grating_um.value() - return freq_mhz - - def _fft_cache_key(self, angle_idx: int) -> tuple: - return (angle_idx, self.chk_bg_sub.isChecked(), self._current_n_fft(), - self.spin_threshold_mv.value()) - - def _aligned_cache_key(self, angle_idx: int, ch_idx: int) -> tuple: - """Mirrors _fft_cache's key granularity so a stale aligned image is - never shown after bg_sub/threshold/pad/grating changes.""" - if ch_idx in CH1_DERIVED_MODES: - return (*self._fft_cache_key(angle_idx), ch_idx, - self.spin_grating_um.value() if ch_idx == VELOCITY_MODE_IDX else None) - return (angle_idx, ch_idx) - - def _aligned_canvas_axes(self) -> tuple[np.ndarray, np.ndarray]: - r = self._alignment_result - n_rows, n_cols = r.canvas_shape - return (r.canvas_origin_mm[0] + np.arange(n_cols) * r.canvas_dx_mm, - r.canvas_origin_mm[1] + np.arange(n_rows) * r.canvas_dy_mm) - - def _get_aligned_display_image(self, raw_img: np.ndarray, angle_idx: int, - ch_idx: int) -> np.ndarray: - key = self._aligned_cache_key(angle_idx, ch_idx) - cached = self._aligned_cache.get(key) - if cached is None: - cached = apply_alignment(self._alignment_result, angle_idx, raw_img) - self._aligned_cache[key] = cached - return cached - - def _refresh_display(self): - """Show the image for the current angle/channel/threshold, using - cached data whenever possible and only falling back to a background - compute (with progress popup) when genuinely nothing is cached yet.""" - if self._sras is None: - return - angle_idx = self.spin_angle.value() - ch_idx = self.combo_channel.currentIndex() - - if ch_idx in CH1_DERIVED_MODES: - raw = self._fft_cache.get(self._fft_cache_key(angle_idx)) - if raw is not None: - self._show_image_now(self._scale_for_display(raw, ch_idx), - angle_idx, ch_idx) - return - else: - cached = self._dc_cache.get((angle_idx, ch_idx)) - if cached is not None: - self._show_image_now(cached, angle_idx, ch_idx) - return - - # Nothing cached for these settings — need a real compute. Changing - # the DC threshold changes *which* pixels get an FFT at all, so it - # can't be satisfied from the cache — but with the DC map already - # known, the recompute skips the FFT for masked-out pixels. - self._start_compute() - - def _show_image_now(self, img: np.ndarray, angle_idx: int, ch_idx: int): - """Display an already-available image with no compute involved.""" - self._current_image = img - self._current_angle = angle_idx - self._current_ch = ch_idx - self.btn_export_csv.setEnabled(ch_idx in CH1_DERIVED_MODES) - self._redraw_image(img) - self._update_roi_ui() - - def _redraw_image(self, img: np.ndarray): - s = self._sras - angle_idx = self._current_angle - ch_idx = self._current_ch - - aligned = (self.chk_aligned_view.isChecked() - and self._alignment_result is not None - and angle_idx in self._alignment_result.per_angle) - if aligned: - display_img = self._get_aligned_display_image(img, angle_idx, ch_idx) - x_axis, y_axis = self._aligned_canvas_axes() - else: - display_img = img - x_axis = s.x_axis_mm(angle_idx) - y_axis = s.y_positions_mm(angle_idx) - - dx = x_axis[1] - x_axis[0] if len(x_axis) > 1 else s.pixel_x_mm - dy = float(y_axis[1] - y_axis[0]) if len(y_axis) > 1 else 1.0 - extent = [x_axis[0] - dx / 2, x_axis[-1] + dx / 2, - y_axis[-1] + dy / 2, y_axis[0] - dy / 2] - - if self.chk_auto.isChecked(): - vmin, vmax = float(display_img.min()), float(display_img.max()) - for spin, val in ((self.spin_vmin, vmin), (self.spin_vmax, vmax)): - spin.blockSignals(True) - spin.setValue(val) - spin.blockSignals(False) - else: - vmin, vmax = self.spin_vmin.value(), self.spin_vmax.value() - - angle_deg = s.angles_deg[angle_idx] - mode_str, unit, colorbar_label = _CHANNEL_DISPLAY[ch_idx] - if ch_idx == VELOCITY_MODE_IDX: - ch_label = f"Velocity [grating={self.spin_grating_um.value():.2f} µm]" - else: - ch_label = CH_LABELS[ch_idx] - - title = f"{CH_NAMES[ch_idx]} | {mode_str} | {angle_deg:.1f}°" - if aligned: - title += " [Aligned]" - - self.image_canvas.show_image( - display_img, extent, - cmap=self.combo_cmap.currentText(), - vmin=vmin, vmax=vmax, - xlabel="X (mm)", ylabel="Y (mm)", - title=title, colorbar_label=colorbar_label, - ) - self.statusBar().showMessage( - f"{s.path.name} | {ch_label} @ {angle_deg:.1f}° " - f"| {display_img.shape[1]} × {display_img.shape[0]} px | {unit}" - f"{' | Aligned' if aligned else ''}" - ) - - # ------------------------------------------------------------------ - # Background compute (only reached on a genuine cache miss) - # ------------------------------------------------------------------ - - def _start_compute(self): - if self._sras is None or self._job_running("compute"): - return # re-checked when the running compute finishes - - angle_idx = self.spin_angle.value() - ch_idx = self.combo_channel.currentIndex() - is_fft = ch_idx in CH1_DERIVED_MODES - - self._pending_angle = angle_idx - self._pending_ch = ch_idx - self._pending_bg_sub = self.chk_bg_sub.isChecked() - self._pending_threshold = self.spin_threshold_mv.value() - self._pending_fft_pad_factor = self._fft_pad_factor - - worker = ComputeWorker( - self._sras, angle_idx, ch_idx, - apply_bg_sub=self._pending_bg_sub, - n_fft=self._current_n_fft(), - dc_threshold_mv=self._pending_threshold, - # Reuse the cached DC4 image (if the precompute has reached this - # angle) so the FFT skips masked-out pixels entirely and doesn't - # need to re-read the CH4 channel from disk. - dc4_mv=self._dc_cache.get((angle_idx, CH4_IDX)), - is_fft_mode=is_fft, - ) - if not self._run_worker( - "compute", worker, - connect=( - ("finished", self._on_compute_done), - ("error", lambda msg: self.statusBar().showMessage( - f"Compute error: {msg}")), - ), - on_done=self._after_compute): - return - - if is_fft: - self.statusBar().showMessage("Computing FFT…") - self._show_progress( - "main", - f"Computing FFT for angle {angle_idx}…\n" - "This can take a while on a large scan — result is cached " - "so revisiting this angle/mode/threshold will be instant.") - else: - self.statusBar().showMessage("Computing DC image…") - self._show_progress("main", f"Computing DC image for angle {angle_idx}…") - - def _after_compute(self): - """If settings changed while the compute was running, re-dispatch - through the cache-aware path — the now-current combination may - already be cached.""" - if (self.spin_angle.value(), self.combo_channel.currentIndex(), - self.chk_bg_sub.isChecked(), self.spin_threshold_mv.value(), - self._fft_pad_factor) != ( - self._pending_angle, self._pending_ch, self._pending_bg_sub, - self._pending_threshold, self._pending_fft_pad_factor): - self._refresh_display() - - def _on_compute_done(self, result): - self._close_progress("main") - if result is None: - return # cancelled mid-compute; the partial image must not cache - angle_idx = self._pending_angle - ch_idx = self._pending_ch - - if ch_idx in CH1_DERIVED_MODES: - self._fft_cache[(angle_idx, self._pending_bg_sub, - self._current_n_fft(), self._pending_threshold)] = result - img = self._scale_for_display(result, ch_idx) - else: - img = result - self._dc_cache[(angle_idx, ch_idx)] = img - - self._show_image_now(img, angle_idx, ch_idx) - - # ------------------------------------------------------------------ - # Background DC precompute (all angles, so switching is fluid) - # ------------------------------------------------------------------ - - def _start_dc_precompute(self): - if self._sras is None: - return - generation = self._dc_generation - n_angles = self._sras.n_angles - - worker = DcPrecomputeWorker(self._sras) - self._run_worker( - "dc_precompute", worker, - connect=( - ("angle_done", lambda a, dc3, dc4, g=generation: - self._on_dc_precompute_angle_done(g, a, dc3, dc4, n_angles)), - ("error", lambda msg: self.statusBar().showMessage( - f"DC precompute error: {msg}", 5000)), - ), - quit_on=("finished", "error"), - ) - - def _on_dc_precompute_angle_done(self, generation: int, angle_idx: int, - dc3_mv: np.ndarray, dc4_mv: np.ndarray, - n_angles: int): - if generation != self._dc_generation: - return # stale result from a previously-loaded file — discard - self._dc_cache[(angle_idx, CH3_IDX)] = dc3_mv - self._dc_cache[(angle_idx, CH4_IDX)] = dc4_mv - - done = sum(1 for a in range(n_angles) if (a, CH4_IDX) in self._dc_cache) - self.lbl_dc_precompute.setText( - f"Precomputing DC images: {done}/{n_angles} angles ready…" - if done < n_angles else "DC images ready for all angles.") - - # If we just finished the angle/channel the user is currently looking - # at and it wasn't shown yet (they switched here before the precompute - # caught up and are still waiting), show it now. - current_ch = self.combo_channel.currentIndex() - if (angle_idx == self.spin_angle.value() - and not self._job_running("compute") - and current_ch in (CH3_IDX, CH4_IDX) - and (self._current_angle != angle_idx or self._current_ch != current_ch)): - self._refresh_display() - - # ------------------------------------------------------------------ - # Pixel inspector - # ------------------------------------------------------------------ - - def _on_pixel_clicked(self, row_idx: int, frame_idx: int): - if self._sras is None or self._current_image is None: - return - angle_idx = self._current_angle - if (self.chk_aligned_view.isChecked() and self._alignment_result is not None - and angle_idx in self._alignment_result.per_angle): - # The click landed on the shared aligned canvas — invert the same - # canvas->raw affine used to display it back to a raw (row, frame) - # index before looking up the waveform. - t = self._alignment_result.per_angle[angle_idx] - raw = t.matrix @ np.array([row_idx, frame_idx], dtype=np.float64) + t.offset - row_idx, frame_idx = int(round(raw[0])), int(round(raw[1])) - n_rows_a, n_frames_a = self._sras.image_shape(angle_idx) - if not (0 <= row_idx < n_rows_a and 0 <= frame_idx < n_frames_a): - self.statusBar().showMessage( - "No source waveform here (padding region of the aligned canvas).") - return - - self.lbl_wave_hint.hide() - if self._current_ch in CH1_DERIVED_MODES: - self.wave_canvas.show_rf_waveform( - self._sras, angle_idx, row_idx, frame_idx, - apply_bg_sub=self.chk_bg_sub.isChecked()) - else: - self.wave_canvas.show_dc_waveform( - self._sras, angle_idx, self._current_ch, row_idx, frame_idx) - - # ------------------------------------------------------------------ - # Progress dialogs - # ------------------------------------------------------------------ - - def _show_progress(self, key: str, message: str, maximum: int = 0): - """Show (or relabel) the progress dialog under *key*. maximum=0 gives - an indeterminate busy indicator.""" - dlg = self._progress_dlgs.get(key) - if dlg is not None: - dlg.setLabelText(message) - return - dlg = QProgressDialog(message, "", 0, maximum, self) - dlg.setWindowTitle("Please wait…") - dlg.setCancelButton(None) - dlg.setWindowModality(Qt.WindowModality.WindowModal) - dlg.setMinimumDuration(300) # only appears if it takes > 300 ms - dlg.show() - self._progress_dlgs[key] = dlg - - def _set_progress(self, key: str, pct: int): - dlg = self._progress_dlgs.get(key) - if dlg is not None: - dlg.setValue(pct) - - def _close_progress(self, key: str): - dlg = self._progress_dlgs.pop(key, None) - if dlg is not None: - dlg.close() - - # ------------------------------------------------------------------ - # Convert menu: batch DC/FFT compute-and-store (v6 -> v7) - # ------------------------------------------------------------------ - - def _on_batch_compute(self, mode: str): - if self._job_running("batch"): - return - label = "DC" if mode == "dc" else "FFT" - paths, _ = QFileDialog.getOpenFileNames( - self, f"Select .sras files to batch-compute {label}", "", - "SRAS files (*.sras);;All files (*)") - if not paths: - return - - self._batch_errors = [] - worker = BatchCacheWorker(paths, mode, self.chk_bg_sub.isChecked()) - started = self._run_worker( - "batch", worker, - connect=( - ("progress", lambda pct: self._set_progress("batch", pct)), - ("file_done", self._on_batch_file_done), - ("finished", lambda p=paths: self._on_batch_finished(p)), - ), - on_done=self._after_batch, - ) - if not started: - return # a second trigger snuck in while the file dialog was open - - self._batch_dc_act.setEnabled(False) - self._batch_fft_act.setEnabled(False) - self._show_progress( - "batch", f"Batch computing {label} for {len(paths)} file(s)…", - maximum=100) - - def _on_batch_file_done(self, path: str, err: str): - if err: - self._batch_errors.append(f"{Path(path).name} — {err}") - self._show_progress("batch", f"Processed {Path(path).name}…") - - def _on_batch_finished(self, paths: list[str]): - self._close_progress("batch") - - n_total = len(paths) - n_failed = len(self._batch_errors) - n_ok = n_total - n_failed - if n_failed: - summary = (f"Batch store: {n_ok}/{n_total} file(s) updated, " - f"{n_failed} failed: {'; '.join(self._batch_errors)}") - else: - summary = f"Batch store: {n_ok}/{n_total} file(s) updated." - self.statusBar().showMessage(summary) - self._batch_errors = [] - - # If the currently-open file was in this batch, reload it so the GUI - # picks up the newly-written v7 cache instead of stale state. - if self._sras is not None and str(self._sras.path) in paths: - self._load_file(str(self._sras.path)) - - def _after_batch(self): - self._batch_dc_act.setEnabled(True) - self._batch_fft_act.setEnabled(True) - - # ------------------------------------------------------------------ - # Fusion: angle alignment - # ------------------------------------------------------------------ - - def _on_angle_alignment(self): - if self._sras is None or self._sras.n_angles <= 1: - return - ref_idx = 0 - threshold_mv = self.spin_threshold_mv.value() - generation = self._alignment_generation - - started = self._run_worker( - "align", AngleAlignmentWorker(self._sras, ref_idx, threshold_mv), - connect=( - ("progress", lambda pct: self._set_progress("main", pct)), - ("finished", lambda result, err, g=generation: - self._on_alignment_done(g, result, err)), - ), - on_done=lambda: self._update_controls_enabled(self._sras is not None), - ) - if not started: - return - - self._alignment_act.setEnabled(False) - self._show_progress( - "main", - f"Computing angle alignment ({self._sras.n_angles} angles, " - f"ref=angle 0, CH4 mask ≥ {threshold_mv:.3f} mV)…", - maximum=100) - - def _on_alignment_done(self, generation: int, result, error_msg: str): - self._close_progress("main") - if generation != self._alignment_generation: - return # a new file was loaded while this was computing — discard - if error_msg: - self.statusBar().showMessage(f"Angle alignment failed: {error_msg}") - return - self._alignment_result = result - self._aligned_cache = {} - self.chk_aligned_view.setEnabled(True) - self.chk_aligned_view.blockSignals(True) - self.chk_aligned_view.setChecked(True) - self.chk_aligned_view.blockSignals(False) - nr, nc = result.canvas_shape - self.statusBar().showMessage( - f"Angle alignment computed ({self._sras.n_angles} angles, " - f"canvas {nc}×{nr} px).") - self._refresh_display() - - # ------------------------------------------------------------------ - # Fusion: manual alignment - # ------------------------------------------------------------------ - - def _on_manual_alignment(self): - if self._sras is None or self._sras.n_angles <= 1: - return - if self._manual_align_dialog is not None: - self._manual_align_dialog.raise_() - self._manual_align_dialog.activateWindow() - return - - ref_idx = 0 - threshold_mv = self.spin_threshold_mv.value() - seed: dict[int, ManualAngleParams] = {} - # Seed only from a previously *saved manual* alignment (this dialog's - # own Save also writes this sidecar) -- never from self._alignment_result - # when it holds the automatic Fusion -> Angle Alignment's output. That - # path's translation comes from FFT phase correlation, which is the - # very thing manual mode exists to work around; inheriting it here - # would silently reintroduce the same bad translations under a - # "manual" label, on top of the (correct) analytic rotation, which is - # exactly what makes manual mode look like it "still does the same - # thing" the automatic one does. - sidecar = load_manual_alignment(self._sras) - if sidecar is not None and sidecar.ref_angle_idx == ref_idx: - seed = dict(sidecar.per_angle) - threshold_mv = sidecar.dc_threshold_mv - - cached_dc4 = {a: img for (a, ch), img in self._dc_cache.items() if ch == CH4_IDX} - dlg = ManualAlignmentDialog( - self, self._sras, ref_angle_idx=ref_idx, dc_threshold_mv=threshold_mv, - seed_per_angle=seed, cached_dc4_mv=cached_dc4) - dlg.alignment_saved.connect(self._on_manual_alignment_saved) - dlg.alignment_cleared.connect(self._on_manual_alignment_cleared) - dlg.finished.connect(self._on_manual_align_dialog_closed) - dlg.setAttribute(Qt.WidgetAttribute.WA_DeleteOnClose) - self._manual_align_dialog = dlg - dlg.show() - - def _on_manual_align_dialog_closed(self, _result_code: int): - self._manual_align_dialog = None - - def _on_manual_alignment_saved(self, result, sidecar_path_str: str): - self._alignment_result = result - self._aligned_cache = {} - self._alignment_generation += 1 - self.chk_aligned_view.setEnabled(True) - self.chk_aligned_view.blockSignals(True) - self.chk_aligned_view.setChecked(True) - self.chk_aligned_view.blockSignals(False) - self._update_controls_enabled(self._sras is not None) - self.statusBar().showMessage( - f"Manual alignment saved to {Path(sidecar_path_str).name}") - if self._current_image is not None: - self._refresh_display() - - def _on_manual_alignment_cleared(self): - self._alignment_result = None - self._aligned_cache = {} - self._alignment_generation += 1 - self.chk_aligned_view.blockSignals(True) - self.chk_aligned_view.setChecked(False) - self.chk_aligned_view.setEnabled(False) - self.chk_aligned_view.blockSignals(False) - self._update_controls_enabled(self._sras is not None) - self.statusBar().showMessage("Manual alignment cleared.") - if self._current_image is not None: - self._refresh_display() - - # ------------------------------------------------------------------ - # FFT Options - # ------------------------------------------------------------------ - - def _on_fft_options(self): - dlg = FftOptionsDialog( - self, - current_backend=compute.get_fft_backend(), - current_pad_factor=self._fft_pad_factor, - samples_per_frame=self._sras.samples_per_frame if self._sras else None, - sample_rate_hz=self._sras.sample_rate_hz if self._sras else None, - grating_um=self.spin_grating_um.value(), - ) - if dlg.exec() != QDialog.DialogCode.Accepted: - return - compute.set_fft_backend(dlg.get_backend()) - self._fft_pad_factor = dlg.get_pad_factor() - self._settings.setValue("fft/backend", compute.get_fft_backend()) - self._settings.setValue("fft/pad_factor", self._fft_pad_factor) - # Pad factor changes the FFT bin count, so it genuinely invalidates - # the cached raw FFT (part of the cache key) — _refresh_display() - # recomputes only on a cache miss. - if self._sras is not None and self.combo_channel.currentIndex() in CH1_DERIVED_MODES: - self._refresh_display() - - # ------------------------------------------------------------------ - - def closeEvent(self, event): - if self._manual_align_dialog is not None: - self._manual_align_dialog.close() - - # Signal every cancellable worker first, then wait. Waiting without - # signalling means sitting out whatever is in flight — on a large - # scan a single angle is ~40 s. - jobs = list(self._jobs.values()) - for _thread, worker, _on_done in jobs: - stop = getattr(worker, "stop", None) - if callable(stop): - stop() - for thread, _worker, _on_done in jobs: - thread.quit() - thread.wait(5000) - super().closeEvent(event) - - -# --------------------------------------------------------------------------- - -def main(): - app = QApplication(sys.argv) - window = SrasViewerWindow( - initial_path=sys.argv[1] if len(sys.argv) > 1 else None) - window.show() - sys.exit(app.exec()) - - -if __name__ == "__main__": - main() diff --git a/sras_viewer/__init__.py b/sras_viewer/__init__.py new file mode 100644 index 0000000..b218d72 --- /dev/null +++ b/sras_viewer/__init__.py @@ -0,0 +1,24 @@ +""" +SRAS Scan File Viewer +PyQt6 application for visualizing channel data from .sras binary scan files. + +Channel semantics (fixed by sc3_aui_app.py acquisition settings): + CH1 — RF Acoustic Packet (AC-coupled, 100 mV/div): FFT → peak frequency + CH3 — Bias A (DC-coupled, 50 mV/div): waveform mean + CH4 — Bias B (DC-coupled, 50 mV/div): waveform mean + +RF images are masked: pixels where CH4_dc < dc_threshold show 0. + +File parsing lives in sras_format, image/alignment math in sras_compute, and +background workers in sras_workers — none of which import Qt or matplotlib, +so multiprocessing children can load them cheaply. +""" + +import faulthandler + +faulthandler.enable() # print a native stack trace on SIGSEGV/SIGABRT/etc. + +from .canvases import ImageCanvas, RoiQuad, WaveformCanvas # noqa: E402,F401 +from .common import CH_LABELS, CMAPS, VELOCITY_MODE_IDX # noqa: E402,F401 +from .dialogs import FftOptionsDialog, ManualAlignmentDialog # noqa: E402,F401 +from .main_window import SrasViewerWindow, main # noqa: E402,F401 diff --git a/sras_viewer/__main__.py b/sras_viewer/__main__.py new file mode 100644 index 0000000..cbd143d --- /dev/null +++ b/sras_viewer/__main__.py @@ -0,0 +1,4 @@ +from .main_window import main + +if __name__ == "__main__": + main() diff --git a/sras_viewer/canvases.py b/sras_viewer/canvases.py new file mode 100644 index 0000000..f7b5e83 --- /dev/null +++ b/sras_viewer/canvases.py @@ -0,0 +1,542 @@ +"""Matplotlib canvases and the ROI primitive.""" + +import numpy as np +from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg +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 + +# --------------------------------------------------------------------------- +# 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): + 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 + + # 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: str, + vmin: float, vmax: float, xlabel: str, ylabel: str, title: str, + colorbar_label: str = ""): + 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 + + im = self.ax.imshow( + img, aspect="auto", origin="upper", + extent=extent, cmap=cmap, vmin=vmin, vmax=vmax, + interpolation="nearest", + ) + cb = self.figure.colorbar(im, ax=self.ax, fraction=0.046, pad=0.04) + 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] + + self._draw_roi() + self.draw_idle() + + 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 ManualAlignOverlayCanvas(FigureCanvasQTAgg): + """Renders ManualAlignmentDialog's multi-angle mask overlay and turns + keyboard input into translate/rotate nudge requests for whichever angle + the dialog currently has active. + + A pure input+render widget — it holds no alignment state and never + touches SrasFile itself; ManualAlignmentDialog 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.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) + diff --git a/sras_viewer/common.py b/sras_viewer/common.py new file mode 100644 index 0000000..7d07a50 --- /dev/null +++ b/sras_viewer/common.py @@ -0,0 +1,115 @@ +"""Shared constants and small layout helpers for the viewer widgets.""" + +from PyQt6.QtCore import Qt +from PyQt6.QtWidgets import ( + QFormLayout, QFrame, QGroupBox, QLabel, QScrollArea, QSizePolicy, + QVBoxLayout, QWidget, +) + +from sras_format import CH1_IDX, CH3_IDX, CH4_IDX + +# --------------------------------------------------------------------------- +# Display constants +# --------------------------------------------------------------------------- + +CH_LABELS = [ + "CH1 — RF (FFT peak freq)", + "CH3 — Bias A (DC mean)", + "CH4 — Bias B (DC mean)", + "CH1 — Velocity (SRAS)", +] + +# Combo index for the derived velocity mode (uses CH1_IDX data) +VELOCITY_MODE_IDX = 3 +# All modes that operate on CH1 waveforms +CH1_DERIVED_MODES = (CH1_IDX, VELOCITY_MODE_IDX) + +CMAPS = ["gray", "viridis", "plasma", "inferno", "hot", "jet", "RdBu_r", "seismic"] + +# (mode_str, status-bar unit, colorbar label) per channel index +_CHANNEL_DISPLAY = { + CH1_IDX: ("RF", "Peak frequency (MHz)", "MHz"), + CH3_IDX: ("DC", "DC mean (mV)", "mV"), + CH4_IDX: ("DC", "DC mean (mV)", "mV"), + VELOCITY_MODE_IDX: ("Velocity", "Velocity (m/s)", "m/s"), +} + +_CSS_HINT = "font-size: 11px; color: #aaa;" +_CSS_INFO = "font-size: 11px;" +_CSS_MUTED = "color: #888; font-size: 11px;" +_CSS_WARN = "color: #e07000; font-size: 11px;" +_CSS_BUSY = "color: #4a90d9; font-size: 11px;" + +# Side-panel column widths (the scroll areas that hold the controls). +_LEFT_PANEL_W = 288 +_RIGHT_PANEL_W = 272 + +# Minimum width for a spin box so its value + suffix are never clipped. +_SPIN_MIN_W = 96 + + +# --------------------------------------------------------------------------- +# Small layout helpers +# --------------------------------------------------------------------------- + +def _wrap_label(text: str = "", css: str | None = None) -> QLabel: + """A word-wrapped QLabel that reports its *wrapped* height to the layout. + + A plain word-wrapped QLabel advertises a single-line minimum height, so in a + fixed-width column the layout happily shrinks it and the extra lines get + clipped. Enabling height-for-width makes the box layout ask for the real + height at the column's width instead. + """ + lbl = QLabel(text) + lbl.setWordWrap(True) + sp = lbl.sizePolicy() + sp.setVerticalPolicy(QSizePolicy.Policy.Minimum) + sp.setHeightForWidth(True) + lbl.setSizePolicy(sp) + if css: + lbl.setStyleSheet(css) + return lbl + + +def _group(title: str) -> tuple[QGroupBox, QVBoxLayout]: + """A group box with consistent, non-cramped internal margins.""" + grp = QGroupBox(title) + lay = QVBoxLayout(grp) + lay.setContentsMargins(10, 8, 10, 10) + lay.setSpacing(6) + return grp, lay + + +def _form() -> QFormLayout: + """A label/field form layout for a narrow side panel.""" + form = QFormLayout() + form.setContentsMargins(0, 0, 0, 0) + form.setHorizontalSpacing(8) + form.setVerticalSpacing(6) + form.setLabelAlignment(Qt.AlignmentFlag.AlignRight + | Qt.AlignmentFlag.AlignVCenter) + form.setFormAlignment(Qt.AlignmentFlag.AlignLeft | Qt.AlignmentFlag.AlignTop) + form.setFieldGrowthPolicy( + QFormLayout.FieldGrowthPolicy.AllNonFixedFieldsGrow) + form.setRowWrapPolicy(QFormLayout.RowWrapPolicy.DontWrapRows) + return form + + +def _scroll_panel(inner: QWidget, width: int) -> QScrollArea: + """Put a side panel in a fixed-width scroll area. + + Without this the panels are sized by the window: a short window squeezes the + controls past their minimum heights, which is what makes text overlap the + widget below it. Scrolling keeps every control at its natural size. + """ + area = QScrollArea() + area.setWidget(inner) + area.setWidgetResizable(True) + area.setFrameShape(QFrame.Shape.NoFrame) + area.setHorizontalScrollBarPolicy(Qt.ScrollBarPolicy.ScrollBarAlwaysOff) + area.setVerticalScrollBarPolicy(Qt.ScrollBarPolicy.ScrollBarAsNeeded) + area.setFixedWidth(width) + area.viewport().setAutoFillBackground(False) + inner.setAutoFillBackground(False) + return area + diff --git a/sras_viewer/dialogs.py b/sras_viewer/dialogs.py new file mode 100644 index 0000000..d1efed7 --- /dev/null +++ b/sras_viewer/dialogs.py @@ -0,0 +1,767 @@ +"""FFT Options and Manual Alignment dialogs.""" + +from typing import TYPE_CHECKING + +import matplotlib as mpl +import numpy as np +from matplotlib.backends.backend_qtagg import NavigationToolbar2QT +from PyQt6.QtCore import pyqtSignal +from PyQt6.QtWidgets import ( + QButtonGroup, QComboBox, QDialog, QDialogButtonBox, QDoubleSpinBox, + QGroupBox, QHBoxLayout, QLabel, QMessageBox, QPushButton, QRadioButton, + QSpinBox, QVBoxLayout, QWidget, +) + +import sras_compute as compute +from sras_compute import ( + PYFFTW_AVAILABLE, ManualAngleParams, build_manual_alignment, + delete_manual_alignment, save_manual_alignment, +) +from sras_format import SrasFile +from sras_workers import Ch4MaskWorker, CrossCorrelateWorker + +from .canvases import ManualAlignOverlayCanvas +from .common import ( + _CSS_HINT, _CSS_MUTED, _CSS_WARN, _SPIN_MIN_W, _form, _group, + _scroll_panel, _wrap_label, +) + +if TYPE_CHECKING: + from .main_window import SrasViewerWindow + + +# --------------------------------------------------------------------------- +# FFT Options dialog +# --------------------------------------------------------------------------- + +class FftOptionsDialog(QDialog): + """Configure FFT backend and zero-padding. + + Changes take effect only when the user clicks Apply. Cancel discards + all pending edits. The live 'frequency resolution' label updates as + the user adjusts the pad factor so they can see the trade-off before + committing. + """ + + def __init__(self, parent=None, *, + current_backend: str, + current_pad_factor: int, + samples_per_frame: int | None, + sample_rate_hz: float | None, + grating_um: float): + super().__init__(parent) + self.setWindowTitle("FFT Options") + self.setModal(True) + self.setMinimumWidth(380) + + self._samples_per_frame = samples_per_frame + self._sample_rate_hz = sample_rate_hz + self._grating_um = grating_um + + layout = QVBoxLayout(self) + + # ---- Backend --------------------------------------------------- + grp_backend = QGroupBox("FFT Backend") + bl = QVBoxLayout(grp_backend) + + self._btn_scipy = QRadioButton("SciPy FFT (pocketfft) (always available)") + self._btn_pyfftw = QRadioButton( + "pyFFTW (faster for large arrays)" if PYFFTW_AVAILABLE + else "pyFFTW (not installed — run: pip install pyfftw)") + self._btn_pyfftw.setEnabled(PYFFTW_AVAILABLE) + + self._backend_group = QButtonGroup(self) + self._backend_group.addButton(self._btn_scipy, id=0) + self._backend_group.addButton(self._btn_pyfftw, id=1) + + if current_backend == "pyfftw" and PYFFTW_AVAILABLE: + self._btn_pyfftw.setChecked(True) + else: + self._btn_scipy.setChecked(True) + + bl.addWidget(self._btn_scipy) + bl.addWidget(self._btn_pyfftw) + layout.addWidget(grp_backend) + + # ---- Zero-padding ---------------------------------------------- + grp_zp = QGroupBox("Zero-Padding") + zl = QVBoxLayout(grp_zp) + + pad_row = QHBoxLayout() + pad_row.addWidget(QLabel("Pad factor:")) + self._spin_pad = QSpinBox() + self._spin_pad.setRange(1, 256) + self._spin_pad.setValue(max(1, current_pad_factor)) + self._spin_pad.setToolTip( + "Multiply the waveform length by this factor via zero-padding\n" + "before computing the FFT.\n" + "1 = no padding (natural length).\n" + "Powers of 2 (2, 4, 8 …) give the best performance." + ) + self._spin_pad.valueChanged.connect(self._update_info) + pad_row.addWidget(self._spin_pad) + zl.addLayout(pad_row) + + self._lbl_nfft = QLabel() + self._lbl_freq_res = QLabel() + self._lbl_vel_res = QLabel() + for lbl in (self._lbl_nfft, self._lbl_freq_res, self._lbl_vel_res): + lbl.setStyleSheet(_CSS_HINT) + zl.addWidget(lbl) + + layout.addWidget(grp_zp) + + # ---- Buttons --------------------------------------------------- + buttons = QDialogButtonBox() + buttons.addButton("Apply", QDialogButtonBox.ButtonRole.AcceptRole + ).clicked.connect(self.accept) + buttons.addButton("Cancel", QDialogButtonBox.ButtonRole.RejectRole + ).clicked.connect(self.reject) + layout.addWidget(buttons) + + self._update_info() + + def _update_info(self): + spf = self._samples_per_frame + sr = self._sample_rate_hz + pad = self._spin_pad.value() + + if spf is None or sr is None: + self._lbl_nfft.setText("Load a file to preview FFT parameters.") + self._lbl_freq_res.setText("") + self._lbl_vel_res.setText("") + return + + n_fft = spf * pad + freq_res_hz = sr / n_fft + freq_res_mhz = freq_res_hz / 1e6 + # v (m/s) = freq (MHz) × grating (µm) + vel_res_ms = freq_res_mhz * self._grating_um + + self._lbl_nfft.setText(f"FFT points: {spf} × {pad} = {n_fft:,}") + self._lbl_freq_res.setText( + f"Frequency bin: {freq_res_mhz:.4f} MHz ({freq_res_hz / 1e3:.2f} kHz)") + self._lbl_vel_res.setText( + f"Velocity bin: {vel_res_ms:.3f} m/s " + f"(at grating = {self._grating_um:.2f} µm)") + + def get_backend(self) -> str: + return "pyfftw" if self._btn_pyfftw.isChecked() and PYFFTW_AVAILABLE else "scipy" + + def get_pad_factor(self) -> int: + return max(1, self._spin_pad.value()) + + + +class ManualAlignmentDialog(QDialog): + """Non-modal manual angle-alignment editor (Fusion -> Manual Alignment...). + + Shows every angle's binarized CH4 (Bias B) mask overlaid in a distinct + color at partial opacity on one shared canvas, so translation/rotation + misalignment is visible by eye. Reference angle (always index 0) is + ground truth and never moves; every other angle is aligned to it. The + user picks an "active" angle and nudges its rotation+translation with + the keyboard; Auto Cross-Correlate finds every non-reference angle's + rotation *and* translation by registering its image against the + reference's (see compute.register_angle_to_reference) — meant to get every + angle stacked on top of each other so keyboard nudging only has to make + small corrections, not find an alignment from scratch; Auto De-rotate is + the weaker fallback that just seeds rotation from the stage's reported + angle, leaving translation alone. Save writes a JSON sidecar next to the + .sras file and hands a freshly-built, full-resolution AlignmentResult back + to the main window — the exact same object shape compute_angle_alignment + produces, so every existing Aligned-View code path (apply_alignment, + _aligned_canvas_axes, the pixel-inspector inverse-transform) works + completely unmodified. + + Non-modal by design (shown via .show(), never .exec() or setModal(True)) + so the user can still interact with the main window. Talks back to + SrasViewerWindow two ways: it reuses parent._run_worker/_jobs directly + for its background mask-fetch and cross-correlate steps, so the main + window's existing shutdown/lifecycle plumbing covers both for free, and + it emits alignment_saved / alignment_cleared signals for the two moments + that should actually mutate the main window's persistent state — + everything else (nudging, Auto De-rotate, Auto Cross-Correlate, threshold + edits) stays purely local to this dialog until Save. + """ + + alignment_saved = pyqtSignal(object, str) # AlignmentResult, sidecar path (str) + alignment_cleared = pyqtSignal() + + _PREVIEW_MARGIN_FRAC = 0.15 + _BASE_ALPHA = 0.42 + _ACTIVE_ALPHA = 0.75 + _MAX_PREVIEW_DIM = 1024 + + # (label, sources passed to compute.register_angle_to_reference). "Both" + # registers on each and keeps whichever scores higher per angle, which + # costs roughly double but removes the failure mode where the single + # chosen source is the one that happens to be uninformative for one angle. + _CORRELATE_SOURCES = ( + ("Both, keep best (recommended)", ("signal", "mask")), + ("Raw signal", ("signal",)), + ("Thresholded mask", ("mask",)), + ) + + def __init__(self, parent: "SrasViewerWindow", sras: SrasFile, *, + ref_angle_idx: int, dc_threshold_mv: float, + seed_per_angle: dict[int, ManualAngleParams] | None, + cached_dc4_mv: dict[int, np.ndarray]): + super().__init__(parent) + self._parent = parent + self._sras = sras + self._ref_angle_idx = ref_angle_idx + self._downsample = (1, 1) # (rows, cols) block-mean factors + self._dc4_mv: dict[int, np.ndarray] = {} + self._masks_small: dict[int, np.ndarray] = {} + self._preview_layers: dict[int, np.ndarray] = {} + self._preview_origin_mm = (0.0, 0.0) + self._preview_shape = (1, 1) + self._preview_pitch_mm = (1.0, 1.0) + self._masks_ready = False + self._fit_notes: dict[int, tuple[float, str]] = {} + self._derotate_sign_flipped = False + + self.setWindowTitle(f"Manual Alignment — {sras.path.name}") + self.resize(1150, 760) + + self._seed_initial_params(seed_per_angle) + n = sras.n_angles + cmap = mpl.colormaps["tab10"] if n <= 10 else mpl.colormaps["tab20"] + self._angle_colors = {a: cmap(a % cmap.N)[:3] for a in range(n)} + self._active_angle = 1 if ref_angle_idx == 0 and n > 1 else 0 + + self._build_ui(dc_threshold_mv) + self._set_controls_enabled(False) # re-enabled once masks are ready + self._start_mask_prep(cached_dc4_mv) + + def showEvent(self, event): + super().showEvent(event) + self.canvas.setFocus() + + # ------------------------------------------------------------------ + # Construction + # ------------------------------------------------------------------ + + def _seed_initial_params(self, seed_per_angle: dict[int, ManualAngleParams] | None): + seed = seed_per_angle or {} + self._angle_params: dict[int, ManualAngleParams] = { + a: (ManualAngleParams(seed[a].rotation_deg, seed[a].shift_mm) + if a in seed else ManualAngleParams()) + for a in range(self._sras.n_angles) + } + self._angle_params[self._ref_angle_idx] = ManualAngleParams() + + def _build_ui(self, dc_threshold_mv: float): + root = QHBoxLayout(self) + + self.canvas = ManualAlignOverlayCanvas() + left = QWidget() + left_l = QVBoxLayout(left) + left_l.setContentsMargins(0, 0, 0, 0) + left_l.setSpacing(4) + left_l.addWidget(NavigationToolbar2QT(self.canvas, left)) + left_l.addWidget(self.canvas) + root.addWidget(left, stretch=1) + + panel = QWidget() + panel_l = QVBoxLayout(panel) + panel_l.setContentsMargins(0, 0, 0, 0) + panel_l.setSpacing(8) + + # ---- Active Angle ------------------------------------------------- + grp_angle, al = _group("Active Angle") + self.combo_active_angle = QComboBox() + for a in range(self._sras.n_angles): + label = f"Angle {a} ({self._sras.angles_deg[a]:.1f}°)" + if a == self._ref_angle_idx: + label += " [reference]" + self.combo_active_angle.addItem(label) + al.addWidget(self.combo_active_angle) + self.lbl_active_note = _wrap_label("", _CSS_WARN) + al.addWidget(self.lbl_active_note) + panel_l.addWidget(grp_angle) + + # ---- Manual Adjustment --------------------------------------------- + self.grp_manual_adjust, mform_box = _group("Manual Adjustment") + mform = _form() + self.spin_active_rotation_deg = QDoubleSpinBox() + self.spin_active_rotation_deg.setRange(-3600.0, 3600.0) + self.spin_active_rotation_deg.setDecimals(3) + self.spin_active_rotation_deg.setSuffix(" °") + self.spin_active_rotation_deg.setMinimumWidth(_SPIN_MIN_W) + mform.addRow("Rotation:", self.spin_active_rotation_deg) + + self.spin_active_shift_x_mm = QDoubleSpinBox() + self.spin_active_shift_x_mm.setRange(-1e5, 1e5) + self.spin_active_shift_x_mm.setDecimals(4) + self.spin_active_shift_x_mm.setSuffix(" mm") + self.spin_active_shift_x_mm.setMinimumWidth(_SPIN_MIN_W) + mform.addRow("Shift X:", self.spin_active_shift_x_mm) + + self.spin_active_shift_y_mm = QDoubleSpinBox() + self.spin_active_shift_y_mm.setRange(-1e5, 1e5) + self.spin_active_shift_y_mm.setDecimals(4) + self.spin_active_shift_y_mm.setSuffix(" mm") + self.spin_active_shift_y_mm.setMinimumWidth(_SPIN_MIN_W) + mform.addRow("Shift Y:", self.spin_active_shift_y_mm) + mform_box.addLayout(mform) + panel_l.addWidget(self.grp_manual_adjust) + + # ---- Nudge Step Sizes ------------------------------------------------ + self.grp_step_sizes, sl = _group("Nudge Step Sizes") + sform = _form() + self.spin_step_translate_mm = QDoubleSpinBox() + self.spin_step_translate_mm.setRange(0.0001, 1000.0) + self.spin_step_translate_mm.setDecimals(4) + self.spin_step_translate_mm.setSuffix(" mm") + self.spin_step_translate_mm.setValue(0.01) + self.spin_step_translate_mm.setMinimumWidth(_SPIN_MIN_W) + sform.addRow("Translate step:", self.spin_step_translate_mm) + + self.spin_step_rotate_deg = QDoubleSpinBox() + self.spin_step_rotate_deg.setRange(0.001, 90.0) + self.spin_step_rotate_deg.setDecimals(3) + self.spin_step_rotate_deg.setSuffix(" °") + self.spin_step_rotate_deg.setValue(0.1) + self.spin_step_rotate_deg.setMinimumWidth(_SPIN_MIN_W) + sform.addRow("Rotate step:", self.spin_step_rotate_deg) + + self.spin_step_multiplier = QDoubleSpinBox() + self.spin_step_multiplier.setRange(1.0, 1000.0) + self.spin_step_multiplier.setDecimals(1) + self.spin_step_multiplier.setValue(10.0) + self.spin_step_multiplier.setMinimumWidth(_SPIN_MIN_W) + sform.addRow("Coarse × (Shift):", self.spin_step_multiplier) + sl.addLayout(sform) + sl.addWidget(_wrap_label( + "Arrow keys nudge X/Y translation; Q/E nudge rotation (CCW/CW). " + "Hold Shift for the coarse step. Click the image once so it has " + "keyboard focus.", _CSS_HINT)) + panel_l.addWidget(self.grp_step_sizes) + + # ---- Mask Threshold --------------------------------------------------- + self.grp_mask_threshold, tl = _group("Mask Threshold") + tform = _form() + self.spin_mask_threshold_mv = QDoubleSpinBox() + self.spin_mask_threshold_mv.setRange(-500.0, 500.0) + self.spin_mask_threshold_mv.setDecimals(3) + self.spin_mask_threshold_mv.setSuffix(" mV") + self.spin_mask_threshold_mv.setValue(dc_threshold_mv) + self.spin_mask_threshold_mv.setMinimumWidth(_SPIN_MIN_W) + tform.addRow("DC threshold:", self.spin_mask_threshold_mv) + tl.addLayout(tform) + panel_l.addWidget(self.grp_mask_threshold) + + # ---- Cross-Correlate (FFT) ----------------------------------------- + self.grp_correlate, cl = _group("Cross-Correlate (FFT)") + cform = _form() + self.combo_correlate_source = QComboBox() + for label, sources in self._CORRELATE_SOURCES: + self.combo_correlate_source.addItem(label, sources) + cform.addRow("Correlate on:", self.combo_correlate_source) + + self.spin_correlate_search_deg = QDoubleSpinBox() + self.spin_correlate_search_deg.setRange(0.0, 180.0) + self.spin_correlate_search_deg.setSingleStep(1.0) + self.spin_correlate_search_deg.setDecimals(1) + self.spin_correlate_search_deg.setSuffix(" °") + self.spin_correlate_search_deg.setValue(6.0) + self.spin_correlate_search_deg.setMinimumWidth(_SPIN_MIN_W) + cform.addRow("Rotation search (±):", self.spin_correlate_search_deg) + cl.addLayout(cform) + self.btn_auto_correlate = QPushButton("Auto Cross-Correlate (vs Reference)") + cl.addWidget(self.btn_auto_correlate) + cl.addWidget(_wrap_label( + "Finds each non-reference angle's rotation *and* translation by " + "cross-correlating its image against the reference's — the stage's " + "reported angle is only the starting point of the search, and both " + "of its signs are tried. Run this first, then nudge only for small " + "corrections.", _CSS_HINT)) + panel_l.addWidget(self.grp_correlate) + + # ---- Actions ------------------------------------------------------ + grp_actions, acl = _group("Actions") + self.btn_auto_derotate = QPushButton("Auto De-rotate (use known angles)") + self.btn_save = QPushButton("Save Alignment") + self.btn_clear = QPushButton("Clear Alignment…") + self.btn_close = QPushButton("Close") + for btn in (self.btn_auto_derotate, self.btn_save, self.btn_clear, self.btn_close): + acl.addWidget(btn) + panel_l.addWidget(grp_actions) + + self.lbl_status = _wrap_label("", _CSS_MUTED) + panel_l.addWidget(self.lbl_status) + panel_l.addStretch() + + root.addWidget(_scroll_panel(panel, 320)) + + self.combo_active_angle.currentIndexChanged.connect(self._on_active_angle_changed) + self.spin_active_rotation_deg.editingFinished.connect(self._on_rotation_spin_edited) + self.spin_active_shift_x_mm.editingFinished.connect(self._on_shift_spin_edited) + self.spin_active_shift_y_mm.editingFinished.connect(self._on_shift_spin_edited) + self.spin_mask_threshold_mv.editingFinished.connect(self._on_mask_threshold_edited) + self.btn_auto_derotate.clicked.connect(self._on_auto_derotate) + self.btn_auto_correlate.clicked.connect(self._on_auto_correlate) + self.btn_save.clicked.connect(self._on_save) + self.btn_clear.clicked.connect(self._on_clear) + self.btn_close.clicked.connect(self.close) + self.canvas.nudge_translate.connect(self._on_nudge_translate) + self.canvas.nudge_rotate.connect(self._on_nudge_rotate) + + self.combo_active_angle.blockSignals(True) + self.combo_active_angle.setCurrentIndex(self._active_angle) + self.combo_active_angle.blockSignals(False) + self._on_active_angle_changed(self._active_angle) + + # ------------------------------------------------------------------ + # Mask preparation (initial CH4 fetch + threshold + downsample) + # ------------------------------------------------------------------ + + def _start_mask_prep(self, cached_dc4_mv: dict[int, np.ndarray]): + self._dc4_mv = dict(cached_dc4_mv) + missing = [a for a in range(self._sras.n_angles) if a not in self._dc4_mv] + if not missing: + self._finish_mask_prep() + return + self.lbl_status.setText(f"Preparing masks: 0/{len(missing)} angle(s) needed…") + started = self._parent._run_worker( + "manual_align_masks", Ch4MaskWorker(self._sras, missing), + connect=( + ("angle_done", self._on_mask_angle_done), + ("error", lambda msg: self.lbl_status.setText(f"Mask prep error: {msg}")), + ), + on_done=self._finish_mask_prep) + if not started: + self.lbl_status.setText( + "Could not start mask preparation (busy) — close and reopen.") + + def _on_mask_angle_done(self, angle_idx: int, dc4_mv: np.ndarray): + self._dc4_mv[angle_idx] = dc4_mv + self.lbl_status.setText( + f"Preparing masks: {len(self._dc4_mv)}/{self._sras.n_angles} ready…") + + def _finish_mask_prep(self): + if len(self._dc4_mv) < self._sras.n_angles: + return # a mask-worker error left some angles unfetched + # Rows and columns get their own factor. A real scan is ~7500 frames + # wide but only ~750 rows tall, so one shared factor sized for the + # frames would throw away 8x more row detail than the preview needs and + # leave the overlay too coarse in y to judge alignment by eye. + max_rows = max(img.shape[0] for img in self._dc4_mv.values()) + max_cols = max(img.shape[1] for img in self._dc4_mv.values()) + self._downsample = ( + max(1, int(np.ceil(max_rows / self._MAX_PREVIEW_DIM))), + max(1, int(np.ceil(max_cols / self._MAX_PREVIEW_DIM)))) + self._recompute_masks_small() + self._rebuild_preview_canvas() + self._set_controls_enabled(True) + self.lbl_status.setText("Ready.") + + def _recompute_masks_small(self): + """Threshold + downsample every angle's already-in-memory full-res + CH4 mV image. Cheap (a compare + block-mean), so this re-runs in + full whenever the mask-threshold spin box changes — no re-fetch. + Purely for the overlay's visuals: no alignment geometry depends on this + threshold, only which pixels the overlay paints.""" + threshold = self.spin_mask_threshold_mv.value() + fy, fx = self._downsample + self._masks_small = { + a: compute.block_mean_2d((img >= threshold).astype(np.float32), fy, fx) + for a, img in self._dc4_mv.items() + } + + # ------------------------------------------------------------------ + # Preview canvas: full rebuild vs. incremental single-layer refresh + # ------------------------------------------------------------------ + + def _rebuild_preview_canvas(self): + """Full geometry rebuild: recomputes the shared preview canvas's + origin/shape (rotation can grow the union bbox — translation alone + cannot, per the padding baked in via _PREVIEW_MARGIN_FRAC) and every + angle's reprojected mask layer. Triggered by: dialog open, + mask-threshold change, Auto De-rotate, a rotation nudge/edit of the + active angle. NOT triggered by a translation-only nudge — see + _refresh_active_preview_layer.""" + dx_ref, dy_ref = compute.pixel_pitch_mm(self._sras, self._ref_angle_idx) + fy, fx = self._downsample + pitch = (dx_ref * fx, dy_ref * fy) + origin, shape = compute.canvas_for_params( + self._sras, self._ref_angle_idx, pitch, self._angle_params, + margin_frac=self._PREVIEW_MARGIN_FRAC, snap=False) + self._preview_origin_mm, self._preview_shape = origin, shape + self._preview_pitch_mm = pitch + self._preview_layers = { + a: self._reproject(a) for a in range(self._sras.n_angles) + } + self._redraw_overlay() + + def _reproject(self, angle_idx: int) -> np.ndarray: + """One angle's downsampled mask on the current preview canvas. + src_downsample must match _masks_small's block-mean factors, or the + layer lands magnified and offset instead of where the alignment + actually puts it.""" + p = self._angle_params[angle_idx] + return compute.reproject_mask( + self._sras, angle_idx, self._ref_angle_idx, + self._masks_small[angle_idx], p.rotation_deg, p.shift_mm, + self._preview_pitch_mm, self._preview_origin_mm, self._preview_shape, + src_downsample=self._downsample) + + def _refresh_active_preview_layer(self): + """Cheap path for a translation-only nudge/edit of the active angle: + reproject just that one angle's downsampled mask onto the *existing* + preview canvas — every other angle's cached layer is untouched.""" + self._preview_layers[self._active_angle] = self._reproject(self._active_angle) + self._redraw_overlay() + + def _redraw_overlay(self): + """Alpha-composite every angle's colored mask layer into one RGBA + image ("all thresholds overlaid with varying opacity"). Each angle + keeps a fixed, distinct color regardless of which is active; the + active angle is drawn last (on top) at a visibly higher alpha so + it's easy to track while nudging.""" + if not self._preview_layers: + return # mask prep hasn't finished yet — nothing to draw + n_rows, n_cols = self._preview_shape + rgba = np.zeros((n_rows, n_cols, 4), dtype=np.float32) + order = sorted(range(self._sras.n_angles), key=lambda a: a == self._active_angle) + for a in order: + layer = self._preview_layers.get(a) + if layer is None: + continue + alpha = self._ACTIVE_ALPHA if a == self._active_angle else self._BASE_ALPHA + color = self._angle_colors[a] + fg_a = layer * alpha + for c in range(3): + rgba[..., c] = color[c] * fg_a + rgba[..., c] * rgba[..., 3] * (1 - fg_a) + rgba[..., 3] = fg_a + rgba[..., 3] * (1 - fg_a) + + x0, y0 = self._preview_origin_mm + dx, dy = self._preview_pitch_mm + x_axis = x0 + np.arange(n_cols) * dx + y_axis = y0 + np.arange(n_rows) * dy + extent = [x_axis[0] - dx / 2, x_axis[-1] + dx / 2, + y_axis[-1] + dy / 2, y_axis[0] - dy / 2] + title = (f"Angle {self._active_angle} active " + f"({self._sras.angles_deg[self._active_angle]:.1f}°)") + self.canvas.show_overlay(rgba, extent, title) + + # ------------------------------------------------------------------ + # Angle selection / nudge / edit handlers + # ------------------------------------------------------------------ + + def _on_active_angle_changed(self, angle_idx: int): + self._active_angle = angle_idx + is_ref = angle_idx == self._ref_angle_idx + self.grp_manual_adjust.setEnabled(self._masks_ready and not is_ref) + self.lbl_active_note.setText( + "Reference angle — defines the shared origin, not adjustable." if is_ref else "") + self._sync_active_spinboxes() + self._redraw_overlay() + + def _sync_active_spinboxes(self): + p = self._angle_params[self._active_angle] + for spin, val in ((self.spin_active_rotation_deg, p.rotation_deg), + (self.spin_active_shift_x_mm, p.shift_mm[0]), + (self.spin_active_shift_y_mm, p.shift_mm[1])): + spin.blockSignals(True) + spin.setValue(val) + spin.blockSignals(False) + + def _on_nudge_translate(self, dir_x: int, dir_y: int, coarse: bool): + if not self._masks_ready or self._active_angle == self._ref_angle_idx: + return + step = self.spin_step_translate_mm.value() + if coarse: + step *= self.spin_step_multiplier.value() + p = self._angle_params[self._active_angle] + p.shift_mm = (p.shift_mm[0] + dir_x * step, p.shift_mm[1] + dir_y * step) + self._sync_active_spinboxes() + self._refresh_active_preview_layer() + + def _on_nudge_rotate(self, direction: int, coarse: bool): + if not self._masks_ready or self._active_angle == self._ref_angle_idx: + return + step = self.spin_step_rotate_deg.value() + if coarse: + step *= self.spin_step_multiplier.value() + self._angle_params[self._active_angle].rotation_deg += direction * step + self._sync_active_spinboxes() + self._rebuild_preview_canvas() + + def _on_rotation_spin_edited(self): + if self._active_angle == self._ref_angle_idx: + return + self._angle_params[self._active_angle].rotation_deg = self.spin_active_rotation_deg.value() + self._rebuild_preview_canvas() + + def _on_shift_spin_edited(self): + if self._active_angle == self._ref_angle_idx: + return + p = self._angle_params[self._active_angle] + p.shift_mm = (self.spin_active_shift_x_mm.value(), self.spin_active_shift_y_mm.value()) + self._refresh_active_preview_layer() + + def _on_mask_threshold_edited(self): + if not self._masks_ready: + return + self._recompute_masks_small() + self._rebuild_preview_canvas() + + # ------------------------------------------------------------------ + # Actions + # ------------------------------------------------------------------ + + def _on_auto_derotate(self): + """Seed every angle's rotation from the stage's reported angle. + + A starting point for nudging by eye, not an alignment: the stage's + sign convention relative to this module's is not knowable from the + file, so the sign that lines the scans up is whichever of the two looks + right in the overlay. Auto Cross-Correlate decides that from the images + instead, and is the button to reach for first. + """ + sign = -1.0 if self._derotate_sign_flipped else 1.0 + self._derotate_sign_flipped = not self._derotate_sign_flipped + n_changed = 0 + for a in range(self._sras.n_angles): + if a == self._ref_angle_idx: + continue + self._angle_params[a].rotation_deg = sign * compute.nominal_delta_deg( + self._sras, a, self._ref_angle_idx) + n_changed += 1 + self._sync_active_spinboxes() + self._rebuild_preview_canvas() + self.lbl_status.setText( + f"Rotation set to the stage angle ({'−' if sign < 0 else '+'}delta) " + f"for {n_changed} angle(s); translation untouched. Click again to " + "try the opposite sign.") + + def _on_auto_correlate(self): + if not self._masks_ready: + return + angles = [a for a in range(self._sras.n_angles) if a != self._ref_angle_idx] + if not angles: + return + worker = CrossCorrelateWorker( + self._sras, self._ref_angle_idx, angles, self._dc4_mv, + sources=self.combo_correlate_source.currentData(), + dc_threshold_mv=self.spin_mask_threshold_mv.value(), + search_deg=self.spin_correlate_search_deg.value()) + self._correlate_done_count = 0 + self._correlate_total = len(angles) + self._fit_notes = {} + self._set_controls_enabled(False) + self.lbl_status.setText(f"Cross-correlating: 0/{self._correlate_total} angle(s)…") + started = self._parent._run_worker( + "manual_align_correlate", worker, + connect=( + ("angle_done", self._on_correlate_angle_done), + ("error", self._on_correlate_error), + ), + on_done=self._finish_auto_correlate) + if not started: + self._set_controls_enabled(True) + self.lbl_status.setText("Could not start cross-correlation (busy) — try again.") + + def _on_correlate_angle_done(self, angle_idx: int, rotation_deg: float, + shift_x_mm: float, shift_y_mm: float, + score: float, source: str): + self._angle_params[angle_idx] = ManualAngleParams(rotation_deg, (shift_x_mm, shift_y_mm)) + self._fit_notes[angle_idx] = (score, source) + self._correlate_done_count += 1 + self.lbl_status.setText( + f"Cross-correlating: {self._correlate_done_count}/{self._correlate_total} angle(s)…") + + def _on_correlate_error(self, msg: str): + self.lbl_status.setText(f"Cross-correlation error: {msg}") + + def _finish_auto_correlate(self): + self._sync_active_spinboxes() + self._rebuild_preview_canvas() + self._set_controls_enabled(True) + self.lbl_status.setText( + f"Cross-correlated {self._correlate_done_count} angle(s) against " + f"Angle {self._ref_angle_idx}.\n" + self._fit_report()) + + def _fit_report(self) -> str: + """Per-angle registration quality, worst first. + + Surfaced rather than buried because a single bad acquisition (stage + glitch, laser dropout) registers poorly and would otherwise be fused in + silently — seeing which angle it is, is what makes dropping it with + sras_edit_scans.py actionable. The deviation from the stage's own + reported angle is shown alongside: a large one means the search and the + stage disagree, which is either a genuine mechanical error or a sign + that this angle's fit is not to be trusted. + """ + if not self._fit_notes: + return "" + rows = sorted(self._fit_notes.items(), key=lambda kv: kv[1][0]) + worst = rows[0] + lines = [f"Worst fit: angle {worst[0]} (score {worst[1][0]:.3f}, " + f"{worst[1][1]})."] + drifted = [] + for a, _note in rows: + nominal = compute.nominal_delta_deg(self._sras, a, self._ref_angle_idx) + got = self._angle_params[a].rotation_deg + dev = min(abs(got - nominal), abs(got + nominal)) + if dev > 1.0: + drifted.append(f"{a} ({dev:.2f}°)") + if drifted: + lines.append("Rotation differs from the stage angle by >1° for " + "angle(s) " + ", ".join(drifted) + ".") + lines.append("Nudge from here for any remaining fine correction.") + return " ".join(lines) + + def _on_save(self): + threshold = self.spin_mask_threshold_mv.value() + resolved = dict(self._angle_params) # already concrete floats + try: + path = save_manual_alignment(self._sras, self._ref_angle_idx, threshold, resolved) + result = build_manual_alignment(self._sras, self._ref_angle_idx, + threshold, resolved) + except OSError as exc: + QMessageBox.warning(self, "Save Alignment Failed", str(exc)) + return + self.lbl_status.setText(f"Saved to {path.name}.") + self.alignment_saved.emit(result, str(path)) + + def _on_clear(self): + reply = QMessageBox.question( + self, "Clear Alignment", + "This resets every angle back to raw/unaligned (0° rotation, no " + "shift) and deletes the saved alignment file for this scan, if " + "any. This cannot be undone. Continue?", + QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No, + QMessageBox.StandardButton.No) + if reply != QMessageBox.StandardButton.Yes: + return + try: + existed = delete_manual_alignment(self._sras) + except OSError as exc: + QMessageBox.warning(self, "Clear Alignment Failed", + f"Could not delete the saved alignment file: {exc}") + return + self._angle_params = {a: ManualAngleParams() for a in range(self._sras.n_angles)} + self._fit_notes = {} + self._sync_active_spinboxes() + self._rebuild_preview_canvas() + self.lbl_status.setText( + "Alignment cleared; saved file removed." if existed + else "Alignment cleared (there was no saved file).") + self.alignment_cleared.emit() + + def _set_controls_enabled(self, enabled: bool): + self._masks_ready = enabled + self.combo_active_angle.setEnabled(enabled) + self.grp_manual_adjust.setEnabled(enabled and self._active_angle != self._ref_angle_idx) + self.grp_step_sizes.setEnabled(enabled) + self.grp_mask_threshold.setEnabled(enabled) + self.grp_correlate.setEnabled(enabled) + self.btn_auto_derotate.setEnabled(enabled) + self.btn_save.setEnabled(enabled) + self.btn_clear.setEnabled(enabled) + + diff --git a/sras_viewer/main_window.py b/sras_viewer/main_window.py new file mode 100644 index 0000000..b0fa6a5 --- /dev/null +++ b/sras_viewer/main_window.py @@ -0,0 +1,1403 @@ +"""The SrasViewerWindow main window and application entry point.""" + +import sys +from pathlib import Path + +import numpy as np +from matplotlib.backends.backend_qtagg import NavigationToolbar2QT +from PyQt6.QtCore import QObject, QSettings, Qt, QThread +from PyQt6.QtGui import QAction +from PyQt6.QtWidgets import ( + QApplication, QCheckBox, QComboBox, QDialog, QDoubleSpinBox, QFileDialog, + QFrame, QHBoxLayout, QLabel, QMainWindow, QProgressDialog, QPushButton, + QSizePolicy, QSpinBox, QSplitter, QVBoxLayout, QWidget, +) + +import sras_compute as compute +from sras_compute import ( + ManualAngleParams, apply_alignment, build_manual_alignment, + load_manual_alignment, sidecar_path, +) +from sras_format import ( + CH3_IDX, CH4_IDX, CH_NAMES, SrasFile, mv_to_adc, + _FALLBACK_YMULT_MV, _FALLBACK_YOFF_ADC, +) +from sras_workers import ( + AngleAlignmentWorker, BatchCacheWorker, ComputeWorker, DcPrecomputeWorker, + LoadWorker, +) + +from .canvases import ImageCanvas, WaveformCanvas +from .common import ( + CH1_DERIVED_MODES, CH_LABELS, CMAPS, VELOCITY_MODE_IDX, _CHANNEL_DISPLAY, + _CSS_BUSY, _CSS_HINT, _CSS_INFO, _CSS_MUTED, _CSS_WARN, _LEFT_PANEL_W, + _RIGHT_PANEL_W, _SPIN_MIN_W, _form, _group, _scroll_panel, _wrap_label, +) +from .dialogs import FftOptionsDialog, ManualAlignmentDialog + +# --------------------------------------------------------------------------- +# Main window +# --------------------------------------------------------------------------- + +class SrasViewerWindow(QMainWindow): + def __init__(self, initial_path: str | None = None): + super().__init__() + self.setWindowTitle("SRAS Scan Viewer") + self.resize(1560, 840) + self.setMinimumSize(960, 560) + self.setAcceptDrops(True) + + self._sras: SrasFile | None = None + self._current_image: np.ndarray | None = None + self._current_angle: int = 0 + self._current_ch: int = 0 + self._pending_angle: int = 0 + self._pending_ch: int = 0 + self._pending_bg_sub: bool = True + self._pending_threshold: float = 50.0 # mV + self._pending_fft_pad_factor: int = 1 + + # Live background jobs, keyed by role — see _run_worker. + self._jobs: dict[str, tuple] = {} + self._progress_dlgs: dict[str, QProgressDialog] = {} + + # FFT settings (configured via FFT Options dialog, persisted across + # sessions). IniFormat: predictable cross-platform and redirectable + # in tests. + self._settings = QSettings(QSettings.Format.IniFormat, + QSettings.Scope.UserScope, + "sras-viewer", "sras-viewer") + compute.set_fft_backend(str(self._settings.value("fft/backend", "scipy"))) + try: + pad = int(self._settings.value("fft/pad_factor", 1)) + except (TypeError, ValueError): + pad = 1 + self._fft_pad_factor: int = max(1, min(256, pad)) # 1 = no padding + + # Convert menu: batch DC/FFT compute-and-store (v6 -> v7) + self._batch_errors: list[str] = [] + + # Display-only settings (colormap, grating) never trigger a + # recompute — they're applied to cached data on redraw. DC images + # (CH3/CH4) are cheap and precomputed for every angle in the + # background right after load. CH1/Velocity FFT images are + # computed lazily (with a progress popup) the first time an + # angle/threshold combination is viewed — using the cached DC4 + # image to skip the FFT entirely for masked-out pixels — and + # cached per (angle, bg_sub, n_fft, threshold) so revisiting the + # same combination is free. + self._dc_cache: dict[tuple[int, int], np.ndarray] = {} + self._fft_cache: dict[tuple[int, bool, int | None, float], np.ndarray] = {} + self._dc_generation: int = 0 + + # Angle alignment ("Fusion" menu) + self._alignment_result = None + self._alignment_generation: int = 0 + self._aligned_cache: dict[tuple, np.ndarray] = {} + self._manual_align_dialog: ManualAlignmentDialog | None = None + + self._build_ui() + + if initial_path: + self._load_file(initial_path) + + # ------------------------------------------------------------------ + # Background job plumbing + # ------------------------------------------------------------------ + + def _run_worker(self, key: str, worker: QObject, *, + connect: tuple = (), quit_on: tuple = ("finished",), + on_done=None) -> bool: + """Move *worker* onto its own QThread and start it. Returns False if + a job under *key* is already running. + + Centralises two lifetime hazards that each cost a process abort: + + 1. The job is claimed in self._jobs *before* start() and before + anything below that can pump the Qt event loop (a + QProgressDialog.show() does on first display). If it weren't, a + re-entrant editingFinished could slip past the busy check, start a + second thread, and then have the first call's own assignment + clobber — and destroy while still running — that second QThread. + + 2. thread.finished fires as the thread winds down but does not + guarantee the OS thread has joined. Dropping the last reference to + a QThread whose thread is still running logs "QThread: Destroyed + while thread is still running" and aborts, so wait() first. + """ + if key in self._jobs: + return False + + thread = QThread() + self._jobs[key] = (thread, worker, on_done) # claim before anything pumps + worker.moveToThread(thread) + thread.started.connect(worker.run) + for signal_name, slot in connect: + getattr(worker, signal_name).connect(slot) + for signal_name in quit_on: + getattr(worker, signal_name).connect(thread.quit) + thread.finished.connect(lambda k=key: self._on_job_finished(k)) + thread.start() + return True + + def _on_job_finished(self, key: str): + job = self._jobs.pop(key, None) + if job is None: + return + thread, _worker, on_done = job + thread.wait() # join before releasing our last reference + if on_done is not None: + on_done() + + def _job_running(self, key: str) -> bool: + return key in self._jobs + + # ------------------------------------------------------------------ + # UI construction + # ------------------------------------------------------------------ + + def _build_ui(self): + central = QWidget() + self.setCentralWidget(central) + root = QHBoxLayout(central) + root.setContentsMargins(8, 8, 8, 8) + root.setSpacing(8) + + root.addWidget(self._build_left_panel()) + root.addWidget(self._build_canvases(), stretch=1) + root.addWidget(self._build_right_panel()) + + self.statusBar().showMessage("Open an .sras file to begin.") + self._build_menus() + + def _build_left_panel(self) -> QWidget: + panel = QWidget() + panel_layout = QVBoxLayout(panel) + panel_layout.setContentsMargins(0, 0, 0, 0) + panel_layout.setSpacing(8) + + # ---- File ------------------------------------------------------- + grp_file, fl = _group("File") + self.btn_open = QPushButton("Open .sras…") + self.btn_open.clicked.connect(self._on_open) + self.lbl_filename = _wrap_label("No file loaded", _CSS_MUTED) + fl.addWidget(self.btn_open) + fl.addWidget(self.lbl_filename) + panel_layout.addWidget(grp_file) + + # ---- Scan info -------------------------------------------------- + grp_info, il = _group("Scan Info") + il.setSpacing(3) + self._info = {} + for key in ("Angles", "Rows", "Frames / row", "Samples / frame", + "Sample rate", "X start", "Pixel Δx", "Laser freq"): + lbl = _wrap_label(f"{key}: —", _CSS_INFO) + il.addWidget(lbl) + self._info[key] = lbl + + # frame-count / format notes + self.lbl_frame_warn = _wrap_label("", _CSS_WARN) + il.addWidget(self.lbl_frame_warn) + + # background DC-precompute progress + self.lbl_dc_precompute = _wrap_label("", _CSS_BUSY) + il.addWidget(self.lbl_dc_precompute) + panel_layout.addWidget(grp_info) + + # ---- View settings ---------------------------------------------- + grp_view, vl = _group("View Settings") + + view_form = _form() + + self.spin_angle = QSpinBox() + self.spin_angle.setRange(0, 0) + self.spin_angle.setEnabled(False) + self.spin_angle.setMinimumWidth(64) + self.spin_angle.editingFinished.connect(self._on_view_changed) + self.lbl_angle_deg = QLabel("—") + angle_field = QWidget() + ar = QHBoxLayout(angle_field) + ar.setContentsMargins(0, 0, 0, 0) + ar.setSpacing(6) + ar.addWidget(self.spin_angle) + ar.addWidget(self.lbl_angle_deg) + ar.addStretch() + view_form.addRow("Angle:", angle_field) + + self.combo_channel = QComboBox() + self.combo_channel.addItems(CH_LABELS) + self.combo_channel.setEnabled(False) + self.combo_channel.setSizePolicy(QSizePolicy.Policy.Expanding, + QSizePolicy.Policy.Fixed) + self.combo_channel.setSizeAdjustPolicy( + QComboBox.SizeAdjustPolicy.AdjustToMinimumContentsLengthWithIcon) + self.combo_channel.setMinimumContentsLength(12) + self.combo_channel.currentIndexChanged.connect(self._on_channel_changed) + view_form.addRow("Channel:", self.combo_channel) + vl.addLayout(view_form) + + sep = QFrame() + sep.setFrameShape(QFrame.Shape.HLine) + sep.setStyleSheet("color: #555;") + vl.addWidget(sep) + + # DC threshold (for RF / CH1 masking) + self.grp_threshold, tl = _group("RF Mask Threshold (CH1 only)") + thr_form = _form() + self.spin_threshold_mv = QDoubleSpinBox() + self.spin_threshold_mv.setRange(-500.0, 500.0) + self.spin_threshold_mv.setDecimals(3) + self.spin_threshold_mv.setSingleStep(0.025) + self.spin_threshold_mv.setSuffix(" mV") + self.spin_threshold_mv.setValue(50.0) + self.spin_threshold_mv.setEnabled(False) + self.spin_threshold_mv.setMinimumWidth(_SPIN_MIN_W) + self.spin_threshold_mv.editingFinished.connect(self._on_threshold_changed) + thr_form.addRow("DC threshold:", self.spin_threshold_mv) + tl.addLayout(thr_form) + self.lbl_threshold_adc = _wrap_label( + f"≈ {mv_to_adc(50.0):.1f} ADC counts", _CSS_MUTED) + tl.addWidget(self.lbl_threshold_adc) + vl.addWidget(self.grp_threshold) + + # Background subtraction (v4+ files only) + self.chk_bg_sub = QCheckBox("Background subtraction (CH1 only)") + self.chk_bg_sub.setChecked(True) + self.chk_bg_sub.setEnabled(False) + self.chk_bg_sub.setToolTip( + "Subtract the stored background waveform from each CH1 frame\n" + "before computing the FFT (v4+ files only)." + ) + self.chk_bg_sub.toggled.connect(self._on_bg_sub_toggled) + vl.addWidget(self.chk_bg_sub) + + # Aligned View (Fusion → Angle Alignment result) + self.chk_aligned_view = QCheckBox("Aligned View (Fusion)") + self.chk_aligned_view.setChecked(False) + self.chk_aligned_view.setEnabled(False) + self.chk_aligned_view.setToolTip( + "Show the current angle/channel resampled onto the shared,\n" + "rotation+translation-aligned canvas from Fusion → Angle\n" + "Alignment. Uncheck to see the raw per-angle scan grid." + ) + self.chk_aligned_view.toggled.connect(self._on_aligned_view_toggled) + vl.addWidget(self.chk_aligned_view) + + self.btn_export_csv = QPushButton("Export Image as CSV…") + self.btn_export_csv.setEnabled(False) + self.btn_export_csv.setToolTip( + "Save the current CH1 image (one scan row per CSV line).") + self.btn_export_csv.clicked.connect(self._on_export_csv) + vl.addWidget(self.btn_export_csv) + + panel_layout.addWidget(grp_view) + + # ---- ROI --------------------------------------------------------- + grp_roi, rl = _group("ROI (Region of Interest)") + + self.btn_draw_roi = QPushButton("Draw ROI") + self.btn_draw_roi.setCheckable(True) + self.btn_draw_roi.setEnabled(False) + self.btn_draw_roi.setToolTip( + "Arm next click+drag on the image to draw a new ROI\n" + "(replaces any existing one). Click again to cancel.\n" + "After drawing, drag inside to move, or grab corners to reshape.\n" + "The ROI is persistent across channels / modes / angles." + ) + self.btn_draw_roi.toggled.connect(self._on_draw_roi_toggled) + rl.addWidget(self.btn_draw_roi) + + self.btn_clear_roi = QPushButton("Clear ROI") + self.btn_clear_roi.setEnabled(False) + self.btn_clear_roi.clicked.connect(self._on_clear_roi) + rl.addWidget(self.btn_clear_roi) + + self.btn_export_roi = QPushButton("Export ROI as CSV…") + self.btn_export_roi.setEnabled(False) + self.btn_export_roi.setToolTip( + "Save every pixel whose centre lies inside the ROI as CSV.\n" + "Columns: row, frame, x_mm, y_mm, value.\n" + "Corner coordinates of the quad are written in the file header." + ) + self.btn_export_roi.clicked.connect(self._on_export_roi_csv) + rl.addWidget(self.btn_export_roi) + + self.lbl_roi_center = _wrap_label("centroid: —", _CSS_HINT) + self.lbl_roi_size = _wrap_label("bbox: —", _CSS_HINT) + self.lbl_roi_npix = _wrap_label("pixels inside: —", _CSS_HINT) + for lbl in (self.lbl_roi_center, self.lbl_roi_size, self.lbl_roi_npix): + rl.addWidget(lbl) + + panel_layout.addWidget(grp_roi) + panel_layout.addStretch() + return _scroll_panel(panel, _LEFT_PANEL_W) + + def _build_canvases(self) -> QWidget: + splitter = QSplitter(Qt.Orientation.Vertical) + splitter.setChildrenCollapsible(False) + + img_widget = QWidget() + img_vl = QVBoxLayout(img_widget) + img_vl.setContentsMargins(0, 0, 0, 0) + img_vl.setSpacing(4) + self.image_canvas = ImageCanvas() + self.image_canvas.setMinimumHeight(220) + self.image_canvas.pixel_clicked.connect(self._on_pixel_clicked) + self.image_canvas.roi_changed.connect(self._update_roi_ui) + self.image_canvas.draw_mode_changed.connect(self._on_draw_mode_changed) + img_vl.addWidget(NavigationToolbar2QT(self.image_canvas, img_widget)) + img_vl.addWidget(self.image_canvas) + splitter.addWidget(img_widget) + + wave_widget = QWidget() + wave_vl = QVBoxLayout(wave_widget) + wave_vl.setContentsMargins(0, 0, 0, 0) + wave_vl.setSpacing(4) + self.lbl_wave_hint = QLabel( + "Click a pixel in the image above to inspect its waveform.") + self.lbl_wave_hint.setAlignment(Qt.AlignmentFlag.AlignCenter) + self.lbl_wave_hint.setStyleSheet(_CSS_MUTED) + self.wave_canvas = WaveformCanvas() + self.wave_canvas.setMinimumHeight(150) + wave_vl.addWidget(self.lbl_wave_hint) + wave_vl.addWidget(self.wave_canvas) + splitter.addWidget(wave_widget) + + splitter.setStretchFactor(0, 3) + splitter.setStretchFactor(1, 1) + splitter.setSizes([580, 250]) + return splitter + + def _build_right_panel(self) -> QWidget: + # Velocity settings (visible only in velocity mode) + self.grp_velocity, vel_l = _group("Velocity Settings (CH1 only)") + vel_form = _form() + self.spin_grating_um = QDoubleSpinBox() + self.spin_grating_um.setRange(0.1, 1000.0) + self.spin_grating_um.setDecimals(2) + self.spin_grating_um.setSingleStep(0.5) + self.spin_grating_um.setSuffix(" µm") + self.spin_grating_um.setValue(25) + self.spin_grating_um.setEnabled(False) + self.spin_grating_um.setMinimumWidth(_SPIN_MIN_W) + self.spin_grating_um.editingFinished.connect(self._on_grating_changed) + vel_form.addRow("Grating size:", self.spin_grating_um) + vel_l.addLayout(vel_form) + vel_l.addWidget(_wrap_label("v (m/s) = freq (MHz) × grating (µm)", + "font-size: 10px; color: #888;")) + self.grp_velocity.setVisible(False) + + grp_display, dl = _group("Display Options") + + cmap_form = _form() + self.combo_cmap = QComboBox() + self.combo_cmap.addItems(CMAPS) + self.combo_cmap.setCurrentText("gray") + self.combo_cmap.setEnabled(False) + self.combo_cmap.setSizePolicy(QSizePolicy.Policy.Expanding, + QSizePolicy.Policy.Fixed) + self.combo_cmap.currentIndexChanged.connect(self._on_cmap_changed) + cmap_form.addRow("Colormap:", self.combo_cmap) + dl.addLayout(cmap_form) + + self.chk_auto = QCheckBox("Auto-scale colormap") + self.chk_auto.setChecked(True) + self.chk_auto.toggled.connect(self._on_autoscale_toggled) + dl.addWidget(self.chk_auto) + + range_form = _form() + for label, attr in (("min:", "spin_vmin"), ("max:", "spin_vmax")): + spin = QDoubleSpinBox() + spin.setRange(-1e9, 1e9) + spin.setDecimals(4) + spin.setEnabled(False) + spin.setMinimumWidth(_SPIN_MIN_W) + spin.editingFinished.connect(self._on_manual_range_changed) + setattr(self, attr, spin) + range_form.addRow(label, spin) + dl.addLayout(range_form) + + right_panel = QWidget() + layout = QVBoxLayout(right_panel) + layout.setContentsMargins(0, 0, 0, 0) + layout.setSpacing(8) + layout.addWidget(self.grp_velocity) + layout.addWidget(grp_display) + layout.addStretch() + return _scroll_panel(right_panel, _RIGHT_PANEL_W) + + def _build_menus(self): + menubar = self.menuBar() + + fft_menu = menubar.addMenu("&FFT") + fft_act = QAction("FFT &Options…", self) + fft_act.setStatusTip("Configure FFT backend and zero-padding") + fft_act.triggered.connect(self._on_fft_options) + fft_menu.addAction(fft_act) + + fusion_menu = menubar.addMenu("&Fusion") + self._alignment_act = QAction("Angle &Alignment", self) + self._alignment_act.setStatusTip( + "Compute a rotation+translation alignment across all angles " + "(from CH4 masks) and enable Aligned View. Requires >1 angle.") + self._alignment_act.setEnabled(False) + self._alignment_act.triggered.connect(self._on_angle_alignment) + fusion_menu.addAction(self._alignment_act) + + self._manual_align_act = QAction("&Manual Alignment…", self) + self._manual_align_act.setStatusTip( + "Open an interactive dialog to align angles by eye: overlaid CH4 " + "threshold masks, keyboard nudge (translate + rotate), auto " + "de-rotate to the known scan angles, and save/clear a persistent " + "alignment.") + self._manual_align_act.setEnabled(False) + self._manual_align_act.triggered.connect(self._on_manual_alignment) + fusion_menu.addAction(self._manual_align_act) + + convert_menu = menubar.addMenu("&Convert") + self._batch_dc_act = QAction("Batch Compute DC and &Store…", self) + self._batch_dc_act.setStatusTip( + "Select .sras files and compute+store DC images (CH3/CH4 mean) " + "for every angle, converting v6 files to v7 in place.") + self._batch_dc_act.triggered.connect(lambda: self._on_batch_compute("dc")) + convert_menu.addAction(self._batch_dc_act) + + self._batch_fft_act = QAction("Batch Compute FFT and Sto&re…", self) + self._batch_fft_act.setStatusTip( + "Select .sras files and compute+store FFT peak-frequency images " + "for every angle, converting v6 files to v7 in place. Stored " + "images are natural-resolution (pad 1); padded views compute live.") + self._batch_fft_act.triggered.connect(lambda: self._on_batch_compute("fft")) + convert_menu.addAction(self._batch_fft_act) + + # ------------------------------------------------------------------ + # Drag-and-drop + # ------------------------------------------------------------------ + + def dragEnterEvent(self, event): + urls = event.mimeData().urls() + if urls and urls[0].toLocalFile().lower().endswith(".sras"): + event.acceptProposedAction() + + def dropEvent(self, event): + self._load_file(event.mimeData().urls()[0].toLocalFile()) + + # ------------------------------------------------------------------ + # File loading + # ------------------------------------------------------------------ + + def _on_open(self): + path, _ = QFileDialog.getOpenFileName( + self, "Open SRAS File", "", "SRAS Files (*.sras);;All Files (*)") + if path: + self._load_file(path) + + def _load_file(self, path: str): + started = self._run_worker( + "load", LoadWorker(path), + connect=( + ("finished", self._on_load_done), + ("error", lambda msg: self.statusBar().showMessage(f"Error: {msg}")), + ), + ) + if not started: + return + self.btn_open.setEnabled(False) + self.statusBar().showMessage(f"Loading {Path(path).name}…") + self._show_progress("main", f"Loading {Path(path).name}…") + + def _on_load_done(self, sras): + self._close_progress("main") + self.btn_open.setEnabled(True) + if sras is None: + return + self._sras = sras + self._current_image = None + + # A manual-alignment dialog bound to the previous file must not + # survive a reload — its per-angle state (and the sras it was + # constructed against) no longer matches the new file's geometry. + if self._manual_align_dialog is not None: + self._manual_align_dialog.close() + self._manual_align_dialog = None + + # Caches (and any in-flight DC precompute) belong to the previous + # file's geometry — discard and start fresh. Bumping the generation + # counters makes any still-running worker's result get dropped when + # it lands. + self._dc_cache = {} + self._fft_cache = {} + self._dc_generation += 1 + self.lbl_dc_precompute.setText("") + + self._alignment_result = None + self._aligned_cache = {} + self._alignment_generation += 1 + self.chk_aligned_view.blockSignals(True) + self.chk_aligned_view.setChecked(False) + self.chk_aligned_view.setEnabled(False) + self.chk_aligned_view.blockSignals(False) + + # Silently restore a previously-saved manual alignment, if any, so + # the work survives closing and reopening the file. + sidecar = load_manual_alignment(sras) + if sidecar is not None: + try: + self._alignment_result = build_manual_alignment( + sras, sidecar.ref_angle_idx, sidecar.dc_threshold_mv, + sidecar.per_angle) + self.chk_aligned_view.blockSignals(True) + self.chk_aligned_view.setChecked(True) + self.chk_aligned_view.blockSignals(False) + self.statusBar().showMessage( + f"Restored saved manual alignment from " + f"{sidecar_path(sras.path).name}") + except Exception as exc: + # A corrupt/foreign sidecar or a rescan that shrank n_angles + # below ref_angle_idx must not block opening the .sras file. + self.statusBar().showMessage( + f"Could not restore saved alignment: {exc}") + + # A ROI from the previous file no longer matches the new scan's + # geometry, so discard it on every load. + self.image_canvas.clear_roi() + + self.lbl_filename.setText(sras.path.name) + + self.spin_angle.blockSignals(True) + self.spin_angle.setRange(0, max(0, sras.n_angles - 1)) + self.spin_angle.setValue(0) + self.spin_angle.blockSignals(False) + + # DC channels are cheap and give an instant, fluid overview of a + # scan; CH1/Velocity require an FFT per pixel that can take minutes + # on a large scan, so don't default to it. + self.combo_channel.blockSignals(True) + self.combo_channel.setCurrentIndex(CH4_IDX) + self.combo_channel.blockSignals(False) + + self._update_controls_enabled(True) + self._on_threshold_changed() # refresh ADC label with file calibration + self._on_view_changed() + self._start_dc_precompute() + + # ------------------------------------------------------------------ + # Scan info panel + # ------------------------------------------------------------------ + + def _update_scan_info_labels(self): + s = self._sras + if s is None: + return + a = self.spin_angle.value() + for key, text in ( + ("Angles", f"{s.n_angles}"), + ("Rows", f"{s.n_rows[a]}"), + ("Frames / row", f"{s.n_frames[a]}"), + ("Samples / frame", f"{s.samples_per_frame}"), + ("Sample rate", f"{s.sample_rate_hz / 1e9:.4g} GS/s"), + ("X start", f"{s.x_start_mm[a]:.4g} mm"), + ("Pixel Δx", f"{s.pixel_x_mm * 1e3:.3g} µm"), + ("Laser freq", f"{s.laser_freq_hz / 1e3:.4g} kHz"), + ): + self._info[key].setText(f"{key}: {text}") + + notes = [] + if s.frame_count_mismatch: + notes.append(f"! Header n_frames={s.n_frames_header}, " + f"actual={s.n_frames[a]} (scanner bug — corrected)") + if s.scan_aborted: + notes.append(f"! Scan aborted: {s.n_angles}/{s.n_angles_declared} " + "angles complete") + if s.background is not None: + notes.append(f"Background waveform: {len(s.background)} samples") + if s.version in (6, 7): + notes.append("v6/v7 format: rows / frames / x_start are per-angle") + + n_dc = sum(1 for x in s.precomputed_dc4_mv if x is not None) + n_fft = sum(1 for x in s.precomputed_freq_mhz if x is not None) + if n_dc or n_fft: + bg_note = " (bg-sub)" if s.precomputed_bg_sub else " (no bg-sub)" + notes.append( + f"Cached images: DC {n_dc}/{s.n_angles} angles, " + f"FFT {n_fft}/{s.n_angles} angles{bg_note if n_fft else ''} " + "— display is instant for cached angles") + elif s.version == 7: + notes.append("v7 format: no cache blocks stored yet") + self.lbl_frame_warn.setText("\n".join(notes)) + + # ------------------------------------------------------------------ + # Controls + # ------------------------------------------------------------------ + + def _update_controls_enabled(self, enabled: bool): + s = self._sras + has_file = enabled and s is not None + ch_idx = self.combo_channel.currentIndex() + is_ch1 = enabled and ch_idx in CH1_DERIVED_MODES + is_vel = enabled and ch_idx == VELOCITY_MODE_IDX + + self.spin_angle.setEnabled(has_file and s.n_angles > 1) + self.combo_channel.setEnabled(enabled) + self.combo_cmap.setEnabled(enabled) + self.chk_auto.setEnabled(enabled) + manual = enabled and not self.chk_auto.isChecked() + self.spin_vmin.setEnabled(manual) + self.spin_vmax.setEnabled(manual) + + # Threshold and bg-sub apply to all CH1 modes + self.spin_threshold_mv.setEnabled(is_ch1) + self.chk_bg_sub.setEnabled(has_file and s.background is not None and is_ch1) + self.spin_grating_um.setEnabled(is_vel) + self.grp_velocity.setVisible(is_vel) + + self.btn_export_csv.setEnabled(is_ch1 and self._current_image is not None) + # ROI: always usable once a file is loaded (independent of channel) + self.btn_draw_roi.setEnabled(has_file) + + # Batch Convert actions pick their own files, independent of + # whatever's currently open — only gated on no batch already running. + can_batch = not self._job_running("batch") + self._batch_dc_act.setEnabled(can_batch) + self._batch_fft_act.setEnabled(can_batch) + + self._alignment_act.setEnabled( + has_file and s.n_angles > 1 and not self._job_running("align")) + self._manual_align_act.setEnabled( + has_file and s.n_angles > 1 and not self._job_running("align")) + self.chk_aligned_view.setEnabled(enabled and self._alignment_result is not None) + self._update_roi_ui() + + def _on_channel_changed(self): + self._update_controls_enabled(self._sras is not None) + self._on_view_changed() + + def _on_bg_sub_toggled(self): + # Background subtraction changes the FFT input, so it genuinely + # invalidates the cached raw FFT (the cache key includes it) — + # _refresh_display() recomputes only on a miss for the new state. + if self._sras is not None and self.combo_channel.currentIndex() in CH1_DERIVED_MODES: + self._refresh_display() + + def _on_grating_changed(self): + # Grating is a pure post-multiply on the cached frequency image — + # never needs a recompute. + if self._sras is not None and self.combo_channel.currentIndex() == VELOCITY_MODE_IDX: + self._refresh_display() + + def _on_threshold_changed(self): + mv = self.spin_threshold_mv.value() + cal = (self._sras.cal(CH4_IDX) if self._sras is not None + else (_FALLBACK_YMULT_MV, _FALLBACK_YOFF_ADC, 0.0)) + self.lbl_threshold_adc.setText(f"≈ {mv_to_adc(mv, *cal):.1f} ADC counts") + # Threshold decides which pixels get an FFT at all, so changing it is + # a genuine cache-key change — but the recompute reuses the cached DC4 + # image to skip masked-out pixels. + if self._sras is not None and self.combo_channel.currentIndex() in CH1_DERIVED_MODES: + self._refresh_display() + + def _on_autoscale_toggled(self, checked: bool): + manual = not checked + self.spin_vmin.setEnabled(manual and self._sras is not None) + self.spin_vmax.setEnabled(manual and self._sras is not None) + if self._sras is not None and self._current_image is not None: + self._redraw_image(self._current_image) + + def _on_manual_range_changed(self): + if not self.chk_auto.isChecked() and self._current_image is not None: + self._redraw_image(self._current_image) + + def _on_cmap_changed(self): + # Colormap is purely how the existing image is rendered. + if self._current_image is not None: + self._redraw_image(self._current_image) + + def _on_view_changed(self): + if self._sras is None: + return + idx = self.spin_angle.value() + self.lbl_angle_deg.setText(f"({self._sras.angles_deg[idx]:.1f}°)") + self._update_scan_info_labels() + self._refresh_display() + + def _on_aligned_view_toggled(self, checked: bool): + if self._current_image is not None: + self._redraw_image(self._current_image) + + # ------------------------------------------------------------------ + # CSV export + # ------------------------------------------------------------------ + + def _on_export_csv(self): + if self._current_image is None or self._sras is None: + return + default_name = (f"{self._sras.path.stem}_angle{self._current_angle}" + f"_{CH_NAMES[self._current_ch]}.csv") + path, _ = QFileDialog.getSaveFileName( + self, "Export Image as CSV", + str(self._sras.path.parent / default_name), + "CSV files (*.csv);;All files (*)") + if not path: + return + np.savetxt(path, self._current_image, delimiter=",", fmt="%.6g") + self.statusBar().showMessage(f"Exported {Path(path).name}") + + def _on_export_roi_csv(self): + if self._current_image is None or self._sras is None: + return + roi = self.image_canvas.get_roi() + if roi is None: + self.statusBar().showMessage("No ROI — draw one first") + return + s = self._sras + x_axis = s.x_axis_mm(self._current_angle) + y_axis = s.y_positions_mm(self._current_angle) + mask = roi.mask_for_grid(x_axis, y_axis) + if not mask.any(): + self.statusBar().showMessage("ROI does not overlap any pixel") + return + img = self._current_image + if img.shape != mask.shape: + self.statusBar().showMessage( + f"ROI shape {mask.shape} does not match image {img.shape}") + return + + X, Y = np.meshgrid(np.asarray(x_axis, dtype=np.float64), + np.asarray(y_axis, dtype=np.float64)) + rows_idx, frames_idx = np.where(mask) + n_pix = int(mask.sum()) + + ch_name = CH_NAMES[self._current_ch] + angle = self._current_angle + default_name = f"{s.path.stem}_angle{angle}_{ch_name}_ROI.csv" + path, _ = QFileDialog.getSaveFileName( + self, "Export ROI as CSV", + str(s.path.parent / default_name), + "CSV files (*.csv);;All files (*)") + if not path: + return + + corners_str = " ".join(f"({p[0]:.6g},{p[1]:.6g})" for p in roi.corners()) + header = ( + f"# ROI quad corners (BL BR TR TL) mm: {corners_str}\n" + f"# source: {s.path.name}, channel={ch_name}, " + f"angle_idx={angle}, angle_deg={s.angles_deg[angle]:.4g}\n" + f"# n_pixels={n_pix}\n" + "row,frame,x_mm,y_mm,value" + ) + data = np.column_stack([ + rows_idx.astype(np.int64), frames_idx.astype(np.int64), + X[mask], Y[mask], img[mask].astype(np.float64), + ]) + # integer columns first, floats after — use a per-column format list + np.savetxt(path, data, delimiter=",", + fmt=["%d", "%d", "%.6g", "%.6g", "%.6g"], + header=header, comments="") + self.statusBar().showMessage( + f"Exported ROI ({n_pix} pixels) to {Path(path).name}") + + # ------------------------------------------------------------------ + # ROI + # ------------------------------------------------------------------ + + def _on_draw_roi_toggled(self, checked: bool): + if checked: + self.image_canvas.start_drawing() + self.statusBar().showMessage( + "Click and drag on the image to draw a new rectangle.") + else: + self.image_canvas.cancel_drawing() + + def _on_draw_mode_changed(self, active: bool): + # Keep the toggle button's visual state in sync with the canvas. + self.btn_draw_roi.blockSignals(True) + self.btn_draw_roi.setChecked(active) + self.btn_draw_roi.blockSignals(False) + + def _on_clear_roi(self): + self.image_canvas.clear_roi() + self.statusBar().showMessage("ROI cleared") + + def _update_roi_ui(self): + roi = self.image_canvas.get_roi() + if roi is None: + self.lbl_roi_center.setText("centroid: —") + self.lbl_roi_size.setText("bbox: —") + self.lbl_roi_npix.setText("pixels inside: —") + self.btn_clear_roi.setEnabled(False) + self.btn_export_roi.setEnabled(False) + return + + cen = roi.centroid() + bbox = roi.bbox_size() + self.lbl_roi_center.setText(f"centroid: ({cen[0]:.3f}, {cen[1]:.3f}) mm") + self.lbl_roi_size.setText(f"bbox: {bbox[0]:.3f} × {bbox[1]:.3f} mm") + + npix = 0 + if self._sras is not None: + try: + # Deliberately always the raw per-angle grid, even when + # Aligned View is on: _on_export_roi_csv also exports on the + # raw grid (never synthetically-resampled pixels), so this + # readout must match what Export ROI actually writes. + mask = roi.mask_for_grid( + self._sras.x_axis_mm(self._current_angle), + self._sras.y_positions_mm(self._current_angle)) + npix = int(mask.sum()) + except Exception: + npix = 0 + self.lbl_roi_npix.setText(f"pixels inside: {npix}") + self.btn_clear_roi.setEnabled(True) + self.btn_export_roi.setEnabled(self._current_image is not None and npix > 0) + + # ------------------------------------------------------------------ + # Display + # ------------------------------------------------------------------ + + def _current_n_fft(self) -> int | None: + if self._fft_pad_factor <= 1 or self._sras is None: + return None + return self._sras.samples_per_frame * self._fft_pad_factor + + def _scale_for_display(self, freq_mhz: np.ndarray, ch_idx: int) -> np.ndarray: + """Velocity is a pure post-multiply of the (already DC-masked) + cached frequency image — never worth a recompute on its own.""" + if ch_idx == VELOCITY_MODE_IDX: + return freq_mhz * self.spin_grating_um.value() + return freq_mhz + + def _fft_cache_key(self, angle_idx: int) -> tuple: + return (angle_idx, self.chk_bg_sub.isChecked(), self._current_n_fft(), + self.spin_threshold_mv.value()) + + def _aligned_cache_key(self, angle_idx: int, ch_idx: int) -> tuple: + """Mirrors _fft_cache's key granularity so a stale aligned image is + never shown after bg_sub/threshold/pad/grating changes.""" + if ch_idx in CH1_DERIVED_MODES: + return (*self._fft_cache_key(angle_idx), ch_idx, + self.spin_grating_um.value() if ch_idx == VELOCITY_MODE_IDX else None) + return (angle_idx, ch_idx) + + def _aligned_canvas_axes(self) -> tuple[np.ndarray, np.ndarray]: + r = self._alignment_result + n_rows, n_cols = r.canvas_shape + return (r.canvas_origin_mm[0] + np.arange(n_cols) * r.canvas_dx_mm, + r.canvas_origin_mm[1] + np.arange(n_rows) * r.canvas_dy_mm) + + def _get_aligned_display_image(self, raw_img: np.ndarray, angle_idx: int, + ch_idx: int) -> np.ndarray: + key = self._aligned_cache_key(angle_idx, ch_idx) + cached = self._aligned_cache.get(key) + if cached is None: + cached = apply_alignment(self._alignment_result, angle_idx, raw_img) + self._aligned_cache[key] = cached + return cached + + def _refresh_display(self): + """Show the image for the current angle/channel/threshold, using + cached data whenever possible and only falling back to a background + compute (with progress popup) when genuinely nothing is cached yet.""" + if self._sras is None: + return + angle_idx = self.spin_angle.value() + ch_idx = self.combo_channel.currentIndex() + + if ch_idx in CH1_DERIVED_MODES: + raw = self._fft_cache.get(self._fft_cache_key(angle_idx)) + if raw is not None: + self._show_image_now(self._scale_for_display(raw, ch_idx), + angle_idx, ch_idx) + return + else: + cached = self._dc_cache.get((angle_idx, ch_idx)) + if cached is not None: + self._show_image_now(cached, angle_idx, ch_idx) + return + + # Nothing cached for these settings — need a real compute. Changing + # the DC threshold changes *which* pixels get an FFT at all, so it + # can't be satisfied from the cache — but with the DC map already + # known, the recompute skips the FFT for masked-out pixels. + self._start_compute() + + def _show_image_now(self, img: np.ndarray, angle_idx: int, ch_idx: int): + """Display an already-available image with no compute involved.""" + self._current_image = img + self._current_angle = angle_idx + self._current_ch = ch_idx + self.btn_export_csv.setEnabled(ch_idx in CH1_DERIVED_MODES) + self._redraw_image(img) + self._update_roi_ui() + + def _redraw_image(self, img: np.ndarray): + s = self._sras + angle_idx = self._current_angle + ch_idx = self._current_ch + + aligned = (self.chk_aligned_view.isChecked() + and self._alignment_result is not None + and angle_idx in self._alignment_result.per_angle) + if aligned: + display_img = self._get_aligned_display_image(img, angle_idx, ch_idx) + x_axis, y_axis = self._aligned_canvas_axes() + else: + display_img = img + x_axis = s.x_axis_mm(angle_idx) + y_axis = s.y_positions_mm(angle_idx) + + dx = x_axis[1] - x_axis[0] if len(x_axis) > 1 else s.pixel_x_mm + dy = float(y_axis[1] - y_axis[0]) if len(y_axis) > 1 else 1.0 + extent = [x_axis[0] - dx / 2, x_axis[-1] + dx / 2, + y_axis[-1] + dy / 2, y_axis[0] - dy / 2] + + if self.chk_auto.isChecked(): + vmin, vmax = float(display_img.min()), float(display_img.max()) + for spin, val in ((self.spin_vmin, vmin), (self.spin_vmax, vmax)): + spin.blockSignals(True) + spin.setValue(val) + spin.blockSignals(False) + else: + vmin, vmax = self.spin_vmin.value(), self.spin_vmax.value() + + angle_deg = s.angles_deg[angle_idx] + mode_str, unit, colorbar_label = _CHANNEL_DISPLAY[ch_idx] + if ch_idx == VELOCITY_MODE_IDX: + ch_label = f"Velocity [grating={self.spin_grating_um.value():.2f} µm]" + else: + ch_label = CH_LABELS[ch_idx] + + title = f"{CH_NAMES[ch_idx]} | {mode_str} | {angle_deg:.1f}°" + if aligned: + title += " [Aligned]" + + self.image_canvas.show_image( + display_img, extent, + cmap=self.combo_cmap.currentText(), + vmin=vmin, vmax=vmax, + xlabel="X (mm)", ylabel="Y (mm)", + title=title, colorbar_label=colorbar_label, + ) + self.statusBar().showMessage( + f"{s.path.name} | {ch_label} @ {angle_deg:.1f}° " + f"| {display_img.shape[1]} × {display_img.shape[0]} px | {unit}" + f"{' | Aligned' if aligned else ''}" + ) + + # ------------------------------------------------------------------ + # Background compute (only reached on a genuine cache miss) + # ------------------------------------------------------------------ + + def _start_compute(self): + if self._sras is None or self._job_running("compute"): + return # re-checked when the running compute finishes + + angle_idx = self.spin_angle.value() + ch_idx = self.combo_channel.currentIndex() + is_fft = ch_idx in CH1_DERIVED_MODES + + self._pending_angle = angle_idx + self._pending_ch = ch_idx + self._pending_bg_sub = self.chk_bg_sub.isChecked() + self._pending_threshold = self.spin_threshold_mv.value() + self._pending_fft_pad_factor = self._fft_pad_factor + + worker = ComputeWorker( + self._sras, angle_idx, ch_idx, + apply_bg_sub=self._pending_bg_sub, + n_fft=self._current_n_fft(), + dc_threshold_mv=self._pending_threshold, + # Reuse the cached DC4 image (if the precompute has reached this + # angle) so the FFT skips masked-out pixels entirely and doesn't + # need to re-read the CH4 channel from disk. + dc4_mv=self._dc_cache.get((angle_idx, CH4_IDX)), + is_fft_mode=is_fft, + ) + if not self._run_worker( + "compute", worker, + connect=( + ("finished", self._on_compute_done), + ("error", lambda msg: self.statusBar().showMessage( + f"Compute error: {msg}")), + ), + on_done=self._after_compute): + return + + if is_fft: + self.statusBar().showMessage("Computing FFT…") + self._show_progress( + "main", + f"Computing FFT for angle {angle_idx}…\n" + "This can take a while on a large scan — result is cached " + "so revisiting this angle/mode/threshold will be instant.") + else: + self.statusBar().showMessage("Computing DC image…") + self._show_progress("main", f"Computing DC image for angle {angle_idx}…") + + def _after_compute(self): + """If settings changed while the compute was running, re-dispatch + through the cache-aware path — the now-current combination may + already be cached.""" + if (self.spin_angle.value(), self.combo_channel.currentIndex(), + self.chk_bg_sub.isChecked(), self.spin_threshold_mv.value(), + self._fft_pad_factor) != ( + self._pending_angle, self._pending_ch, self._pending_bg_sub, + self._pending_threshold, self._pending_fft_pad_factor): + self._refresh_display() + + def _on_compute_done(self, result): + self._close_progress("main") + if result is None: + return # cancelled mid-compute; the partial image must not cache + angle_idx = self._pending_angle + ch_idx = self._pending_ch + + if ch_idx in CH1_DERIVED_MODES: + self._fft_cache[(angle_idx, self._pending_bg_sub, + self._current_n_fft(), self._pending_threshold)] = result + img = self._scale_for_display(result, ch_idx) + else: + img = result + self._dc_cache[(angle_idx, ch_idx)] = img + + self._show_image_now(img, angle_idx, ch_idx) + + # ------------------------------------------------------------------ + # Background DC precompute (all angles, so switching is fluid) + # ------------------------------------------------------------------ + + def _start_dc_precompute(self): + if self._sras is None: + return + generation = self._dc_generation + n_angles = self._sras.n_angles + + worker = DcPrecomputeWorker(self._sras) + self._run_worker( + "dc_precompute", worker, + connect=( + ("angle_done", lambda a, dc3, dc4, g=generation: + self._on_dc_precompute_angle_done(g, a, dc3, dc4, n_angles)), + ("error", lambda msg: self.statusBar().showMessage( + f"DC precompute error: {msg}", 5000)), + ), + quit_on=("finished", "error"), + ) + + def _on_dc_precompute_angle_done(self, generation: int, angle_idx: int, + dc3_mv: np.ndarray, dc4_mv: np.ndarray, + n_angles: int): + if generation != self._dc_generation: + return # stale result from a previously-loaded file — discard + self._dc_cache[(angle_idx, CH3_IDX)] = dc3_mv + self._dc_cache[(angle_idx, CH4_IDX)] = dc4_mv + + done = sum(1 for a in range(n_angles) if (a, CH4_IDX) in self._dc_cache) + self.lbl_dc_precompute.setText( + f"Precomputing DC images: {done}/{n_angles} angles ready…" + if done < n_angles else "DC images ready for all angles.") + + # If we just finished the angle/channel the user is currently looking + # at and it wasn't shown yet (they switched here before the precompute + # caught up and are still waiting), show it now. + current_ch = self.combo_channel.currentIndex() + if (angle_idx == self.spin_angle.value() + and not self._job_running("compute") + and current_ch in (CH3_IDX, CH4_IDX) + and (self._current_angle != angle_idx or self._current_ch != current_ch)): + self._refresh_display() + + # ------------------------------------------------------------------ + # Pixel inspector + # ------------------------------------------------------------------ + + def _on_pixel_clicked(self, row_idx: int, frame_idx: int): + if self._sras is None or self._current_image is None: + return + angle_idx = self._current_angle + if (self.chk_aligned_view.isChecked() and self._alignment_result is not None + and angle_idx in self._alignment_result.per_angle): + # The click landed on the shared aligned canvas — invert the same + # canvas->raw affine used to display it back to a raw (row, frame) + # index before looking up the waveform. + t = self._alignment_result.per_angle[angle_idx] + raw = t.matrix @ np.array([row_idx, frame_idx], dtype=np.float64) + t.offset + row_idx, frame_idx = int(round(raw[0])), int(round(raw[1])) + n_rows_a, n_frames_a = self._sras.image_shape(angle_idx) + if not (0 <= row_idx < n_rows_a and 0 <= frame_idx < n_frames_a): + self.statusBar().showMessage( + "No source waveform here (padding region of the aligned canvas).") + return + + self.lbl_wave_hint.hide() + if self._current_ch in CH1_DERIVED_MODES: + self.wave_canvas.show_rf_waveform( + self._sras, angle_idx, row_idx, frame_idx, + apply_bg_sub=self.chk_bg_sub.isChecked()) + else: + self.wave_canvas.show_dc_waveform( + self._sras, angle_idx, self._current_ch, row_idx, frame_idx) + + # ------------------------------------------------------------------ + # Progress dialogs + # ------------------------------------------------------------------ + + def _show_progress(self, key: str, message: str, maximum: int = 0): + """Show (or relabel) the progress dialog under *key*. maximum=0 gives + an indeterminate busy indicator.""" + dlg = self._progress_dlgs.get(key) + if dlg is not None: + dlg.setLabelText(message) + return + dlg = QProgressDialog(message, "", 0, maximum, self) + dlg.setWindowTitle("Please wait…") + dlg.setCancelButton(None) + dlg.setWindowModality(Qt.WindowModality.WindowModal) + dlg.setMinimumDuration(300) # only appears if it takes > 300 ms + dlg.show() + self._progress_dlgs[key] = dlg + + def _set_progress(self, key: str, pct: int): + dlg = self._progress_dlgs.get(key) + if dlg is not None: + dlg.setValue(pct) + + def _close_progress(self, key: str): + dlg = self._progress_dlgs.pop(key, None) + if dlg is not None: + dlg.close() + + # ------------------------------------------------------------------ + # Convert menu: batch DC/FFT compute-and-store (v6 -> v7) + # ------------------------------------------------------------------ + + def _on_batch_compute(self, mode: str): + if self._job_running("batch"): + return + label = "DC" if mode == "dc" else "FFT" + paths, _ = QFileDialog.getOpenFileNames( + self, f"Select .sras files to batch-compute {label}", "", + "SRAS files (*.sras);;All files (*)") + if not paths: + return + + self._batch_errors = [] + worker = BatchCacheWorker(paths, mode, self.chk_bg_sub.isChecked()) + started = self._run_worker( + "batch", worker, + connect=( + ("progress", lambda pct: self._set_progress("batch", pct)), + ("file_done", self._on_batch_file_done), + ("finished", lambda p=paths: self._on_batch_finished(p)), + ), + on_done=self._after_batch, + ) + if not started: + return # a second trigger snuck in while the file dialog was open + + self._batch_dc_act.setEnabled(False) + self._batch_fft_act.setEnabled(False) + self._show_progress( + "batch", f"Batch computing {label} for {len(paths)} file(s)…", + maximum=100) + + def _on_batch_file_done(self, path: str, err: str): + if err: + self._batch_errors.append(f"{Path(path).name} — {err}") + self._show_progress("batch", f"Processed {Path(path).name}…") + + def _on_batch_finished(self, paths: list[str]): + self._close_progress("batch") + + n_total = len(paths) + n_failed = len(self._batch_errors) + n_ok = n_total - n_failed + if n_failed: + summary = (f"Batch store: {n_ok}/{n_total} file(s) updated, " + f"{n_failed} failed: {'; '.join(self._batch_errors)}") + else: + summary = f"Batch store: {n_ok}/{n_total} file(s) updated." + self.statusBar().showMessage(summary) + self._batch_errors = [] + + # If the currently-open file was in this batch, reload it so the GUI + # picks up the newly-written v7 cache instead of stale state. + if self._sras is not None and str(self._sras.path) in paths: + self._load_file(str(self._sras.path)) + + def _after_batch(self): + self._batch_dc_act.setEnabled(True) + self._batch_fft_act.setEnabled(True) + + # ------------------------------------------------------------------ + # Fusion: angle alignment + # ------------------------------------------------------------------ + + def _on_angle_alignment(self): + if self._sras is None or self._sras.n_angles <= 1: + return + ref_idx = 0 + threshold_mv = self.spin_threshold_mv.value() + generation = self._alignment_generation + + started = self._run_worker( + "align", AngleAlignmentWorker(self._sras, ref_idx, threshold_mv), + connect=( + ("progress", lambda pct: self._set_progress("main", pct)), + ("finished", lambda result, err, g=generation: + self._on_alignment_done(g, result, err)), + ), + on_done=lambda: self._update_controls_enabled(self._sras is not None), + ) + if not started: + return + + self._alignment_act.setEnabled(False) + self._show_progress( + "main", + f"Computing angle alignment ({self._sras.n_angles} angles, " + f"ref=angle 0, CH4 mask ≥ {threshold_mv:.3f} mV)…", + maximum=100) + + def _on_alignment_done(self, generation: int, result, error_msg: str): + self._close_progress("main") + if generation != self._alignment_generation: + return # a new file was loaded while this was computing — discard + if error_msg: + self.statusBar().showMessage(f"Angle alignment failed: {error_msg}") + return + self._alignment_result = result + self._aligned_cache = {} + self.chk_aligned_view.setEnabled(True) + self.chk_aligned_view.blockSignals(True) + self.chk_aligned_view.setChecked(True) + self.chk_aligned_view.blockSignals(False) + nr, nc = result.canvas_shape + self.statusBar().showMessage( + f"Angle alignment computed ({self._sras.n_angles} angles, " + f"canvas {nc}×{nr} px).") + self._refresh_display() + + # ------------------------------------------------------------------ + # Fusion: manual alignment + # ------------------------------------------------------------------ + + def _on_manual_alignment(self): + if self._sras is None or self._sras.n_angles <= 1: + return + if self._manual_align_dialog is not None: + self._manual_align_dialog.raise_() + self._manual_align_dialog.activateWindow() + return + + ref_idx = 0 + threshold_mv = self.spin_threshold_mv.value() + seed: dict[int, ManualAngleParams] = {} + # Seed only from a previously *saved manual* alignment (this dialog's + # own Save also writes this sidecar) -- never from self._alignment_result + # when it holds the automatic Fusion -> Angle Alignment's output. That + # path's translation comes from FFT phase correlation, which is the + # very thing manual mode exists to work around; inheriting it here + # would silently reintroduce the same bad translations under a + # "manual" label, on top of the (correct) analytic rotation, which is + # exactly what makes manual mode look like it "still does the same + # thing" the automatic one does. + sidecar = load_manual_alignment(self._sras) + if sidecar is not None and sidecar.ref_angle_idx == ref_idx: + seed = dict(sidecar.per_angle) + threshold_mv = sidecar.dc_threshold_mv + + cached_dc4 = {a: img for (a, ch), img in self._dc_cache.items() if ch == CH4_IDX} + dlg = ManualAlignmentDialog( + self, self._sras, ref_angle_idx=ref_idx, dc_threshold_mv=threshold_mv, + seed_per_angle=seed, cached_dc4_mv=cached_dc4) + dlg.alignment_saved.connect(self._on_manual_alignment_saved) + dlg.alignment_cleared.connect(self._on_manual_alignment_cleared) + dlg.finished.connect(self._on_manual_align_dialog_closed) + dlg.setAttribute(Qt.WidgetAttribute.WA_DeleteOnClose) + self._manual_align_dialog = dlg + dlg.show() + + def _on_manual_align_dialog_closed(self, _result_code: int): + self._manual_align_dialog = None + + def _on_manual_alignment_saved(self, result, sidecar_path_str: str): + self._alignment_result = result + self._aligned_cache = {} + self._alignment_generation += 1 + self.chk_aligned_view.setEnabled(True) + self.chk_aligned_view.blockSignals(True) + self.chk_aligned_view.setChecked(True) + self.chk_aligned_view.blockSignals(False) + self._update_controls_enabled(self._sras is not None) + self.statusBar().showMessage( + f"Manual alignment saved to {Path(sidecar_path_str).name}") + if self._current_image is not None: + self._refresh_display() + + def _on_manual_alignment_cleared(self): + self._alignment_result = None + self._aligned_cache = {} + self._alignment_generation += 1 + self.chk_aligned_view.blockSignals(True) + self.chk_aligned_view.setChecked(False) + self.chk_aligned_view.setEnabled(False) + self.chk_aligned_view.blockSignals(False) + self._update_controls_enabled(self._sras is not None) + self.statusBar().showMessage("Manual alignment cleared.") + if self._current_image is not None: + self._refresh_display() + + # ------------------------------------------------------------------ + # FFT Options + # ------------------------------------------------------------------ + + def _on_fft_options(self): + dlg = FftOptionsDialog( + self, + current_backend=compute.get_fft_backend(), + current_pad_factor=self._fft_pad_factor, + samples_per_frame=self._sras.samples_per_frame if self._sras else None, + sample_rate_hz=self._sras.sample_rate_hz if self._sras else None, + grating_um=self.spin_grating_um.value(), + ) + if dlg.exec() != QDialog.DialogCode.Accepted: + return + compute.set_fft_backend(dlg.get_backend()) + self._fft_pad_factor = dlg.get_pad_factor() + self._settings.setValue("fft/backend", compute.get_fft_backend()) + self._settings.setValue("fft/pad_factor", self._fft_pad_factor) + # Pad factor changes the FFT bin count, so it genuinely invalidates + # the cached raw FFT (part of the cache key) — _refresh_display() + # recomputes only on a cache miss. + if self._sras is not None and self.combo_channel.currentIndex() in CH1_DERIVED_MODES: + self._refresh_display() + + # ------------------------------------------------------------------ + + def closeEvent(self, event): + if self._manual_align_dialog is not None: + self._manual_align_dialog.close() + + # Signal every cancellable worker first, then wait. Waiting without + # signalling means sitting out whatever is in flight — on a large + # scan a single angle is ~40 s. + jobs = list(self._jobs.values()) + for _thread, worker, _on_done in jobs: + stop = getattr(worker, "stop", None) + if callable(stop): + stop() + for thread, _worker, _on_done in jobs: + thread.quit() + thread.wait(5000) + super().closeEvent(event) + + +# --------------------------------------------------------------------------- + +def main(): + app = QApplication(sys.argv) + window = SrasViewerWindow( + initial_path=sys.argv[1] if len(sys.argv) > 1 else None) + window.show() + sys.exit(app.exec()) + diff --git a/tests/test_gui.py b/tests/test_gui.py index cb0bc28..ada0b73 100644 --- a/tests/test_gui.py +++ b/tests/test_gui.py @@ -186,7 +186,7 @@ def test_roi_and_csv_export(ctx): assert win.btn_export_roi.isEnabled(), "Export ROI enabled" csv_path = ctx.tmpdir / "roi.csv" - with patch("sras_viewer.QFileDialog.getSaveFileName", + with patch("sras_viewer.main_window.QFileDialog.getSaveFileName", return_value=(str(csv_path), "")): win._on_export_roi_csv() assert csv_path.exists(), "ROI CSV written" @@ -195,7 +195,7 @@ def test_roi_and_csv_export(ctx): f"ROI CSV has {len(body)} lines for {npix} pixels (want header + one per pixel)" img_csv = ctx.tmpdir / "img.csv" - with patch("sras_viewer.QFileDialog.getSaveFileName", + with patch("sras_viewer.main_window.QFileDialog.getSaveFileName", return_value=(str(img_csv), "")): win._on_export_csv() assert img_csv.exists(), "image CSV written" @@ -426,7 +426,7 @@ def test_stale_schema_sidecar_ignored(ctx): def test_clear_with_confirmation(ctx): win, dlg, s = ctx.win, ctx.dlg, ctx.s - with patch("sras_viewer.QMessageBox.question", + with patch("sras_viewer.dialogs.QMessageBox.question", return_value=QMessageBox.StandardButton.Yes): dlg._on_clear() assert not ctx.sidecar.exists(), "sidecar file deleted" From f0f622b9abf5a2c66a97b85ef2766d068f1005df Mon Sep 17 00:00:00 2001 From: Thomas Ales Date: Thu, 6 Aug 2026 10:56:50 -0500 Subject: [PATCH 09/10] Dedup viewer Qt boilerplate after the split - One _apply_alignment_result() replaces the five copies of the cache-clear / generation-bump / Aligned-View-checkbox dance, and _on_manual_alignment_saved/_cleared collapse into a shared _manual_alignment_changed. - QSignalBlocker context managers replace every hand-rolled blockSignals(True)/set/blockSignals(False) triple (11 sites). - _make_dspin() replaces the 11 four-to-seven-line QDoubleSpinBox constructions; _axes_extent() replaces the duplicated imshow-extent formula; _is_fft_mode() replaces the repeated CH1-mode guard. - Jobs constants replace the stringly-typed background-job keys. - The two 160-line widget builders split along their section banners: _build_left_panel -> file/info/view/roi groups, the manual-alignment _build_ui -> six per-group builders + _connect_controls. Co-Authored-By: Claude Fable 5 --- sras_viewer/common.py | 40 +++++++- sras_viewer/dialogs.py | 127 ++++++++++-------------- sras_viewer/main_window.py | 191 +++++++++++++++++-------------------- 3 files changed, 172 insertions(+), 186 deletions(-) diff --git a/sras_viewer/common.py b/sras_viewer/common.py index 7d07a50..5793e90 100644 --- a/sras_viewer/common.py +++ b/sras_viewer/common.py @@ -2,8 +2,8 @@ from PyQt6.QtCore import Qt from PyQt6.QtWidgets import ( - QFormLayout, QFrame, QGroupBox, QLabel, QScrollArea, QSizePolicy, - QVBoxLayout, QWidget, + QDoubleSpinBox, QFormLayout, QFrame, QGroupBox, QLabel, QScrollArea, + QSizePolicy, QVBoxLayout, QWidget, ) from sras_format import CH1_IDX, CH3_IDX, CH4_IDX @@ -52,6 +52,42 @@ _SPIN_MIN_W = 96 # Small layout helpers # --------------------------------------------------------------------------- +class Jobs: + """Keys for SrasViewerWindow's background-job registry (_run_worker / + _job_running) and its progress dialogs — one place instead of string + literals scattered across window and dialogs.""" + LOAD = "load" + COMPUTE = "compute" + DC_PRECOMPUTE = "dc_precompute" + BATCH = "batch" + ALIGN = "align" + MANUAL_ALIGN_MASKS = "manual_align_masks" + MANUAL_ALIGN_CORRELATE = "manual_align_correlate" + + +def _make_dspin(lo: float, hi: float, decimals: int, *, suffix: str = "", + value: float | None = None, step: float | None = None) -> QDoubleSpinBox: + """A QDoubleSpinBox with the panel-standard construction.""" + spin = QDoubleSpinBox() + spin.setRange(lo, hi) + spin.setDecimals(decimals) + if suffix: + spin.setSuffix(suffix) + if step is not None: + spin.setSingleStep(step) + if value is not None: + spin.setValue(value) + spin.setMinimumWidth(_SPIN_MIN_W) + return spin + + +def _axes_extent(x_axis, y_axis, dx: float, dy: float) -> list[float]: + """Matplotlib imshow extent with half-pixel margins, Y flipped so row 0 + renders at the top.""" + return [x_axis[0] - dx / 2, x_axis[-1] + dx / 2, + y_axis[-1] + dy / 2, y_axis[0] - dy / 2] + + def _wrap_label(text: str = "", css: str | None = None) -> QLabel: """A word-wrapped QLabel that reports its *wrapped* height to the layout. diff --git a/sras_viewer/dialogs.py b/sras_viewer/dialogs.py index d1efed7..5589698 100644 --- a/sras_viewer/dialogs.py +++ b/sras_viewer/dialogs.py @@ -5,11 +5,11 @@ from typing import TYPE_CHECKING import matplotlib as mpl import numpy as np from matplotlib.backends.backend_qtagg import NavigationToolbar2QT -from PyQt6.QtCore import pyqtSignal +from PyQt6.QtCore import QSignalBlocker, pyqtSignal from PyQt6.QtWidgets import ( - QButtonGroup, QComboBox, QDialog, QDialogButtonBox, QDoubleSpinBox, - QGroupBox, QHBoxLayout, QLabel, QMessageBox, QPushButton, QRadioButton, - QSpinBox, QVBoxLayout, QWidget, + QButtonGroup, QComboBox, QDialog, QDialogButtonBox, QGroupBox, + QHBoxLayout, QLabel, QMessageBox, QPushButton, QRadioButton, QSpinBox, + QVBoxLayout, QWidget, ) import sras_compute as compute @@ -22,7 +22,7 @@ from sras_workers import Ch4MaskWorker, CrossCorrelateWorker from .canvases import ManualAlignOverlayCanvas from .common import ( - _CSS_HINT, _CSS_MUTED, _CSS_WARN, _SPIN_MIN_W, _form, _group, + _CSS_HINT, _CSS_MUTED, _CSS_WARN, Jobs, _axes_extent, _form, _group, _make_dspin, _scroll_panel, _wrap_label, ) @@ -268,8 +268,20 @@ class ManualAlignmentDialog(QDialog): panel_l = QVBoxLayout(panel) panel_l.setContentsMargins(0, 0, 0, 0) panel_l.setSpacing(8) + panel_l.addWidget(self._build_angle_group()) + panel_l.addWidget(self._build_adjust_group()) + panel_l.addWidget(self._build_step_group()) + panel_l.addWidget(self._build_threshold_group(dc_threshold_mv)) + panel_l.addWidget(self._build_correlate_group()) + panel_l.addWidget(self._build_actions_group()) + self.lbl_status = _wrap_label("", _CSS_MUTED) + panel_l.addWidget(self.lbl_status) + panel_l.addStretch() - # ---- Active Angle ------------------------------------------------- + root.addWidget(_scroll_panel(panel, 320)) + self._connect_controls() + + def _build_angle_group(self) -> QWidget: grp_angle, al = _group("Active Angle") self.combo_active_angle = QComboBox() for a in range(self._sras.n_angles): @@ -280,80 +292,52 @@ class ManualAlignmentDialog(QDialog): al.addWidget(self.combo_active_angle) self.lbl_active_note = _wrap_label("", _CSS_WARN) al.addWidget(self.lbl_active_note) - panel_l.addWidget(grp_angle) + return grp_angle - # ---- Manual Adjustment --------------------------------------------- + def _build_adjust_group(self) -> QWidget: self.grp_manual_adjust, mform_box = _group("Manual Adjustment") mform = _form() - self.spin_active_rotation_deg = QDoubleSpinBox() - self.spin_active_rotation_deg.setRange(-3600.0, 3600.0) - self.spin_active_rotation_deg.setDecimals(3) - self.spin_active_rotation_deg.setSuffix(" °") - self.spin_active_rotation_deg.setMinimumWidth(_SPIN_MIN_W) + self.spin_active_rotation_deg = _make_dspin(-3600.0, 3600.0, 3, suffix=" °") mform.addRow("Rotation:", self.spin_active_rotation_deg) - self.spin_active_shift_x_mm = QDoubleSpinBox() - self.spin_active_shift_x_mm.setRange(-1e5, 1e5) - self.spin_active_shift_x_mm.setDecimals(4) - self.spin_active_shift_x_mm.setSuffix(" mm") - self.spin_active_shift_x_mm.setMinimumWidth(_SPIN_MIN_W) + self.spin_active_shift_x_mm = _make_dspin(-1e5, 1e5, 4, suffix=" mm") mform.addRow("Shift X:", self.spin_active_shift_x_mm) - self.spin_active_shift_y_mm = QDoubleSpinBox() - self.spin_active_shift_y_mm.setRange(-1e5, 1e5) - self.spin_active_shift_y_mm.setDecimals(4) - self.spin_active_shift_y_mm.setSuffix(" mm") - self.spin_active_shift_y_mm.setMinimumWidth(_SPIN_MIN_W) + self.spin_active_shift_y_mm = _make_dspin(-1e5, 1e5, 4, suffix=" mm") mform.addRow("Shift Y:", self.spin_active_shift_y_mm) mform_box.addLayout(mform) - panel_l.addWidget(self.grp_manual_adjust) + return self.grp_manual_adjust - # ---- Nudge Step Sizes ------------------------------------------------ + def _build_step_group(self) -> QWidget: self.grp_step_sizes, sl = _group("Nudge Step Sizes") sform = _form() - self.spin_step_translate_mm = QDoubleSpinBox() - self.spin_step_translate_mm.setRange(0.0001, 1000.0) - self.spin_step_translate_mm.setDecimals(4) - self.spin_step_translate_mm.setSuffix(" mm") - self.spin_step_translate_mm.setValue(0.01) - self.spin_step_translate_mm.setMinimumWidth(_SPIN_MIN_W) + self.spin_step_translate_mm = _make_dspin(0.0001, 1000.0, 4, + suffix=" mm", value=0.01) sform.addRow("Translate step:", self.spin_step_translate_mm) - self.spin_step_rotate_deg = QDoubleSpinBox() - self.spin_step_rotate_deg.setRange(0.001, 90.0) - self.spin_step_rotate_deg.setDecimals(3) - self.spin_step_rotate_deg.setSuffix(" °") - self.spin_step_rotate_deg.setValue(0.1) - self.spin_step_rotate_deg.setMinimumWidth(_SPIN_MIN_W) + self.spin_step_rotate_deg = _make_dspin(0.001, 90.0, 3, + suffix=" °", value=0.1) sform.addRow("Rotate step:", self.spin_step_rotate_deg) - self.spin_step_multiplier = QDoubleSpinBox() - self.spin_step_multiplier.setRange(1.0, 1000.0) - self.spin_step_multiplier.setDecimals(1) - self.spin_step_multiplier.setValue(10.0) - self.spin_step_multiplier.setMinimumWidth(_SPIN_MIN_W) + self.spin_step_multiplier = _make_dspin(1.0, 1000.0, 1, value=10.0) sform.addRow("Coarse × (Shift):", self.spin_step_multiplier) sl.addLayout(sform) sl.addWidget(_wrap_label( "Arrow keys nudge X/Y translation; Q/E nudge rotation (CCW/CW). " "Hold Shift for the coarse step. Click the image once so it has " "keyboard focus.", _CSS_HINT)) - panel_l.addWidget(self.grp_step_sizes) + return self.grp_step_sizes - # ---- Mask Threshold --------------------------------------------------- + def _build_threshold_group(self, dc_threshold_mv: float) -> QWidget: self.grp_mask_threshold, tl = _group("Mask Threshold") tform = _form() - self.spin_mask_threshold_mv = QDoubleSpinBox() - self.spin_mask_threshold_mv.setRange(-500.0, 500.0) - self.spin_mask_threshold_mv.setDecimals(3) - self.spin_mask_threshold_mv.setSuffix(" mV") - self.spin_mask_threshold_mv.setValue(dc_threshold_mv) - self.spin_mask_threshold_mv.setMinimumWidth(_SPIN_MIN_W) + self.spin_mask_threshold_mv = _make_dspin(-500.0, 500.0, 3, + suffix=" mV", value=dc_threshold_mv) tform.addRow("DC threshold:", self.spin_mask_threshold_mv) tl.addLayout(tform) - panel_l.addWidget(self.grp_mask_threshold) + return self.grp_mask_threshold - # ---- Cross-Correlate (FFT) ----------------------------------------- + def _build_correlate_group(self) -> QWidget: self.grp_correlate, cl = _group("Cross-Correlate (FFT)") cform = _form() self.combo_correlate_source = QComboBox() @@ -361,13 +345,8 @@ class ManualAlignmentDialog(QDialog): self.combo_correlate_source.addItem(label, sources) cform.addRow("Correlate on:", self.combo_correlate_source) - self.spin_correlate_search_deg = QDoubleSpinBox() - self.spin_correlate_search_deg.setRange(0.0, 180.0) - self.spin_correlate_search_deg.setSingleStep(1.0) - self.spin_correlate_search_deg.setDecimals(1) - self.spin_correlate_search_deg.setSuffix(" °") - self.spin_correlate_search_deg.setValue(6.0) - self.spin_correlate_search_deg.setMinimumWidth(_SPIN_MIN_W) + self.spin_correlate_search_deg = _make_dspin(0.0, 180.0, 1, suffix=" °", + value=6.0, step=1.0) cform.addRow("Rotation search (±):", self.spin_correlate_search_deg) cl.addLayout(cform) self.btn_auto_correlate = QPushButton("Auto Cross-Correlate (vs Reference)") @@ -378,9 +357,9 @@ class ManualAlignmentDialog(QDialog): "reported angle is only the starting point of the search, and both " "of its signs are tried. Run this first, then nudge only for small " "corrections.", _CSS_HINT)) - panel_l.addWidget(self.grp_correlate) + return self.grp_correlate - # ---- Actions ------------------------------------------------------ + def _build_actions_group(self) -> QWidget: grp_actions, acl = _group("Actions") self.btn_auto_derotate = QPushButton("Auto De-rotate (use known angles)") self.btn_save = QPushButton("Save Alignment") @@ -388,14 +367,9 @@ class ManualAlignmentDialog(QDialog): self.btn_close = QPushButton("Close") for btn in (self.btn_auto_derotate, self.btn_save, self.btn_clear, self.btn_close): acl.addWidget(btn) - panel_l.addWidget(grp_actions) - - self.lbl_status = _wrap_label("", _CSS_MUTED) - panel_l.addWidget(self.lbl_status) - panel_l.addStretch() - - root.addWidget(_scroll_panel(panel, 320)) + return grp_actions + def _connect_controls(self): self.combo_active_angle.currentIndexChanged.connect(self._on_active_angle_changed) self.spin_active_rotation_deg.editingFinished.connect(self._on_rotation_spin_edited) self.spin_active_shift_x_mm.editingFinished.connect(self._on_shift_spin_edited) @@ -409,9 +383,8 @@ class ManualAlignmentDialog(QDialog): self.canvas.nudge_translate.connect(self._on_nudge_translate) self.canvas.nudge_rotate.connect(self._on_nudge_rotate) - self.combo_active_angle.blockSignals(True) - self.combo_active_angle.setCurrentIndex(self._active_angle) - self.combo_active_angle.blockSignals(False) + with QSignalBlocker(self.combo_active_angle): + self.combo_active_angle.setCurrentIndex(self._active_angle) self._on_active_angle_changed(self._active_angle) # ------------------------------------------------------------------ @@ -426,7 +399,7 @@ class ManualAlignmentDialog(QDialog): return self.lbl_status.setText(f"Preparing masks: 0/{len(missing)} angle(s) needed…") started = self._parent._run_worker( - "manual_align_masks", Ch4MaskWorker(self._sras, missing), + Jobs.MANUAL_ALIGN_MASKS, Ch4MaskWorker(self._sras, missing), connect=( ("angle_done", self._on_mask_angle_done), ("error", lambda msg: self.lbl_status.setText(f"Mask prep error: {msg}")), @@ -541,8 +514,7 @@ class ManualAlignmentDialog(QDialog): dx, dy = self._preview_pitch_mm x_axis = x0 + np.arange(n_cols) * dx y_axis = y0 + np.arange(n_rows) * dy - extent = [x_axis[0] - dx / 2, x_axis[-1] + dx / 2, - y_axis[-1] + dy / 2, y_axis[0] - dy / 2] + extent = _axes_extent(x_axis, y_axis, dx, dy) title = (f"Angle {self._active_angle} active " f"({self._sras.angles_deg[self._active_angle]:.1f}°)") self.canvas.show_overlay(rgba, extent, title) @@ -565,9 +537,8 @@ class ManualAlignmentDialog(QDialog): for spin, val in ((self.spin_active_rotation_deg, p.rotation_deg), (self.spin_active_shift_x_mm, p.shift_mm[0]), (self.spin_active_shift_y_mm, p.shift_mm[1])): - spin.blockSignals(True) - spin.setValue(val) - spin.blockSignals(False) + with QSignalBlocker(spin): + spin.setValue(val) def _on_nudge_translate(self, dir_x: int, dir_y: int, coarse: bool): if not self._masks_ready or self._active_angle == self._ref_angle_idx: @@ -655,7 +626,7 @@ class ManualAlignmentDialog(QDialog): self._set_controls_enabled(False) self.lbl_status.setText(f"Cross-correlating: 0/{self._correlate_total} angle(s)…") started = self._parent._run_worker( - "manual_align_correlate", worker, + Jobs.MANUAL_ALIGN_CORRELATE, worker, connect=( ("angle_done", self._on_correlate_angle_done), ("error", self._on_correlate_error), diff --git a/sras_viewer/main_window.py b/sras_viewer/main_window.py index b0fa6a5..07c3c71 100644 --- a/sras_viewer/main_window.py +++ b/sras_viewer/main_window.py @@ -5,11 +5,11 @@ from pathlib import Path import numpy as np from matplotlib.backends.backend_qtagg import NavigationToolbar2QT -from PyQt6.QtCore import QObject, QSettings, Qt, QThread +from PyQt6.QtCore import QObject, QSettings, QSignalBlocker, Qt, QThread from PyQt6.QtGui import QAction from PyQt6.QtWidgets import ( - QApplication, QCheckBox, QComboBox, QDialog, QDoubleSpinBox, QFileDialog, - QFrame, QHBoxLayout, QLabel, QMainWindow, QProgressDialog, QPushButton, + QApplication, QCheckBox, QComboBox, QDialog, QFileDialog, QFrame, + QHBoxLayout, QLabel, QMainWindow, QProgressDialog, QPushButton, QSizePolicy, QSpinBox, QSplitter, QVBoxLayout, QWidget, ) @@ -30,8 +30,8 @@ from sras_workers import ( from .canvases import ImageCanvas, WaveformCanvas from .common import ( CH1_DERIVED_MODES, CH_LABELS, CMAPS, VELOCITY_MODE_IDX, _CHANNEL_DISPLAY, - _CSS_BUSY, _CSS_HINT, _CSS_INFO, _CSS_MUTED, _CSS_WARN, _LEFT_PANEL_W, - _RIGHT_PANEL_W, _SPIN_MIN_W, _form, _group, _scroll_panel, _wrap_label, + _CSS_BUSY, _axes_extent, _CSS_HINT, _CSS_INFO, _CSS_MUTED, _CSS_WARN, _LEFT_PANEL_W, + _RIGHT_PANEL_W, Jobs, _form, _group, _make_dspin, _scroll_panel, _wrap_label, ) from .dialogs import FftOptionsDialog, ManualAlignmentDialog @@ -175,17 +175,23 @@ class SrasViewerWindow(QMainWindow): panel_layout = QVBoxLayout(panel) panel_layout.setContentsMargins(0, 0, 0, 0) panel_layout.setSpacing(8) + panel_layout.addWidget(self._build_file_group()) + panel_layout.addWidget(self._build_info_group()) + panel_layout.addWidget(self._build_view_group()) + panel_layout.addWidget(self._build_roi_group()) + panel_layout.addStretch() + return _scroll_panel(panel, _LEFT_PANEL_W) - # ---- File ------------------------------------------------------- + def _build_file_group(self) -> QWidget: grp_file, fl = _group("File") self.btn_open = QPushButton("Open .sras…") self.btn_open.clicked.connect(self._on_open) self.lbl_filename = _wrap_label("No file loaded", _CSS_MUTED) fl.addWidget(self.btn_open) fl.addWidget(self.lbl_filename) - panel_layout.addWidget(grp_file) + return grp_file - # ---- Scan info -------------------------------------------------- + def _build_info_group(self) -> QWidget: grp_info, il = _group("Scan Info") il.setSpacing(3) self._info = {} @@ -202,9 +208,9 @@ class SrasViewerWindow(QMainWindow): # background DC-precompute progress self.lbl_dc_precompute = _wrap_label("", _CSS_BUSY) il.addWidget(self.lbl_dc_precompute) - panel_layout.addWidget(grp_info) + return grp_info - # ---- View settings ---------------------------------------------- + def _build_view_group(self) -> QWidget: grp_view, vl = _group("View Settings") view_form = _form() @@ -244,14 +250,9 @@ class SrasViewerWindow(QMainWindow): # DC threshold (for RF / CH1 masking) self.grp_threshold, tl = _group("RF Mask Threshold (CH1 only)") thr_form = _form() - self.spin_threshold_mv = QDoubleSpinBox() - self.spin_threshold_mv.setRange(-500.0, 500.0) - self.spin_threshold_mv.setDecimals(3) - self.spin_threshold_mv.setSingleStep(0.025) - self.spin_threshold_mv.setSuffix(" mV") - self.spin_threshold_mv.setValue(50.0) + self.spin_threshold_mv = _make_dspin(-500.0, 500.0, 3, suffix=" mV", + value=50.0, step=0.025) self.spin_threshold_mv.setEnabled(False) - self.spin_threshold_mv.setMinimumWidth(_SPIN_MIN_W) self.spin_threshold_mv.editingFinished.connect(self._on_threshold_changed) thr_form.addRow("DC threshold:", self.spin_threshold_mv) tl.addLayout(thr_form) @@ -289,10 +290,9 @@ class SrasViewerWindow(QMainWindow): "Save the current CH1 image (one scan row per CSV line).") self.btn_export_csv.clicked.connect(self._on_export_csv) vl.addWidget(self.btn_export_csv) + return grp_view - panel_layout.addWidget(grp_view) - - # ---- ROI --------------------------------------------------------- + def _build_roi_group(self) -> QWidget: grp_roi, rl = _group("ROI (Region of Interest)") self.btn_draw_roi = QPushButton("Draw ROI") @@ -327,10 +327,7 @@ class SrasViewerWindow(QMainWindow): self.lbl_roi_npix = _wrap_label("pixels inside: —", _CSS_HINT) for lbl in (self.lbl_roi_center, self.lbl_roi_size, self.lbl_roi_npix): rl.addWidget(lbl) - - panel_layout.addWidget(grp_roi) - panel_layout.addStretch() - return _scroll_panel(panel, _LEFT_PANEL_W) + return grp_roi def _build_canvases(self) -> QWidget: splitter = QSplitter(Qt.Orientation.Vertical) @@ -372,14 +369,9 @@ class SrasViewerWindow(QMainWindow): # Velocity settings (visible only in velocity mode) self.grp_velocity, vel_l = _group("Velocity Settings (CH1 only)") vel_form = _form() - self.spin_grating_um = QDoubleSpinBox() - self.spin_grating_um.setRange(0.1, 1000.0) - self.spin_grating_um.setDecimals(2) - self.spin_grating_um.setSingleStep(0.5) - self.spin_grating_um.setSuffix(" µm") - self.spin_grating_um.setValue(25) + self.spin_grating_um = _make_dspin(0.1, 1000.0, 2, suffix=" µm", + value=25, step=0.5) self.spin_grating_um.setEnabled(False) - self.spin_grating_um.setMinimumWidth(_SPIN_MIN_W) self.spin_grating_um.editingFinished.connect(self._on_grating_changed) vel_form.addRow("Grating size:", self.spin_grating_um) vel_l.addLayout(vel_form) @@ -407,11 +399,8 @@ class SrasViewerWindow(QMainWindow): range_form = _form() for label, attr in (("min:", "spin_vmin"), ("max:", "spin_vmax")): - spin = QDoubleSpinBox() - spin.setRange(-1e9, 1e9) - spin.setDecimals(4) + spin = _make_dspin(-1e9, 1e9, 4) spin.setEnabled(False) - spin.setMinimumWidth(_SPIN_MIN_W) spin.editingFinished.connect(self._on_manual_range_changed) setattr(self, attr, spin) range_form.addRow(label, spin) @@ -494,7 +483,7 @@ class SrasViewerWindow(QMainWindow): def _load_file(self, path: str): started = self._run_worker( - "load", LoadWorker(path), + Jobs.LOAD, LoadWorker(path), connect=( ("finished", self._on_load_done), ("error", lambda msg: self.statusBar().showMessage(f"Error: {msg}")), @@ -530,25 +519,18 @@ class SrasViewerWindow(QMainWindow): self._dc_generation += 1 self.lbl_dc_precompute.setText("") - self._alignment_result = None - self._aligned_cache = {} - self._alignment_generation += 1 - self.chk_aligned_view.blockSignals(True) - self.chk_aligned_view.setChecked(False) - self.chk_aligned_view.setEnabled(False) - self.chk_aligned_view.blockSignals(False) + self._apply_alignment_result(None, view_checked=False) # Silently restore a previously-saved manual alignment, if any, so # the work survives closing and reopening the file. sidecar = load_manual_alignment(sras) if sidecar is not None: try: - self._alignment_result = build_manual_alignment( - sras, sidecar.ref_angle_idx, sidecar.dc_threshold_mv, - sidecar.per_angle) - self.chk_aligned_view.blockSignals(True) - self.chk_aligned_view.setChecked(True) - self.chk_aligned_view.blockSignals(False) + self._apply_alignment_result( + build_manual_alignment( + sras, sidecar.ref_angle_idx, sidecar.dc_threshold_mv, + sidecar.per_angle), + view_checked=True) self.statusBar().showMessage( f"Restored saved manual alignment from " f"{sidecar_path(sras.path).name}") @@ -564,17 +546,15 @@ class SrasViewerWindow(QMainWindow): self.lbl_filename.setText(sras.path.name) - self.spin_angle.blockSignals(True) - self.spin_angle.setRange(0, max(0, sras.n_angles - 1)) - self.spin_angle.setValue(0) - self.spin_angle.blockSignals(False) + with QSignalBlocker(self.spin_angle): + self.spin_angle.setRange(0, max(0, sras.n_angles - 1)) + self.spin_angle.setValue(0) # DC channels are cheap and give an instant, fluid overview of a # scan; CH1/Velocity require an FFT per pixel that can take minutes # on a large scan, so don't default to it. - self.combo_channel.blockSignals(True) - self.combo_channel.setCurrentIndex(CH4_IDX) - self.combo_channel.blockSignals(False) + with QSignalBlocker(self.combo_channel): + self.combo_channel.setCurrentIndex(CH4_IDX) self._update_controls_enabled(True) self._on_threshold_changed() # refresh ADC label with file calibration @@ -657,14 +637,14 @@ class SrasViewerWindow(QMainWindow): # Batch Convert actions pick their own files, independent of # whatever's currently open — only gated on no batch already running. - can_batch = not self._job_running("batch") + can_batch = not self._job_running(Jobs.BATCH) self._batch_dc_act.setEnabled(can_batch) self._batch_fft_act.setEnabled(can_batch) self._alignment_act.setEnabled( - has_file and s.n_angles > 1 and not self._job_running("align")) + has_file and s.n_angles > 1 and not self._job_running(Jobs.ALIGN)) self._manual_align_act.setEnabled( - has_file and s.n_angles > 1 and not self._job_running("align")) + has_file and s.n_angles > 1 and not self._job_running(Jobs.ALIGN)) self.chk_aligned_view.setEnabled(enabled and self._alignment_result is not None) self._update_roi_ui() @@ -676,7 +656,7 @@ class SrasViewerWindow(QMainWindow): # Background subtraction changes the FFT input, so it genuinely # invalidates the cached raw FFT (the cache key includes it) — # _refresh_display() recomputes only on a miss for the new state. - if self._sras is not None and self.combo_channel.currentIndex() in CH1_DERIVED_MODES: + if self._is_fft_mode(): self._refresh_display() def _on_grating_changed(self): @@ -693,7 +673,7 @@ class SrasViewerWindow(QMainWindow): # Threshold decides which pixels get an FFT at all, so changing it is # a genuine cache-key change — but the recompute reuses the cached DC4 # image to skip masked-out pixels. - if self._sras is not None and self.combo_channel.currentIndex() in CH1_DERIVED_MODES: + if self._is_fft_mode(): self._refresh_display() def _on_autoscale_toggled(self, checked: bool): @@ -810,9 +790,8 @@ class SrasViewerWindow(QMainWindow): def _on_draw_mode_changed(self, active: bool): # Keep the toggle button's visual state in sync with the canvas. - self.btn_draw_roi.blockSignals(True) - self.btn_draw_roi.setChecked(active) - self.btn_draw_roi.blockSignals(False) + with QSignalBlocker(self.btn_draw_roi): + self.btn_draw_roi.setChecked(active) def _on_clear_roi(self): self.image_canvas.clear_roi() @@ -859,6 +838,11 @@ class SrasViewerWindow(QMainWindow): return None return self._sras.samples_per_frame * self._fft_pad_factor + def _is_fft_mode(self) -> bool: + """Is the selected channel an FFT-derived (CH1/Velocity) mode?""" + return (self._sras is not None + and self.combo_channel.currentIndex() in CH1_DERIVED_MODES) + def _scale_for_display(self, freq_mhz: np.ndarray, ch_idx: int) -> np.ndarray: """Velocity is a pure post-multiply of the (already DC-masked) cached frequency image — never worth a recompute on its own.""" @@ -947,15 +931,13 @@ class SrasViewerWindow(QMainWindow): dx = x_axis[1] - x_axis[0] if len(x_axis) > 1 else s.pixel_x_mm dy = float(y_axis[1] - y_axis[0]) if len(y_axis) > 1 else 1.0 - extent = [x_axis[0] - dx / 2, x_axis[-1] + dx / 2, - y_axis[-1] + dy / 2, y_axis[0] - dy / 2] + extent = _axes_extent(x_axis, y_axis, dx, dy) if self.chk_auto.isChecked(): vmin, vmax = float(display_img.min()), float(display_img.max()) for spin, val in ((self.spin_vmin, vmin), (self.spin_vmax, vmax)): - spin.blockSignals(True) - spin.setValue(val) - spin.blockSignals(False) + with QSignalBlocker(spin): + spin.setValue(val) else: vmin, vmax = self.spin_vmin.value(), self.spin_vmax.value() @@ -988,7 +970,7 @@ class SrasViewerWindow(QMainWindow): # ------------------------------------------------------------------ def _start_compute(self): - if self._sras is None or self._job_running("compute"): + if self._sras is None or self._job_running(Jobs.COMPUTE): return # re-checked when the running compute finishes angle_idx = self.spin_angle.value() @@ -1013,7 +995,7 @@ class SrasViewerWindow(QMainWindow): is_fft_mode=is_fft, ) if not self._run_worker( - "compute", worker, + Jobs.COMPUTE, worker, connect=( ("finished", self._on_compute_done), ("error", lambda msg: self.statusBar().showMessage( @@ -1073,7 +1055,7 @@ class SrasViewerWindow(QMainWindow): worker = DcPrecomputeWorker(self._sras) self._run_worker( - "dc_precompute", worker, + Jobs.DC_PRECOMPUTE, worker, connect=( ("angle_done", lambda a, dc3, dc4, g=generation: self._on_dc_precompute_angle_done(g, a, dc3, dc4, n_angles)), @@ -1101,7 +1083,7 @@ class SrasViewerWindow(QMainWindow): # caught up and are still waiting), show it now. current_ch = self.combo_channel.currentIndex() if (angle_idx == self.spin_angle.value() - and not self._job_running("compute") + and not self._job_running(Jobs.COMPUTE) and current_ch in (CH3_IDX, CH4_IDX) and (self._current_angle != angle_idx or self._current_ch != current_ch)): self._refresh_display() @@ -1171,7 +1153,7 @@ class SrasViewerWindow(QMainWindow): # ------------------------------------------------------------------ def _on_batch_compute(self, mode: str): - if self._job_running("batch"): + if self._job_running(Jobs.BATCH): return label = "DC" if mode == "dc" else "FFT" paths, _ = QFileDialog.getOpenFileNames( @@ -1183,9 +1165,9 @@ class SrasViewerWindow(QMainWindow): self._batch_errors = [] worker = BatchCacheWorker(paths, mode, self.chk_bg_sub.isChecked()) started = self._run_worker( - "batch", worker, + Jobs.BATCH, worker, connect=( - ("progress", lambda pct: self._set_progress("batch", pct)), + ("progress", lambda pct: self._set_progress(Jobs.BATCH, pct)), ("file_done", self._on_batch_file_done), ("finished", lambda p=paths: self._on_batch_finished(p)), ), @@ -1197,16 +1179,16 @@ class SrasViewerWindow(QMainWindow): self._batch_dc_act.setEnabled(False) self._batch_fft_act.setEnabled(False) self._show_progress( - "batch", f"Batch computing {label} for {len(paths)} file(s)…", + Jobs.BATCH, f"Batch computing {label} for {len(paths)} file(s)…", maximum=100) def _on_batch_file_done(self, path: str, err: str): if err: self._batch_errors.append(f"{Path(path).name} — {err}") - self._show_progress("batch", f"Processed {Path(path).name}…") + self._show_progress(Jobs.BATCH, f"Processed {Path(path).name}…") def _on_batch_finished(self, paths: list[str]): - self._close_progress("batch") + self._close_progress(Jobs.BATCH) n_total = len(paths) n_failed = len(self._batch_errors) @@ -1240,7 +1222,7 @@ class SrasViewerWindow(QMainWindow): generation = self._alignment_generation started = self._run_worker( - "align", AngleAlignmentWorker(self._sras, ref_idx, threshold_mv), + Jobs.ALIGN, AngleAlignmentWorker(self._sras, ref_idx, threshold_mv), connect=( ("progress", lambda pct: self._set_progress("main", pct)), ("finished", lambda result, err, g=generation: @@ -1265,12 +1247,9 @@ class SrasViewerWindow(QMainWindow): if error_msg: self.statusBar().showMessage(f"Angle alignment failed: {error_msg}") return - self._alignment_result = result - self._aligned_cache = {} - self.chk_aligned_view.setEnabled(True) - self.chk_aligned_view.blockSignals(True) - self.chk_aligned_view.setChecked(True) - self.chk_aligned_view.blockSignals(False) + # No generation bump: this result *is* the current generation's. + self._apply_alignment_result(result, view_checked=True, + bump_generation=False) nr, nc = result.canvas_shape self.statusBar().showMessage( f"Angle alignment computed ({self._sras.n_angles} angles, " @@ -1320,32 +1299,32 @@ class SrasViewerWindow(QMainWindow): def _on_manual_align_dialog_closed(self, _result_code: int): self._manual_align_dialog = None - def _on_manual_alignment_saved(self, result, sidecar_path_str: str): + def _apply_alignment_result(self, result, *, view_checked: bool, + bump_generation: bool = True): + """Install (or clear, with result=None) the active alignment: reset + the aligned-image cache and set the Aligned View checkbox without + firing its change signal.""" self._alignment_result = result self._aligned_cache = {} - self._alignment_generation += 1 - self.chk_aligned_view.setEnabled(True) - self.chk_aligned_view.blockSignals(True) - self.chk_aligned_view.setChecked(True) - self.chk_aligned_view.blockSignals(False) + if bump_generation: + self._alignment_generation += 1 + with QSignalBlocker(self.chk_aligned_view): + self.chk_aligned_view.setChecked(view_checked) + self.chk_aligned_view.setEnabled(result is not None) + + def _manual_alignment_changed(self, result, message: str): + self._apply_alignment_result(result, view_checked=result is not None) self._update_controls_enabled(self._sras is not None) - self.statusBar().showMessage( - f"Manual alignment saved to {Path(sidecar_path_str).name}") + self.statusBar().showMessage(message) if self._current_image is not None: self._refresh_display() + def _on_manual_alignment_saved(self, result, sidecar_path_str: str): + self._manual_alignment_changed( + result, f"Manual alignment saved to {Path(sidecar_path_str).name}") + def _on_manual_alignment_cleared(self): - self._alignment_result = None - self._aligned_cache = {} - self._alignment_generation += 1 - self.chk_aligned_view.blockSignals(True) - self.chk_aligned_view.setChecked(False) - self.chk_aligned_view.setEnabled(False) - self.chk_aligned_view.blockSignals(False) - self._update_controls_enabled(self._sras is not None) - self.statusBar().showMessage("Manual alignment cleared.") - if self._current_image is not None: - self._refresh_display() + self._manual_alignment_changed(None, "Manual alignment cleared.") # ------------------------------------------------------------------ # FFT Options @@ -1369,7 +1348,7 @@ class SrasViewerWindow(QMainWindow): # Pad factor changes the FFT bin count, so it genuinely invalidates # the cached raw FFT (part of the cache key) — _refresh_display() # recomputes only on a cache miss. - if self._sras is not None and self.combo_channel.currentIndex() in CH1_DERIVED_MODES: + if self._is_fft_mode(): self._refresh_display() # ------------------------------------------------------------------ From 8154c57066f75d88a714b43befeb7a7cce560946 Mon Sep 17 00:00:00 2001 From: Thomas Ales Date: Thu, 6 Aug 2026 11:01:25 -0500 Subject: [PATCH 10/10] Move design essays to docs/design.md, leave pointers MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The three block essays in sras_compute.py — memory budget / row chunking, angle-alignment coordinate frames, and the sidecar placement + schema history — move to docs/design.md (joined by a new section on the zoom FFT peak search), each replaced by a 2-4 line pointer. The save_manual_alignment docstring no longer restates the JSON schema its own eight lines of code construct. Load-bearing trap notes (normalization=None, subpixel-score mixing, memmap lazy reads) stay in place. Zero code changes — golden-hash diff verified empty. Co-Authored-By: Claude Fable 5 --- .gitignore | 1 + docs/design.md | 137 ++++++++++++++++++++++++++++++++++++++++++++++++ sras_compute.py | 103 ++++++------------------------------ 3 files changed, 155 insertions(+), 86 deletions(-) create mode 100644 docs/design.md diff --git a/.gitignore b/.gitignore index 2694b86..33ad664 100644 --- a/.gitignore +++ b/.gitignore @@ -6,3 +6,4 @@ __pycache__/ *.sras baseline*.txt after*.txt +sras_viewer.egg-info/ diff --git a/docs/design.md b/docs/design.md new file mode 100644 index 0000000..3ea672b --- /dev/null +++ b/docs/design.md @@ -0,0 +1,137 @@ +# sras-viewer design notes + +Rationale that outgrew code comments. Each section is referenced by a short +pointer comment at the relevant definition, so the code stays scannable and +the reasoning stays findable. + +## Memory budget and row chunking (`sras_compute.py`) + +DC images are computed over row chunks so the float32 working buffers for one +chunk stay under a memory budget. A fixed row count (the original design) +works fine for small legacy scans but is catastrophic for a v6 scan with a +large per-angle frame/sample count — e.g. a 7500-frame × 2500-sample angle +needs ~2.4 GB for a single 32-row chunk. + +With chunks running concurrently the budget has to cover *all* live chunks at +once. On a large scan `chunk_rows` is already clamped to its floor of one row +(one row alone is ~75 MB of float32 at 7507×2500), so shrinking the per-chunk +size cannot buy more concurrency — the worker count must be derived from the +budget instead: `_plan_chunks` picks the worker count *first* and sizes the +chunk to it. Sizing the chunk first is the trap: a single chunk would always +consume the whole budget and leave room for exactly one worker, precisely on +the large scans that need concurrency most. + +The 1024 MB default (`SRAS_MEM_BUDGET_MB`) is the measured knee on a 16-core +machine against a 7507-frame × 2500-sample angle: 512 MB left ~20% of the +speedup on the table, and 1536+ MB cost ~0.4 GB more resident memory for no +further gain. + +A caller that itself runs several computations concurrently (angle-level +parallelism, `plan_angle_level`) must pass *both* `max_workers=1` and its +share of the budget. Capping the workers alone is not enough: the chunk would +still be sized against the whole budget, and N concurrent callers would each +allocate all of it. + +## FFT peak search: block-parallel zoom refinement (`sras_compute.py`) + +The displayed RF value per pixel is the argmax of the zero-padded power +spectrum of that pixel's CH1 waveform. At the pad factor of 40 needed for +mapping resolution, materialising padded spectra is hopeless: ~9 GB per scan +row, which is what used to collapse the old row-chunk planner to one worker +and make synthesis single-threaded. + +`_peak_bins_zoom` never materialises the padded spectrum: + +1. a coarse rfft at `next_fast_len(2*spf)` — 2× oversampled, so the padded + power spectrum (a trig polynomial of degree spf−1) cannot hide its global + max between coarse samples; +2. every coarse bin within `_ZOOM_CAND_RATIO` (0.7) of its row's coarse max + becomes a refinement candidate. Quarter-natural-bin scalloping at the 2× + grid can understate a peak's power by at most ~19%, so 0.7 keeps a wide + margin. The DC-adjacent window is always refined too: the coarse DC bin + is zeroed for suppression, which would otherwise blind the scan to fine + bins closer to DC than the first coarse sample (where the leakage skirt + of an un-subtracted offset peaks); +3. each candidate window (±`_ZOOM_HALFWIDTH` = 0.75 coarse spacings; every + fine bin lies within 0.5 spacings of its nearest coarse bin) is evaluated + on the exact `n_fft` grid by one small complex gemm, with np.argmax's + lowest-bin tie-break preserved across windows. + +The selected bin is bit-identical to the full padded argmax — enforced by +`tests/test_compute.py::test_zoom_identity`, a fuzz test over adversarial +spectra, and the golden-hash harness (`tools/check_equivalence.py`), whose +baseline was captured on the old full-padded path. + +Work fans out over a persistent thread pool in `_FFT_BLOCK` = 512-waveform +tasks: smaller blocks serialise on GIL-held numpy dispatch, larger ones lose +cache residency and task granularity (measured on a 16-core machine, where +this path runs ~35× faster than the old serial padded transform at pad 40). +pyFFTW runs through per-thread `builders` plans (FFTW_MEASURE, wisdom +persisted under `~/.cache/sras-viewer/`), and `threadpoolctl` clamps BLAS to +one thread under the pool so the refinement gemm cannot oversubscribe. +`compute_rf_image(exact=True)` (or `SRAS_FFT_EXACT=1`) keeps the reference +full-padded path for audits. + +## Angle alignment coordinate frames (`sras_compute.py`) + +Alignment puts every angle's images onto one shared, zero-padded pixel grid +using a rigid transform only — rotation + translation, never scale. + +Angle 0 (the reference) is the sole coordinate authority: it is the only +angle whose stage XY (`x_start_mm` / `y_positions_mm`) is ever read, and the +shared canvas is literally an extension of angle 0's own pixel grid, so the +aligned view carries angle 0's real X/Y axes. Every *other* angle is placed +purely by content — its rotation and translation come from cross-correlating +its CH4 image against angle 0's (`register_angle_to_reference`) — and its own +stage XY is deliberately never consulted. That is not an oversight: the +rotation stage moves the sample relative to the scan window, so where a +window sat in stage coordinates says nothing about where the sample is, and +an earlier design that pivoted each angle on a signal-weighted centroid of +its own window put every angle on a ~20 mm circle around the optical center +instead of stacking them into one shape. + +Only two coordinate frames exist: + +* **local mm** — one angle's own physical frame: origin at the *center of its + own pixel array*, x along +column, y along +row, scaled by that angle's own + pitches. Carries no stage position whatsoever. +* **ref mm** — the reference angle's local mm. A registration result + `(rotation_deg, shift_mm)` is exactly the rigid map from an angle's local + mm to ref mm: `q = R(rotation_deg) @ l + shift_mm`. Stage coordinates + re-enter once, at the very end, when the canvas origin is converted to + angle 0's stage mm (`AlignmentResult.canvas_origin_mm`). + +Rotation is done in mm, never on raw pixel indices: the x pitch +(`SrasFile.pixel_x_mm`, 5 µm on a real scan) and the y/row pitch (50 µm) +differ by 10×, so rotating the raw index grid would shear the image — an +unwanted anisotropic scale. Registration runs on a resampled *isotropic* grid +for the same reason, and every affine maps shared-grid index → mm → undo +rotation/shift → that angle's own local mm → that angle's own raw index, +matching the output→input convention `scipy.ndimage.affine_transform` wants. + +## Manual-alignment sidecar (`sras_compute.py`) + +`.sras.align.json` lives next to the scan file. The code lives in +`sras_compute`, not `sras_format`: `sras_format` is scoped to the versioned +binary .sras spec itself (see `scan_format.md`), while a manual alignment is +a viewer-computed *derived* artifact, analogous in kind to `AlignmentResult` +— so it belongs with the alignment math it serialises. json + pathlib are +stdlib, so this adds no dependency to a module whose load-bearing constraint +is staying free of Qt/matplotlib for cheap multiprocessing-child imports. + +### Schema history + +The stored `rotation_deg`/`shift_mm` are meaningless without the frame they +were measured in, so `_SIDECAR_SCHEMA_VERSION` is bumped whenever that frame +changes. Each bump makes older files describe a different (and, for the bugs +each bump fixed, actively wrong) transform than the same numbers would today; +loading one unchanged would silently reproduce the very "scans show up +everywhere" symptom the bump fixed — so older sidecars are treated as absent +rather than migrated. + +* **1 → 2** — pivot moved from the scan-window bbox center to a + content-derived centroid, and the rotation sign convention was corrected. +* **2 → 3** — the content centroid was abandoned entirely: rotation is now + about each angle's own array center, mapped onto the reference's array + center, with `shift_mm` in the reference's local mm frame. No angle but the + reference contributes stage coordinates any more. diff --git a/sras_compute.py b/sras_compute.py index 0305c54..9120959 100644 --- a/sras_compute.py +++ b/sras_compute.py @@ -276,22 +276,10 @@ def _peak_bins_direct(waves: np.ndarray, n_len: int) -> np.ndarray: # Chunking / parallel budget # --------------------------------------------------------------------------- # -# Rows are batched so the float32 working buffers for one chunk stay under a -# memory budget. A fixed row count (the original design) works fine for small -# legacy scans but is catastrophic for a v6 scan with a large per-angle -# frame/sample count — e.g. a 7500-frame x 2500-sample angle needs ~2.4 GB for -# a single 32-row chunk. -# -# With chunks running concurrently the budget has to cover *all* live chunks at -# once. Note that on a large scan chunk_rows is already clamped to its floor of -# 1 row (one row alone is ~75 MB of float32 at 7507x2500), so shrinking the -# per-chunk size cannot buy more concurrency — the worker count must be derived -# from the budget instead. See _plan_chunks. - -# 1024 MB is the measured knee on a 16-core machine against a 7507-frame x -# 2500-sample angle: 512 MB leaves ~20% of the FFT speedup on the table, and -# 1536+ MB costs ~0.4 GB more resident for no further gain. Override with -# SRAS_MEM_BUDGET_MB on a smaller machine. +# Row chunks are budgeted so all concurrently-live working buffers fit in +# memory; the worker count is derived from the budget, not vice versa. +# Rationale and the measured 1024 MB default: docs/design.md ("Memory budget +# and row chunking"). _TOTAL_BYTES_BUDGET = int(os.environ.get("SRAS_MEM_BUDGET_MB", 1024)) * 1024 * 1024 _CHUNK_ROWS_MAX = 32 # cap for small scans (original behavior) _MAX_WORKERS = int(os.environ.get("SRAS_MAX_WORKERS", 0)) or (os.cpu_count() or 4) @@ -659,41 +647,12 @@ def cache_file(path: str, mode: str, apply_bg_sub: bool, # --------------------------------------------------------------------------- # Angle alignment (Fusion menu) # -# Puts every angle's images onto one shared, zero-padded pixel grid using a -# rigid transform only — rotation + translation, never scale. -# -# Angle 0 (the reference) is the sole coordinate authority: it is the only -# angle whose stage XY (x_start_mm / y_positions_mm) is ever read, and the -# shared canvas is literally an extension of angle 0's own pixel grid, so the -# aligned view carries angle 0's real X/Y axes. Every *other* angle is placed -# purely by content — its rotation and translation come from cross-correlating -# its CH4 image against angle 0's (register_angle_to_reference) — and its own -# stage XY is deliberately never consulted. That is not an oversight: the -# rotation stage moves the sample relative to the scan window, so where a -# window sat in stage coordinates says nothing about where the sample is, and -# an earlier design that pivoted each angle on a signal-weighted centroid of -# its own window put every angle on a ~20 mm circle around the optical center -# instead of stacking them into one shape. -# -# Only two coordinate frames exist here: -# -# local mm — one angle's own physical frame: origin at the *center of its own -# pixel array*, x along +column, y along +row, scaled by that angle's own -# pitches. Carries no stage position whatsoever. -# -# ref mm — the reference angle's local mm. A registration result -# (rotation_deg, shift_mm) is exactly the rigid map from an angle's local -# mm to ref mm: q = R(rotation_deg) @ l + shift_mm. Stage coordinates -# re-enter once, at the very end, when the canvas origin is converted to -# angle 0's stage mm (AlignmentResult.canvas_origin_mm). -# -# Rotation is done in mm, never on raw pixel indices: the x pitch -# (SrasFile.pixel_x_mm, 5 µm on a real scan) and the y/row pitch (50 µm) differ -# by 10x, so rotating the raw index grid would shear the image — an unwanted -# anisotropic scale. Registration runs on a resampled *isotropic* grid for the -# same reason, and every affine here maps shared-grid index -> mm -> undo -# rotation/shift -> that angle's own local mm -> that angle's own raw index, -# matching the output->input convention scipy.ndimage.affine_transform wants. +# Rigid transforms only (rotation + translation, never scale), computed in mm +# on two frames: each angle's "local mm" (origin at its own array center) and +# the reference angle's local mm. Angle 0 is the sole coordinate authority; +# every other angle is placed purely by image content. Why, and the full +# frame/affine conventions: docs/design.md ("Angle alignment coordinate +# frames"). # --------------------------------------------------------------------------- @dataclass @@ -1555,14 +1514,8 @@ def build_manual_alignment(sras: SrasFile, ref_angle_idx: int, # ---- Sidecar persistence (.sras.align.json) ------------------------- -# -# Lives here, not sras_format.py: sras_format.py is scoped to the versioned -# binary .sras spec itself (see scan_format.md); a manual alignment is a -# viewer-computed *derived* artifact, analogous in kind to AlignmentResult — -# so it belongs with the alignment math it serialises, which already lives -# in this module. json + pathlib are both stdlib, so this doesn't add a new -# dependency to a module whose only load-bearing constraint is staying free -# of Qt/matplotlib for cheap multiprocessing-child imports. +# A viewer-computed derived artifact, so it lives with the alignment math +# rather than in sras_format (see docs/design.md, "Manual-alignment sidecar"). @dataclass class ManualAlignmentSidecar: @@ -1579,18 +1532,9 @@ def sidecar_path(sras_path) -> Path: return p.with_name(p.name + ".align.json") -# The stored rotation_deg/shift_mm are meaningless without the frame they were -# measured in, so this is bumped whenever that frame changes. Each bump makes -# older files describe a different (and, for the bugs each bump fixed, actively -# wrong) transform than the same numbers would today, and loading one unchanged -# would silently reproduce the very "scans show up everywhere" symptom the bump -# fixed — so older sidecars are treated as absent rather than migrated. -# 1 -> 2 pivot moved from the scan-window bbox center to a content-derived -# centroid, and the rotation sign convention was corrected. -# 2 -> 3 the content centroid was abandoned entirely: rotation is now about -# each angle's own array center, mapped onto the reference's array -# center, with shift_mm in the reference's local mm frame. No angle -# but the reference contributes stage coordinates any more. +# Bumped whenever the frame the stored numbers are measured in changes; older +# sidecars are treated as absent, never migrated. Bump history: +# docs/design.md ("Schema history"). _SIDECAR_SCHEMA_VERSION = 3 @@ -1598,21 +1542,8 @@ def save_manual_alignment(sras: SrasFile, ref_angle_idx: int, dc_threshold_mv: float, per_angle: dict[int, ManualAngleParams]) -> Path: """Write the sidecar JSON for sras.path (overwriting any existing one) - and return the path written. - - Schema (schema_version 3): - { - "schema_version": 3, - "ref_angle_idx": , - "dc_threshold_mv": , - "per_angle": { - "": {"rotation_deg": , "shift_mm": [, ]}, - ... - } - } - Angle indices are JSON object keys, so they round-trip as strings — - load_manual_alignment converts them back to int. - """ + and return the path written. Angle indices become JSON object keys, so + they round-trip as strings — load_manual_alignment converts them back.""" path = sidecar_path(sras.path) payload = { "schema_version": _SIDECAR_SCHEMA_VERSION,