// SPDX-FileCopyrightText: 2026 Filip Leonarski, Paul Scherrer Institute // SPDX-License-Identifier: GPL-3.0-only #include "PostRefine.h" #include #include #include "../../common/JFJochMath.h" // PI #include "XtalResidual.h" // XtalResidual (the positional detector<->reciprocal residual, step B) #include "LatticeReduction.h" #include "ceres/ceres.h" #include "ceres/rotation.h" namespace { // One integrated partial, flattened across all images. struct Partial { int h, k, l; float img; double I, sigma, zeta; double angle_rad; // frame mid-exposure goniometer angle double obs_x, obs_y; // observed spot centroid (pixels); NAN if the box sum found no centroid }; // A rocking event and its precomputed reference reciprocal vector (phi=0 frame, from the indexed lattice). struct Event { double phi_obs; // rad, intensity-weighted rocking centroid double weight; // sqrt(sum I / sum sigma) double e_ref[3]; // h*a* + k*b* + l*c* at the reference (unrefined) cell/orientation int h, k, l; }; // Distance-INDEPENDENT Ewald excitation residual for a uniform cell-scale parameter s and a refined // goniometer axis (3-vector). The header-distance miscalibration leaves a uniform cell scale; the axis is // the other phi_obs lever. Both are phi_obs-constrained (distance-independent). e_ref is the reference // reciprocal (h*a* + k*b* + l*c* at the indexed cell). On the Ewald sphere <=> |p|^2 + 2 p_z/lambda == 0. struct ScaleAxisExcitationResidual { ScaleAxisExcitationResidual(double lambda, double angle_rad, double weight, const double e_ref[3]) : inv_lambda(1.0 / lambda), angle_rad(angle_rad), weight(weight), ex(e_ref[0]), ey(e_ref[1]), ez(e_ref[2]) {} template bool operator()(const T *const s, const T *const axis, T *residual) const { const T inv_s = T(1) / s[0]; const T p_ref[3] = {T(ex) * inv_s, T(ey) * inv_s, T(ez) * inv_s}; const T aa[3] = {T(-angle_rad) * axis[0], T(-angle_rad) * axis[1], T(-angle_rad) * axis[2]}; T p_lab[3]; ceres::AngleAxisRotatePoint(aa, p_ref, p_lab); const T zeta = p_lab[0] * p_lab[0] + p_lab[1] * p_lab[1] + p_lab[2] * p_lab[2] + T(2.0) * p_lab[2] * T(inv_lambda); residual[0] = T(weight) * zeta * T(0.5) / T(inv_lambda); return true; } const double inv_lambda, angle_rad, weight, ex, ey, ez; }; } // namespace PostRefineResult PostRefineRotationGeometry(const std::vector &outcomes, const GoniometerAxis &axis, const DiffractionGeometry &nominal_geom, const CrystalLattice &reference_latt, const PostRefineSettings &settings, Logger &logger) { PostRefineResult result; result.geom = nominal_geom; result.cell = reference_latt.GetUnitCell(); result.distance_before_mm = nominal_geom.GetDetectorDistance_mm(); result.distance_after_mm = nominal_geom.GetDetectorDistance_mm(); try { const double wedge_half = axis.GetWedge_deg() / 2.0; const double lambda = nominal_geom.GetWavelength_A(); const Coord ax = axis.GetAxis(); const Coord Astar = reference_latt.Astar(), Bstar = reference_latt.Bstar(), Cstar = reference_latt.Cstar(); std::vector pts; for (const auto &o : outcomes) for (const auto &r : o.reflections) { if (!std::isfinite(r.I) || !std::isfinite(r.sigma) || r.sigma <= 0.0f) continue; const double mid_deg = axis.GetAngle_deg(r.image_number) + wedge_half; const double ox = std::isfinite(r.observed_x) ? r.observed_x : NAN; const double oy = std::isfinite(r.observed_y) ? r.observed_y : NAN; pts.push_back(Partial{r.h, r.k, r.l, r.image_number, r.I, r.sigma, std::isfinite(r.zeta) ? r.zeta : 0.0, mid_deg * PI / 180.0, ox, oy}); } logger.Info("Post-refine: {} partials gathered", pts.size()); if (pts.size() < static_cast(settings.min_events)) return result; std::sort(pts.begin(), pts.end(), [](const Partial &a, const Partial &b) { if (a.h != b.h) return a.h < b.h; if (a.k != b.k) return a.k < b.k; if (a.l != b.l) return a.l < b.l; return a.img < b.img; }); // Split into rocking events (same raw hkl, adjacent frames). Only >=2-frame events carry an // unbiased phi_obs (a single-frame centroid is just the frame centre); precompute e_ref per event. constexpr float MAX_FRAME_GAP = 2.0f; std::vector events; std::vector rock_std; // per-event observed rocking width (rad); task-3 mosaicity signal std::vector zeta_ev; // per-event intensity-weighted zeta (aligned with rock_std) size_t i = 0; while (i < pts.size()) { size_t j = i + 1; while (j < pts.size() && pts[j].h == pts[i].h && pts[j].k == pts[i].k && pts[j].l == pts[i].l && pts[j].img - pts[j - 1].img <= MAX_FRAME_GAP) ++j; if (j - i >= 2) { double sumI = 0, sumIphi = 0, sumIphi2 = 0, sumSig = 0, sumIzeta = 0; for (size_t m = i; m < j; ++m) { const double Ipos = std::max(0.0, pts[m].I); sumI += Ipos; sumIphi += Ipos * pts[m].angle_rad; sumIphi2 += Ipos * pts[m].angle_rad * pts[m].angle_rad; sumSig += pts[m].sigma; sumIzeta += Ipos * pts[m].zeta; } if (sumI > 0.0 && sumSig > 0.0) { const double phi = sumIphi / sumI; // Intensity-weighted rocking width (rad) - the observed spread of the reflection over the // frames it spans (mosaicity + beam-divergence + oscillation, sampled by phi_obs). const double rvar = std::max(0.0, sumIphi2 / sumI - phi * phi); rock_std.push_back(std::sqrt(rvar)); zeta_ev.push_back(sumIzeta / sumI); const Coord e = Astar * static_cast(pts[i].h) + Bstar * static_cast(pts[i].k) + Cstar * static_cast(pts[i].l); events.push_back(Event{phi, std::sqrt(sumI / sumSig), {e.x, e.y, e.z}, pts[i].h, pts[i].k, pts[i].l}); } } i = j; } logger.Info("Post-refine: {} multi-frame rocking events", events.size()); if (static_cast(events.size()) < settings.min_events) return result; // Task 3 signal: the observed rocking width (median over events) vs the frame oscillation. If the // rocking width materially exceeds the oscillation, mosaicity/beam-divergence is resolved and could // be refined from it; if it is ~ the oscillation, the reflections are barely rocking (under-sampled). if (!rock_std.empty()) { // Deconvolve the oscillation from the observed rocking width and convert the residual angular // spread to a mosaicity: sigma_mos_angular = sqrt(rock_std^2 - (osc/sqrt12)^2), mosaicity = it*zeta. const double osc_rad = axis.GetWedge_deg() * PI / 180.0; const double osc_var = (osc_rad / std::sqrt(12.0)) * (osc_rad / std::sqrt(12.0)); std::vector mos; mos.reserve(rock_std.size()); for (size_t e = 0; e < rock_std.size(); ++e) { const double sig_ang = std::sqrt(std::max(0.0, rock_std[e] * rock_std[e] - osc_var)); mos.push_back(sig_ang * zeta_ev[e] * 180.0 / PI); } std::nth_element(rock_std.begin(), rock_std.begin() + rock_std.size() / 2, rock_std.end()); const double med_rock_deg = rock_std[rock_std.size() / 2] * 180.0 / PI; std::nth_element(mos.begin(), mos.begin() + mos.size() / 2, mos.end()); result.est_mosaicity_deg = mos[mos.size() / 2]; logger.Info("Post-refine mosaicity signal: median rocking width {:.4f} deg (osc {:.4f}) -> " "estimated mosaicity {:.4f} deg", med_rock_deg, axis.GetWedge_deg(), result.est_mosaicity_deg); } constexpr size_t MAX_EVENTS = 20000; if (events.size() > MAX_EVENTS) { std::nth_element(events.begin(), events.begin() + MAX_EVENTS, events.end(), [](const Event &a, const Event &b) { return a.weight > b.weight; }); events.resize(MAX_EVENTS); } // ---- GEOMETRY REFINEMENT: the XtalOptimizer-equivalent, done as TWO SEPARATE // cross-validated steps rather than one joint fit (the same lesson as integration: refining the // profile width and the scale jointly fails, refining them separately works). Each step is committed // only if it lowers a HELD-OUT (deterministic split-half) residual - otherwise that part of the // geometry is left at nominal ("quit when things go wrong"): // Step A: cell scale + rotation axis from phi_obs (distance-independent excitation residual). // Step B: detector distance + beam centre from the observed spot positions, with the cell FIXED at // step A (so the positional residual is no longer degenerate with the cell scale). // Detector tilt is held fixed (gauge-coupled to orientation on a single crystal). ---- if (settings.refine_geometry) { const gemmi::CrystalSystem sys = (settings.crystal_system == gemmi::CrystalSystem::Trigonal) ? gemmi::CrystalSystem::Hexagonal : settings.crystal_system; const double ax0[3] = {ax.x, ax.y, ax.z}; const double lambda_l = lambda; const double rot3 = nominal_geom.GetPoniRot3_rad(); const double pixel_mm = nominal_geom.GetPixelSize_mm(); const double det_rot[2] = {nominal_geom.GetPoniRot1_rad(), nominal_geom.GetPoniRot2_rad()}; const UnitCell r0 = reference_latt.GetUnitCell(); // Deterministic split of the reflections into a fit half and a held-out half. Avalanche-mix the // hkl hash so the split bit is decorrelated from the LSB - a plain h+k+l parity collides with the // lattice centering condition (e.g. an I-centred lattice has h+k+l even for EVERY present // reflection, so a parity split would leave the validation half empty). auto is_val = [](int h, int k, int l) { unsigned u = static_cast(h) * 2654435761u + static_cast(k) * 2246822519u + static_cast(l) * 3266489917u; u ^= u >> 15; u *= 2246822519u; u ^= u >> 13; return (u & 1u) != 0u; }; enum Subset { FIT, VAL, ALL }; auto in = [&](int h, int k, int l, Subset s) { return s == ALL || (is_val(h, k, l) == (s == VAL)); }; // ===== Step A: cell scale s + rotation axis from phi_obs ===== auto excit_cost = [&](Subset s, double sc, const double axv[3]) { double c = 0.0; int n = 0; for (const auto &ev : events) { if (!in(ev.h, ev.k, ev.l, s)) continue; ScaleAxisExcitationResidual r(lambda_l, ev.phi_obs, 1.0, ev.e_ref); double sd = sc, av[3] = {axv[0], axv[1], axv[2]}, resid = 0.0; r(&sd, av, &resid); c += resid * resid; ++n; } return n ? c / n : 0.0; }; auto solve_scale_axis = [&](Subset s, double &s_out, double ax_out[3]) { double sc = 1.0, axv[3] = {ax0[0], ax0[1], ax0[2]}; ceres::Problem p; for (const auto &ev : events) { if (!in(ev.h, ev.k, ev.l, s)) continue; p.AddResidualBlock(new ceres::AutoDiffCostFunction( new ScaleAxisExcitationResidual(lambda_l, ev.phi_obs, settings.excitation_weight, ev.e_ref)), new ceres::CauchyLoss(0.02), &sc, axv); } p.SetParameterLowerBound(&sc, 0, 0.9); p.SetParameterUpperBound(&sc, 0, 1.1); for (int j = 0; j < 3; ++j) { p.SetParameterLowerBound(axv, j, ax0[j] - 0.05); p.SetParameterUpperBound(axv, j, ax0[j] + 0.05); } ceres::Solver::Options o; o.linear_solver_type = ceres::DENSE_QR; o.max_num_iterations = 50; o.num_threads = std::max(1, settings.num_threads); o.logging_type = ceres::LoggingType::SILENT; ceres::Solver::Summary sum; ceres::Solve(o, &p, &sum); s_out = sc; ax_out[0] = axv[0]; ax_out[1] = axv[1]; ax_out[2] = axv[2]; return sum.IsSolutionUsable(); }; double s_fit = 1.0, ax_fit[3]; const bool convA = solve_scale_axis(FIT, s_fit, ax_fit); const double cvA_nom = excit_cost(VAL, 1.0, ax0); const double cvA_ref = excit_cost(VAL, s_fit, ax_fit); // Commit the cell scale only for a small, credible move: a well-calibrated header needs < ~0.6 %, // so a > 1 % scale is a red flag (on multi-lattice / noisy data the excitation fit is biased the // same way in every cross-validation fold, so the relative-improvement gate cannot catch it). result.cell_refined = convA && cvA_ref < 0.98 * cvA_nom && std::fabs(s_fit - 1.0) < 0.01; double s = 1.0, axv[3] = {ax0[0], ax0[1], ax0[2]}; if (result.cell_refined) solve_scale_axis(ALL, s, axv); // commit: re-fit on all data const double axlen = std::sqrt(axv[0]*axv[0] + axv[1]*axv[1] + axv[2]*axv[2]); const double axdev = std::acos(std::clamp((axv[0]*ax0[0]+axv[1]*ax0[1]+axv[2]*ax0[2]) / std::max(1e-9, axlen), -1.0, 1.0)) * 180.0 / PI; logger.Info("Post-refine GEOM step A (cell/axis): s = {:.5f}, rot-axis {:.3f} deg, held-out excit " "{:.3e} -> {:.3e} => {}", s, axdev, cvA_nom, cvA_ref, result.cell_refined ? "COMMIT" : "reject (kept nominal cell)"); // Cell (scale s, shape fixed) as the XtalResidual parameter blocks p0/p1/p2, held CONSTANT in step B. double p0[3] = {0, 0, 0}, p1[3] = {0, 0, 0}, p2[3] = {0, 0, 0}; double beta = r0.beta; switch (sys) { case gemmi::CrystalSystem::Tetragonal: LatticeToRodriguesAndLengths_GS(reference_latt, p0, p1); p1[0] = (p1[0] + p1[1]) / 2.0; break; case gemmi::CrystalSystem::Cubic: LatticeToRodriguesAndLengths_GS(reference_latt, p0, p1); p1[0] = (p1[0] + p1[1] + p1[2]) / 3.0; break; case gemmi::CrystalSystem::Hexagonal: LatticeToRodriguesAndLengths_Hex(reference_latt, p0, p1); break; case gemmi::CrystalSystem::Monoclinic: LatticeToRodriguesLengthsBeta_Mono(reference_latt, p0, p1, beta); p2[0] = beta; break; case gemmi::CrystalSystem::Orthorhombic: LatticeToRodriguesAndLengths_GS(reference_latt, p0, p1); break; default: LatticeToRodriguesAndLengths_GS(reference_latt, p0, p1); p2[0] = r0.alpha * PI / 180.0; p2[1] = r0.beta * PI / 180.0; p2[2] = r0.gamma * PI / 180.0; break; } for (int j = 0; j < 3; ++j) p1[j] *= s; // apply the committed cell scale double rot_vec[3] = {axv[0], axv[1], axv[2]}; // committed (or nominal) axis // ===== Step B: detector distance + beam from the observed positions, cell fixed ===== std::vector obs; for (const auto &pp : pts) if (std::isfinite(pp.obs_x) && std::isfinite(pp.obs_y)) obs.push_back(&pp); constexpr size_t MAX_OBS = 20000; if (obs.size() > MAX_OBS) { std::nth_element(obs.begin(), obs.begin() + MAX_OBS, obs.end(), [](const Partial *a, const Partial *b) { return a->I / std::max(1e-9, a->sigma) > b->I / std::max(1e-9, b->sigma); }); obs.resize(MAX_OBS); } result.obs_used = static_cast(obs.size()); const double beam_x0 = nominal_geom.GetBeamX_pxl(), beam_y0 = nominal_geom.GetBeamY_pxl(); const double dist0 = nominal_geom.GetDetectorDistance_mm(); auto pos_cost = [&](Subset s, const double beam[2], const double dist[1]) { double c = 0.0; int n = 0; for (const Partial *pp : obs) { if (!in(pp->h, pp->k, pp->l, s)) continue; XtalResidual r(pp->obs_x, pp->obs_y, lambda_l, pixel_mm, rot3, pp->angle_rad, pp->h, pp->k, pp->l, sys); double resid[3] = {0, 0, 0}; r(beam, dist, det_rot, rot_vec, p0, p1, p2, resid); c += resid[0]*resid[0] + resid[1]*resid[1] + resid[2]*resid[2]; ++n; } return n ? c / n : 0.0; }; auto solve_detector = [&](Subset s, double beam_out[2], double &dist_out) { double beam[2] = {beam_x0, beam_y0}, dist[1] = {dist0}; ceres::Problem p; for (const Partial *pp : obs) { if (!in(pp->h, pp->k, pp->l, s)) continue; p.AddResidualBlock(new ceres::AutoDiffCostFunction( new XtalResidual(pp->obs_x, pp->obs_y, lambda_l, pixel_mm, rot3, pp->angle_rad, pp->h, pp->k, pp->l, sys)), new ceres::CauchyLoss(0.02), beam, dist, const_cast(det_rot), rot_vec, p0, p1, p2); p.SetParameterBlockConstant(const_cast(det_rot)); p.SetParameterBlockConstant(rot_vec); p.SetParameterBlockConstant(p0); p.SetParameterBlockConstant(p1); p.SetParameterBlockConstant(p2); } if (p.NumResidualBlocks() == 0) { beam_out[0] = beam_x0; beam_out[1] = beam_y0; dist_out = dist0; return false; } p.SetParameterLowerBound(dist, 0, dist0 * 0.95); p.SetParameterUpperBound(dist, 0, dist0 * 1.05); for (int j = 0; j < 2; ++j) { p.SetParameterLowerBound(beam, j, beam[j] - 15.0); p.SetParameterUpperBound(beam, j, beam[j] + 15.0); } ceres::Solver::Options o; o.linear_solver_type = ceres::DENSE_QR; o.max_num_iterations = 60; o.num_threads = std::max(1, settings.num_threads); o.logging_type = ceres::LoggingType::SILENT; ceres::Solver::Summary sum; ceres::Solve(o, &p, &sum); beam_out[0] = beam[0]; beam_out[1] = beam[1]; dist_out = dist[0]; return sum.IsSolutionUsable(); }; double beam[2] = {beam_x0, beam_y0}, dist = dist0; if (obs.size() >= static_cast(settings.min_events)) { double beam_fit[2], dist_fit; const bool convB = solve_detector(FIT, beam_fit, dist_fit); const double b_nom[2] = {beam_x0, beam_y0}, d_nom[1] = {dist0}; const double b_ref[2] = {beam_fit[0], beam_fit[1]}, d_ref[1] = {dist_fit}; const double cvB_nom = pos_cost(VAL, b_nom, d_nom); const double cvB_ref = pos_cost(VAL, b_ref, d_ref); // Commit the detector geometry only for a small, credible move: distance < 1 % (a calibrated // header needs < ~0.6 %). A larger move is the red flag for an unreliable fit - typically a // second lattice whose spots bias every cross-validation fold identically, so the relative // "it improved" gate is blind to it and pulls a spurious distance<->cell pair (the radial // degeneracy) far off. The absolute size of the move discriminates a genuine header correction // from that failure far better than the absolute residual, which real marginal (noisy / iced) // data shares with the multi-lattice case. const bool in_bounds = std::fabs(dist_fit - dist0) < 0.01 * dist0 && std::hypot(beam_fit[0] - beam_x0, beam_fit[1] - beam_y0) < 15.0; result.detector_refined = convB && cvB_ref < 0.98 * cvB_nom && in_bounds; if (result.detector_refined) { double bo[2]; solve_detector(ALL, bo, dist); beam[0] = bo[0]; beam[1] = bo[1]; } logger.Info("Post-refine GEOM step B (distance/beam): dist {:.3f} -> {:.3f} mm, beam " "({:.2f},{:.2f}) -> ({:.2f},{:.2f}), held-out pos {:.3e} -> {:.3e} => {}", dist0, result.detector_refined ? dist : dist0, beam_x0, beam_y0, result.detector_refined ? beam[0] : beam_x0, result.detector_refined ? beam[1] : beam_y0, cvB_nom, cvB_ref, result.detector_refined ? "COMMIT" : "reject (kept nominal detector)"); } else { logger.Info("Post-refine GEOM step B: only {} positional observations - skipped", obs.size()); } // Assemble the committed geometry. UnitCell cellA = r0; if (result.cell_refined) { cellA.a = static_cast(r0.a * s); cellA.b = static_cast(r0.b * s); cellA.c = static_cast(r0.c * s); } result.cell = cellA; result.distance_after_mm = dist; result.beam_x_before_px = beam_x0; result.beam_x_after_px = beam[0]; result.beam_y_before_px = beam_y0; result.beam_y_after_px = beam[1]; result.events_used = static_cast(events.size()); result.ok = result.cell_refined || result.detector_refined; if (!result.ok) logger.Info("Post-refine GEOM: neither step passed cross-validation - geometry left at nominal"); return result; } return result; // refine_geometry is the only supported mode; nothing refined otherwise } catch (...) { result.ok = false; return result; } }