Repository navigation
Cut peak memory of the robust convolution by ~3.4x - #91
Merged
Merged
Conversation
A single 19533x17000 plane needed ~17 GB to convolve, which is more than twice the per-worker share of a typical 128 GB / 16-process dask job, so wide-field cubes were being OOM killed. Two things drove that, both incidental rather than needed by the algorithm: - `convolve` cast the u/v grids to complex64 before `gaussft`, so every intermediate in there, and the returned taper, was complex for a quantity that is real throughout. numba then promoted the whole expression to complex128 against the float64 scalars, at 16 bytes per pixel. - `np.fft.fft2` promotes a float32 image to complex128, holds one of those per transform axis, then `.astype(np.complex64)` copies it back down. The NaN mask got a second full complex transform at the same cost. The image is real, so half its spectrum is redundant. Take the taper real and in the image's own precision, fill it in a single pass so only the taper itself is ever allocated at image size, and transform with `scipy.fft.rfft2`/`irfft2`, which preserves single precision and stores the last axis once. The taper is applied in place and the spectrum is consumed by its own inverse, so neither is duplicated. Measured on a 9000x7800 float32 plane with 42.6% NaNs, peak RSS goes from 13.5x to 4.0x the image, and the call is 3.4x faster. Double-precision input still transforms at double precision, so nothing loses accuracy. Output through `beamcon_2D` is bit-identical for all four conv modes, the NaN mask and scaling factor are unchanged, and odd axis lengths round-trip. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01AN9y36hf41ZmfkLQPnd4AN
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Why
A single 19533x17000 plane needs ~17 GB to convolve through
conv_mode="robust". That is more than twice the per-worker share of a 128 GB / 16-process dask job, so wide-field polarisation cubes get OOM killed part way through aflintrun.Peak RSS was measured by replaying
convolve's hot path on a float32 plane with 42.6% NaNs (the real blanking fraction), sampling/proc/self/statm:Linear in pixel count, so 332 Mpix lands at ~17 GB.
What was costing it
Two things, both incidental rather than required by the algorithm:
complex64beforegaussft. Every intermediate in there (ur,vr,g_arg,ur_in,vr_in,dg_arg,g_final) became complex, for a quantity that is real throughout. numba then promoted the whole expression to complex128 against the float64 scalars, at 16 bytes per pixel, andconvolvecast the result back down tocomplex64afterwards.np.fft.fft2promotes a float32 image tocomplex128. numpy's fft API is defined at double precision, so a 1.24 GB plane became a 5.3 GB spectrum, with one of those live per transform axis, then.astype(np.complex64)copied it back down. The NaN mask got a second full complex transform pair at the same cost.What this does
The image is real, so half its spectrum is redundant. So:
gaussftnow returns a real taper in the dtype ofu/v, filled in a single pass. Only the taper itself is ever allocated at image size, instead of up to four full-size intermediates.convolvetransforms withscipy.fft.rfft2/irfft2, which preserves single precision and stores the last axis once rather than twice. The taper is applied in place and the spectrum is consumed by its own inverse, so neither is duplicated.Double-precision input still transforms at double precision, so nothing loses accuracy.
real_dtypeis matched on dtype kind and width rather than identity, becausefits.getdatahands back big-endian>f4andnp.dtype(">f4") == np.float32isFalse— an identity check there would send every real image down the double path and undo the saving. There is a test pinning exactly that.Result
Same 9000x7800 float32 plane, 42.6% NaNs, through
convolve:So ~17 GB becomes ~5 GB per plane, which fits an 8 GB-per-worker job without changing the dask config, and the call is 3.4x faster as a side effect.
This looks to be at the floor for a whole-image FFT: I also tried fusing the taper into the multiply so the full-size taper array never exists, and peak RSS did not move (3.0x either way on the image-only path) because the peak is set by the FFT working set, not the taper. Going lower would need a different algorithm (overlap-add or tiled), since the convolving kernel is tiny next to the image. Not attempted here.
Correctness
Checked against a pristine 4.3.2 copy loaded side by side, over
{128x96, 129x97} x {float32, float64} x {NaNs, no NaNs}:gaussftagrees with a double-precision reference of the same expression tortol=1e-12.The miriad-gated tests (
test_robust,test_scipy,test_astropy,test_smooth,test_robust_3d) cannot run without miriad installed, so instead pristine 4.3.2 and this branch were both driven throughbeamcon_2D.smooth_fits_fileson the same fixture fromtests/test_2d.py. Output is bit-identical for all four conv modes on both fixture beams —beamcon_2Dwrites float32, and the double-precision difference is exactly one float32 ULP (1.192e-07), so it rounds away on write. Those tests assertatol=1e-3, so their margin is untouched. They will still run for real in CI.21 passedlocally (12 pre-existing + 9 new), with the same 5 miriad errors as onmaster.ruff0.15.11 check and format both clean.New tests
tests/test_convolve_uv.pycovers the taper's dtype and realness, agreement with a double-precision reference, cross-precision agreement ofconvolveincluding odd shapes and NaN handling, that a>f4FITS array stays in single precision (asserting on the actual spectrum dtype and half-spectrum shape), and that the caller's array is not consumed.Open question: should
workersbe exposed?scipy.ffttakes aworkersargument. It is a real speedup when cores are idle, but it is worth nothing inflint's currentcores: 16/processes: 16config, where every allocated core already has its own task. Measured with N concurrent processes on an N-core box, all starting their transform at the same instant:3.1x per plane when cores are free, but only +9% throughput at
cores == processes, which is within run-to-run noise. Task-level parallelism has already saturated the cores. Note also thatworkersis pocketfft's own thread pool and does not readOMP_NUM_THREADS, andworkers=-1resolves toos.cpu_count(), i.e. the whole node rather than the SLURM allocation — so it would need to default to 1 regardless.Left out of this PR on that basis. It only starts paying if
processesdrops belowcores; happy to plumb it throughsmooth()and both CLIs as a follow-up if you want it.🤖 Generated with Claude Code
https://claude.ai/code/session_01AN9y36hf41ZmfkLQPnd4AN
Generated by Claude Code