Files
Jungfraujoch/image_analysis/spot_finding/ImageSpotFinderGPU.cu
T
leonarski_fandClaude Opus 5.5 aac1c54fb7 GPU spot finder: warp-per-32-columns local test, second pass only at ring pixels
The local-box SNR kernel (analyze_pixel) was 55% of all GPU kernel time on a
rotation run and held the image loop GPU-bound. Three exact changes:

- analyze_pixel is rewritten as one warp per 32 output columns, each lane
  holding the vertical sums of two input columns in registers and the 31-wide
  horizontal window taken from warp prefix scans. No shared memory and no
  block-wide synchronisation; the sums are modular 64-bit integers, so the
  result is the same bits as before.
- The adaptive finder keeps only (local-test & ring threshold), and a pixel's
  second-pass result depends only on its own window, so the second pass is
  evaluated at the ring pixels alone (analyze_candidates) instead of densely
  followed by and_bits.
- The first pass is then only read within NBX of a ring pixel, so a warp whose
  tile no ring pixel can reach skips it; it also no longer reads an all-zero
  previous-pass buffer.

Checked bit for bit against the old kernel on every frame of a 16M rotation
run, and p.hkl / p_unmerged.mtz md5-identical on two inhouse EIGER2 16M sets.
Total kernel time 19.8 -> 12.1 s, image loop 2.69 -> 1.37 ms/image; wall
44.0 -> 37.2 s and 56.0 -> 48.7 s.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01D1G8gJVAy6gp1K5Dz3NE5C
2026-09-26 17:16:06 +02:00

368 lines
18 KiB
Plaintext

