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
368 lines
18 KiB
Plaintext
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());
|
|
}
|