Port Combine() (partials->fulls) to CUDA, mirroring process_rawrun bit-for-bit: one thread per raw-hkl run splits its usable partials into rocking events (frame gap <= 2), pools background, seeds F, runs 3 de-biased Poisson reweights and the capture-uncertainty term. Emission is deterministic - a count pass, a host exclusive prefix sum for per-run offsets, then an emit pass at those offsets - so fulls come out in raw-run-major/event order, identical to the CPU path; both pass instantiations share the same arithmetic so count == emit exactly. Dmax/Dmin/Fmax reproduce std::max/min NaN semantics (not fmax) for parity. Validated across the 18-crystal rotation battery: all 15 deterministic crystals (P1/P2/C2/H3/I23/P41212/P222/P422) match the CPU combine exactly on SG/ISa/CC1.2/ completeness and run-to-run (fulls count bit-identical); the 3 upstream-nondet crystals vary from GPU-prediction overflow, not the combine. Gated opt-in behind JFJOCH_RSM_GPU_COMBINE (default = CPU combine): combine alone is timing-neutral because the shared 1.2M SortFullsByFrame std::sort dominates and the fulls round-trip adds a copy - it only pays off once the fulls stay resident for scale-fulls + merge. Also add JFJOCH_RSM_NO_GPU master switch to force the CPU fallback (incl. phase-1 scaling) from one binary for A/B parity. SortFullsByFrame extracted from the Combine tail and shared by both paths. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
87 lines
5.0 KiB
C++
87 lines
5.0 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;
|
|
|
|
// --- 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). Uploaded once, alongside SetPartials.
|
|
void SetCombineInputs(const float *bkg, const float *image_number, const float *d);
|
|
|
|
// 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);
|
|
|
|
// Download the combined fulls SoA (length = Combine()'s return). corr/partiality/rlp are 1 by
|
|
// construction and not returned; the caller sets them.
|
|
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;
|
|
|
|
private:
|
|
struct Impl;
|
|
std::unique_ptr<Impl> impl_;
|
|
};
|