diff --git a/pyproject.toml b/pyproject.toml index ca708ca..ae73ea4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -16,6 +16,8 @@ dependencies = [ "scikit-image==0.26.0", # Faster rfft backend; the viewer falls back to scipy.fft without it. "pyFFTW==0.15.1", + # Clamps BLAS threading under the FFT worker pool. + "threadpoolctl==3.6.0", ] [project.optional-dependencies] diff --git a/sras_compute.py b/sras_compute.py index bf48c6b..5700b02 100644 --- a/sras_compute.py +++ b/sras_compute.py @@ -6,8 +6,10 @@ multiprocessing child can import it without loading Qt or matplotlib — which matters because Python 3.14 on macOS spawns rather than forks. """ +import atexit import json import os +import threading from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from pathlib import Path @@ -24,12 +26,16 @@ from sras_format import CH1_IDX, CH3_IDX, CH4_IDX, SrasFile, adc_to_mv try: import pyfftw - pyfftw.interfaces.cache.enable() PYFFTW_AVAILABLE = True except ImportError: PYFFTW_AVAILABLE = False -_fft_backend = "numpy" # "numpy" or "pyfftw"; set via set_fft_backend() +try: + from threadpoolctl import threadpool_limits +except ImportError: + threadpool_limits = None + +_fft_backend = "numpy" # "numpy" (scipy.fft) or "pyfftw"; set via set_fft_backend() def set_fft_backend(name: str): @@ -52,6 +58,218 @@ def _do_rfft(x: np.ndarray, n: int | None = None, axis: int = -1, return scipy_fft.rfft(x, n=n, axis=axis, workers=workers) +# --------------------------------------------------------------------------- +# FFT worker pool and per-thread pyFFTW plans +# --------------------------------------------------------------------------- + +_FFT_BLOCK = 512 # waveforms per FFT task. The knee on a 16-core machine: + # smaller blocks serialise on GIL-held numpy dispatch, + # larger ones lose cache residency and task granularity. +_ZOOM_MIN_PAD = 4 # zoom refinement engages at n_fft >= _ZOOM_MIN_PAD * spf +_FFT_EXACT_ENV = bool(int(os.environ.get("SRAS_FFT_EXACT", "0") or 0)) + +_pool_lock = threading.Lock() +_pool: ThreadPoolExecutor | None = None + + +def _fft_pool() -> ThreadPoolExecutor: + """The persistent process-wide pool for FFT block tasks.""" + global _pool + with _pool_lock: + if _pool is None: + _pool = ThreadPoolExecutor(max_workers=_MAX_WORKERS, + thread_name_prefix="sras-fft") + atexit.register(_pool.shutdown, wait=False, cancel_futures=True) + return _pool + + +_WISDOM_PATH = Path.home() / ".cache" / "sras-viewer" / "fftw_wisdom" +_wisdom_lock = threading.Lock() +_wisdom_loaded = False +_fftw_local = threading.local() + + +def _load_wisdom_once(): + """Import saved FFTW wisdom so FFTW_MEASURE planning is a one-time cost + per machine. Purely an optimisation: failures are ignored.""" + global _wisdom_loaded + with _wisdom_lock: + if _wisdom_loaded: + return + _wisdom_loaded = True + try: + pyfftw.import_wisdom(_WISDOM_PATH.read_bytes().split(b"\x00\n")) + except Exception: + pass + + +def _save_wisdom(): + with _wisdom_lock: + try: + _WISDOM_PATH.parent.mkdir(parents=True, exist_ok=True) + _WISDOM_PATH.write_bytes(b"\x00\n".join(pyfftw.export_wisdom())) + except Exception: + pass + + +def _fftw_block_rfft(waves: np.ndarray, n: int) -> np.ndarray: + """rfft of a (B, spf) float32 block via a cached per-thread FFTW plan. + + Plans have a fixed (_FFT_BLOCK, spf) input shape so each worker thread + plans once per transform length; a remainder block runs through the same + plan with its tail rows ignored. The returned array is the plan's output + buffer — consume it before the next call on the same thread. + """ + n_wf, spf = waves.shape + plans = getattr(_fftw_local, "plans", None) + if plans is None: + plans = _fftw_local.plans = {} + key = (_FFT_BLOCK, spf, n) + plan = plans.get(key) + if plan is None: + _load_wisdom_once() + buf = pyfftw.empty_aligned((_FFT_BLOCK, spf), dtype="float32") + # No overwrite_input: FFTW must not scribble on input_array, whose + # zero-padded tail (columns spf..n) is zeroed exactly once here. + plan = pyfftw.builders.rfft(buf, n=n, axis=-1, threads=1, + planner_effort="FFTW_MEASURE") + plan.input_array[:] = 0.0 + plans[key] = plan + _save_wisdom() + inp = plan.input_array + inp[:n_wf, :spf] = waves + return plan()[:n_wf] + + +def _block_rfft(waves: np.ndarray, n: int) -> np.ndarray: + """Single-threaded rfft of one block; outer parallelism comes from the + pool, so the transform itself must not spin up threads.""" + if _fft_backend == "pyfftw" and PYFFTW_AVAILABLE: + return _fftw_block_rfft(waves, n) + return scipy_fft.rfft(waves, n=n, axis=-1, workers=1) + + +# --------------------------------------------------------------------------- +# Zoom peak search: coarse rfft + local fine DFT around the winning bin +# --------------------------------------------------------------------------- + +_ZOOM_HALFWIDTH = 0.75 # refinement window half-width, in coarse spacings. + # Every fine bin lies within 0.5 spacings of its + # nearest coarse bin, and that bin is guaranteed to + # be a candidate (see _ZOOM_CAND_RATIO), so 0.5 + # suffices; 0.75 adds rounding margin. +_ZOOM_CAND_RATIO = 0.7 # refine every coarse bin within this power ratio of + # its row's coarse maximum. Quarter-natural-bin + # scalloping at the 2x-oversampled coarse grid can + # understate a peak by at most ~19% in power, so 0.7 + # keeps a wide margin — near-contenders are resolved + # on the fine grid, never ranked from coarse samples. + + +@dataclass +class _ZoomPlan: + """Constants for the coarse+refine peak search, built once per + compute_rf_image call and shared across worker threads. *phases* is a + lazily-filled window-start -> phase-vector cache; a benign duplicate + compute under concurrency is cheaper than locking.""" + n_fft: int + n_coarse: int + m: int # fine bins per refinement window + n_bins_fine: int + E: np.ndarray # (spf, m) complex64 fine-DFT matrix, relative bins + j: np.ndarray # (spf,) float64 sample indices + phases: dict + + +def _zoom_plan(spf: int, n_fft: int) -> _ZoomPlan: + # 2x-oversampled coarse grid: the padded power spectrum is a trig + # polynomial of degree spf-1, so at 2x sampling its global max cannot + # hide between coarse bins. + n_coarse = scipy_fft.next_fast_len(2 * spf, real=True) + n_bins_fine = n_fft // 2 + 1 + m = int(np.ceil(2 * _ZOOM_HALFWIDTH * n_fft / n_coarse)) + 1 + m = min(m, n_bins_fine - 1) + j = np.arange(spf, dtype=np.float64) + E = np.exp((-2j * np.pi / n_fft) * np.outer(j, np.arange(m))).astype(np.complex64) + return _ZoomPlan(n_fft, n_coarse, m, n_bins_fine, E, j, {}) + + +def _window_start(k_c: np.ndarray, zp: _ZoomPlan) -> np.ndarray: + """First fine bin of the refinement window around each coarse bin. + Clipped to [1, ...] so fine bin 0 stays excluded (DC suppression).""" + k0 = np.floor((k_c - _ZOOM_HALFWIDTH) * (zp.n_fft / zp.n_coarse)).astype(np.int64) + return np.clip(k0, 1, max(1, zp.n_bins_fine - zp.m)) + + +def _refine_window(waves: np.ndarray, rows: np.ndarray, k0: int, zp: _ZoomPlan, + best_pow: np.ndarray, best_bin: np.ndarray): + """Evaluate fine bins k0..k0+m-1 for the given rows and fold the result + into the per-row best (power, bin), preserving np.argmax's lowest-bin + tie-break.""" + ph = zp.phases.get(k0) + if ph is None: + ph = np.exp((-2j * np.pi * k0 / zp.n_fft) * zp.j).astype(np.complex64) + zp.phases[k0] = ph + F = (waves[rows] * ph) @ zp.E + q = F.real ** 2 + q += F.imag ** 2 + i = np.argmax(q, axis=1) + p = q[np.arange(len(rows)), i] + b = k0 + i + upd = (p > best_pow[rows]) | ((p == best_pow[rows]) & (b < best_bin[rows])) + ridx = rows[upd] + best_pow[ridx] = p[upd] + best_bin[ridx] = b[upd] + + +def _peak_bins_zoom(waves: np.ndarray, zp: _ZoomPlan) -> np.ndarray: + """Peak fine-bin per waveform without materialising the padded spectrum: + coarse rfft, then a small fine DFT (one gemm per shared window) on the + exact n_fft bin grid. Identity with the full padded argmax is enforced + by tests/test_compute.py::test_zoom_identity and the golden-hash sweep.""" + n_wf = waves.shape[0] + S = _block_rfft(waves, zp.n_coarse) + P = S.real ** 2 + P += S.imag ** 2 + P[:, 0] = 0.0 + p_c = np.max(P, axis=1) + + best_pow = np.full(n_wf, -1.0, dtype=np.float32) + best_bin = np.full(n_wf, np.iinfo(np.int64).max, dtype=np.int64) + + # For a clean signal this yields one or two windows; a noise spectrum + # (many near-equal peaks) yields a dozen or so — still a tiny fraction + # of the padded grid. + thr = np.where(p_c > 0, np.float32(_ZOOM_CAND_RATIO) * p_c, + np.float32(np.inf)) + rows_c, bins_c = np.nonzero(P >= thr[:, None]) + k0_c = _window_start(bins_c, zp) + # Zeroing the coarse DC bin (suppression) blinds the candidate scan to + # fine bins closer to DC than the first coarse sample — where the + # DC-leakage skirt of an un-subtracted offset peaks. Always refine the + # DC-adjacent window too. + rows_c = np.concatenate([rows_c, np.arange(n_wf)]) + k0_c = np.concatenate([k0_c, np.ones(n_wf, dtype=np.int64)]) + pair = np.unique(np.stack([rows_c, k0_c], axis=1), axis=0) + for k0 in np.unique(pair[:, 1]): + _refine_window(waves, pair[pair[:, 1] == k0, 0], int(k0), zp, + best_pow, best_bin) + + # An all-zero spectrum must reproduce argmax-of-zeros = bin 0. + best_bin[p_c == 0.0] = 0 + return best_bin + + +def _peak_bins_direct(waves: np.ndarray, n_len: int) -> np.ndarray: + """Peak bin per waveform via the full transform — for pad factors too + small for the zoom search to pay.""" + S = _block_rfft(waves, n_len) + power = S.real ** 2 + power += S.imag ** 2 + power[:, 0] = 0.0 + return np.argmax(power, axis=1) + + # --------------------------------------------------------------------------- # Chunking / parallel budget # --------------------------------------------------------------------------- @@ -76,17 +294,21 @@ _TOTAL_BYTES_BUDGET = int(os.environ.get("SRAS_MEM_BUDGET_MB", 1024)) * 1024 * 1 _CHUNK_ROWS_MAX = 32 # cap for small scans (original behavior) _MAX_WORKERS = int(os.environ.get("SRAS_MAX_WORKERS", 0)) or (os.cpu_count() or 4) -# An rfft chunk holds, live at once: the float32 input, the complex64 -# transform, and the float32 power spectrum — roughly 3x the input buffer. -_FFT_LIVE_MULTIPLIER = 3 - - def _chunk_rows_for(n_frames: int, samples_per_frame: int, budget: int = _TOTAL_BYTES_BUDGET) -> int: bytes_per_row = max(1, n_frames * samples_per_frame * 4) # float32 return int(max(1, min(_CHUNK_ROWS_MAX, budget // bytes_per_row))) +def _plan_fft_rows(n_frames: int, samples_per_frame: int, budget: int) -> int: + """Rows per outer chunk for the block-FFT path. Only the float32 + waveform buffer scales with the chunk (in-flight block spectra total a + few MB across the whole pool), so budget it with 2x slack and let the + block fan-out saturate the pool regardless of pad factor.""" + bytes_per_row = max(1, n_frames * samples_per_frame * 4 * 2) + return int(max(1, min(_CHUNK_ROWS_MAX, budget // bytes_per_row))) + + def _plan_chunks(n_rows: int, n_frames: int, samples_per_frame: int, live_multiplier: int = 1, max_workers: int | None = None, @@ -215,7 +437,8 @@ def compute_rf_image(sras: SrasFile, angle_idx: int, dc4_mv: np.ndarray | None = None, max_workers: int | None = None, budget: int | None = None, - should_stop=None) -> np.ndarray: + should_stop=None, + exact: bool = False) -> np.ndarray: """FFT of each CH1 waveform; pixel = peak frequency in MHz. Pixels where CH4_dc < dc_threshold_mv are set to 0 — and the FFT is @@ -233,6 +456,13 @@ def compute_rf_image(sras: SrasFile, angle_idx: int, DC-channel precompute cache), pass it as *dc4_mv* (mV, shape (n_rows, n_frames)) to reuse it instead of re-reading CH4 here. + At pad factors >= _ZOOM_MIN_PAD the padded spectrum is never + materialised: a coarse rfft finds each peak and a local fine DFT + resolves it on the exact n_fft bin grid (_peak_bins_zoom). *exact* + forces the reference full-padded transform instead — it exists for + tests, audits, and the SRAS_FFT_EXACT=1 escape hatch, and is slow and + memory-hungry at high pad. + Fast path: if the file has a precomputed peak-frequency image for this angle (v5 PREC or v7 CACH), zero-padding is off, and the bg-sub flag matches, the stored image is used directly — no FFT is run. @@ -257,17 +487,31 @@ def compute_rf_image(sras: SrasFile, angle_idx: int, return freq_img # ---- Chunked FFT path -------------------------------------------------- - freq_axis = sras.freq_axis_mhz(n_fft) + exact = exact or _FFT_EXACT_ENV + spf = sras.samples_per_frame + n_len = n_fft if n_fft is not None else spf + freq32 = sras.freq_axis_mhz(n_fft).astype(np.float32) img = np.zeros((n_rows, n_frames), dtype=np.float32) - n_fft_bins = n_fft if n_fft is not None else sras.samples_per_frame - chunk_rows, n_workers = _plan_chunks( - n_rows, n_frames, max(sras.samples_per_frame, n_fft_bins), - live_multiplier=_FFT_LIVE_MULTIPLIER, max_workers=max_workers, - budget=budget) background = sras.background if (apply_bg_sub and sras.background is not None) else None cal4 = sras.cal(CH4_IDX) - def chunk(r0: int, r1: int): + zp = (_zoom_plan(spf, n_fft) + if not exact and n_fft is not None and n_fft >= _ZOOM_MIN_PAD * spf + else None) + total = _TOTAL_BYTES_BUDGET if budget is None else max(1, budget) + cap = max_workers if max_workers is not None else _MAX_WORKERS + if exact: + # The reference path materialises the full padded spectrum, so rows + # are budgeted against it (complex64 + power + temp per bin) and the + # chunk runs as one serial transform. + bytes_per_row = max(1, n_frames * (4 * spf + 16 * (n_len // 2 + 1))) + chunk_rows = int(max(1, min(_CHUNK_ROWS_MAX, total // bytes_per_row))) + cap = 1 + else: + chunk_rows = _plan_fft_rows(n_frames, spf, total) + pool = _fft_pool() if cap > 1 else None + + def process(r0: int, r1: int): if dc_threshold_mv is None: valid = None else: @@ -280,36 +524,73 @@ def compute_rf_image(sras: SrasFile, angle_idx: int, if not valid.any(): return - # Index the raw memmap slice with the boolean mask *before* - # converting dtype — this is a lazy view until touched, so only the - # selected elements are actually read from disk; masked-out pixels' - # pages are never paged in at all. - raw = data[r0:r1, CH1_IDX, :, :] - waves = (raw[valid] if valid is not None else raw).astype(np.float32) + counts = (valid.sum(axis=1) if valid is not None + else np.full(r1 - r0, n_frames, dtype=np.int64)) + offs = np.concatenate(([0], np.cumsum(counts))) + n_wf = int(offs[-1]) + waves = np.empty((n_wf, spf), dtype=np.float32) - if background is not None: - waves -= background # background is 1-D (spf,) + def read_row(i: int): + if should_stop is not None and should_stop(): + return + # Index the raw memmap slice with the boolean mask *before* + # converting dtype — this is a lazy view until touched, so only + # the selected elements are actually read from disk; masked-out + # pixels' pages are never paged in at all. + raw = data[r0 + i, CH1_IDX] + dst = waves[offs[i]:offs[i + 1]] + dst[:] = raw[valid[i]] if valid is not None else raw + if background is not None: + dst -= background # background is 1-D (spf,) - # Parallelism comes from the outer chunk loop, so keep the inner - # transform single-threaded to avoid oversubscribing the machine. - spectrum = _do_rfft(waves, n=n_fft, axis=-1, workers=1) - del waves - # |z|^2 without np.abs()'s extra full-size temporary. - power = spectrum.real ** 2 - power += spectrum.imag ** 2 - del spectrum - # [..., 0] not [:, 0]: the unmasked path keeps the (rows, frames, - # bins) shape, where [:, 0] would blank a whole frame. - power[..., 0] = 0.0 # suppress DC bin - peak_bins = np.argmax(power, axis=-1) - del power - - if valid is not None: - img[r0:r1][valid] = freq_axis[peak_bins] + if pool is None: + for i in range(r1 - r0): + read_row(i) else: - img[r0:r1] = freq_axis[peak_bins] + list(pool.map(read_row, range(r1 - r0))) - _map_row_chunks(n_rows, chunk_rows, n_workers, chunk, should_stop=should_stop) + out = np.empty(n_wf, dtype=np.float32) + + def fft_block(b0: int): + if should_stop is not None and should_stop(): + return + b1 = min(b0 + _FFT_BLOCK, n_wf) + w = waves[b0:b1] + bins = (_peak_bins_zoom(w, zp) if zp is not None + else _peak_bins_direct(w, n_len)) + out[b0:b1] = freq32[bins] + + if exact: + spectrum = _do_rfft(waves, n=n_fft, axis=-1, workers=1) + power = spectrum.real ** 2 + power += spectrum.imag ** 2 + power[:, 0] = 0.0 # suppress DC bin + out[:] = freq32[np.argmax(power, axis=1)] + elif pool is None: + for b0 in range(0, n_wf, _FFT_BLOCK): + fft_block(b0) + else: + list(pool.map(fft_block, range(0, n_wf, _FFT_BLOCK))) + + # Boolean scatter is row-major, matching the read pass's + # concatenation order. + if valid is not None: + img[r0:r1][valid] = out + else: + img[r0:r1] = out.reshape(r1 - r0, n_frames) + + # BLAS must not thread under the pool (the fine-DFT gemm would multiply + # against the pool's own workers). + limiter = (threadpool_limits(limits=1) + if pool is not None and threadpool_limits is not None else None) + try: + for r0 in range(0, n_rows, chunk_rows): + if should_stop is not None and should_stop(): + break + process(r0, min(r0 + chunk_rows, n_rows)) + finally: + if limiter is not None: + limiter.unregister() return img diff --git a/tests/test_compute.py b/tests/test_compute.py index e08af39..408d79f 100644 --- a/tests/test_compute.py +++ b/tests/test_compute.py @@ -10,6 +10,7 @@ import sys from pathlib import Path import numpy as np +import pytest import sras_compute as compute from sras_compute import ( @@ -117,9 +118,17 @@ def test_parallel_identity(tmp_path, monkeypatch): geometry=[(n_rows, n_frames)]) sras = SrasFile(str(path)) - # Shrink the budget so chunk_rows collapses to 1 and every row is - # its own chunk — the worst case for boundary bugs. + # Shrink the budget so the outer row loop splits into many chunks, and + # the block size so every chunk splits into many FFT tasks — the worst + # case for boundary bugs. monkeypatch.setattr(compute, "_TOTAL_BYTES_BUDGET", 8 * n_frames * spf * 4) + monkeypatch.setattr(compute, "_FFT_BLOCK", 4) + fft_rows = compute._plan_fft_rows(n_frames, spf, compute._TOTAL_BYTES_BUDGET) + assert fft_rows < n_rows, \ + f"FFT work actually splits into multiple chunks ({fft_rows} of {n_rows})" + dc_rows = compute._chunk_rows_for(n_frames, spf, compute._TOTAL_BYTES_BUDGET) + assert dc_rows < n_rows, \ + f"DC work actually splits into multiple chunks ({dc_rows} of {n_rows})" monkeypatch.setattr(compute, "_MAX_WORKERS", 1) dc_serial = compute_dc_image(sras, 0, CH4_IDX) @@ -128,27 +137,97 @@ def test_parallel_identity(tmp_path, monkeypatch): thr = float(np.median(dc4)) rf_masked_serial = compute_rf_image(sras, 0, dc_threshold_mv=thr, apply_bg_sub=True) - - chunk_rows, n_workers = compute._plan_chunks( - n_rows, n_frames, spf, live_multiplier=compute._FFT_LIVE_MULTIPLIER) - assert n_workers == 1, f"serial plan uses 1 worker (chunk_rows={chunk_rows})" - assert chunk_rows < n_rows, \ - f"work actually splits into multiple chunks ({chunk_rows} of {n_rows} rows)" + rf_pad_serial = compute_rf_image(sras, 0, dc_threshold_mv=thr, + apply_bg_sub=True, n_fft=spf * 8) monkeypatch.setattr(compute, "_MAX_WORKERS", 8) - chunk_rows, n_workers = compute._plan_chunks( - n_rows, n_frames, spf, live_multiplier=compute._FFT_LIVE_MULTIPLIER) - assert n_workers > 1, \ - f"parallel plan uses >1 worker (chunk_rows={chunk_rows} workers={n_workers})" - dc_par = compute_dc_image(sras, 0, CH4_IDX) rf_par = compute_rf_image(sras, 0, dc_threshold_mv=None, apply_bg_sub=True) rf_masked_par = compute_rf_image(sras, 0, dc_threshold_mv=thr, apply_bg_sub=True) + rf_pad_par = compute_rf_image(sras, 0, dc_threshold_mv=thr, + apply_bg_sub=True, n_fft=spf * 8) assert np.array_equal(dc_serial, dc_par), "dc image identical" assert np.array_equal(rf_serial, rf_par), "rf image identical (unmasked)" assert np.array_equal(rf_masked_serial, rf_masked_par), \ "rf image identical (masked)" + assert np.array_equal(rf_pad_serial, rf_pad_par), \ + "rf image identical (masked, padded/zoom)" + + +@pytest.mark.parametrize("spf,bps", [(64, 2), (37, 1)]) +def test_zoom_identity(tmp_path, monkeypatch, spf, bps): + """The zoom peak search must reproduce the full padded-rfft argmax + bit-for-bit, across pad factors, masking, bg-sub, dtype, and backend.""" + path = tmp_path / f"zoom_{spf}.sras" + gen.write(path, n_angles=2, seed=6, samples_per_frame=spf, bps=bps) + sras = SrasFile(str(path)) + dc4 = dc_image_mv(sras, 0, CH4_IDX) + thr = float(np.median(dc4)) + + backends = ["numpy"] + (["pyfftw"] if compute.PYFFTW_AVAILABLE else []) + for backend in backends: + monkeypatch.setattr(compute, "_fft_backend", backend) + for pad in (4, 8, 40): + n_fft = spf * pad + for thr_v in (None, thr): + for bg in (False, True): + ref = compute_rf_image(sras, 0, dc_threshold_mv=thr_v, + apply_bg_sub=bg, n_fft=n_fft, + exact=True) + zoom = compute_rf_image(sras, 0, dc_threshold_mv=thr_v, + apply_bg_sub=bg, n_fft=n_fft) + diff = int((ref != zoom).sum()) + assert diff == 0, \ + (f"{diff} px differ: backend={backend} pad={pad} " + f"thr={thr_v} bg={bg} spf={spf}") + + # A threshold above every pixel masks everything: both paths must agree + # on an all-zero image. + all_masked = compute_rf_image(sras, 0, dc_threshold_mv=1e9, n_fft=spf * 8) + assert not all_masked.any() + + +def test_zoom_identity_fuzz(): + """Hammer _peak_bins_zoom directly with adversarial spectra: noise, + un-subtracted DC offsets, on-bin and off-bin tones, near-tie tone pairs, + and all-zero rows.""" + import scipy.fft as scipy_fft + + rng = np.random.default_rng(42) + for _ in range(25): + spf = int(rng.integers(16, 220)) + pad = int(rng.choice([4, 5, 8, 16, 40])) + n_fft = spf * pad + n_wf = 24 + w = rng.normal(scale=20.0, size=(n_wf, spf)) + t = np.arange(spf) + # rows 0-5: pure/noisy tones (some off-bin), row 6-7: near-tie pair, + # row 8: big DC offset, row 9: all zeros, rest: plain noise. + for r in range(6): + f = rng.uniform(1.0, spf / 2 - 1) + w[r] = 60 * np.sin(2 * np.pi * f * t / spf) + w[r] * (r % 2) + f1, f2 = rng.uniform(2.0, spf / 2 - 2, size=2) + w[6] = 50 * np.sin(2 * np.pi * f1 * t / spf) \ + + 49.9 * np.sin(2 * np.pi * f2 * t / spf) + w[7] = 50 * np.sin(2 * np.pi * f1 * t / spf) \ + + 50 * np.cos(2 * np.pi * f2 * t / spf) + w[8] = 90 + rng.normal(scale=5.0, size=spf) + w[9] = 0.0 + w = w.astype(np.float32) + + S = scipy_fft.rfft(w, n=n_fft, axis=-1, workers=1) + P = S.real ** 2 + P += S.imag ** 2 + P[:, 0] = 0.0 + ref = np.argmax(P, axis=1) + + zp = compute._zoom_plan(spf, n_fft) + got = compute._peak_bins_zoom(w, zp) + bad = np.nonzero(ref != got)[0] + assert not len(bad), \ + (f"spf={spf} pad={pad}: rows {bad.tolist()} picked " + f"{got[bad].tolist()} instead of {ref[bad].tolist()}") def test_nomask_equals_low_threshold(tmp_path): diff --git a/tools/bench_fft.py b/tools/bench_fft.py new file mode 100644 index 0000000..d71285f --- /dev/null +++ b/tools/bench_fft.py @@ -0,0 +1,106 @@ +#!/usr/bin/env python3 +"""Benchmark the FFT peak-search path: exact vs zoom, serial vs pooled. + +Reports wall time, waveforms/s, CPU utilization (utime+stime over wall, in +cores), and verifies every variant against the exact reference image. + +Usage: + python tools/bench_fft.py # synthetic, pads 1/8/40 + python tools/bench_fft.py --pads 40 --spf 2500 --rows 8 --frames 1024 + python tools/bench_fft.py --real /path/big.sras --real-rows 32 --pads 40 +""" + +import argparse +import resource +import sys +import tempfile +import time +from pathlib import Path + +import numpy as np + +sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) + +import sras_compute as compute # noqa: E402 +from sras_compute import compute_rf_image, set_fft_backend # noqa: E402 +from sras_format import SrasFile # noqa: E402 +import tools.make_test_sras as gen # noqa: E402 +from tools.check_equivalence import row_slice # noqa: E402 + + +def _timed(fn): + r0 = resource.getrusage(resource.RUSAGE_SELF) + t0 = time.perf_counter() + out = fn() + wall = time.perf_counter() - t0 + r1 = resource.getrusage(resource.RUSAGE_SELF) + cpu = (r1.ru_utime - r0.ru_utime) + (r1.ru_stime - r0.ru_stime) + return out, wall, cpu / max(wall, 1e-9) + + +def bench(sras, pads, backends): + n_wf = sum(int(sras.n_rows[a]) * int(sras.n_frames[a]) + for a in range(sras.n_angles)) + spf = sras.samples_per_frame + print(f"{n_wf} waveforms x {spf} samples, {sras.n_angles} angle(s)") + print(f"{'pad':>4} {'backend':>8} {'variant':>16} {'wall':>9} " + f"{'wf/s':>10} {'util':>6} match") + + for pad in pads: + n_fft = spf * pad if pad > 1 else None + for backend in backends: + set_fft_backend(backend) + + def run(**kw): + imgs = [compute_rf_image(sras, a, dc_threshold_mv=None, + apply_bg_sub=True, n_fft=n_fft, **kw) + for a in range(sras.n_angles)] + return np.concatenate([i.ravel() for i in imgs]) + + ref, wall, util = _timed(lambda: run(exact=True)) + rows = [("exact(serial)", ref, wall, util, True)] + for label, kw in (("zoom(serial)", dict(max_workers=1)), + ("zoom(pool)", {})): + img, wall, util = _timed(lambda: run(**kw)) + rows.append((label, img, wall, util, bool(np.array_equal(img, ref)))) + for label, img, wall, util, ok in rows: + print(f"{pad:>4} {backend:>8} {label:>16} {wall:>8.2f}s " + f"{n_wf / wall:>10.0f} {util:>5.1f}x " + f"{'OK' if ok else 'MISMATCH'}") + + +def main(): + p = argparse.ArgumentParser(description=__doc__) + p.add_argument("--pads", default="1,8,40", + help="comma-separated pad factors (default 1,8,40)") + p.add_argument("--spf", type=int, default=2500) + p.add_argument("--rows", type=int, default=8) + p.add_argument("--frames", type=int, default=1024) + p.add_argument("--backends", default=None, + help="comma-separated (default: numpy,pyfftw if available)") + p.add_argument("--real", help="path to a real .sras file") + p.add_argument("--real-rows", type=int, default=32, + help="rows of angle 0 to use from the real file") + args = p.parse_args() + + pads = [int(x) for x in args.pads.split(",")] + if args.backends: + backends = args.backends.split(",") + else: + backends = ["numpy"] + (["pyfftw"] if compute.PYFFTW_AVAILABLE else []) + + if args.real: + sras = row_slice(SrasFile(args.real), 0, args.real_rows) + sras.data = [sras.data[0]] + sras.n_angles = 1 + bench(sras, pads, backends) + else: + with tempfile.TemporaryDirectory(prefix="sras_bench_") as tmp: + path = Path(tmp) / "bench.sras" + gen.write(path, n_angles=1, seed=0, samples_per_frame=args.spf, + geometry=[(args.rows, args.frames)]) + bench(SrasFile(str(path)), pads, backends) + + +if __name__ == "__main__": + main()