Bound nested parallelism, make compute cancellable, harden batch pool
Three defects found by testing the parallel paths end to end.
Nested parallelism was unbounded. DcPrecomputeWorker and the alignment
driver both parallelise over angles, and each angle's compute_dc_image
then parallelised over rows again: 14 x 14 threads, and — worse — every
angle sized its chunk against the *whole* memory budget, so worst-case
live buffers were ~13.9 GB against a 1 GB budget. Capping the inner
worker count alone does not fix it; the chunk is still sized against the
full budget. plan_angle_level now hands each angle both max_workers=1
and its share of the budget, which bounds the real file at 1002 MB.
Chunk planning could not produce concurrency at all. _plan_chunks sized
chunk_rows first, and _chunk_rows_for spends whatever budget it is given
— so a single chunk consumed the lot and budget//chunk_bytes came back
as one worker. It now picks the worker count first and sizes the chunk to
it, giving 7-16 workers on the real file instead of 1.
Cancellation was too coarse to shut down. stop() was only checked between
angles, and one angle of the real scan is ~40 s, so closing the window
sat through it and returned with threads still running. should_stop is
now polled per row chunk and closeEvent signals every worker before
waiting: close went from 4.0 s with two live threads to 1.04 s with both
finished. A cancelled ComputeWorker emits None so a half-filled image is
never cached.
Also: BatchCacheWorker no longer reports every file as failed when the
process pool itself dies (unguarded __main__, frozen build, sandbox) —
it retries those files in-process. Batches below 512 MB skip the pool
entirely, since interpreter startup would otherwise make small jobs
slower. cache_file rejects an unknown mode instead of silently treating
it as "fft".
Measured on the real 496 GB scan, angle 3, 32 rows:
DC 2.94s -> 0.75s RF 6.10s -> 1.99s peak RSS 1.40 -> 1.93 GB
Batch convert, 24 files / 1 GB: 4.03s -> 2.36s.
Memory budget defaults to 1024 MB (SRAS_MEM_BUDGET_MB), the measured knee.
Testing: tools/test_refactor.py (70 checks: cache round-trip and block
carry-forward, partial-cache handling, parallel-vs-serial identity,
no-mask path, ROI mask vs full-grid, v2/v3/v4 parsing, sras_average
round-trip incl. remainder handling) and tools/test_gui.py (46 checks
driving the real widgets and workers headlessly). check_equivalence.py
still reports all 167 outputs identical to ed0eba4.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
+53
-4
@@ -41,9 +41,11 @@ def _preamble(ymult_v: float, yoff_adc: float, yzero_v: float) -> bytes:
|
||||
).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)
|
||||
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
|
||||
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,
|
||||
samples_per_frame: int = 64) -> dict:
|
||||
payload, meta = build(n_angles, seed, samples_per_frame)
|
||||
samples_per_frame: int = 64,
|
||||
geometry: list[tuple[int, int]] | None = None) -> dict:
|
||||
payload, meta = build(n_angles, seed, samples_per_frame, geometry)
|
||||
path.write_bytes(payload)
|
||||
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():
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
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