Rewrite the FFT peak search: block-parallel zoom refinement
At pad 40 the old path materialised a ~9 GB padded spectrum per row,
which collapsed the chunk planner to one worker and one rfft call with
workers=1 — synthesis ran single-threaded, ~1 hour per angle on real
files.
The padded spectrum is never materialised now. Each block of 512
waveforms gets a coarse rfft at next_fast_len(2*spf); every coarse bin
within 0.7 of its row's max (plus the DC-adjacent window, which coarse
DC suppression would otherwise blind) is refined onto the exact n_fft
grid by a small complex gemm. The selected bin is bit-identical to the
full padded argmax — enforced by test_zoom_identity, a 25-seed fuzz
test over adversarial spectra, and a clean golden-hash diff against the
pre-rewrite baseline across pads {1,2,4,8,40}, masked/unmasked, bg
on/off, int8/int16, and both backends.
Blocks fan out over a persistent thread pool; pyFFTW runs through
per-thread FFTW_MEASURE builder plans with wisdom persisted to
~/.cache/sras-viewer, and threadpoolctl clamps BLAS under the pool.
compute_rf_image(exact=True) (or SRAS_FFT_EXACT=1) keeps the reference
padded path for audits.
tools/bench_fft.py measures: pad 40, 16 cores, 8192x2500 synthetic —
exact serial 717 wf/s -> zoom pool 25100 wf/s (35x, pyFFTW backend;
19x scipy), every variant verified equal to the reference.
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
+92
-13
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user