Files
Jungfraujoch/image_analysis/scale_merge/RotationScaleMergeGPU.h
T
leonarski_fandClaude Opus 5 d8029524e7 Scaling: never let an observation's own fluctuation set its weight
A weighted mean is only unbiased while the weights are independent of the values
being averaged. The IUCr's own nomenclature report (Schwarzenbach et al., Acta
Cryst A45 (1989) 63-75) puts it directly: weights in averaging "should not be
based on the counting statistics of the individual observations whose estimated
variances are biased and result in larger weights for accidentally low
intensities". Two places in the rotation pipeline were doing exactly that, and
between them they drove whole resolution shells of merged intensity negative.

1. The profile fit computed its non-signal variance as

       var_bkg = max(0, 1/den - max(0, I) + bkg-estimate term)

   The point of a separate var_bkg is that it does NOT move with the
   reflection's own fluctuation, and 1/den - I is the quantity that does not:
   1/den is the fit variance taken at the fitted intensity and grows with it
   roughly one for one. Clamping the subtrahend at zero left a down-fluctuated
   reflection's own deflated variance standing as its background variance.
   Measured over 6.9 M partials of one weak rotation dataset, var_bkg/bkg came
   out at 3.7-5.4 for observations with I < 0 against 11.4-13.7 for I > 0 - the
   down-fluctuated half of every reflection carried a variance ~2.7x too small
   and was weighted up by the same factor, first in the 3D combine and then
   again in the merge. Removing the clamp makes var_bkg flat in I (~13 x bkg
   across the whole range).

2. The merge then weighted each combined full by 1/sigma_full^2, and sigma_full
   is by construction a function of the full's own answer: the combine's
   variance carries a corr*max(0, F) signal term, so every full with F <= 0 got
   the smallest variance the model allows while the strongest quartile got
   2.26x more. The merge now rebuilds that variance at the reflection's mean
   instead, from a linear model var(I) = var_bkg + var_per_I * I that the
   combine measures and stores on the full. This mirrors
   MergeOnTheFly::CorrectedSigma, whose comment already claimed to mirror the
   rotation combine.

Verified against an estimator that cannot see the fluctuation - summing the
partials and dividing by the summed partiality, the classical construction every
other program uses (Greenhough & Suddath, J. Appl. Cryst. 19 (1986) 400-409, via
Leslie, Acta Cryst D55 (1999) 1696-1702: profile fitting biases the individual
partials but not their sum). Reproducing the merge on dumped observations, the
shipped weighting sat ~1.9 sigma below that reference in the noise shells; the
two changes recover most of it, and every intensity-independent weighting
scheme agrees with the reference once (1) is in.

Four-crystal probe, XDS resolution limits, branch fingerprint identical on all
four (so none of these is a two-pass branch flip):

  weak cubic case   last shell <I/sig> -1.6 -> +0.2 (XDS +0.10), last shell
                    R_meas 478% -> 250% (XDS 246%), overall <I/sig> 6.1 -> 7.5
                    (XDS 7.18), R_meas 18.3% -> 18.1%, CC1/2_hi 38.2% -> 43.7%
  tetragonal case   outer shells <I/sig> -0.4/-0.8/-0.9/-1.0 -> +1.8/+1.2/
                    +0.9/+0.4, R_meas 184%/595%/7614%/nan -> 95%/119%/135%/232%
                    (the nan was the shell mean crossing zero), R_meas 33.3% ->
                    32.9%, CC1/2_hi 38.3% -> 56.5%
  trigonal case     R_meas 13.0% -> 12.5%, CC1/2_hi 14.4% -> 16.5%
  strong control    unchanged to every printed digit but ISa

Cost: ISa falls (17.2 -> 14.0 and 16.7 -> 14.9 on the two mid-strength cases,
28.3 -> 27.8 on the control). Strong reflections are untouched by (1) - their
partials are all positive, so var_bkg is bit-identical - but the joint a/b fit
redistributes: honest weak sigmas lower a, and b rises to keep the strong bins
fitted. The median reduced chi^2 improves (1.25 -> 1.14, 1.35 -> 1.28) so the
new split describes the scatter better, but ISa is the one headline metric that
moves the wrong way and it should be watched over the full battery.

The integrator change is shared, so the stills merge sees it too; there it feeds
GetExpectedVarianceMerge, which had been handed the same contaminated var_bkg.
That path is untested here.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-08-11 17:40:15 +02:00

161 lines
9.8 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);
// --- 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_;
};