The multiplicative outlier band (c8de5f8d6) took its width from the error model's b. b is a Gaussian width in I and grows with whatever tail the fulls have: on a low-ISa hexagonal small-molecule set with a population of fulls lost to near zero, b = 0.54 against a core ln-spread of 0.25, and exp(6 b) = 25 let the high outliers of the weak reflections through (SHELXL R1 0.191 -> 0.333 there, 0.111 -> 0.124 on its 25 keV sweep). The spread is now measured on the merge's own fulls: the weighted median of |ln(I / median)|, weighted by (<I>/sigma_counting)^2 so the fulls whose ratio counting noise does not blur carry it; b remains the fallback where nothing can be measured. Battery (29 sets: 13 small-molecule, 16 protein), SHELXL R1 against the median fix alone: the strongly absorbing cubic set 0.116 -> 0.0998 (b-band 0.103); the low-ISa hexagonal set 0.191 -> 0.174 at 20 keV, 0.111 -> 0.101 at 25 keV; organics within +-0.0003 or better (cytidine 0.0617 -> 0.0613, lalanine 0.0541 -> 0.0533). Proteins unchanged except two low-ISa sets whose CC1/2 cutoff moves: 6yqf 3.32 -> 2.89 A (placement R-free at a fixed 3.32 A 0.4329 -> 0.4326, at 2.89 A 0.4642 -> 0.4630), and the split myoglobin set 1.77 -> 1.94 A (low-resolution R_meas 33.0% -> 27.3%). Private subset: two sets' R_meas lower, the rest unchanged. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01K5K8jvPPbmCrbqnWkddTuB
232 lines
15 KiB
C++
232 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 (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;
|
|
// 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_;
|
|
};
|