The delta-CC1/2 disposition keeps a weak stretch in the merge (downgraded) wherever removing it would not raise CC1/2, and the merge then carries each of its observations at the small 1/sigma^2 its scaled-up counting error gives it. R_meas counted those observations at full weight, so the frames that add next to nothing to the intensities set the number. hq-pool battery: 8xtf kept 115 weak frames (scale 0.1-0.2 of the run's) with the same CC1/2, <I/sigma> and CC_model per shell as rc173 with 84 frames rejected, and R_meas was 1.3-1.5x higher in every shell; 9w3y (33 deg rejected -> 0, every model metric better) and lyso_x10sa_strong read the same way. The disposition itself is not segmentation-dependent: conviction is on the batch grid, not the ledger ranges, and those sets' dispositions changed because the corrected data changed. R_meas now weights each observation by its merge weight v = 1/sigma^2 under the error model (corrected_sigma on the host, ModelSigma on the GPU, with the error model of the last MergeAccum), normalised per reflection to Kish's effective count, so equal sigmas give the ordinary formula (WeightedRmeasTerms). Where the proportional term dominates - strong reflections - frames are weighted alike, as in the merge: 8xtf's lowest shell reads 15.0%, the rc173 run with 84 frames rejected 15.2%. The per-hand table is weighted the same way. R_MEAS_UNWEIGHTED / REFRES_R_MEAS_UNWEIGHTED keep the XDS/AIMLESS convention and are what to set beside XDS; the battery scorer records them. MULTIPLICITY stays a count. Effect (weighted / unweighted): lyso_x06da_ref 0.0454 / 0.0479, 8xtf 0.225 / 0.907, insu_I_x06da_ref REFRES 0.072 / 0.232. GPU and CPU paths agree to the last printed digit on lyso_x06da_ref. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01D1G8gJVAy6gp1K5Dz3NE5C
199 lines
12 KiB
C++
199 lines
12 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 (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 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;
|
|
|
|
// 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.
|
|
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 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).
|
|
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, 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. 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);
|
|
|
|
// Start the Unity scaling of the resident fulls: corr, partiality, prescaling and zeta all 1, no
|
|
// frame fitted yet. Once before the ScaleFulls iterations.
|
|
void ResetFullsScale();
|
|
|
|
// Run `iters` of the Unity scaling loop on the resident fulls (reduce group means -> per-frame LS G
|
|
// -> update corr), in place on the fulls' working corr. Requires SetFullsFrameCSR + SetFullsGroups
|
|
// and ResetFullsScale.
|
|
void ScaleFulls(int iters, 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 ScaleFulls.
|
|
void GetFullsCorr(float *corr) const;
|
|
|
|
private:
|
|
struct Impl;
|
|
std::unique_ptr<Impl> impl_;
|
|
};
|