Run downloaded the correction factor for every partial, scattered it into an eight-and-a-half-million element array of eighty-byte records, kept a copy of it, filtered it twice on the host, gathered it back and uploaded it again. Five passes over six hundred and eighty megabytes, on the path where the data was already resident on the card. A comment above it warned that three host readers needed the scattered copy, and an earlier attempt read that as a reason to leave the whole thing alone. Taken one at a time the readers fall: both pass filters test quantities the card already holds - zeta and the frame index - so they become kernels over the correction array in place; the saved copy is now the download itself, thirty-four megabytes rather than a gather of the whole record; and the combine, the only genuine host reader, gets its own scatter immediately before it, which matters solely when observations are dumped. Zeta is compared in double on the device so the promotion matches the host's comparison exactly, and the drop count is an atomic add. The merge's own sweeps had the same shape: they walked the eighty-byte record to reach twenty bytes of it. They now build those twenty bytes once per merge, on all threads, and stream them. The reject median's first walk over every full goes entirely - the counts it was accumulating are the ones the error-model pass has already produced. The group histogram is one flat uninitialised buffer whose rows are cleared by the threads that use them, in place of a vector of vectors cleared twice, and the ingest no longer zeroes four hundred and fifty megabytes of staging that the following line overwrites. Faster on seventeen of seventeen matched pairs across two alternating sessions; eight to ten per cent of whole-run wall clock on the datasets where the tail dominates. The reflection files are byte-identical, including with a frame correlation cut and with an observation dump, which are what exercise the two filters and the host combine. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_016NNnL26LAvruQ9eLUUWvrJ
168 lines
10 KiB
C++
168 lines
10 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;
|
|
|
|
// 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). 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,
|
|
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 *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 per-obs fields the combine needs on top of the scaling inputs (image-local bkg, fractional
|
|
// frame position for event contiguity, resolution, and the predicted detector position px/py carried
|
|
// to the full for the absorption surface). Uploaded once, alongside SetPartials.
|
|
void SetCombineInputs(const float *bkg, const float *var_bkg,
|
|
const float *image_number, const float *d,
|
|
const float *px, const float *py);
|
|
|
|
// 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,
|
|
double min_captured_fraction);
|
|
|
|
// 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_;
|
|
};
|