Files
Jungfraujoch/image_analysis/scale_merge/RotationScaleMergeGPU.h
T
leonarski_fandClaude Opus 4.8 2c928b27cd RotationScaleMerge: GPU 3D-combine (CUDA port, phase 2 step 1)
Port Combine() (partials->fulls) to CUDA, mirroring process_rawrun bit-for-bit:
one thread per raw-hkl run splits its usable partials into rocking events (frame
gap <= 2), pools background, seeds F, runs 3 de-biased Poisson reweights and the
capture-uncertainty term. Emission is deterministic - a count pass, a host
exclusive prefix sum for per-run offsets, then an emit pass at those offsets - so
fulls come out in raw-run-major/event order, identical to the CPU path; both pass
instantiations share the same arithmetic so count == emit exactly. Dmax/Dmin/Fmax
reproduce std::max/min NaN semantics (not fmax) for parity.

Validated across the 18-crystal rotation battery: all 15 deterministic crystals
(P1/P2/C2/H3/I23/P41212/P222/P422) match the CPU combine exactly on SG/ISa/CC1.2/
completeness and run-to-run (fulls count bit-identical); the 3 upstream-nondet
crystals vary from GPU-prediction overflow, not the combine.

Gated opt-in behind JFJOCH_RSM_GPU_COMBINE (default = CPU combine): combine alone
is timing-neutral because the shared 1.2M SortFullsByFrame std::sort dominates and
the fulls round-trip adds a copy - it only pays off once the fulls stay resident
for scale-fulls + merge. Also add JFJOCH_RSM_NO_GPU master switch to force the CPU
fallback (incl. phase-1 scaling) from one binary for A/B parity. SortFullsByFrame
extracted from the Combine tail and shared by both paths.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-03 07:27:33 +02:00

87 lines
5.0 KiB
C++

// SPDX-FileCopyrightText: 2026 Filip Leonarski, Paul Scherrer Institute <filip.leonarski@psi.ch>
// SPDX-License-Identifier: GPL-3.0-only
#pragma once
#include <cstdint>
#include <memory>
#include <optional>
#include <vector>
// GPU engine for the RotationScaleMerge hot loops. The class keeps the per-observation data resident on
// the device as a structure-of-arrays (coalesced) and runs the scaling loop there. The host keeps the
// one-time raw-hkl sort and the per-space-group ASU keying (gemmi); it hands the GPU the dense group ids
// plus a group-ordered permutation so the per-group reduction is a deterministic segmented reduction
// (one block per group, fixed order, no atomics) - matching the run-to-run determinism of the CPU path.
//
// Only compiled when CUDA is available; the header is safe to include unconditionally (the impl behind
// the pimpl is null without CUDA, and RotationScaleMerge falls back to the CPU loops).
class RotationScaleMergeGPU {
public:
RotationScaleMergeGPU();
~RotationScaleMergeGPU();
RotationScaleMergeGPU(const RotationScaleMergeGPU &) = delete;
RotationScaleMergeGPU &operator=(const RotationScaleMergeGPU &) = delete;
// True if a GPU was found and the engine is usable.
[[nodiscard]] bool Available() const;
// Upload the immutable per-observation fields (once). Arrays are length n_obs unless noted; the
// frame CSR (frame_start/frame_count) is length n_frames and indexes the obs arrays in frame order.
void SetPartials(int n_obs, int n_frames,
const float *I, const float *sigma, const float *rlp, const float *partiality,
const float *zeta, const uint8_t *on_ice, const int32_t *frame,
const float *corr0,
const int32_t *frame_start, const int32_t *frame_count);
// Per space group: the dense ASU-group id per obs, and a group-ordered permutation of the obs whose
// group >= 0 (group_perm), with its CSR (group_start/group_count, length n_groups) - so each group's
// observations are a contiguous, fixed-order segment for the reduction.
void SetGroups(int n_groups, const int32_t *group,
const int32_t *group_perm, int n_group_perm,
const int32_t *group_start, const int32_t *group_count);
// Re-upload the working corr (length n_obs) before a scaling pass (the host mutates it via smooth-G
// between passes). SetPartials uploads the initial corr; this refreshes it.
void SetCorr(const float *corr);
// Run `iters` of {reduce group means -> per-frame robust IRLS G -> update corr} on the device,
// in place on the resident corr. Rotation model (partiality folded via the stored partiality).
void ScalePartials(int iters, double robust_k, double min_partiality, bool has_d_min);
// Copy the updated corr back to the host (length n_obs), and the fitted per-frame G (length
// n_frames, as double) plus the per-frame "was fitted" flag.
void GetCorr(float *corr_out) const;
void GetG(double *g_out, uint8_t *scaled_out) const;
// --- 3D combine (partials -> fulls), all on the device ---
// The per-obs fields the combine needs on top of the scaling inputs (image-local bkg, fractional
// frame position for event contiguity, resolution). Uploaded once, alongside SetPartials.
void SetCombineInputs(const float *bkg, const float *image_number, const float *d);
// The one-time raw-hkl run layout (space-group-independent): the (raw h,k,l, image_number)-sorted
// permutation of the obs, split into contiguous per-raw-hkl runs. Uploaded once in Ingest.
void SetRawRuns(int n_runs, int n_perm, const int32_t *perm,
const int32_t *rawrun_start, const int32_t *rawrun_count,
const int32_t *rawrun_h, const int32_t *rawrun_k, const int32_t *rawrun_l);
// Combine the resident partials (reading the current resident corr) into fulls on the device,
// mirroring RotationScaleMerge::Combine: one thread per raw-hkl run splits its usable partials into
// rocking events (frame gap <= 2), pools background, seeds F, does 3 de-biased Poisson reweights and
// adds the capture-uncertainty term. rawrun_group (length n_runs) is the current space group's ASU
// id per raw hkl (it becomes the full's group). Deterministic: fulls are emitted in raw-run-major,
// event order (a count pass -> host prefix sum -> emit-at-offset), matching the CPU path. Returns the
// number of fulls (call GetFulls with buffers of that length).
int Combine(const int32_t *rawrun_group, double min_partiality, double capture_uncertainty_coeff);
// Download the combined fulls SoA (length = Combine()'s return). corr/partiality/rlp are 1 by
// construction and not returned; the caller sets them.
void GetFulls(int32_t *h, int32_t *k, int32_t *l, float *I, float *sigma, float *d,
float *image_number, int32_t *frame, uint8_t *on_ice, int32_t *group) const;
private:
struct Impl;
std::unique_ptr<Impl> impl_;
};