Refactor/parallel and dedup #1
+98
-25
@@ -66,7 +66,11 @@ def _do_rfft(x: np.ndarray, n: int | None = None, axis: int = -1,
|
|||||||
# per-chunk size cannot buy more concurrency — the worker count must be derived
|
# per-chunk size cannot buy more concurrency — the worker count must be derived
|
||||||
# from the budget instead. See _plan_chunks.
|
# from the budget instead. See _plan_chunks.
|
||||||
|
|
||||||
_TOTAL_BYTES_BUDGET = int(os.environ.get("SRAS_MEM_BUDGET_MB", 1536)) * 1024 * 1024
|
# 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.
|
||||||
|
_TOTAL_BYTES_BUDGET = int(os.environ.get("SRAS_MEM_BUDGET_MB", 1024)) * 1024 * 1024
|
||||||
_CHUNK_ROWS_MAX = 32 # cap for small scans (original behavior)
|
_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)
|
_MAX_WORKERS = int(os.environ.get("SRAS_MAX_WORKERS", 0)) or (os.cpu_count() or 4)
|
||||||
|
|
||||||
@@ -83,21 +87,38 @@ def _chunk_rows_for(n_frames: int, samples_per_frame: int,
|
|||||||
|
|
||||||
def _plan_chunks(n_rows: int, n_frames: int, samples_per_frame: int,
|
def _plan_chunks(n_rows: int, n_frames: int, samples_per_frame: int,
|
||||||
live_multiplier: int = 1,
|
live_multiplier: int = 1,
|
||||||
max_workers: int | None = None) -> tuple[int, int]:
|
max_workers: int | None = None,
|
||||||
"""(chunk_rows, n_workers) sized so live_multiplier x chunk_rows x
|
budget: int | None = None) -> tuple[int, int]:
|
||||||
n_frames x samples x 4 bytes x n_workers stays within the budget."""
|
"""(chunk_rows, n_workers) such that all concurrent chunks together fit
|
||||||
per_chunk_budget = max(1, _TOTAL_BYTES_BUDGET // live_multiplier)
|
the budget: n_workers x live_multiplier x chunk_rows x n_frames x
|
||||||
chunk_rows = _chunk_rows_for(n_frames, samples_per_frame, per_chunk_budget)
|
samples x 4 bytes <= budget.
|
||||||
|
|
||||||
chunk_bytes = max(1, live_multiplier * chunk_rows * n_frames * samples_per_frame * 4)
|
The worker count is chosen *first* and the chunk sized to it. Sizing the
|
||||||
|
chunk first is the trap: _chunk_rows_for spends whatever budget it is
|
||||||
|
given, so a single chunk would always consume the lot and leave room for
|
||||||
|
exactly one worker — no concurrency, on precisely the large scans that
|
||||||
|
need it most.
|
||||||
|
|
||||||
|
A caller that is itself running several of these concurrently 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.
|
||||||
|
"""
|
||||||
|
total = _TOTAL_BYTES_BUDGET if budget is None else max(1, budget)
|
||||||
cap = max_workers if max_workers is not None else _MAX_WORKERS
|
cap = max_workers if max_workers is not None else _MAX_WORKERS
|
||||||
n_workers = int(max(1, min(cap,
|
bytes_per_row = max(1, live_multiplier * n_frames * samples_per_frame * 4)
|
||||||
_TOTAL_BYTES_BUDGET // chunk_bytes,
|
|
||||||
-(-n_rows // chunk_rows))))
|
# A chunk is at least one row, so that alone caps how many can be live.
|
||||||
|
n_workers = int(max(1, min(cap, n_rows, total // bytes_per_row)))
|
||||||
|
chunk_rows = _chunk_rows_for(n_frames, samples_per_frame,
|
||||||
|
max(1, total // (n_workers * live_multiplier)))
|
||||||
|
# With chunk_rows known, more workers than chunks buys nothing.
|
||||||
|
n_workers = int(max(1, min(n_workers, -(-n_rows // chunk_rows))))
|
||||||
return chunk_rows, n_workers
|
return chunk_rows, n_workers
|
||||||
|
|
||||||
|
|
||||||
def _map_row_chunks(n_rows: int, chunk_rows: int, n_workers: int, fn):
|
def _map_row_chunks(n_rows: int, chunk_rows: int, n_workers: int, fn,
|
||||||
|
should_stop=None):
|
||||||
"""Apply fn(r0, r1) over row chunks, in parallel when it pays.
|
"""Apply fn(r0, r1) over row chunks, in parallel when it pays.
|
||||||
|
|
||||||
Chunks write to disjoint output slices, so no locking is needed. numpy
|
Chunks write to disjoint output slices, so no locking is needed. numpy
|
||||||
@@ -105,48 +126,94 @@ def _map_row_chunks(n_rows: int, chunk_rows: int, n_workers: int, fn):
|
|||||||
threads give real parallelism here — and on a very large file they also
|
threads give real parallelism here — and on a very large file they also
|
||||||
keep many more page-fault requests in flight, which is what the I/O path
|
keep many more page-fault requests in flight, which is what the I/O path
|
||||||
wants.
|
wants.
|
||||||
|
|
||||||
|
*should_stop* is polled per chunk so a cancelled job abandons work within
|
||||||
|
one chunk rather than one whole angle — on a large scan an angle is
|
||||||
|
~40 s, which is far too long to block application shutdown.
|
||||||
"""
|
"""
|
||||||
bounds = [(r0, min(r0 + chunk_rows, n_rows))
|
bounds = [(r0, min(r0 + chunk_rows, n_rows))
|
||||||
for r0 in range(0, n_rows, chunk_rows)]
|
for r0 in range(0, n_rows, chunk_rows)]
|
||||||
|
|
||||||
|
def run(b):
|
||||||
|
if should_stop is not None and should_stop():
|
||||||
|
return
|
||||||
|
fn(*b)
|
||||||
|
|
||||||
if n_workers <= 1 or len(bounds) == 1:
|
if n_workers <= 1 or len(bounds) == 1:
|
||||||
for r0, r1 in bounds:
|
for b in bounds:
|
||||||
fn(r0, r1)
|
run(b)
|
||||||
return
|
return
|
||||||
with ThreadPoolExecutor(max_workers=min(n_workers, len(bounds))) as pool:
|
with ThreadPoolExecutor(max_workers=min(n_workers, len(bounds))) as pool:
|
||||||
list(pool.map(lambda b: fn(*b), bounds))
|
list(pool.map(run, bounds))
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Image computation
|
# Image computation
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
def compute_dc_image(sras: SrasFile, angle_idx: int, ch_idx: int) -> np.ndarray:
|
def plan_angle_level(sras: SrasFile,
|
||||||
"""Mean of each waveform → (n_rows, n_frames) float32, in ADC counts."""
|
live_multiplier: int = 1) -> tuple[int, int]:
|
||||||
|
"""(n_angle_workers, per_angle_budget) for a caller that parallelises
|
||||||
|
over angles instead of over rows.
|
||||||
|
|
||||||
|
Callers that parallelise over angles must not also let each angle
|
||||||
|
parallelise over rows: the two levels multiply, in threads and in memory.
|
||||||
|
Each angle is therefore given max_workers=1 and the returned budget share,
|
||||||
|
which together keep total live buffers within _TOTAL_BYTES_BUDGET.
|
||||||
|
"""
|
||||||
|
biggest = max(range(sras.n_angles), key=lambda a: int(sras.n_frames[a]))
|
||||||
|
_, n_workers = _plan_chunks(int(sras.n_rows[biggest]),
|
||||||
|
int(sras.n_frames[biggest]),
|
||||||
|
sras.samples_per_frame,
|
||||||
|
live_multiplier=live_multiplier)
|
||||||
|
n_workers = max(1, min(n_workers, sras.n_angles))
|
||||||
|
return n_workers, max(1, _TOTAL_BYTES_BUDGET // n_workers)
|
||||||
|
|
||||||
|
|
||||||
|
def compute_dc_image(sras: SrasFile, angle_idx: int, ch_idx: int,
|
||||||
|
max_workers: int | None = None,
|
||||||
|
budget: int | None = None,
|
||||||
|
should_stop=None) -> np.ndarray:
|
||||||
|
"""Mean of each waveform → (n_rows, n_frames) float32, in ADC counts.
|
||||||
|
|
||||||
|
Pass max_workers=1 when the caller is already parallelising over angles.
|
||||||
|
If *should_stop* ever returns True the result is incomplete — the caller
|
||||||
|
is expected to be abandoning it.
|
||||||
|
"""
|
||||||
n_rows, n_frames = sras.image_shape(angle_idx)
|
n_rows, n_frames = sras.image_shape(angle_idx)
|
||||||
data = sras.data[angle_idx]
|
data = sras.data[angle_idx]
|
||||||
chunk_rows, n_workers = _plan_chunks(n_rows, n_frames, sras.samples_per_frame)
|
chunk_rows, n_workers = _plan_chunks(n_rows, n_frames, sras.samples_per_frame,
|
||||||
|
max_workers=max_workers, budget=budget)
|
||||||
img = np.empty((n_rows, n_frames), dtype=np.float32)
|
img = np.empty((n_rows, n_frames), dtype=np.float32)
|
||||||
|
|
||||||
def chunk(r0: int, r1: int):
|
def chunk(r0: int, r1: int):
|
||||||
img[r0:r1] = data[r0:r1, ch_idx, :, :].astype(np.float32).mean(axis=-1)
|
img[r0:r1] = data[r0:r1, ch_idx, :, :].astype(np.float32).mean(axis=-1)
|
||||||
|
|
||||||
_map_row_chunks(n_rows, chunk_rows, n_workers, chunk)
|
_map_row_chunks(n_rows, chunk_rows, n_workers, chunk, should_stop=should_stop)
|
||||||
return img
|
return img
|
||||||
|
|
||||||
|
|
||||||
def dc_image_mv(sras: SrasFile, angle_idx: int, ch_idx: int) -> np.ndarray:
|
def dc_image_mv(sras: SrasFile, angle_idx: int, ch_idx: int,
|
||||||
|
max_workers: int | None = None, budget: int | None = None,
|
||||||
|
should_stop=None) -> np.ndarray:
|
||||||
"""DC image for (angle, channel) in mV, preferring a stored v5/v7 cache."""
|
"""DC image for (angle, channel) in mV, preferring a stored v5/v7 cache."""
|
||||||
cached = sras.cached_dc_mv(angle_idx, ch_idx)
|
cached = sras.cached_dc_mv(angle_idx, ch_idx)
|
||||||
if cached is not None:
|
if cached is not None:
|
||||||
return cached
|
return cached
|
||||||
return adc_to_mv(compute_dc_image(sras, angle_idx, ch_idx), *sras.cal(ch_idx))
|
return adc_to_mv(
|
||||||
|
compute_dc_image(sras, angle_idx, ch_idx, max_workers=max_workers,
|
||||||
|
budget=budget, should_stop=should_stop),
|
||||||
|
*sras.cal(ch_idx))
|
||||||
|
|
||||||
|
|
||||||
def compute_rf_image(sras: SrasFile, angle_idx: int,
|
def compute_rf_image(sras: SrasFile, angle_idx: int,
|
||||||
dc_threshold_mv: float | None,
|
dc_threshold_mv: float | None,
|
||||||
apply_bg_sub: bool = True,
|
apply_bg_sub: bool = True,
|
||||||
n_fft: int | None = None,
|
n_fft: int | None = None,
|
||||||
dc4_mv: np.ndarray | None = None) -> np.ndarray:
|
dc4_mv: np.ndarray | None = None,
|
||||||
|
max_workers: int | None = None,
|
||||||
|
budget: int | None = None,
|
||||||
|
should_stop=None) -> np.ndarray:
|
||||||
"""FFT of each CH1 waveform; pixel = peak frequency in MHz.
|
"""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
|
Pixels where CH4_dc < dc_threshold_mv are set to 0 — and the FFT is
|
||||||
@@ -193,7 +260,8 @@ def compute_rf_image(sras: SrasFile, angle_idx: int,
|
|||||||
n_fft_bins = n_fft if n_fft is not None else sras.samples_per_frame
|
n_fft_bins = n_fft if n_fft is not None else sras.samples_per_frame
|
||||||
chunk_rows, n_workers = _plan_chunks(
|
chunk_rows, n_workers = _plan_chunks(
|
||||||
n_rows, n_frames, max(sras.samples_per_frame, n_fft_bins),
|
n_rows, n_frames, max(sras.samples_per_frame, n_fft_bins),
|
||||||
live_multiplier=_FFT_LIVE_MULTIPLIER)
|
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
|
background = sras.background if (apply_bg_sub and sras.background is not None) else None
|
||||||
cal4 = sras.cal(CH4_IDX)
|
cal4 = sras.cal(CH4_IDX)
|
||||||
|
|
||||||
@@ -239,7 +307,7 @@ def compute_rf_image(sras: SrasFile, angle_idx: int,
|
|||||||
else:
|
else:
|
||||||
img[r0:r1] = freq_axis[peak_bins]
|
img[r0:r1] = freq_axis[peak_bins]
|
||||||
|
|
||||||
_map_row_chunks(n_rows, chunk_rows, n_workers, chunk)
|
_map_row_chunks(n_rows, chunk_rows, n_workers, chunk, should_stop=should_stop)
|
||||||
return img
|
return img
|
||||||
|
|
||||||
|
|
||||||
@@ -258,6 +326,8 @@ def cache_file(path: str, mode: str, apply_bg_sub: bool,
|
|||||||
"""
|
"""
|
||||||
global _MAX_WORKERS
|
global _MAX_WORKERS
|
||||||
try:
|
try:
|
||||||
|
if mode not in ("dc", "fft"):
|
||||||
|
return f"unknown cache mode {mode!r} (expected 'dc' or 'fft')"
|
||||||
set_fft_backend(fft_backend)
|
set_fft_backend(fft_backend)
|
||||||
if max_workers:
|
if max_workers:
|
||||||
_MAX_WORKERS = max_workers
|
_MAX_WORKERS = max_workers
|
||||||
@@ -493,7 +563,8 @@ def compute_angle_alignment(sras: SrasFile, ref_angle_idx: int,
|
|||||||
background-thread workers must not touch GUI-thread-owned caches."""
|
background-thread workers must not touch GUI-thread-owned caches."""
|
||||||
n = sras.n_angles # already the *complete*-angle count for aborted v6 scans
|
n = sras.n_angles # already the *complete*-angle count for aborted v6 scans
|
||||||
dx_ref, dy_ref = _pixel_pitch_mm(sras, ref_angle_idx)
|
dx_ref, dy_ref = _pixel_pitch_mm(sras, ref_angle_idx)
|
||||||
n_workers = max(1, min(_MAX_WORKERS, n))
|
# Parallel over angles, serial within each — see plan_angle_level.
|
||||||
|
n_workers, angle_budget = plan_angle_level(sras)
|
||||||
|
|
||||||
# progress_cb fires from pool threads, so the counter behind it must be
|
# progress_cb fires from pool threads, so the counter behind it must be
|
||||||
# atomic. list.append is, and len() of a list is a consistent read.
|
# atomic. list.append is, and len() of a list is a consistent read.
|
||||||
@@ -506,7 +577,9 @@ def compute_angle_alignment(sras: SrasFile, ref_angle_idx: int,
|
|||||||
|
|
||||||
# ---- Step 1: binarized CH4 mask per angle, native per-angle grid -----
|
# ---- Step 1: binarized CH4 mask per angle, native per-angle grid -----
|
||||||
def mask_for(a: int) -> np.ndarray:
|
def mask_for(a: int) -> np.ndarray:
|
||||||
dc4 = adc_to_mv(compute_dc_image(sras, a, CH4_IDX), *sras.cal(CH4_IDX))
|
dc4 = adc_to_mv(
|
||||||
|
compute_dc_image(sras, a, CH4_IDX, max_workers=1, budget=angle_budget),
|
||||||
|
*sras.cal(CH4_IDX))
|
||||||
m = (dc4 >= dc_threshold_mv).astype(np.float32)
|
m = (dc4 >= dc_threshold_mv).astype(np.float32)
|
||||||
tick(0, 25)
|
tick(0, 25)
|
||||||
return m
|
return m
|
||||||
|
|||||||
+12
-4
@@ -1624,6 +1624,8 @@ class SrasViewerWindow(QMainWindow):
|
|||||||
|
|
||||||
def _on_compute_done(self, result):
|
def _on_compute_done(self, result):
|
||||||
self._close_progress("main")
|
self._close_progress("main")
|
||||||
|
if result is None:
|
||||||
|
return # cancelled mid-compute; the partial image must not cache
|
||||||
angle_idx = self._pending_angle
|
angle_idx = self._pending_angle
|
||||||
ch_idx = self._pending_ch
|
ch_idx = self._pending_ch
|
||||||
|
|
||||||
@@ -1883,11 +1885,17 @@ class SrasViewerWindow(QMainWindow):
|
|||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
def closeEvent(self, event):
|
def closeEvent(self, event):
|
||||||
if self._dc_precompute_worker is not None:
|
# Signal every cancellable worker first, then wait. Waiting without
|
||||||
self._dc_precompute_worker.stop()
|
# signalling means sitting out whatever is in flight — on a large
|
||||||
for thread, _worker, _on_done in list(self._jobs.values()):
|
# 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.quit()
|
||||||
thread.wait(2000)
|
thread.wait(5000)
|
||||||
super().closeEvent(event)
|
super().closeEvent(event)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+113
-31
@@ -9,6 +9,7 @@ everything they need through their constructor and hand results back by signal.
|
|||||||
|
|
||||||
import os
|
import os
|
||||||
from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor, as_completed
|
from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor, as_completed
|
||||||
|
from concurrent.futures.process import BrokenProcessPool
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from PyQt6.QtCore import QObject, pyqtSignal
|
from PyQt6.QtCore import QObject, pyqtSignal
|
||||||
@@ -27,6 +28,31 @@ from sras_format import CH3_IDX, CH4_IDX, SrasFile
|
|||||||
_BATCH_MAX_PROCS = int(os.environ.get("SRAS_BATCH_PROCS", 0)) or min(
|
_BATCH_MAX_PROCS = int(os.environ.get("SRAS_BATCH_PROCS", 0)) or min(
|
||||||
4, os.cpu_count() or 2)
|
4, os.cpu_count() or 2)
|
||||||
|
|
||||||
|
# Spawning a pool costs roughly a second of interpreter startup (each child
|
||||||
|
# re-imports the entry module). That is noise against a multi-GB scan but
|
||||||
|
# dominates a batch of small files, where it would make the job *slower* —
|
||||||
|
# so below this total size the batch just runs in the worker thread.
|
||||||
|
_BATCH_POOL_MIN_BYTES = int(os.environ.get("SRAS_BATCH_POOL_MIN_MB", 512)) * 1024 * 1024
|
||||||
|
|
||||||
|
|
||||||
|
class CancellableWorker(QObject):
|
||||||
|
"""A worker whose compute polls stop() between row chunks.
|
||||||
|
|
||||||
|
Without this a shutdown has to wait out whatever is in flight, and on a
|
||||||
|
large scan a single angle is ~40 s — far too long to block closing the
|
||||||
|
window. Chunk-level polling bounds the wait to one chunk instead.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self._stop = False
|
||||||
|
|
||||||
|
def stop(self):
|
||||||
|
self._stop = True
|
||||||
|
|
||||||
|
def _stopped(self) -> bool:
|
||||||
|
return self._stop
|
||||||
|
|
||||||
|
|
||||||
class LoadWorker(QObject):
|
class LoadWorker(QObject):
|
||||||
finished = pyqtSignal(object) # SrasFile | None
|
finished = pyqtSignal(object) # SrasFile | None
|
||||||
@@ -44,7 +70,7 @@ class LoadWorker(QObject):
|
|||||||
self.finished.emit(None)
|
self.finished.emit(None)
|
||||||
|
|
||||||
|
|
||||||
class ComputeWorker(QObject):
|
class ComputeWorker(CancellableWorker):
|
||||||
"""Computes one displayable image for (angle, channel).
|
"""Computes one displayable image for (angle, channel).
|
||||||
|
|
||||||
For CH1/Velocity (FFT-derived) channels, the FFT is only run for pixels
|
For CH1/Velocity (FFT-derived) channels, the FFT is only run for pixels
|
||||||
@@ -79,15 +105,19 @@ class ComputeWorker(QObject):
|
|||||||
img = compute_rf_image(
|
img = compute_rf_image(
|
||||||
self._sras, self._angle, dc_threshold_mv=self._dc_threshold,
|
self._sras, self._angle, dc_threshold_mv=self._dc_threshold,
|
||||||
apply_bg_sub=self._apply_bg_sub, n_fft=self._n_fft,
|
apply_bg_sub=self._apply_bg_sub, n_fft=self._n_fft,
|
||||||
dc4_mv=self._dc4_mv)
|
dc4_mv=self._dc4_mv, should_stop=self._stopped)
|
||||||
else:
|
else:
|
||||||
img = dc_image_mv(self._sras, self._angle, self._ch)
|
img = dc_image_mv(self._sras, self._angle, self._ch,
|
||||||
self.finished.emit(img)
|
should_stop=self._stopped)
|
||||||
|
# On cancellation the image is only partly filled, so hand back
|
||||||
|
# None rather than something that would be cached as real. The
|
||||||
|
# signal still fires either way — it is what quits the thread.
|
||||||
|
self.finished.emit(None if self._stop else img)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
self.error.emit(str(exc))
|
self.error.emit(str(exc))
|
||||||
|
|
||||||
|
|
||||||
class DcPrecomputeWorker(QObject):
|
class DcPrecomputeWorker(CancellableWorker):
|
||||||
"""Computes CH3/CH4 DC images for every angle in the background.
|
"""Computes CH3/CH4 DC images for every angle in the background.
|
||||||
|
|
||||||
DC images are cheap (a per-waveform mean, no FFT) compared to the
|
DC images are cheap (a per-waveform mean, no FFT) compared to the
|
||||||
@@ -108,30 +138,34 @@ class DcPrecomputeWorker(QObject):
|
|||||||
def __init__(self, sras: SrasFile):
|
def __init__(self, sras: SrasFile):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self._sras = sras
|
self._sras = sras
|
||||||
self._stop = False
|
|
||||||
|
|
||||||
def stop(self):
|
|
||||||
self._stop = True
|
|
||||||
|
|
||||||
def _one_angle(self, a: int) -> tuple[int, np.ndarray, np.ndarray]:
|
def _one_angle(self, a: int) -> tuple[int, np.ndarray, np.ndarray]:
|
||||||
return a, dc_image_mv(self._sras, a, CH3_IDX), dc_image_mv(self._sras, a, CH4_IDX)
|
# 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).
|
||||||
|
kw = dict(max_workers=1, budget=self._angle_budget,
|
||||||
|
should_stop=self._stopped)
|
||||||
|
return (a,
|
||||||
|
dc_image_mv(self._sras, a, CH3_IDX, **kw),
|
||||||
|
dc_image_mv(self._sras, a, CH4_IDX, **kw))
|
||||||
|
|
||||||
def run(self):
|
def run(self):
|
||||||
try:
|
try:
|
||||||
n = self._sras.n_angles
|
n = self._sras.n_angles
|
||||||
n_workers = max(1, min(compute._MAX_WORKERS, n))
|
n_workers, self._angle_budget = compute.plan_angle_level(self._sras)
|
||||||
with ThreadPoolExecutor(max_workers=n_workers) as pool:
|
pool = ThreadPoolExecutor(max_workers=n_workers)
|
||||||
futures = {pool.submit(self._one_angle, a): a for a in range(n)}
|
|
||||||
try:
|
try:
|
||||||
|
futures = {pool.submit(self._one_angle, a): a for a in range(n)}
|
||||||
for fut in as_completed(futures):
|
for fut in as_completed(futures):
|
||||||
if self._stop:
|
if self._stop:
|
||||||
break
|
break
|
||||||
a, dc3, dc4 = fut.result()
|
a, dc3, dc4 = fut.result()
|
||||||
self.angle_done.emit(a, dc3, dc4)
|
self.angle_done.emit(a, dc3, dc4)
|
||||||
finally:
|
finally:
|
||||||
if self._stop:
|
# cancel_futures drops the queued angles; should_stop lets the
|
||||||
for fut in futures:
|
# in-flight ones bail within a chunk. Not waiting here is what
|
||||||
fut.cancel()
|
# keeps closing the window responsive on a large scan.
|
||||||
|
pool.shutdown(wait=not self._stop, cancel_futures=True)
|
||||||
self.finished.emit()
|
self.finished.emit()
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
self.error.emit(str(exc))
|
self.error.emit(str(exc))
|
||||||
@@ -163,24 +197,27 @@ class BatchCacheWorker(QObject):
|
|||||||
self._mode = mode
|
self._mode = mode
|
||||||
self._apply_bg_sub = apply_bg_sub
|
self._apply_bg_sub = apply_bg_sub
|
||||||
|
|
||||||
def _pool(self, n_files: int):
|
def _report(self, path: str, err: str, done: int, total: int):
|
||||||
"""A process pool, falling back to threads if the platform refuses to
|
self.file_done.emit(path, err)
|
||||||
spawn (still a win: the compute releases the GIL)."""
|
self.progress.emit(int(done / max(1, total) * 100))
|
||||||
n = max(1, min(_BATCH_MAX_PROCS, n_files))
|
|
||||||
try:
|
|
||||||
return ProcessPoolExecutor(max_workers=n), n
|
|
||||||
except (OSError, ValueError):
|
|
||||||
return ThreadPoolExecutor(max_workers=n), n
|
|
||||||
|
|
||||||
def run(self):
|
def _run_pooled(self, paths: list[str], n_procs: int) -> list[str]:
|
||||||
paths = self._paths
|
"""Process the batch across *n_procs* subprocesses. Returns the paths
|
||||||
executor, n_procs = self._pool(len(paths))
|
that never got a real answer because the pool itself died, so the
|
||||||
|
caller can retry them in-process.
|
||||||
|
|
||||||
|
Under spawn each child re-imports the entry module, so the batch must
|
||||||
|
survive that going wrong (an unguarded __main__, a frozen build, a
|
||||||
|
sandbox that forbids subprocesses) rather than reporting every file as
|
||||||
|
failed — hence the retry list instead of a per-file error.
|
||||||
|
"""
|
||||||
# Each child threads internally; divide the machine rather than
|
# Each child threads internally; divide the machine rather than
|
||||||
# letting every process claim every core.
|
# letting every process claim every core.
|
||||||
per_proc_workers = max(1, (os.cpu_count() or 4) // n_procs)
|
per_proc_workers = max(1, (os.cpu_count() or 4) // n_procs)
|
||||||
|
unresolved: list[str] = []
|
||||||
done = 0
|
done = 0
|
||||||
with executor:
|
|
||||||
|
with ProcessPoolExecutor(max_workers=n_procs) as executor:
|
||||||
futures = {
|
futures = {
|
||||||
executor.submit(cache_file, p, self._mode, self._apply_bg_sub,
|
executor.submit(cache_file, p, self._mode, self._apply_bg_sub,
|
||||||
compute.get_fft_backend(), per_proc_workers): p
|
compute.get_fft_backend(), per_proc_workers): p
|
||||||
@@ -190,11 +227,56 @@ class BatchCacheWorker(QObject):
|
|||||||
path = futures[fut]
|
path = futures[fut]
|
||||||
try:
|
try:
|
||||||
err = fut.result()
|
err = fut.result()
|
||||||
|
except BrokenProcessPool:
|
||||||
|
unresolved.append(path)
|
||||||
|
continue
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
err = str(exc)
|
err = str(exc)
|
||||||
done += 1
|
done += 1
|
||||||
self.file_done.emit(path, err)
|
self._report(path, err, done, len(paths))
|
||||||
self.progress.emit(int(done / max(1, len(paths)) * 100))
|
|
||||||
|
return unresolved
|
||||||
|
|
||||||
|
def _run_inline(self, paths: list[str], done: int, total: int):
|
||||||
|
"""Fallback / single-file path: compute in this thread. Still uses the
|
||||||
|
full core count internally, since nothing else is competing."""
|
||||||
|
for path in paths:
|
||||||
|
try:
|
||||||
|
err = cache_file(path, self._mode, self._apply_bg_sub,
|
||||||
|
compute.get_fft_backend(), compute._MAX_WORKERS)
|
||||||
|
except Exception as exc:
|
||||||
|
err = str(exc)
|
||||||
|
done += 1
|
||||||
|
self._report(path, err, done, total)
|
||||||
|
|
||||||
|
def _worth_pooling(self, paths: list[str]) -> bool:
|
||||||
|
if len(paths) < 2:
|
||||||
|
return False
|
||||||
|
total = 0
|
||||||
|
for p in paths:
|
||||||
|
try:
|
||||||
|
total += os.path.getsize(p)
|
||||||
|
except OSError:
|
||||||
|
pass # unreadable files are reported by cache_file
|
||||||
|
return total >= _BATCH_POOL_MIN_BYTES
|
||||||
|
|
||||||
|
def run(self):
|
||||||
|
paths = self._paths
|
||||||
|
n_procs = max(1, min(_BATCH_MAX_PROCS, len(paths)))
|
||||||
|
|
||||||
|
if not self._worth_pooling(paths):
|
||||||
|
self._run_inline(paths, 0, len(paths))
|
||||||
|
self.finished.emit()
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
unresolved = self._run_pooled(paths, n_procs)
|
||||||
|
except Exception:
|
||||||
|
# The pool could not be created or collapsed wholesale.
|
||||||
|
unresolved = list(paths)
|
||||||
|
|
||||||
|
if unresolved:
|
||||||
|
self._run_inline(unresolved, len(paths) - len(unresolved), len(paths))
|
||||||
|
|
||||||
self.finished.emit()
|
self.finished.emit()
|
||||||
|
|
||||||
|
|||||||
+53
-4
@@ -41,9 +41,11 @@ def _preamble(ymult_v: float, yoff_adc: float, yzero_v: float) -> bytes:
|
|||||||
).encode("utf-8")
|
).encode("utf-8")
|
||||||
|
|
||||||
|
|
||||||
def build(n_angles: int, seed: int, samples_per_frame: int) -> tuple[bytes, dict]:
|
def build(n_angles: int, seed: int, samples_per_frame: int,
|
||||||
|
geometry: list[tuple[int, int]] | None = None) -> tuple[bytes, dict]:
|
||||||
rng = np.random.default_rng(seed)
|
rng = np.random.default_rng(seed)
|
||||||
geom = [_GEOMETRY[a % len(_GEOMETRY)] for a in range(n_angles)]
|
src_geom = geometry or _GEOMETRY
|
||||||
|
geom = [src_geom[a % len(src_geom)] for a in range(n_angles)]
|
||||||
n_ch = 3
|
n_ch = 3
|
||||||
bps = 1
|
bps = 1
|
||||||
|
|
||||||
@@ -116,12 +118,59 @@ def build(n_angles: int, seed: int, samples_per_frame: int) -> tuple[bytes, dict
|
|||||||
|
|
||||||
|
|
||||||
def write(path: Path, n_angles: int = 3, seed: int = 0,
|
def write(path: Path, n_angles: int = 3, seed: int = 0,
|
||||||
samples_per_frame: int = 64) -> dict:
|
samples_per_frame: int = 64,
|
||||||
payload, meta = build(n_angles, seed, samples_per_frame)
|
geometry: list[tuple[int, int]] | None = None) -> dict:
|
||||||
|
payload, meta = build(n_angles, seed, samples_per_frame, geometry)
|
||||||
path.write_bytes(payload)
|
path.write_bytes(payload)
|
||||||
return meta
|
return meta
|
||||||
|
|
||||||
|
|
||||||
|
HDR_FMT_LEGACY = ">4sBHHffffIIdBB"
|
||||||
|
|
||||||
|
|
||||||
|
def write_legacy(path: Path, version: int = 4, n_angles: int = 2,
|
||||||
|
n_rows: int = 4, n_frames: int = 10,
|
||||||
|
samples_per_frame: int = 32, seed: int = 0) -> dict:
|
||||||
|
"""Write a v2/v3/v4 file: uniform geometry, one flat waveform block.
|
||||||
|
|
||||||
|
Used to exercise sras_average.py, which only handles the legacy formats.
|
||||||
|
"""
|
||||||
|
rng = np.random.default_rng(seed)
|
||||||
|
n_ch, bps = 3, 1
|
||||||
|
|
||||||
|
out = bytearray()
|
||||||
|
out += struct.pack(
|
||||||
|
HDR_FMT_LEGACY, b"SRAS", version, n_angles, n_rows,
|
||||||
|
-0.5, 1.0, _VELOCITY_MM_S, _LASER_FREQ_HZ,
|
||||||
|
n_frames, samples_per_frame, _SAMPLE_RATE_HZ, bps, n_ch,
|
||||||
|
)
|
||||||
|
angles = np.linspace(0.0, 45.0, n_angles, dtype=np.float32)
|
||||||
|
out += angles.astype(">f4").tobytes()
|
||||||
|
y = (np.arange(n_rows) * _ROW_SPACING_MM).astype(np.float32)
|
||||||
|
out += y.astype(">f4").tobytes()
|
||||||
|
|
||||||
|
if version >= 3:
|
||||||
|
for ymult_v, yoff, yzero_v in ((1.5625e-3, -87.04, 0.0),
|
||||||
|
(2.0e-3, -60.0, 1.0e-3),
|
||||||
|
(2.5e-3, -40.0, -2.0e-3))[:n_ch]:
|
||||||
|
p = _preamble(ymult_v, yoff, yzero_v)
|
||||||
|
out += struct.pack(">H", len(p)) + p
|
||||||
|
|
||||||
|
background = rng.integers(-8, 9, size=samples_per_frame, dtype=np.int8)
|
||||||
|
if version >= 4:
|
||||||
|
out += struct.pack(">I", samples_per_frame) + background.tobytes()
|
||||||
|
|
||||||
|
data = rng.integers(-100, 101,
|
||||||
|
size=(n_angles, n_rows, n_ch, n_frames, samples_per_frame),
|
||||||
|
dtype=np.int8)
|
||||||
|
out += data.tobytes()
|
||||||
|
path.write_bytes(bytes(out))
|
||||||
|
return {"version": version, "n_angles": n_angles, "n_rows": n_rows,
|
||||||
|
"n_frames": n_frames, "samples_per_frame": samples_per_frame,
|
||||||
|
"n_channels": n_ch, "data": data, "angles_deg": angles,
|
||||||
|
"y_positions": y, "background": background}
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
p = argparse.ArgumentParser(description=__doc__)
|
p = argparse.ArgumentParser(description=__doc__)
|
||||||
p.add_argument("output")
|
p.add_argument("output")
|
||||||
|
|||||||
@@ -0,0 +1,270 @@
|
|||||||
|
#!/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, aligned view, ROI draw/move, and CSV export.
|
||||||
|
|
||||||
|
Usage: QT_QPA_PLATFORM=offscreen python tools/test_gui.py [file.sras]
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import tempfile
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
os.environ.setdefault("QT_QPA_PLATFORM", "offscreen")
|
||||||
|
|
||||||
|
import numpy as np # noqa: E402
|
||||||
|
from PyQt6.QtCore import QEventLoop, QTimer # noqa: E402
|
||||||
|
from PyQt6.QtWidgets import QApplication # noqa: E402
|
||||||
|
|
||||||
|
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||||
|
|
||||||
|
from sras_format import CH1_IDX, CH3_IDX, CH4_IDX # 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"
|
||||||
|
from unittest.mock import patch
|
||||||
|
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("\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())
|
||||||
@@ -0,0 +1,358 @@
|
|||||||
|
#!/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()
|
||||||
Reference in New Issue
Block a user