// SPDX-FileCopyrightText: 2025 Filip Leonarski, Paul Scherrer Institute <filip.leonarski@psi.ch>
// SPDX-License-Identifier: GPL-3.0-only
// GPU Spot finding developed by Hans-Christian Stadler (PSI)
// Copyright (2019-2023) Paul Scherrer Institute
#include "ImageSpotFinderGPU.h"
#include "../../common/JFJochException.h"
struct spot_parameters {
int32_t width;
int32_t height;
float strong_pixel_threshold2;
int32_t count_threshold;
};
// input X x Y pixels array
// output X x Y bit array
static constexpr int WARP_SIZE = 32; // assume warp size of 32 cuda threads per warp
inline void cuda_err(cudaError_t val) {
if (val != cudaSuccess)
throw JFJochException(JFJochExceptionCategory::GPUCUDAError, cudaGetErrorString(val));
}
// Write pixel results to bit array
// params: spot finding parameters
// out: pixel result bit array
// pixel: flat pixel index = bit index into bit array
// val: pixel result
// **NOTE**: assumes sizeof(*out) * 8 == WARP_SIZE
__device__ __forceinline__ void write_result(const spot_parameters& params, uint32_t* out, int32_t pixel, uint8_t val)
{
static_assert(sizeof(*out) * 8 == WARP_SIZE, "Violation of essential implementation assumption: WARP_SIZE must match output array element type bit size!");
static constexpr unsigned ALL_THREADS = UINT32_MAX;
const int32_t laneid = threadIdx.x & (WARP_SIZE - 1);
unsigned result = __ballot_sync(ALL_THREADS, val);
const int32_t idx = pixel / WARP_SIZE; // global uint32_t index
const int32_t bit = pixel % WARP_SIZE; // local bit index
if ((bit >= laneid) && (laneid == 0)) { // write to upper part of uint32_t
result <<= bit;
if (result)
atomicOr(&out[idx], result);
} else if ((bit < laneid) && (bit == 0)) { // write to lower part of uint32_t
result >>= laneid;
if (result)
atomicOr(&out[idx], result);
}
}
// Determine if pixel could be a spot
// params: spot finding parameters
// val: pixel value
// sum: window sum
// sum2: window sum of squares
// count: window valid pixels count
// return the pixel result: 0-no spot / 1-spot candidate
__device__ __forceinline__ uint8_t pixel_result(const spot_parameters& params, const int64_t val, int64_t sum, int64_t sum2, int64_t count)
{
sum -= val;
sum2 -= val * val;
count -= 1;
const int64_t var = count * sum2 - (sum * sum); // This should be divided by ((2*NBX+1) * (2*NBY+1)-1)*((2*NBX+1) * (2*NBY+1))
const int64_t in_minus_mean = val * count - sum; // Should be divided by ((2*NBX+1) * (2*NBY+1));
const int64_t tmp1 = in_minus_mean * in_minus_mean;
const float tmp2 = var * params.strong_pixel_threshold2;
bool snr_criterion;
if (params.strong_pixel_threshold2 == 0.0f)
snr_criterion = true;
else
snr_criterion = (count > ImageSpotFinder::MIN_VALID_PIXELS) && (in_minus_mean > 0) && (tmp1 > tmp2);
bool count_criterion = (params.count_threshold == 0.0f) || (val > params.count_threshold);
bool strong_pixel = snr_criterion && count_criterion;
if (val == INT32_MAX)
strong_pixel = true;
else if (val == INT32_MIN)
strong_pixel = false;
return strong_pixel ? 1 : 0;
}
// Inclusive prefix sum over the lanes of a warp (modular arithmetic).
__device__ __forceinline__ uint64_t warp_inclusive_scan(uint64_t v) {
const int32_t lane = threadIdx.x & (WARP_SIZE - 1);
for (int offset = 1; offset < WARP_SIZE; offset *= 2) {
const uint64_t t = __shfl_up_sync(UINT32_MAX, v, offset);
if (lane >= offset)
v += t;
}
return v;
}
// Sum over the 2 * NBX + 1 = 31 columns of each lane's window, the columns held two per lane: a from
// lane i is column i of the warp's 64, b from lane i is column 32 + i (lane 31 holds no b). Lane l's
// window is columns l + 1 .. l + 31: a of lanes l+1..31 plus b of lanes 0..l-1.
__device__ __forceinline__ uint64_t warp_window_sum(uint64_t a, uint64_t b) {
static_assert(2 * ImageSpotFinder::NBX + 1 == WARP_SIZE - 1, "the two-columns-per-lane layout needs a 31-pixel window");
const int32_t lane = threadIdx.x & (WARP_SIZE - 1);
const uint64_t pa = warp_inclusive_scan(a);
const uint64_t pb = warp_inclusive_scan(b);
const uint64_t total_a = __shfl_sync(UINT32_MAX, pa, WARP_SIZE - 1);
uint64_t pb_before = __shfl_up_sync(UINT32_MAX, pb, 1);
if (lane < 1)
pb_before = 0;
return total_a - pa + pb_before;
}
// Find pixels that could be spots
// in: image input values
// out: pixel result bit array, 1 bit per pixel (0:no/1:candidate spot)
// prev_out: result of a previous pass; its pixels read as INT32_MAX (not counted, and strong).
// nullptr = no previous pass.
// needed: nullptr, or a bit buffer of the pixels whose result will actually be read by someone who
// looks within NBX pixels of them (analyze_candidates). A warp whose tile holds no pixel that
// any of them can reach leaves its tile unwritten.
// params: spot finding parameters
//
// One warp per 32 output columns and a band of rows (a "wave", gridDim-independent: the warps are
// numbered group-fastest). Each lane keeps the vertical sums (sum, sum2, count) of two input columns
// over the window's rows in registers - 64 columns for the warp, enough for 32 windows of 31 - and
// slides them one row at a time: the row leaving the window is read again from global memory (it was
// read 31 rows earlier and sits in cache) and the row entering it is added. The horizontal window
// sums come from warp prefix scans, so there is no shared memory and no block-wide synchronisation.
// All sums are integer (modular 64-bit), so the result is exact and independent of the order.
__global__ void analyze_pixel(const int32_t *in, const uint32_t *prev_out, uint32_t *out,
const uint32_t *needed, const spot_parameters params, int32_t nWaves)
{
constexpr int32_t NBX = ImageSpotFinder::NBX;
const int32_t lane = threadIdx.x & (WARP_SIZE - 1);
const int32_t warp = (blockIdx.x * blockDim.x + threadIdx.x) / WARP_SIZE;
const int32_t ngroups = (params.width + WARP_SIZE - 1) / WARP_SIZE;
const int32_t group = warp % ngroups;
const int32_t wave = warp / ngroups;
const int32_t rowsPerWave = (params.height + nWaves - 1) / nWaves;
const int32_t rmin = wave * rowsPerWave;
const int32_t rmax = min(rmin + rowsPerWave, params.height);
// Warp-uniform, so the shuffles and ballots below stay collective.
if (wave >= nWaves || rmin >= params.height)
return;
const int32_t out_col = group * WARP_SIZE + lane; // the column this lane writes
const int32_t col_a = group * WARP_SIZE - NBX - 1 + lane; // the two input columns this lane sums
const int32_t col_b = col_a + WARP_SIZE;
const bool use_a = col_a >= 0 && col_a < params.width;
const bool use_b = lane < WARP_SIZE - 1 && col_b >= 0 && col_b < params.width;
if (needed != nullptr) {
// Is any needed pixel within NBX of the tile, rows [rmin, rmax) x columns [c0, c0 + 32)?
// Each lane checks rows of the dilated band, the columns [c0 - NBX, c0 + 32 + NBX) as the
// (at most three) words they fall in, masked to exactly that range.
const int64_t c_lo = max(group * WARP_SIZE - NBX, 0);
const int64_t c_hi = min(group * WARP_SIZE + WARP_SIZE + NBX, params.width); // past the last
bool any = false;
for (int32_t r = max(rmin - NBX, 0) + lane; r < min(rmax + NBX, params.height) && !any; r += WARP_SIZE) {
const int64_t first = static_cast<int64_t>(r) * params.width + c_lo;
const int64_t last = static_cast<int64_t>(r) * params.width + c_hi - 1;
for (int64_t w = first / 32; w <= last / 32; w++) {
uint32_t bits = needed[w];
if (w == first / 32)
bits &= UINT32_MAX << (first % 32);
if (w == last / 32 && last % 32 != 31)
bits &= (1U << (last % 32 + 1)) - 1;
any = any || (bits != 0);
}
}
if (!__any_sync(UINT32_MAX, any))
return;
}
const auto value_at = [&](int32_t row, int32_t col) -> int32_t {
const int32_t npixel = row * params.width + col;
const bool sat = prev_out != nullptr && ((prev_out[npixel / 32] & (1U << (npixel % 32))) != 0);
return sat ? INT32_MAX : in[npixel];
};
uint64_t sum_a = 0, sum2_a = 0, count_a = 0;
uint64_t sum_b = 0, sum2_b = 0, count_b = 0;
// Add (sign +1) or remove (-1) row r of this lane's two columns; sentinels count as nothing.
const auto slide = [&](int32_t row, uint64_t sign) {
if (use_a) {
const int32_t v = value_at(row, col_a);
if (v != INT32_MAX && v != INT32_MIN) {
sum_a += sign * static_cast<uint64_t>(static_cast<int64_t>(v));
sum2_a += sign * static_cast<uint64_t>(static_cast<int64_t>(v) * v);
count_a += sign;
}
}
if (use_b) {
const int32_t v = value_at(row, col_b);
if (v != INT32_MAX && v != INT32_MIN) {
sum_b += sign * static_cast<uint64_t>(static_cast<int64_t>(v));
sum2_b += sign * static_cast<uint64_t>(static_cast<int64_t>(v) * v);
count_b += sign;
}
}
};
for (int32_t r = max(rmin - NBX, 0); r < min(rmin + NBX, params.height - 1) + 1; r++)
slide(r, 1);
for (int32_t row = rmin; row < rmax; row++) {
const auto sum = static_cast<int64_t>(warp_window_sum(sum_a, sum_b));
const auto sum2 = static_cast<int64_t>(warp_window_sum(sum2_a, sum2_b));
const auto count = static_cast<int64_t>(warp_window_sum(count_a, count_b));
uint8_t val = 0;
if (out_col < params.width)
val = pixel_result(params, value_at(row, out_col), sum, sum2, count);
write_result(params, out, row * params.width + out_col, val);
if (row - NBX >= 0)
slide(row - NBX, UINT64_MAX); // -1: the row leaving the window
if (row + NBX + 1 < params.height)
slide(row + NBX + 1, 1);
}
}
// The second pass of analyze_pixel, evaluated only at the candidate pixels (bit set in cand) instead
// of over the whole image. Used where the caller keeps only (pass-2 result & cand), as the adaptive
// finder does with its ring mask: every other pixel of the dense pass-2 result is discarded, and a
// pixel's pass-2 result depends only on the image and the pass-1 bits inside its own window, so this
// yields exactly those bits. The window is summed directly - the same clipped 31 x 31 box, the same
// exclusion of sentinel and pass-1 pixels, integer sums, the same pixel_result - so the result is
// bit-identical to the dense pass followed by the AND.
//
// One warp per candidate: lane dx (0..30) sums column c + dx - NBX over the window's rows, and the
// warp reduces. The candidate words are scanned 32 at a time, one per lane; a word is owned by one
// warp, which writes it whole, so out needs no atomics (it must be zeroed beforehand).
__global__ void analyze_candidates(const int32_t *in, const uint32_t *prev_out, const uint32_t *cand,
uint32_t *out, const spot_parameters params, size_t nwords)
{
constexpr unsigned ALL_THREADS = UINT32_MAX;
constexpr int NBX = ImageSpotFinder::NBX;
const int lane = threadIdx.x & (WARP_SIZE - 1);
const size_t warp = (static_cast<size_t>(blockIdx.x) * blockDim.x + threadIdx.x) / WARP_SIZE;
const size_t nwarps = static_cast<size_t>(gridDim.x) * blockDim.x / WARP_SIZE;
const int32_t npix = params.width * params.height;
const auto value_at = [&](int32_t pixel) -> int32_t {
const bool sat = ((prev_out[pixel / 32] & (1U << (pixel % 32))) != 0);
return sat ? INT32_MAX : in[pixel];
};
for (size_t base = warp * WARP_SIZE; base < nwords; base += nwarps * WARP_SIZE) {
const uint32_t my_word = (base + lane < nwords) ? cand[base + lane] : 0;
unsigned words = __ballot_sync(ALL_THREADS, my_word != 0);
while (words) {
const int src = __ffs(words) - 1;
words &= words - 1;
uint32_t bits = __shfl_sync(ALL_THREADS, my_word, src);
const size_t word = base + src;
uint32_t result = 0;
while (bits) {
const int b = __ffs(bits) - 1;
bits &= bits - 1;
const int32_t pixel = static_cast<int32_t>(word * 32 + b);
if (pixel >= npix)
break;
const int32_t row = pixel / params.width;
const int32_t col = pixel % params.width + lane - NBX;
int64_t sum = 0, sum2 = 0;
int32_t count = 0;
if (lane < 2 * NBX + 1 && col >= 0 && col < params.width) {
const int32_t rlo = max(row - NBX, 0);
const int32_t rhi = min(row + NBX, params.height - 1);
for (int32_t r = rlo; r <= rhi; r++) {
const int32_t val = value_at(r * params.width + col);
if (val != INT32_MAX && val != INT32_MIN) {
sum += val;
sum2 += static_cast<int64_t>(val) * val;
count += 1;
}
}
}
for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
sum += __shfl_down_sync(ALL_THREADS, sum, offset);
sum2 += __shfl_down_sync(ALL_THREADS, sum2, offset);
count += __shfl_down_sync(ALL_THREADS, count, offset);
}
if (lane == 0 && pixel_result(params, value_at(pixel), sum, sum2, count))
result |= 1u << b;
}
if (lane == 0)
out[word] = result;
}
}
}
ImageSpotFinderGPU::ImageSpotFinderGPU(int32_t in_width, int32_t in_height,
std::shared_ptr<CudaStream> in_stream)
: ImageSpotFinder(in_width, in_height, false),
stream(in_stream),
extractor(in_width, in_height, std::move(in_stream)) {
gpu_out_0 = CudaDevicePtr<uint32_t>(OutputSize());
gpu_out_1 = CudaDevicePtr<uint32_t>(OutputSize());
int device;
cuda_err(cudaGetDevice(&device));
int sm_count;
cuda_err(cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, device));
candidateBlocks = 8 * sm_count;
}
void ImageSpotFinderGPU::SetResolutionMaskBits(const std::vector<uint32_t> &packed_mask) {
ImageSpotFinder::SetResolutionMaskBits(packed_mask);
extractor.SetResolutionMask(res_mask_bits);
}
const std::vector<DiffractionSpot> &ImageSpotFinderGPU::ExtractComponents(const ImagePreprocessorBuffer &image,
const SpotFindingSettings &settings) {
extractor.Extract(gpu_out_1, image.getGPUBuffer(), settings, components);
return components;
}
void ImageSpotFinderGPU::Detect(const ImagePreprocessorBuffer &image, const SpotFindingSettings &settings) {
RunDetect(image, settings, nullptr);
// The bit buffer stays on the device - ExtractComponents reads it there.
cuda_err(cudaStreamSynchronize(*stream));
}
void ImageSpotFinderGPU::DetectAt(const ImagePreprocessorBuffer &image, const SpotFindingSettings &settings,
const uint32_t *gpu_candidates) {
RunDetect(image, settings, gpu_candidates);
}
void ImageSpotFinderGPU::RunDetect(const ImagePreprocessorBuffer &image, const SpotFindingSettings &settings,
const uint32_t *gpu_candidates) {
spot_parameters spot_params{};
spot_params.height = height;
spot_params.width = width;
spot_params.strong_pixel_threshold2 = settings.signal_to_noise_threshold * settings.signal_to_noise_threshold;
spot_params.count_threshold = settings.photon_count_threshold;
if (2 * NBX + 1 > windowSizeLimit)
throw JFJochException(JFJochExceptionCategory::SpotFinderError, "nbx exceeds window size limit");
if (2 * NBX + 1 > windowSizeLimit)
throw JFJochException(JFJochExceptionCategory::SpotFinderError, "nby exceeds window size limit");
if (windowSizeLimit > numberOfCudaThreads)
throw JFJochException(JFJochExceptionCategory::SpotFinderError, "window size limit exceeds number of cuda threads");
if (windowSizeLimit > spot_params.width)
throw JFJochException(JFJochExceptionCategory::SpotFinderError, "window size limit exceeds number of columns");
if (windowSizeLimit > spot_params.height)
throw JFJochException(JFJochExceptionCategory::SpotFinderError, "window size limit exceeds number of height");
const int32_t ngroups = (spot_params.width + 31) / 32;
const int32_t warpsPerBlock = numberOfCudaThreads / 32;
const int32_t nBlocks = (ngroups * numberOfWaves + warpsPerBlock - 1) / warpsPerBlock;
cuda_err(cudaMemsetAsync(gpu_out_0, 0, OutputSize() * sizeof(uint32_t), *stream));
cuda_err(cudaMemsetAsync(gpu_out_1, 0, OutputSize() * sizeof(uint32_t), *stream));
// The first pass has no previous one (the classic engine used to hand it the zeroed gpu_out_1).
analyze_pixel<<<nBlocks, numberOfCudaThreads, 0, *stream>>>
(image.getGPUBuffer(), nullptr, gpu_out_0, gpu_candidates, spot_params, numberOfWaves);
cuda_err(cudaGetLastError());
if (gpu_candidates == nullptr)
analyze_pixel<<<nBlocks, numberOfCudaThreads, 0, *stream>>>
(image.getGPUBuffer(), gpu_out_0, gpu_out_1, nullptr, spot_params, numberOfWaves);
else
analyze_candidates<<<candidateBlocks, 256, 0, *stream>>>
(image.getGPUBuffer(), gpu_out_0, gpu_candidates, gpu_out_1, spot_params, OutputSize());
cuda_err(cudaGetLastError());
}