The crystal refinement (XtalOptimizer, both the seven-block and the reduced beam+orientation form, and XtalOptimizerRotationOnly) no longer builds a ceres::Problem. XtalRefine holds the problem as data and solves it with LMSolver, which follows Ceres' trust-region LM step for step - Jacobi scaling, damping and radius updates, stopping rules, box projection, the projected Armijo line search with cubic interpolation on bounded problems, the SphereManifold for the spindle - but takes J^T J and J^T r directly instead of a Jacobian. The residual is the same XtalResidual code, now Ceres-free and evaluated on a forward-mode Dual (Dual.h); everything that depends on parameters alone (detector-angle trig, per-frame back-rotation, reciprocal basis, orientation rotation) is worked out once per evaluation, and the observed and predicted halves carry 6 and 9 derivative lanes rather than 16. The sums are cut into blocks that depend on the residual count alone, so the answer does not depend on the thread count. Because the line-search trial point is the candidate point, a bounded iteration costs one evaluation instead of Ceres' three. Validation (rc174 + this, -march=x86-64-v3): - p.mtz md5 identical to the Ceres build on myob/cytc/thau x10sa, GPU and CPU builds, and on the lyso8 stills reference. - Solve corpus (every 16-parameter solve and every 10th per-image solve of the three sets, 8.3k problems, inputs and Ceres results dumped from a run that reproduced the md5s): usable/failed agree on all, iteration counts identical on all, parameters agree to <2e-11 (in px / rad / 0.01 A units), costs to 1e-13. - Same process, same threads: 7-9x faster per solve than Ceres. - In-run (GPU, loaded box): xtal 16-parameter solves myob 22.2 -> 6.4 core-s, cytc 88 -> 26 core-s; per-image solves 5.4 -> 1.2 core-s (myob); cytc first pass indexing windows 1.9 -> 0.85 s, myob 1.1 -> 0.45 s; solver share of the whole cytc run 17% -> 4% of CPU samples. New tests compare the solver with Ceres on synthetic rotation problems (full/weighted/reduced) and check thread-count independence. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01K5K8jvPPbmCrbqnWkddTuB
128 lines
5.5 KiB
C++
128 lines
5.5 KiB
C++
// SPDX-FileCopyrightText: 2026 Filip Leonarski, Paul Scherrer Institute <filip.leonarski@psi.ch>
|
|
// SPDX-License-Identifier: GPL-3.0-only
|
|
|
|
// Ceres first: its logging header defines a CHECK macro of its own, which Catch's must replace here.
|
|
#include "XtalRefineCeres.h"
|
|
#undef CHECK
|
|
#include <catch2/catch_all.hpp>
|
|
|
|
#include "../image_analysis/geom_refinement/LatticeReduction.h"
|
|
#include "../image_analysis/bragg_prediction/BraggPrediction.h"
|
|
|
|
// The LM solver of XtalRefine is meant to take Ceres' path to Ceres' answer: same steps accepted, same
|
|
// stopping rule, same point. These cases hand one problem to both and compare.
|
|
namespace {
|
|
// A monoclinic crystal rotated about X over ten 3-degree frames, its predicted spots as observations;
|
|
// the refinement starts from a perturbed beam, tilt and cell.
|
|
XtalRefineProblem RotationProblem(bool reduced, bool weighted) {
|
|
DiffractionExperiment exp;
|
|
exp.IncidentEnergy_keV(WVL_1A_IN_KEV).BeamX_pxl(1000).BeamY_pxl(1000)
|
|
.PoniRot1_rad(0.01).PoniRot2_rad(0.02).DetectorDistance_mm(200);
|
|
const auto geom = exp.GetDiffractionGeometry();
|
|
const CrystalLattice latt(40, 50, 80, 90, 95, 90);
|
|
const GoniometerAxis axis("omega", 0.0f, 3.0f, Coord(1, 0, 0), std::nullopt);
|
|
const gemmi::CrystalSystem sys = gemmi::CrystalSystem::Monoclinic;
|
|
|
|
XtalRefineProblem p;
|
|
p.crystal_system = sys;
|
|
p.beam_and_orientation_only = reduced;
|
|
p.distance_mm = 200;
|
|
|
|
BraggPrediction prediction;
|
|
const BraggPredictionSettings settings{.high_res_A = 1.5, .ewald_dist_cutoff = 0.002};
|
|
for (int img = 0; img < 10; img++) {
|
|
const float angle_deg = axis.GetAngle_deg(img) + axis.GetWedge_deg() / 2.0f;
|
|
const auto n = prediction.Calc(exp, latt.Multiply(axis.GetTransformationAngle(angle_deg).transpose()),
|
|
settings);
|
|
p.frame_angle_rad.push_back(angle_deg * PI / 180.0);
|
|
for (int i = 0; i < n; i++) {
|
|
const auto &r = prediction.GetReflections().at(i);
|
|
p.residuals.emplace_back(r.predicted_x + 0.3 * std::sin(i), r.predicted_y + 0.3 * std::cos(i),
|
|
geom.GetWavelength_A(), geom.GetPixelSize_mm(), 1.0, 0.0,
|
|
angle_deg * PI / 180.0, r.h, r.k, r.l, sys);
|
|
p.frame.push_back(img);
|
|
if (weighted)
|
|
p.weight_sq.push_back(0.2 + 0.6 * (i % 5) / 4.0);
|
|
}
|
|
}
|
|
|
|
p.beam[0] = 1000.0;
|
|
p.beam[1] = 997.0;
|
|
p.detector_rot[0] = 0.012;
|
|
p.detector_rot[1] = 0.018;
|
|
p.rot_vec[0] = 1.0;
|
|
p.rot_vec[1] = 0.0;
|
|
p.rot_vec[2] = 0.0;
|
|
double beta = 0;
|
|
LatticeToRodriguesLengthsBeta_Mono(CrystalLattice(39.7f, 50.6f, 79.6f, 90.0f, 94.5f, 90.0f),
|
|
p.latt_vec0, p.latt_vec1, beta);
|
|
p.latt_vec2[0] = beta;
|
|
|
|
if (!reduced) {
|
|
p.detector_rot_constant = false;
|
|
p.rot_vec_constant = false;
|
|
p.latt_vec1_constant = false;
|
|
p.latt_vec2_constant = false;
|
|
for (int i = 0; i < 2; i++) {
|
|
p.detector_rot_lower[i] = p.detector_rot[i] - 0.05;
|
|
p.detector_rot_upper[i] = p.detector_rot[i] + 0.05;
|
|
}
|
|
for (int i = 0; i < 3; i++) {
|
|
p.latt_vec1_lower[i] = 5.0;
|
|
p.latt_vec1_upper[i] = 100.0;
|
|
}
|
|
p.latt_vec2_lower[0] = PI / 3;
|
|
p.latt_vec2_upper[0] = 2 * PI / 3;
|
|
p.priors.push_back({XtalRefinePrior::Block::Beam, 1.0, 0.0, p.beam[0], 0.5});
|
|
p.priors.push_back({XtalRefinePrior::Block::DetectorRot, 0.0, 1.0, p.detector_rot[1], 50.0});
|
|
}
|
|
p.options.max_iterations = 50;
|
|
return p;
|
|
}
|
|
|
|
void CompareWithCeres(XtalRefineProblem p) {
|
|
XtalRefineProblem q = p;
|
|
const LMSummary lm = SolveXtalRefine(p, 4);
|
|
const ceres::Solver::Summary ref = SolveXtalRefineCeres(q, 4);
|
|
|
|
REQUIRE(lm.IsSolutionUsable() == ref.IsSolutionUsable());
|
|
CHECK(lm.iterations == static_cast<int>(ref.iterations.size()));
|
|
CHECK(lm.final_cost == Catch::Approx(ref.final_cost).epsilon(1e-9));
|
|
const auto same = [](const double *a, const double *b, int n, double tol) {
|
|
for (int i = 0; i < n; i++)
|
|
CHECK(a[i] == Catch::Approx(b[i]).margin(tol));
|
|
};
|
|
same(p.beam, q.beam, 2, 1e-7);
|
|
same(p.detector_rot, q.detector_rot, 2, 1e-10);
|
|
same(p.rot_vec, q.rot_vec, 3, 1e-10);
|
|
same(p.latt_vec0, q.latt_vec0, 3, 1e-10);
|
|
same(p.latt_vec1, q.latt_vec1, 3, 1e-8);
|
|
same(p.latt_vec2, q.latt_vec2, 3, 1e-10);
|
|
}
|
|
}
|
|
|
|
TEST_CASE("XtalRefine_matches_Ceres_full", "[XtalOptimizer]") {
|
|
CompareWithCeres(RotationProblem(false, false));
|
|
}
|
|
|
|
TEST_CASE("XtalRefine_matches_Ceres_full_weighted", "[XtalOptimizer]") {
|
|
CompareWithCeres(RotationProblem(false, true));
|
|
}
|
|
|
|
TEST_CASE("XtalRefine_matches_Ceres_beam_orientation", "[XtalOptimizer]") {
|
|
CompareWithCeres(RotationProblem(true, false));
|
|
}
|
|
|
|
TEST_CASE("XtalRefine_same_answer_at_any_thread_count", "[XtalOptimizer]") {
|
|
XtalRefineProblem a = RotationProblem(false, false);
|
|
XtalRefineProblem b = a;
|
|
SolveXtalRefine(a, 1);
|
|
SolveXtalRefine(b, 7);
|
|
for (int i = 0; i < 3; i++) {
|
|
CHECK(a.latt_vec0[i] == b.latt_vec0[i]);
|
|
CHECK(a.latt_vec1[i] == b.latt_vec1[i]);
|
|
}
|
|
CHECK(a.beam[0] == b.beam[0]);
|
|
CHECK(a.beam[1] == b.beam[1]);
|
|
}
|