PERF: ~1.5x faster denoise_image on MPS/CPU via separable box-sum filter - #70
Merged
Merged
Conversation
…noise_image
antstorch.denoise_image's core SANLM loop used F.conv2d/conv3d with an
all-ones kernel to compute local box-sums (mean/variance, and the per
search-offset patch-similarity sum). PyTorch's MPS conv3d backend has
very high per-call overhead for this kernel/volume size (~40ms/call),
and this convolution is invoked once per search offset (124x for the
default r=2), making it the dominant cost of the whole filter.
Since the "convolution" is just a box-sum (all-ones kernel), it is
separable: replace it with a new _box_filter_sum() that sums 2r+1
shifted tensor views per axis directly, with no convolution op. A
cumsum/prefix-sum formulation was tried first but rejected: it
differences two large accumulated sums and loses float32 precision
(subtly larger RMSE vs. the ANTs reference); direct sliding-window
summation avoids that cancellation entirely.
Benchmarked against ants.denoise_image on a 136x176x176 volume
(r=2, p=1, Rician), --ants-threads 1 for a deterministic reference:
- MPS: 10.78s -> 7.0s (ANTs/ANTsTorch speed ratio 2.48x -> 3.77x)
- CPU: also improved
- Accuracy vs. ANTs unchanged (RMSE 0.00836, correlation 0.99999),
output remains bit-exact deterministic run-to-run.
Also includes the benchmark script's existing pipeline warm-up,
per-run timing breakdown, and HTML comparison report generation
(tools/benchmarks/compare_denoise_image.py), used throughout this
work to profile and validate the change against the ANTs reference.
Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Member
|
Brilliant. Thanks @stnava |
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.
Summary
F.conv2d/F.conv3d-based box-sum filtering inantstorch.denoise_image's core SANLM loop with a new_box_filter_sum()helper that computes the same box-sum via direct sliding-window addition (sum of2r+1shifted views per axis).tools/benchmarks/compare_denoise_image.py(pipeline warm-up timing, per-run timing breakdown, HTML comparison report) that were used throughout to profile and validate this change against the ANTs reference.Why
Profiling
antstorch.denoise_imageon a real 136x176x176 volume (r=2,p=1, Rician) showed the core SANLM loop's per-search-offset box-sum convolution (F.conv3dwith a 3x3x3 ones kernel, called 124 times for the default search radius) was the dominant cost — PyTorch's MPS backend has very high per-call overhead forconv3dat this kernel/volume size (~40ms/call), independent of the actual FLOP count.Results
Benchmarked with
ants.denoise_imageas reference (--ants-threads 1for a deterministic single-threaded baseline):CPU also benefits from the same change. Accuracy is unchanged (differences are float32 noise-level, well within existing ANTs vs. ANTsTorch parity).
Test plan
pytest tests/test_denoise_image.py— 13/13 passedtools/benchmarks/compare_denoise_image.pyagainstants.denoise_imageon a real 3-D volume (MPS and CPU), confirming speedup and unchanged accuracy/determinism vs. the ANTs reference