Files
Jungfraujoch/image_analysis/scale_merge/ErrorModel.cpp
T
leonarski_fandClaude Opus 5.5 87f0bdc870 RotationScaleMerge: fit the error model on the device in the GPU build (tier F)
FitErrorModelGPU puts the samples in the host's rank order with four stable
radix passes, so the equal-count bins hold the same samples; computes every
normalised deviation with the host's rounding, so every bin median is the
host's to the bit; and hands the per-bin terms to the host's own update step,
now ErrorModelUpdate (a pure extraction from FitErrorModel). The one thing
that differs is the order each bin's two sums are added in: a fixed tree on the
device, the order nth_element happened to leave the bin in on the host. a and
b^2 can therefore differ from the host fit in their last bits - measured on two
synthetic pools: relative difference 2e-16 and 3e-13 in a, 5e-16 in b^2. The
device fit is deterministic (same pool, same bits).

p.mtz md5 nevertheless unchanged on myob, cytc, 8a1a and 8qaw (GPU build): the
difference does not survive the float rounding of the merged sigmas there. A
knife-edge decision elsewhere could still see it, which is why this is kept as
its own commit.

The host fit took 0.4-0.6 s per fit on the large merges, about half of it the
initial equal-count split, with only 16-way parallelism in the iterations.
Device peak about 72 bytes per sample (the rank sort; each stage's buffers are
freed when it is done); the free memory is checked first and a device without
it stops with a message naming the CPU build.

Kept last on the branch so that it can be dropped on its own.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01SVmAWnzCmRKAXVUCdc4iNi
2026-10-08 08:52:04 +02:00

144 lines
7.1 KiB
C++

