Exact (tier E): the device gets the same group ids and the same group-ordered permutation (ascending observation index within a group) the host's counting sort produced. ComputeAsuGroups still finds the groups on the host (one ASU reduction per raw-hkl run, the (key, run) sort, the dense ids, the representatives - now unpacked from the packed key instead of a second ASU reduction). With the GPU resident it then hands the device only the runs in group order (two arrays as long as the runs, not the observations): BuildGroups stamps every observation's group from its run (the device evaluates the same finiteness test on the same uploaded floats), counts each group's observations, and fills each group's segment and puts it in index order. The host no longer builds or uploads the two observation-length arrays; the device keeps its per-observation group array across Runs instead of reallocating it beside the old one. SetGroups (host-built upload) is gone; the CPU path keeps the host histogram CSR. Measured (prototype, GPU RTX 5080, Run-span sections): ComputeAsuGroups summed over a run's merges 8tyy 9.0 -> 4.2 s, 8a1a 2.2 -> 0.85 s, cytc 0.51 -> 0.19 s. md5 of p.mtz identical to production on myob/cytc/thau/8a1a/8tyy and cytc -N 4. Clean branch rebuilt from scratch (GPU and CPU) and re-verified: p.mtz md5 identical to production on myob/cytc/thau/8a1a/8tyy/8qaw/9gdj (GPU), cytc -N 4, myob/cytc and myob -N 4 (CPU). Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01SVmAWnzCmRKAXVUCdc4iNi
234 lines
15 KiB
C++
234 lines
15 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);
|
|
void SetObsClipped(int offset, int count, const uint8_t *clipped);
|
|
|
|
// Per space group: the dense ASU-group id per obs, and a group-ordered permutation of the obs whose
|
|
// group >= 0, with its CSR - so each group's observations are a contiguous, fixed-order segment for
|
|
// the reduction. Built on the device from the raw-hkl runs uploaded by SetRawRuns: rawrun_group is
|
|
// each run's group (-1 = none), and group g holds the runs sorted_run[group_first[g] ..
|
|
// group_first[g + 1]). An observation is in its run's group when it passes the finiteness test
|
|
// Ingest applies, and each group lists its observations in ascending index order.
|
|
void BuildGroups(int n_groups, const int32_t *rawrun_group, int n_sorted, const int32_t *sorted_run,
|
|
const int32_t *group_first);
|
|
|
|
// 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 weighted-LS 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 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;
|
|
// The information each frame's last fit had on G, sum w^2 c^2 (length n_frames, 0 = not fitted).
|
|
void GetInfo(double *info_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.
|
|
// hand / has_hands are the per-full Bijvoet hand and the per-group "this group pools two hands",
|
|
// which the error-model samples are formed on; both null leaves the fit on the pooled group.
|
|
void MergeEmSamples(bool for_search, double min_partiality,
|
|
const uint8_t *hand, const uint8_t *has_hands,
|
|
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.
|
|
// The per-group results stay on the device for MergeAccumRange to download; reject_median is
|
|
// uploaded (NAN where none). rejected_obs is the per-full flag (length n_fulls): on entry the
|
|
// observations the host already rejected (the Wilson test), on return those plus the median test's;
|
|
// the host needs it for the reductions it still does itself, above all the anomalous I(+)/I(-) split.
|
|
// frame_cc_factor (length n_frames) is each frame's (G_ref/G)^2 floored at 1; swh_typ0/1 are the
|
|
// half-set weights multiplied by it. Requires MergeEmSamples first (em_mean resident).
|
|
// reject_var_add (n_groups) widens the pooled cut by the shell's own measured Bijvoet
|
|
// variance; null leaves the plain n-sigma test.
|
|
// reject_mult_sd is the multiplicative spread of the outlier band (OutlierBand.h).
|
|
void MergeAccum(double error_model_a, double error_model_b, bool error_model_active,
|
|
bool reject_outliers, double reject_nsigma, double reject_mult_sd, const float *reject_median,
|
|
const float *reject_var_add,
|
|
const uint8_t *half, const double *frame_cc_factor, uint8_t *rejected_obs);
|
|
// Download groups [g0, g0 + n) of what MergeAccum left on the device; every output has length n.
|
|
// rejected[g] counts the outliers dropped from the group.
|
|
void MergeAccumRange(int g0, int n, double *swI, double *sw, double *swIh0, double *swIh1,
|
|
double *swh0, double *swh1, double *swh_typ0, double *swh_typ1,
|
|
int32_t *nh0, int32_t *nh1, double *d_out,
|
|
int32_t *rejected, uint8_t *on_ice_out);
|
|
|
|
// Per-group R_meas accumulators: sum|I_corr-merged_I| and sum I, the same with each observation
|
|
// weighted by its merge weight v = 1/sigma^2 under the error model of the last MergeAccum,
|
|
// sum v, sum v^2, 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, double *wabsdev, double *wsumI,
|
|
double *sumv, double *sumv2, 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); n_overloaded receives how many rocking
|
|
// events were dropped for a saturated pixel (see the host combine).
|
|
int Combine(const int32_t *rawrun_group, double min_partiality, double capture_uncertainty_coeff,
|
|
double min_captured_fraction, float max_frame_gap, int64_t &n_overloaded);
|
|
|
|
// Download the combined fulls SoA (length = Combine()'s return). The working corr is downloaded
|
|
// separately by GetFullsCorr (it is only meaningful after the fulls scaling; 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, uint8_t *clipped,
|
|
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 + (capture * I)^2. Length = n_fulls.
|
|
void GetFullsVariance(float *var_bkg, float *var_per_I, float *capture) 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);
|
|
|
|
// Start the Unity scaling of the resident fulls: corr, partiality, prescaling and zeta all 1, no
|
|
// frame fitted yet. Once before the FitFullsScale rounds.
|
|
void ResetFullsScale();
|
|
|
|
// One Unity fit of the resident fulls (reduce group means -> per-frame LS G and its information,
|
|
// GetG / GetInfo), leaving the fulls' working corr as it is: the host smooths the fitted scales
|
|
// and applies them with SmoothFullsCorr. Requires SetFullsFrameCSR + SetFullsGroups and
|
|
// ResetFullsScale.
|
|
void FitFullsScale(double min_partiality);
|
|
|
|
// The fulls' counterpart of SmoothCorr: f_corr[i] *= ratio[f_frame[i]] where apply[f].
|
|
void SmoothFullsCorr(const uint8_t *apply, const double *ratio);
|
|
|
|
// Download the fulls' working corr (length = n_fulls), valid after the fulls scaling.
|
|
void GetFullsCorr(float *corr) const;
|
|
|
|
// --- correction-surface fit (RotationScaleMerge::ApplyCellSurface) ---
|
|
// The two passes every round of the fit makes over its terms, on the device; the host keeps the
|
|
// per-cell step, the gauge and the cross-validation. Both sums are formed in exactly the order the
|
|
// host forms them, so the fitted surface is the host's to the last bit (see the kernels).
|
|
|
|
// One observation as the surface fit sees it: the host's term, uploaded as it stands.
|
|
struct SurfaceTerm { float I, sigma, corr, d; int32_t cell, group; };
|
|
|
|
// The terms (in fulls order) with each one's frame parity, the ASU-group CSR over them (gperm lists
|
|
// the terms of group g at [gstart[g], gstart[g+1]), in fulls order) and the cell count.
|
|
void SurfaceSetTerms(int n_terms, const SurfaceTerm *term, const uint8_t *parity,
|
|
int n_groups, const int32_t *gperm, const int32_t *gstart, int ncell);
|
|
|
|
// One subset of the terms (0 = even frames, 1 = odd, 2 = all), cut into the host's n_blocks
|
|
// reduction blocks and ordered within each block by cell, keeping term order inside a cell:
|
|
// block b, cell c is perm[seg_start[b * ncell + c], seg_start[b * ncell + c + 1]).
|
|
void SurfaceSetSubset(int subset, int n_blocks, const int32_t *perm, const int32_t *seg_start);
|
|
|
|
// The per-group reference sums sw / swI over the terms of frame parity `parity` (< 0 = all) with
|
|
// the surface A (length ncell) applied. They stay on the device for SurfaceFitSums;
|
|
// SurfaceGetReference downloads them (length n_groups each).
|
|
void SurfaceReference(int parity, const double *A);
|
|
void SurfaceGetReference(double *sw, double *swI) const;
|
|
|
|
// The fit's per-cell sums over one subset against the last SurfaceReference and its A:
|
|
// cross = sum w Is Iref and ref2 = sum w Iref^2 (length ncell each).
|
|
void SurfaceFitSums(int subset, double *cross, double *ref2);
|
|
|
|
private:
|
|
struct Impl;
|
|
std::unique_ptr<Impl> impl_;
|
|
};
|