Files
Jungfraujoch/image_analysis/scale_merge/RotationScaleMergeGPU.h
T
leonarski_f cb5a2f032a
Build Packages / Create release (push) Successful in 23s
Build Packages / build:rugnux-tgz (x86_64) (push) Successful in 10m6s
Build Packages / build:rugnux:aarch64 (cross) (push) Successful in 9m6s
Build Packages / build:viewer-tgz:cpu (push) Successful in 11m15s
Build Packages / build:viewer-tgz:cuda (push) Successful in 12m21s
Build Packages / build:windows:nocuda (push) Successful in 17m9s
Build Packages / build:windows:cuda (push) Successful in 19m49s
Build Packages / HDF5 consumer tests (DIALS, XDS) (push) Successful in 25m42s
Build Packages / build:rpm (ubuntu2204_nocuda) (push) Successful in 16m0s
Build Packages / build:rpm (ubuntu2404_nocuda) (push) Successful in 14m54s
Build Packages / build:rpm (rocky8_nocuda) (push) Successful in 17m7s
Build Packages / build:rugnux:windows (push) Successful in 10m47s
Build Packages / build:rpm (rocky9_nocuda) (push) Successful in 17m4s
Build Packages / build:rpm (rocky8_sls9) (push) Successful in 17m8s
Build Packages / Generate python client (push) Successful in 45s
Build Packages / Build documentation (push) Successful in 1m45s
Build Packages / build:rpm (rocky9_sls9) (push) Successful in 19m23s
Build Packages / build:rpm (ubuntu2204) (push) Successful in 19m20s
Build Packages / build:rpm (ubuntu2404) (push) Successful in 18m43s
Build Packages / build:rpm (rocky8) (push) Successful in 19m31s
Build Packages / build:rpm (rocky9) (push) Successful in 20m16s
Build Packages / Unit tests (push) Successful in 1h41m19s
v1.0.0-rc.170 (#80)
* Fixed a `jfjoch_broker` crash during indexing: sorting no longer misbehaves on non-finite values, and GPU FFT indexer kernel launches are now error-checked.
* rugnux needs about a third less peak memory to scale, merge and post-refine rotation data, with identical results.
* `rugnux --model`: the placed coordinate file carries the space group its own coordinates obey, and says so when that is not the group the reflection files beside it carry.
* `jfjoch_viewer`: fixes in the dataset plots, inspector and layout; spot markers lose their black outline by default (a checkbox under "Image features" restores it) and the highest-pixel markers are white boxes around the pixel.

Reviewed-on: #80
Co-authored-by: Filip Leonarski <filip.leonarski@psi.ch>
2026-09-16 18:17:46 +02:00

175 lines
11 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;
// The observation count and the frame CSR (frame_start/frame_count, length n_frames, indexing
// the obs arrays in frame order), plus every device-side observation buffer, allocated here so
// the chunked uploads below have somewhere to land. Call once, before them.
void SetPartialsLayout(int n_obs, int n_frames,
const int32_t *frame_start, const int32_t *frame_count);
// One slice ([offset, offset+count) of n_obs) of one immutable per-observation field, so the
// host can stage a bounded chunk at a time out of its AoS record: staged whole and all at once,
// the fourteen arrays are more than half the entire observation payload held live again.
// Corr0 seeds the resident, mutable corr (SetCorr refreshes it). Bkg/VarBkg/ImageNumber/D/Px/Py
// are the combine inputs (image-local background, fractional frame position for event
// contiguity, resolution, and the predicted detector position carried to the full for the
// absorption surface).
enum class ObsField { I, Sigma, PrescalingCorr, Partiality, Zeta, Corr0,
Bkg, VarBkg, ImageNumber, D, Px, Py };
void SetObsField(ObsField f, int offset, int count, const float *v);
void SetObsFrame(int offset, int count, const int32_t *frame);
void SetObsOnIce(int offset, int count, const uint8_t *on_ice);
// 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). SetObsField(Corr0) 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;
// Apply the smooth-G correction to the resident corr: corr[i] *= ratio[frame[i]] for frames with
// apply[f] (both length n_frames), matching CPU SmoothG. Keeps corr resident (no round-trip).
void SmoothCorr(const uint8_t *apply, const double *ratio);
// The two pass filters, on the resident corr (both mirror the host loops in Run and keep corr on the
// device). By zeta: zero corr wherever the rocking geometry fails the de-novo search threshold,
// returning how many observations that removed from the merge, i.e. how many had a finite, positive
// corr. By frame: zero corr on the frames flagged in `reject` (length n_frames).
int64_t FilterCorrByZeta(double min_zeta);
void FilterCorrByFrame(const uint8_t *reject);
// --- merge + error-model reductions over the resident, scaled fulls (reuse the fulls group CSR) ---
// The per-frame cell-consistency mask (length n_frames) used by the merge filter. Uploaded once.
void SetFrameCellOk(const uint8_t *frame_cell_ok);
// Per-group inv-var mean (em_mean, length n_groups) + per-full leverage-corrected error-model samples
// (s2/I2/dev2 + valid flag, length n_fulls), mirroring MergeAndStats' first two error-model loops.
// Stashes (for_search, min_partiality) for the MergeAccum/MergeRmeas calls that follow.
void MergeEmSamples(bool for_search, double min_partiality,
double *em_mean_out, int32_t *cnt_out, double *s2_out, double *I2_out,
double *dev2_out, uint8_t *valid_out);
// Per-group merge accumulators (inv-var sums + deterministic half-sets, error-model-corrected sigma
// from a/b). `half` is the per-full CC1/2 half-set (length n_fulls) assigned on the host, so this
// kernel and the host merge loop cannot disagree about it.
// Outputs length n_groups; rejected[g] counts outliers dropped (reject_median uploaded, NAN
// where none). rejected_obs is the per-full flag (length n_fulls): the host needs it for the reductions
// it still does itself, above all the anomalous I(+)/I(-) split.
// Requires MergeEmSamples first (em_mean resident).
void MergeAccum(double error_model_a, double error_model_b, bool error_model_active,
bool reject_outliers, double reject_nsigma, const float *reject_median,
const uint8_t *half,
double *swI, double *sw, double *swIh0, double *swIh1,
double *swh0, double *swh1, int32_t *nh0, int32_t *nh1, double *d_out,
int32_t *rejected, uint8_t *on_ice_out, uint8_t *rejected_obs);
// Per-group R_meas accumulators (sum|I_corr-merged_I|, sum_I, n, and the count this looser walk
// accepted - which the host uses only to skip empty groups, NOT as the per-shell
// total_observations); merged_I is uploaded. All arrays length n_groups.
void MergeRmeas(const double *merged_I, double *absdev, double *sumI, int32_t *n, int32_t *nusable);
// Post-smooth per-frame diagnostic CC: recompute the group means from the resident (smoothed) corr
// and the Pearson CC of each frame's I*corr vs its group mean, downloading only the per-frame cc /
// cc_n (length n_frames). Mirrors ReduceGroupMeans(partials) + FinalizePerFrameScale's CC loop.
void ComputePartialCC(double min_partiality, double *cc_out, int64_t *cc_n_out);
// --- 3D combine (partials -> fulls), all on the device ---
// 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 <= max_frame_gap, the host's RockingEventFrameGap), 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,
double min_captured_fraction, float max_frame_gap);
// Download the combined fulls SoA (length = Combine()'s return). The working corr is downloaded
// separately by GetFullsCorr (it is only meaningful after ScaleFulls; otherwise the caller sets it).
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;
// Download the fulls' predicted detector position (peak partial's px/py), for the host absorption
// surface. Length = n_fulls.
void GetFullsPxPy(float *px, float *py) const;
// Download the fulls' variance model, var(I) = var_bkg + var_per_I * I. Length = n_fulls.
void GetFullsVariance(float *var_bkg, float *var_per_I) const;
// Re-upload the fulls' working corr (length n_fulls) after the host correction surfaces (decay /
// absorption) mutate it, so the resident merge reads the corrected scale.
void SetFullsCorr(const float *corr);
// --- scale the resident fulls on the device (Unity model), no round-trip ---
// Download the fulls' frame and ASU-group keys (emit order) so the host can build the frame/group CSRs
// with a counting sort (deterministic, no GPU stable-sort) and hand them back below.
void GetFullsKeys(int32_t *frame, int32_t *group) const;
// The fulls' per-frame CSR: frame_perm groups the emit-ordered fulls by frame (frame_start/count length
// n_frames index it), so FitPerFrameG can scale the fulls without physically reordering them.
void SetFullsFrameCSR(const int32_t *frame_perm, int n_perm,
const int32_t *frame_start, const int32_t *frame_count);
// The fulls' per-ASU-group CSR (group-ordered permutation of the fulls with group>=0, + its CSR).
void SetFullsGroups(const int32_t *gperm, int n_gperm,
const int32_t *gstart, const int32_t *gcount);
// Run `iters` of the Unity scaling loop on the resident fulls (reduce group means -> per-frame IRLS G
// -> update corr), in place on the fulls' working corr. Requires SetFullsFrameCSR + SetFullsGroups.
void ScaleFulls(int iters, double robust_k, double min_partiality);
// Download the fulls' working corr (length = n_fulls), valid after ScaleFulls.
void GetFullsCorr(float *corr) const;
private:
struct Impl;
std::unique_ptr<Impl> impl_;
};