// SPDX-FileCopyrightText: 2026 Filip Leonarski, Paul Scherrer Institute <filip.leonarski@psi.ch>
// SPDX-License-Identifier: GPL-3.0-only
#include "ErrorModel.h"
#include <algorithm>
#include <cmath>
#include <functional>
#include <future>
#include "../../common/ParallelFor.h"
ErrorModelFit FitErrorModel(const std::vector<ErrorModelSample> &pool, std::vector<ErrorModelBinned> &smp,
size_t nthreads) {
constexpr int n_bins = ERROR_MODEL_BINS;
ErrorModelFit fit;
if (pool.size() < static_cast<size_t>(8 * n_bins))
return fit;
// The rank key is computed once, here, and carried with the sample: a key recomputed inside the
// comparator can round differently at different inlined call sites, and nth_element then sees an
// inconsistent order.
smp.resize(pool.size());
ParallelChunks(static_cast<int>(pool.size()), ThreadsForWork(pool.size(), nthreads), [&](int lo, int hi) {
for (int i = lo; i < hi; ++i)
smp[i] = {pool[i].I2 / pool[i].s2, pool[i].s2, pool[i].I2, pool[i].dev2, 0.0};
});
// Equal-count bins, each boundary put in place with nth_element; the split runs down the middle so
// each level halves the range it works on. Total order, so which of two equal keys lands in which bin
// is a property of the samples, not of the order they arrived in. The two halves of a split are
// disjoint ranges, so they can run side by side while the range is still worth a thread.
const size_t per = smp.size() / n_bins;
const auto by_snr = [](const ErrorModelBinned &a, const ErrorModelBinned &b) {
if (a.snr2 != b.snr2) return a.snr2 < b.snr2;
if (a.s2 != b.s2) return a.s2 < b.s2;
if (a.I2 != b.I2) return a.I2 < b.I2;
return a.dev2 < b.dev2;
};
constexpr size_t SPLIT_MIN_PARALLEL = 1u << 17;
const std::function<void(size_t, size_t, int, int, int)> split =
[&](size_t lo, size_t hi, int b0, int b1, int depth) {
if (b0 >= b1) return;
const int mid = (b0 + b1) / 2;
const size_t k = static_cast<size_t>(mid) * per;
std::nth_element(smp.begin() + lo, smp.begin() + k, smp.begin() + hi, by_snr);
if (depth > 0 && hi - lo > SPLIT_MIN_PARALLEL) {
auto left = std::async(std::launch::async,
[&, lo, k, b0, mid, depth] { split(lo, k, b0, mid, depth - 1); });
split(k, hi, mid + 1, b1, depth - 1);
left.get();
} else {
split(lo, k, b0, mid, 0);
split(k, hi, mid + 1, b1, 0);
}
};
split(0, smp.size(), 1, n_bins, nthreads > 1 ? 3 : 0);
auto bin_lo = [&](int bin) { return static_cast<size_t>(bin) * per; };
auto bin_hi = [&](int bin) { return bin == n_bins - 1 ? smp.size() : static_cast<size_t>(bin + 1) * per; };
// b is identified only if the strongest bin reaches where b*<I> can outgrow the counting term: when
// fewer reflections are strong than one bin holds, that bin sits at I/sigma ~ 1 and a two-term fit
// hands b the bins' own noise instead (measured: b = 5.6, ISa 0.5, on a crystal XDS gives 6.1). Then
// fit a alone and leave b at zero. Over the rotation battery the two crystals this fires on sit at
// (I/sigma)^2 0.22 and 0.84 while the next is 31.7, so the threshold is not delicate.
{
const auto top_mid = smp.begin() + (bin_lo(n_bins - 1) + bin_hi(n_bins - 1)) / 2;
std::nth_element(smp.begin() + bin_lo(n_bins - 1), top_mid, smp.end(), by_snr);
fit.b_measured = top_mid->snr2 > ERROR_MODEL_B_LEVER_MIN;
}
std::vector<double> bs2(n_bins), bI2(n_bins), bd2(n_bins);
double a = 1.0, b2 = 0.0;
for (int iter = 0; iter < 30; ++iter) {
// Per bin: the terms of var = a*s2 + b^2*I2, each divided by the current model's variance, and
// the median normalised squared deviation, expressed as the sum it stands for. Reordering inside
// a bin does not change which samples it holds. One bin per task, so the result does not depend
// on threads.
ParallelChunks(n_bins, ThreadsForWork(smp.size(), nthreads), [&](int blo, int bhi) {
for (int bin = blo; bin < bhi; ++bin) {
double s = 0.0, I = 0.0;
for (size_t i = bin_lo(bin); i < bin_hi(bin); ++i) {
auto &x = smp[i];
const double v = a * x.s2 + b2 * x.I2;
s += x.s2 / v; I += x.I2 / v;
x.z2 = x.dev2 / v;
}
const size_t n = bin_hi(bin) - bin_lo(bin);
const auto mid = smp.begin() + bin_lo(bin) + n / 2;
std::nth_element(smp.begin() + bin_lo(bin), mid, smp.begin() + bin_hi(bin),
[](const ErrorModelBinned &p, const ErrorModelBinned &q) { return p.z2 < q.z2; });
bs2[bin] = s; bI2[bin] = I; bd2[bin] = mid->z2 / ERROR_MODEL_CHI2_1_MEDIAN * static_cast<double>(n);
}
});
if (!ErrorModelUpdate(bs2, bI2, bd2, a, b2, fit))
break;
}
return fit;
}
bool ErrorModelUpdate(const std::vector<double> &bs2, const std::vector<double> &bI2, const std::vector<double> &bd2,
double &a, double &b2, ErrorModelFit &fit) {
const int n_bins = static_cast<int>(bs2.size());
// Every bin weighs the same in RELATIVE terms, so the weak bins still constrain a.
double Ass = 0, AsI = 0, AII = 0, Bs = 0, BI = 0;
for (int bin = 0; bin < n_bins; ++bin) {
if (!(bd2[bin] > 0.0)) continue;
const double w = 1.0 / (bd2[bin] * bd2[bin]);
Ass += w * bs2[bin] * bs2[bin]; AsI += w * bs2[bin] * bI2[bin]; AII += w * bI2[bin] * bI2[bin];
Bs += w * bs2[bin] * bd2[bin]; BI += w * bI2[bin] * bd2[bin];
}
const double det = Ass * AII - AsI * AsI;
double a_new, b2_new;
if (!fit.b_measured) {
if (!(Ass > 0.0)) return false;
a_new = std::clamp(Bs / Ass, 0.25, 100.0);
b2_new = 0.0;
} else {
if (!(det > 1e-10 * Ass * AII)) return false;
const double a_fit = (Bs * AII - BI * AsI) / det;
const double b2_fit = (Ass * BI - AsI * Bs) / det;
a_new = std::clamp(a_fit, 0.25, 100.0);
b2_new = std::max(b2_fit, 0.0);
// Leverage is not enough: b^2 can have it and still come out at zero within its own error,
// and 1/b then reads the noise in b as an I/sigma. The standard error of b^2 comes from the
// bins' own scatter about the fitted line; this only decides what is printed.
double chi = 0.0;
for (int bin = 0; bin < n_bins; ++bin) {
if (!(bd2[bin] > 0.0)) continue;
const double r = (bd2[bin] - a_fit * bs2[bin] - b2_fit * bI2[bin]) / bd2[bin];
chi += r * r;
}
fit.b_resolved = b2_fit >= 2.0 * std::sqrt(chi / (n_bins - 2) * Ass / det);
}
const bool settled = std::fabs(a_new - a) <= 1e-4 * a && std::fabs(b2_new - b2) <= 1e-4 * b2;
a = a_new;
b2 = b2_new;
fit.active = true;
fit.a = a;
fit.b2 = b2;
return !settled;
}