midas-diffract 0.7.0__tar.gz → 0.8.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/PKG-INFO +1 -1
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/midas_diffract/__init__.py +3 -2
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/midas_diffract/forward.py +105 -70
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/midas_diffract/optimize.py +84 -21
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/midas_diffract.egg-info/PKG-INFO +1 -1
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/midas_diffract.egg-info/SOURCES.txt +2 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/pyproject.toml +1 -1
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/tests/test_distortion_layer.py +1 -1
- midas_diffract-0.8.0/tests/test_omega_solver_singularities.py +205 -0
- midas_diffract-0.8.0/tests/test_optimize_no_match.py +92 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/LICENSE +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/README.md +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/midas_diffract/hkls.py +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/midas_diffract/losses.py +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/midas_diffract/simulate_panel_zarrs.py +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/midas_diffract.egg-info/dependency_links.txt +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/midas_diffract.egg-info/requires.txt +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/midas_diffract.egg-info/top_level.txt +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/setup.cfg +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/tests/test_c_comparison.py +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/tests/test_forward.py +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/tests/test_hkls.py +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/tests/test_losses.py +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/tests/test_multi_detector.py +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/tests/test_omega_box_filter.py +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/tests/test_strain_tensor.py +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/tests/test_tilts.py +0 -0
- {midas_diffract-0.7.0 → midas_diffract-0.8.0}/tests/test_wedge.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: midas-diffract
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.8.0
|
|
4
4
|
Summary: End-to-end differentiable forward model for High-Energy Diffraction Microscopy (FF, NF, pf-HEDM)
|
|
5
5
|
Author-email: Hemant Sharma <hsharma@anl.gov>
|
|
6
6
|
License-Expression: BSD-3-Clause
|
|
@@ -28,7 +28,7 @@ Quick start
|
|
|
28
28
|
loss.backward()
|
|
29
29
|
"""
|
|
30
30
|
|
|
31
|
-
__version__ = "0.
|
|
31
|
+
__version__ = "0.8.0"
|
|
32
32
|
|
|
33
33
|
from .forward import (
|
|
34
34
|
HEDMForwardModel,
|
|
@@ -39,7 +39,7 @@ from .forward import (
|
|
|
39
39
|
)
|
|
40
40
|
from .hkls import hkls_for_forward_model
|
|
41
41
|
from .losses import SpotMatchingLoss
|
|
42
|
-
from .optimize import optimize_single_grain, evaluate_recovery
|
|
42
|
+
from .optimize import optimize_single_grain, evaluate_recovery, NoMatchError
|
|
43
43
|
|
|
44
44
|
__all__ = [
|
|
45
45
|
"HEDMForwardModel",
|
|
@@ -50,6 +50,7 @@ __all__ = [
|
|
|
50
50
|
"SpotMatchingLoss",
|
|
51
51
|
"hkls_for_forward_model",
|
|
52
52
|
"optimize_single_grain",
|
|
53
|
+
"NoMatchError",
|
|
53
54
|
"evaluate_recovery",
|
|
54
55
|
"__version__",
|
|
55
56
|
]
|
|
@@ -224,6 +224,90 @@ class SpotDescriptors:
|
|
|
224
224
|
det_id: Optional[torch.Tensor] = None # (..., K, M) int64 panel index (panel mode only)
|
|
225
225
|
|
|
226
226
|
|
|
227
|
+
# ---------------------------------------------------------------------------
|
|
228
|
+
# Omega solver
|
|
229
|
+
# ---------------------------------------------------------------------------
|
|
230
|
+
|
|
231
|
+
def solve_omega(Gx: torch.Tensor, Gy: torch.Tensor, v: torch.Tensor):
|
|
232
|
+
"""Solve ``-Gx cos(w) + Gy sin(w) = v`` for both omega branches.
|
|
233
|
+
|
|
234
|
+
Returns ``(omega_p, omega_n, valid)`` with omega in radians on ``(-pi, pi]``
|
|
235
|
+
and ``valid`` True where a solution exists.
|
|
236
|
+
|
|
237
|
+
Closed form, no quadratic::
|
|
238
|
+
|
|
239
|
+
-Gx cos w + Gy sin w == rho cos(w - phi),
|
|
240
|
+
rho = |(Gx, Gy)|, phi = atan2(Gy, -Gx)
|
|
241
|
+
|
|
242
|
+
so the constraint is ``rho cos(w - phi) = v`` and the solutions are
|
|
243
|
+
``w = phi +- acos(v / rho)``.
|
|
244
|
+
|
|
245
|
+
This replaces a Gy^2-divided quadratic together with everything that route
|
|
246
|
+
dragged in: ``+ 1e-7`` in each denominator, a discriminant, a clamped
|
|
247
|
+
``safe_arccos``, a +-w branch chosen by comparing residuals, and a separate
|
|
248
|
+
``|Gy| < 1e-12`` special case. All were consequences of solving for cos w
|
|
249
|
+
first rather than for w.
|
|
250
|
+
|
|
251
|
+
Why it is not merely tidier. MEASURED over 200k random coefficient sets in
|
|
252
|
+
float64, as the residual of the defining equation above:
|
|
253
|
+
|
|
254
|
+
=========================== ==================== =================
|
|
255
|
+
regime old (eps = 1e-7) this form
|
|
256
|
+
=========================== ==================== =================
|
|
257
|
+
ordinary ``|Gy| ~ 0.35`` median 1.2e-07 median 2.8e-17
|
|
258
|
+
small ``|Gy| <= 3e-4`` median 1.7e-04 median 1.4e-17
|
|
259
|
+
tiny ``|Gy| <= 1e-8`` median 3.0e-04 median 1.4e-17
|
|
260
|
+
``Gy`` exactly 0 median 3.0e-04 median 1.4e-17
|
|
261
|
+
=========================== ==================== =================
|
|
262
|
+
|
|
263
|
+
and, recovering a KNOWN omega within 0.02 deg of zero -- the band the old
|
|
264
|
+
``safe_arccos`` clamp froze -- median error 1.56e-02 deg before, 7.2e-15
|
|
265
|
+
deg now. Note the epsilon was never confined to the small-|Gy| spots the
|
|
266
|
+
audit identified: ``y2 + 1e-7`` perturbs every coefficient, so essentially
|
|
267
|
+
every spot carried a small bias.
|
|
268
|
+
|
|
269
|
+
The one remaining singularity is ``acos`` at ``|v/rho| = 1``: exact
|
|
270
|
+
tangency of the Ewald sphere, a genuinely grazing reflection and genuinely
|
|
271
|
+
non-differentiable. That is where a singularity belongs. The old form
|
|
272
|
+
additionally put one at ``omega ~ 0``, which is an ordinary spot.
|
|
273
|
+
|
|
274
|
+
Verified against the C reference (``ForwardSimulationCompressed`` and
|
|
275
|
+
``simulateNF``) via ``tests/test_c_comparison.py``.
|
|
276
|
+
"""
|
|
277
|
+
rho = torch.sqrt(Gx * Gx + Gy * Gy + 1e-300)
|
|
278
|
+
phi = torch.atan2(Gy, -Gx)
|
|
279
|
+
|
|
280
|
+
# A solution exists iff |v| <= rho -- exactly the old condition: the old
|
|
281
|
+
# discriminant was 4 Gy^2 (rho^2 - v^2), so ``disc >= 0`` meant
|
|
282
|
+
# ``|v| <= rho`` for Gy != 0, and the Gy ~ 0 branch's ``|-v/Gx| <= 1`` is
|
|
283
|
+
# the same statement when rho = |Gx|.
|
|
284
|
+
cos_off = v / rho
|
|
285
|
+
valid = torch.abs(cos_off) <= 1.0
|
|
286
|
+
|
|
287
|
+
# acos'(+-1) is infinite. Guard the INPUT: torch sums per-spot gradients,
|
|
288
|
+
# so one tangential spot would otherwise NaN the whole batch. The boundary
|
|
289
|
+
# value (0 or pi) is restored exactly by the outer where.
|
|
290
|
+
at_edge = torch.abs(cos_off) >= 1.0
|
|
291
|
+
cos_safe = torch.where(at_edge, torch.zeros_like(cos_off),
|
|
292
|
+
cos_off.clamp(-1.0, 1.0))
|
|
293
|
+
d_off = torch.acos(cos_safe)
|
|
294
|
+
d_off = torch.where(
|
|
295
|
+
at_edge,
|
|
296
|
+
torch.where(cos_off > 0, torch.zeros_like(d_off),
|
|
297
|
+
torch.full_like(d_off, math.pi)),
|
|
298
|
+
d_off,
|
|
299
|
+
)
|
|
300
|
+
|
|
301
|
+
zero_w = torch.zeros_like(phi)
|
|
302
|
+
omega_p = torch.where(valid, phi + d_off, zero_w)
|
|
303
|
+
omega_n = torch.where(valid, phi - d_off, zero_w)
|
|
304
|
+
# Wrap into (-pi, pi], the range the old acos-based form produced.
|
|
305
|
+
two_pi = 2.0 * math.pi
|
|
306
|
+
omega_p = omega_p - two_pi * torch.round(omega_p / two_pi)
|
|
307
|
+
omega_n = omega_n - two_pi * torch.round(omega_n / two_pi)
|
|
308
|
+
return omega_p, omega_n, valid
|
|
309
|
+
|
|
310
|
+
|
|
227
311
|
# ---------------------------------------------------------------------------
|
|
228
312
|
# Forward model
|
|
229
313
|
# ---------------------------------------------------------------------------
|
|
@@ -570,8 +654,21 @@ class HEDMForwardModel(nn.Module):
|
|
|
570
654
|
# ------------------------------------------------------------------
|
|
571
655
|
|
|
572
656
|
def safe_arccos(self, x: torch.Tensor) -> torch.Tensor:
|
|
573
|
-
"""Numerically stable arccos: clamp to [-1+eps, 1-eps].
|
|
574
|
-
|
|
657
|
+
"""Numerically stable arccos: clamp to [-1+eps, 1-eps].
|
|
658
|
+
|
|
659
|
+
The clamp sits at the float64 boundary, NOT at ``self.epsilon``.
|
|
660
|
+
``self.epsilon`` is 1e-7, and ``acos(1 - 1e-7) = 4.47e-4 rad``, so
|
|
661
|
+
clamping there froze both the value and the gradient in a +-0.0256 deg
|
|
662
|
+
band -- 0.029% of valid spots pinned to the band edge. At 1e-15 the
|
|
663
|
+
residual band is 2.6e-6 deg, four orders smaller and well below any
|
|
664
|
+
angular resolution this model is used at.
|
|
665
|
+
|
|
666
|
+
Where a sine is available, prefer ``atan2``: it has no band at all.
|
|
667
|
+
The omega solver does exactly that; this is kept for the ``Gy ~ 0``
|
|
668
|
+
branch and for eta, where the constraint supplies no sine.
|
|
669
|
+
"""
|
|
670
|
+
eps = 1e-15 if x.dtype == torch.float64 else 1e-7
|
|
671
|
+
return torch.acos(torch.clamp(x, -1.0 + eps, 1.0 - eps))
|
|
575
672
|
|
|
576
673
|
# ------------------------------------------------------------------
|
|
577
674
|
# strain_as_voigt (accept full 3x3 tensor OR plain-Voigt 6-vector)
|
|
@@ -885,67 +982,9 @@ class HEDMForwardModel(nn.Module):
|
|
|
885
982
|
v = v_no_wedge + sin_W * Gz_p
|
|
886
983
|
# ---------------------------------------------------------------
|
|
887
984
|
|
|
888
|
-
#
|
|
889
|
-
|
|
890
|
-
|
|
891
|
-
# C uses almostzero=1e-12 for the Gy≈0 branch (see
|
|
892
|
-
# NF_HEDM/src/CalcDiffractionSpots.c:96 and
|
|
893
|
-
# FF_HEDM/src/ForwardSimulationCompressed.c:168). Match exactly.
|
|
894
|
-
almostzero = 1e-12
|
|
895
|
-
x2 = Gx * Gx
|
|
896
|
-
y2 = Gy * Gy
|
|
897
|
-
a = 1.0 + x2 / (y2 + self.epsilon)
|
|
898
|
-
b_coeff = 2.0 * v * Gx / (y2 + self.epsilon)
|
|
899
|
-
c_coeff = v * v / (y2 + self.epsilon) - 1.0
|
|
900
|
-
discriminant = b_coeff * b_coeff - 4.0 * a * c_coeff
|
|
901
|
-
|
|
902
|
-
sqrt_disc = torch.sqrt(torch.abs(discriminant))
|
|
903
|
-
|
|
904
|
-
coswp = (-b_coeff + sqrt_disc) / (2.0 * a)
|
|
905
|
-
coswn = (-b_coeff - sqrt_disc) / (2.0 * a)
|
|
906
|
-
|
|
907
|
-
wap = self.safe_arccos(coswp)
|
|
908
|
-
wan = self.safe_arccos(coswn)
|
|
909
|
-
wbp = -wap
|
|
910
|
-
wbn = -wan
|
|
911
|
-
|
|
912
|
-
# Select correct branch: the one satisfying -Gx*cos(w)+Gy*sin(w)=v
|
|
913
|
-
eqap = -Gx * torch.cos(wap) + Gy * torch.sin(wap)
|
|
914
|
-
eqbp = -Gx * torch.cos(wbp) + Gy * torch.sin(wbp)
|
|
915
|
-
eqan = -Gx * torch.cos(wan) + Gy * torch.sin(wan)
|
|
916
|
-
eqbn = -Gx * torch.cos(wbn) + Gy * torch.sin(wbn)
|
|
917
|
-
|
|
918
|
-
Dap = torch.abs(eqap - v)
|
|
919
|
-
Dbp = torch.abs(eqbp - v)
|
|
920
|
-
Dan = torch.abs(eqan - v)
|
|
921
|
-
Dbn = torch.abs(eqbn - v)
|
|
922
|
-
|
|
923
|
-
all_wp = torch.where(Dap < Dbp, wap, wbp)
|
|
924
|
-
all_wn = torch.where(Dan < Dbn, wan, wbn)
|
|
925
|
-
|
|
926
|
-
# Special case: Gy ~ 0 (C uses almostzero=1e-12)
|
|
927
|
-
# C code (CalcDiffractionSpots.c:97-106):
|
|
928
|
-
# cosome1 = -v / x;
|
|
929
|
-
# if (|cosome1| <= 1) { ome = acos(cosome1); solutions: +ome, -ome }
|
|
930
|
-
gy_zero = torch.abs(Gy) < almostzero
|
|
931
|
-
cosome_special = -v / (Gx + self.epsilon)
|
|
932
|
-
cosome_special_valid = (torch.abs(cosome_special) <= 1.0) & gy_zero & (torch.abs(Gx) > self.epsilon)
|
|
933
|
-
special_w = self.safe_arccos(cosome_special) # positive omega solution
|
|
934
|
-
# Two solutions: +ome and -ome
|
|
935
|
-
special_wp = special_w # positive
|
|
936
|
-
special_wn = -special_w # negative
|
|
937
|
-
|
|
938
|
-
# When |Gy| < almostzero, use the special case; otherwise use the quadratic
|
|
939
|
-
disc_valid = (discriminant >= 0) & (~gy_zero)
|
|
940
|
-
coswp_valid = (coswp >= -1.0) & (coswp <= 1.0)
|
|
941
|
-
coswn_valid = (coswn >= -1.0) & (coswn <= 1.0)
|
|
942
|
-
|
|
943
|
-
omega_p = torch.where(cosome_special_valid, special_wp,
|
|
944
|
-
torch.where(disc_valid & coswp_valid, all_wp,
|
|
945
|
-
torch.zeros_like(all_wp)))
|
|
946
|
-
omega_n = torch.where(cosome_special_valid, special_wn,
|
|
947
|
-
torch.where(disc_valid & coswn_valid, all_wn,
|
|
948
|
-
torch.zeros_like(all_wn)))
|
|
985
|
+
# Omega solver -- see module-level ``solve_omega``.
|
|
986
|
+
omega_p, omega_n, sol_valid = solve_omega(Gx, Gy, v)
|
|
987
|
+
|
|
949
988
|
|
|
950
989
|
# Concatenate two solutions: (..., 2N, M)
|
|
951
990
|
all_omega = torch.cat([omega_p, omega_n], dim=-2)
|
|
@@ -983,13 +1022,9 @@ class HEDMForwardModel(nn.Module):
|
|
|
983
1022
|
tt = tt.expand_as(omega_p)
|
|
984
1023
|
two_theta = torch.cat([tt, tt], dim=-2)
|
|
985
1024
|
|
|
986
|
-
# Validity mask
|
|
987
|
-
|
|
988
|
-
|
|
989
|
-
# For gy_zero special case, valid only if cosome is in [-1, 1]
|
|
990
|
-
valid_p = valid_p | cosome_special_valid
|
|
991
|
-
valid_n = valid_n | cosome_special_valid
|
|
992
|
-
valid = torch.cat([valid_p, valid_n], dim=-2).float()
|
|
1025
|
+
# Validity mask. Both omega solutions exist under the same condition,
|
|
1026
|
+
# |v| <= rho -- the closed form has no separate Gy ~ 0 case to union in.
|
|
1027
|
+
valid = torch.cat([sol_valid, sol_valid], dim=-2).float()
|
|
993
1028
|
|
|
994
1029
|
# Eta bounds
|
|
995
1030
|
eta_ok = (torch.abs(eta) >= self.min_eta) & \
|
|
@@ -36,6 +36,19 @@ DEG2RAD = math.pi / 180.0
|
|
|
36
36
|
RAD2DEG = 180.0 / math.pi
|
|
37
37
|
|
|
38
38
|
|
|
39
|
+
class NoMatchError(RuntimeError):
|
|
40
|
+
"""Too few observed spots matched for the optimisation to be defined.
|
|
41
|
+
|
|
42
|
+
Raised instead of returning a sentinel loss: with no matched spots there is
|
|
43
|
+
no residual, hence no gradient and no information, so any number returned
|
|
44
|
+
would be a fiction that a caller could compare against a real fit.
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
def __init__(self, message: str, *, n_matched: int):
|
|
48
|
+
super().__init__(message)
|
|
49
|
+
self.n_matched = n_matched
|
|
50
|
+
|
|
51
|
+
|
|
39
52
|
def _associate(pred_valid: torch.Tensor, observed: torch.Tensor,
|
|
40
53
|
max_dist: float) -> "tuple[torch.Tensor, torch.Tensor]":
|
|
41
54
|
"""Nearest-neighbour observed->predicted association.
|
|
@@ -97,9 +110,10 @@ def optimize_single_grain(
|
|
|
97
110
|
Discard observed spots whose nearest predicted neighbour is
|
|
98
111
|
farther than this in the angular metric (radians).
|
|
99
112
|
min_matches : int
|
|
100
|
-
If fewer than this many spots are matched at any iteration,
|
|
101
|
-
|
|
102
|
-
|
|
113
|
+
If fewer than this many spots are matched at any iteration, the fit is
|
|
114
|
+
not defined and the function returns a record with ``success=False``
|
|
115
|
+
and ``loss_history=[inf]`` rather than a sentinel value that could be
|
|
116
|
+
mistaken for a converged fit.
|
|
103
117
|
phase1_steps, phase2_steps, phase3_steps : int
|
|
104
118
|
Outer-loop step counts per phase. Each step is one L-BFGS call
|
|
105
119
|
with up to ``lbfgs_max_iter`` inner iterations.
|
|
@@ -118,6 +132,10 @@ def optimize_single_grain(
|
|
|
118
132
|
when no ground truth is supplied; see
|
|
119
133
|
:func:`evaluate_recovery` for ground-truth eval.
|
|
120
134
|
``loss_history`` -- list of phase-final loss values.
|
|
135
|
+
``success`` -- False if too few spots ever matched; the parameters
|
|
136
|
+
are then the caller's own seed, ``loss_history`` is
|
|
137
|
+
``[inf]``, and ``failure_reason`` says why.
|
|
138
|
+
``n_matched`` -- spots matched at the last evaluated state.
|
|
121
139
|
"""
|
|
122
140
|
if loss is None:
|
|
123
141
|
loss = SpotMatchingLoss(metric="l2")
|
|
@@ -130,6 +148,7 @@ def optimize_single_grain(
|
|
|
130
148
|
opt_euler = init_euler.clone().requires_grad_(True)
|
|
131
149
|
opt_latc = init_lattice.clone().requires_grad_(False)
|
|
132
150
|
loss_history: list = []
|
|
151
|
+
_last_n_matched = [0]
|
|
133
152
|
|
|
134
153
|
def make_closure(params):
|
|
135
154
|
def closure():
|
|
@@ -143,17 +162,31 @@ def optimize_single_grain(
|
|
|
143
162
|
pred_flat = coords.squeeze().reshape(-1, 3)
|
|
144
163
|
valid_flat = valid.squeeze().reshape(-1)
|
|
145
164
|
pred_valid = pred_flat[valid_flat > 0.5]
|
|
165
|
+
# These two early-outs used to return torch.tensor(1e6,
|
|
166
|
+
# requires_grad=True) -- a FRESH LEAF, disconnected from opt_euler,
|
|
167
|
+
# on which .backward() was never called. So opt_euler.grad stayed
|
|
168
|
+
# None, L-BFGS took no step, and optimize_single_grain returned
|
|
169
|
+
# loss_history = [1e6, 1e6, 1e6] with the parameters moved exactly
|
|
170
|
+
# 0.0 and no error raised. A constant is the worst possible signal
|
|
171
|
+
# here: it looks like a fit that plateaued.
|
|
172
|
+
#
|
|
173
|
+
# There is nothing to restore -- with no matched spots there are no
|
|
174
|
+
# observations and hence no gradient. Raise instead, and let the
|
|
175
|
+
# caller record the failure.
|
|
146
176
|
if pred_valid.shape[0] == 0:
|
|
147
|
-
|
|
148
|
-
|
|
177
|
+
raise NoMatchError(
|
|
178
|
+
"no valid predicted spots at the current state", n_matched=0
|
|
149
179
|
)
|
|
150
180
|
pred_match, obs_match = _associate(
|
|
151
181
|
pred_valid, observed_spots, max_match_distance
|
|
152
182
|
)
|
|
153
183
|
if pred_match.shape[0] < min_matches:
|
|
154
|
-
|
|
155
|
-
|
|
184
|
+
raise NoMatchError(
|
|
185
|
+
f"only {pred_match.shape[0]} spots matched, need "
|
|
186
|
+
f"{min_matches}",
|
|
187
|
+
n_matched=int(pred_match.shape[0]),
|
|
156
188
|
)
|
|
189
|
+
_last_n_matched[0] = int(pred_match.shape[0])
|
|
157
190
|
l = loss(pred_match, obs_match)
|
|
158
191
|
l.backward()
|
|
159
192
|
return l
|
|
@@ -172,6 +205,25 @@ def optimize_single_grain(
|
|
|
172
205
|
lat_err = (opt_latc[:3] - init_lattice[:3]).abs().max().item()
|
|
173
206
|
print(f"{step:5d} {l.item():12.6e} {misori:12.6f} {lat_err:10.6f}")
|
|
174
207
|
|
|
208
|
+
def _failed(exc: "NoMatchError") -> Dict[str, Any]:
|
|
209
|
+
"""Report a fit that never had enough data to be defined.
|
|
210
|
+
|
|
211
|
+
``loss_history`` is inf rather than a finite constant so a caller
|
|
212
|
+
cannot compare it against a real fit and prefer it, and ``success`` is
|
|
213
|
+
False so "the parameters did not move" is distinguishable from "the fit
|
|
214
|
+
converged".
|
|
215
|
+
"""
|
|
216
|
+
return {
|
|
217
|
+
"euler_rad": opt_euler.detach().clone(),
|
|
218
|
+
"euler_deg": opt_euler.detach().clone() * RAD2DEG,
|
|
219
|
+
"lattice": opt_latc.detach().clone(),
|
|
220
|
+
"misori_deg": current_misori_deg(),
|
|
221
|
+
"loss_history": [float("inf")],
|
|
222
|
+
"success": False,
|
|
223
|
+
"n_matched": exc.n_matched,
|
|
224
|
+
"failure_reason": str(exc),
|
|
225
|
+
}
|
|
226
|
+
|
|
175
227
|
if verbose:
|
|
176
228
|
print(f"{'Step':>5} {'Loss':>12} {'dMisori(deg)':>12} {'dLat':>10}")
|
|
177
229
|
print("-" * 55)
|
|
@@ -181,11 +233,14 @@ def optimize_single_grain(
|
|
|
181
233
|
[opt_euler], lr=1.0, max_iter=lbfgs_max_iter,
|
|
182
234
|
line_search_fn="strong_wolfe",
|
|
183
235
|
)
|
|
184
|
-
|
|
185
|
-
|
|
186
|
-
|
|
187
|
-
|
|
188
|
-
|
|
236
|
+
try:
|
|
237
|
+
for step in range(phase1_steps):
|
|
238
|
+
l = optimizer.step(make_closure([opt_euler]))
|
|
239
|
+
log(step, l)
|
|
240
|
+
if current_misori_deg() < convergence_misori_deg and step > 0:
|
|
241
|
+
break
|
|
242
|
+
except NoMatchError as exc:
|
|
243
|
+
return _failed(exc)
|
|
189
244
|
loss_history.append(float(l.detach()))
|
|
190
245
|
|
|
191
246
|
if verbose:
|
|
@@ -196,12 +251,15 @@ def optimize_single_grain(
|
|
|
196
251
|
[opt_latc], lr=1.0, max_iter=lbfgs_max_iter,
|
|
197
252
|
line_search_fn="strong_wolfe",
|
|
198
253
|
)
|
|
199
|
-
|
|
200
|
-
|
|
201
|
-
|
|
202
|
-
|
|
203
|
-
|
|
204
|
-
|
|
254
|
+
try:
|
|
255
|
+
for step in range(phase2_steps):
|
|
256
|
+
l = optimizer.step(make_closure([opt_latc]))
|
|
257
|
+
log(step + phase1_steps, l)
|
|
258
|
+
if (opt_latc[:3] - init_lattice[:3]).abs().max().item() > 0 \
|
|
259
|
+
and abs(float(l.detach()) - loss_history[-1]) < convergence_lattice_err:
|
|
260
|
+
break
|
|
261
|
+
except NoMatchError as exc:
|
|
262
|
+
return _failed(exc)
|
|
205
263
|
loss_history.append(float(l.detach()))
|
|
206
264
|
|
|
207
265
|
if verbose:
|
|
@@ -212,9 +270,12 @@ def optimize_single_grain(
|
|
|
212
270
|
[opt_euler, opt_latc], lr=0.5, max_iter=lbfgs_max_iter,
|
|
213
271
|
line_search_fn="strong_wolfe",
|
|
214
272
|
)
|
|
215
|
-
|
|
216
|
-
|
|
217
|
-
|
|
273
|
+
try:
|
|
274
|
+
for step in range(phase3_steps):
|
|
275
|
+
l = optimizer.step(make_closure([opt_euler, opt_latc]))
|
|
276
|
+
log(step + phase1_steps + phase2_steps, l)
|
|
277
|
+
except NoMatchError as exc:
|
|
278
|
+
return _failed(exc)
|
|
218
279
|
loss_history.append(float(l.detach()))
|
|
219
280
|
|
|
220
281
|
return {
|
|
@@ -223,6 +284,8 @@ def optimize_single_grain(
|
|
|
223
284
|
"lattice": opt_latc.detach().clone(),
|
|
224
285
|
"misori_deg": current_misori_deg(),
|
|
225
286
|
"loss_history": loss_history,
|
|
287
|
+
"success": True,
|
|
288
|
+
"n_matched": _last_n_matched[0],
|
|
226
289
|
}
|
|
227
290
|
|
|
228
291
|
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: midas-diffract
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.8.0
|
|
4
4
|
Summary: End-to-end differentiable forward model for High-Energy Diffraction Microscopy (FF, NF, pf-HEDM)
|
|
5
5
|
Author-email: Hemant Sharma <hsharma@anl.gov>
|
|
6
6
|
License-Expression: BSD-3-Clause
|
|
@@ -19,6 +19,8 @@ tests/test_hkls.py
|
|
|
19
19
|
tests/test_losses.py
|
|
20
20
|
tests/test_multi_detector.py
|
|
21
21
|
tests/test_omega_box_filter.py
|
|
22
|
+
tests/test_omega_solver_singularities.py
|
|
23
|
+
tests/test_optimize_no_match.py
|
|
22
24
|
tests/test_strain_tensor.py
|
|
23
25
|
tests/test_tilts.py
|
|
24
26
|
tests/test_wedge.py
|
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "midas-diffract"
|
|
7
|
-
version = "0.
|
|
7
|
+
version = "0.8.0"
|
|
8
8
|
description = "End-to-end differentiable forward model for High-Energy Diffraction Microscopy (FF, NF, pf-HEDM)"
|
|
9
9
|
readme = "README.md"
|
|
10
10
|
license = "BSD-3-Clause"
|
|
@@ -19,7 +19,7 @@ _GD = dict(
|
|
|
19
19
|
omega_start=180.0, omega_step=-0.25, n_frames=1440,
|
|
20
20
|
n_pixels_y=1679, n_pixels_z=1679, min_eta=6.0, wavelength=0.172979,
|
|
21
21
|
)
|
|
22
|
-
# A representative calibrated v2 coefficient set (
|
|
22
|
+
# A representative calibrated v2 coefficient set (datasetD Pilatus 2M CeO2).
|
|
23
23
|
_COEFFS = [0.00707, -0.01, 0.00624, 0.01, -34.76, 0.00234, 81.47,
|
|
24
24
|
-0.00369, -12.29, -0.00727, -5.29, -0.00863, -1.51, -0.00446, -7.79]
|
|
25
25
|
|
|
@@ -0,0 +1,205 @@
|
|
|
1
|
+
"""The omega solver: accuracy and gradients where the old form failed.
|
|
2
|
+
|
|
3
|
+
The solver satisfies -Gx cos(w) + Gy sin(w) = v, so the residual of THAT
|
|
4
|
+
identity is the arbiter -- not agreement with the previous implementation,
|
|
5
|
+
which was itself biased. The closed form
|
|
6
|
+
|
|
7
|
+
rho = |(Gx, Gy)|, phi = atan2(Gy, -Gx), w = phi +- acos(v/rho)
|
|
8
|
+
|
|
9
|
+
replaced a Gy^2-divided quadratic whose ``+ 1e-7`` denominators perturbed every
|
|
10
|
+
spot, and whose ``safe_arccos`` froze both value and gradient in a +-0.0256 deg
|
|
11
|
+
band around omega = 0.
|
|
12
|
+
|
|
13
|
+
C parity is covered separately by tests/test_c_comparison.py (FF and NF); note
|
|
14
|
+
its tolerance is 0.5 deg, which confirms no regression but cannot resolve the
|
|
15
|
+
1e-2 deg effects tested here.
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
import math
|
|
21
|
+
|
|
22
|
+
import numpy as np
|
|
23
|
+
import pytest
|
|
24
|
+
import torch
|
|
25
|
+
|
|
26
|
+
from midas_diffract.forward import HEDMForwardModel, HEDMGeometry, solve_omega
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
@pytest.fixture(autouse=True)
|
|
30
|
+
def _float64_default():
|
|
31
|
+
"""float64 for this module only.
|
|
32
|
+
|
|
33
|
+
Setting the default dtype at import time leaks into every test module
|
|
34
|
+
imported afterwards -- it broke 49 tests in test_forward.py, which expects
|
|
35
|
+
the float32 default.
|
|
36
|
+
"""
|
|
37
|
+
prev = torch.get_default_dtype()
|
|
38
|
+
torch.set_default_dtype(torch.float64)
|
|
39
|
+
yield
|
|
40
|
+
torch.set_default_dtype(prev)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def _solve(Gx, Gy, v):
|
|
44
|
+
"""The SHIPPED solver -- not a copy of it.
|
|
45
|
+
|
|
46
|
+
An earlier draft of this file reimplemented the closed form here, which
|
|
47
|
+
tested the copy rather than the code and silently dropped the tangency
|
|
48
|
+
guard. ``solve_omega`` is module-level precisely so the test can call it.
|
|
49
|
+
"""
|
|
50
|
+
wp, wn, _ = solve_omega(Gx, Gy, v)
|
|
51
|
+
return wp, wn
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _residual(w, Gx, Gy, v):
|
|
55
|
+
return torch.abs(-Gx * torch.cos(w) + Gy * torch.sin(w) - v)
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def _coeffs(n, gy_scale, seed):
|
|
59
|
+
rng = np.random.default_rng(seed)
|
|
60
|
+
Gx = torch.tensor(rng.normal(0, 0.35, n))
|
|
61
|
+
Gy = torch.tensor(rng.uniform(-gy_scale, gy_scale, n)) if gy_scale < 0.1 \
|
|
62
|
+
else torch.tensor(rng.normal(0, gy_scale, n))
|
|
63
|
+
v = torch.tensor(rng.normal(0, 0.20, n))
|
|
64
|
+
keep = (v * v) <= (Gx * Gx + Gy * Gy)
|
|
65
|
+
return Gx[keep], Gy[keep], v[keep]
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
@pytest.mark.parametrize("gy_scale,label", [
|
|
69
|
+
(0.35, "ordinary"), (3e-4, "small |Gy|"), (1e-8, "tiny |Gy|"),
|
|
70
|
+
])
|
|
71
|
+
def test_constraint_residual_is_at_machine_precision(gy_scale, label):
|
|
72
|
+
"""The old form gave median 1.2e-07 (ordinary) to 3.0e-04 (tiny |Gy|)."""
|
|
73
|
+
Gx, Gy, v = _coeffs(50_000, gy_scale, seed=0)
|
|
74
|
+
wp, wn = _solve(Gx, Gy, v)
|
|
75
|
+
r = torch.minimum(_residual(wp, Gx, Gy, v), _residual(wn, Gx, Gy, v))
|
|
76
|
+
assert float(r.median()) < 1e-15, f"{label}: median residual {float(r.median()):.3e}"
|
|
77
|
+
assert float(r.max()) < 1e-12, f"{label}: max residual {float(r.max()):.3e}"
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def test_gy_exactly_zero_needs_no_special_branch():
|
|
81
|
+
"""The old code carried a separate |Gy| < 1e-12 branch. Gy = 0 is now
|
|
82
|
+
just phi = atan2(0, -Gx), which is 0 or pi. Nothing special about it."""
|
|
83
|
+
rng = np.random.default_rng(1)
|
|
84
|
+
Gx = torch.tensor(rng.normal(0, 0.35, 20_000))
|
|
85
|
+
Gy = torch.zeros_like(Gx)
|
|
86
|
+
v = torch.tensor(rng.normal(0, 0.20, 20_000))
|
|
87
|
+
keep = (v * v) <= (Gx * Gx)
|
|
88
|
+
Gx, Gy, v = Gx[keep], Gy[keep], v[keep]
|
|
89
|
+
wp, wn = _solve(Gx, Gy, v)
|
|
90
|
+
r = torch.minimum(_residual(wp, Gx, Gy, v), _residual(wn, Gx, Gy, v))
|
|
91
|
+
assert float(r.max()) < 1e-12
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def test_no_frozen_band_near_omega_zero():
|
|
95
|
+
"""Recover a KNOWN omega inside the band the old clamp pinned.
|
|
96
|
+
|
|
97
|
+
Old: median error 1.56e-02 deg, i.e. the answer was the band edge rather
|
|
98
|
+
than the spot's actual omega.
|
|
99
|
+
"""
|
|
100
|
+
rng = np.random.default_rng(2)
|
|
101
|
+
n = 50_000
|
|
102
|
+
Gx = torch.tensor(rng.normal(0, 0.35, n))
|
|
103
|
+
Gy = torch.tensor(rng.normal(0, 0.35, n))
|
|
104
|
+
w_true = torch.tensor(rng.uniform(-0.02, 0.02, n)) * math.pi / 180.0
|
|
105
|
+
v = -Gx * torch.cos(w_true) + Gy * torch.sin(w_true) # exact by construction
|
|
106
|
+
wp, wn = _solve(Gx, Gy, v)
|
|
107
|
+
# Either branch may carry the true root; take the nearer.
|
|
108
|
+
err = torch.minimum((wp - w_true).abs(), (wn - w_true).abs())
|
|
109
|
+
err_deg = err * 180.0 / math.pi
|
|
110
|
+
assert float(err_deg.median()) < 1e-10, (
|
|
111
|
+
f"median error {float(err_deg.median()):.3e} deg "
|
|
112
|
+
f"(old form: 1.56e-02 deg -- the frozen band)")
|
|
113
|
+
assert float(np.percentile(err_deg.numpy(), 99.9)) < 1e-6
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
def test_no_epsilon_bias_at_moderate_gy():
|
|
117
|
+
"""``y2 + 1e-7`` perturbed EVERY coefficient, not only small-|Gy| spots.
|
|
118
|
+
|
|
119
|
+
At |Gy| ~ 0.35 the old relative perturbation was ~1e-6, which is what put
|
|
120
|
+
a ~1e-7 residual on ordinary spots. Nothing here may depend on a tolerance
|
|
121
|
+
that loose.
|
|
122
|
+
"""
|
|
123
|
+
Gx, Gy, v = _coeffs(20_000, 0.35, seed=3)
|
|
124
|
+
wp, wn = _solve(Gx, Gy, v)
|
|
125
|
+
r = torch.minimum(_residual(wp, Gx, Gy, v), _residual(wn, Gx, Gy, v))
|
|
126
|
+
assert float((r > 1e-14).double().mean()) < 1e-3
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
# ------------------------------------------------------------- gradients
|
|
130
|
+
|
|
131
|
+
def test_gradient_finite_and_fd_matching_at_omega_zero():
|
|
132
|
+
"""omega = 0 is an ORDINARY spot. Its gradient must be right, not merely
|
|
133
|
+
finite -- the old clamp made it exactly 0."""
|
|
134
|
+
GX, GY = 0.3, 0.4
|
|
135
|
+
# w = 0 => -Gx cos 0 + Gy sin 0 = -Gx, so v = -Gx puts a root at omega=0.
|
|
136
|
+
# It lands on the MINUS branch here (phi + acos(v/rho) is the other root).
|
|
137
|
+
v = torch.tensor(-GX)
|
|
138
|
+
|
|
139
|
+
def f(gx, gy):
|
|
140
|
+
return _solve(gx, gy, v)[1]
|
|
141
|
+
|
|
142
|
+
Gx = torch.tensor(GX, requires_grad=True)
|
|
143
|
+
Gy = torch.tensor(GY, requires_grad=True)
|
|
144
|
+
out = f(Gx, Gy)
|
|
145
|
+
assert abs(float(out.detach())) < 1e-12, "setup should put omega at 0"
|
|
146
|
+
g = torch.autograd.grad(out, [Gx, Gy])
|
|
147
|
+
assert all(torch.isfinite(x).all() for x in g)
|
|
148
|
+
h = 1e-7
|
|
149
|
+
for i, (val, gi) in enumerate(zip((GX, GY), g)):
|
|
150
|
+
a = [torch.tensor(GX), torch.tensor(GY)]
|
|
151
|
+
a[i] = torch.tensor(val + h); up = float(f(*a))
|
|
152
|
+
a[i] = torch.tensor(val - h); dn = float(f(*a))
|
|
153
|
+
fd = (up - dn) / (2 * h)
|
|
154
|
+
assert abs(float(gi) - fd) < 1e-4, f"param {i}: {float(gi)} vs FD {fd}"
|
|
155
|
+
# The old clamp made this exactly 0. At least one component must be alive.
|
|
156
|
+
assert max(abs(float(x)) for x in g) > 1e-6, "gradient is dead at omega = 0"
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
def test_one_tangential_spot_does_not_poison_the_batch():
|
|
160
|
+
"""acos'(+-1) is infinite at exact tangency; torch sums gradients, so an
|
|
161
|
+
unguarded singular spot would NaN every other spot's gradient."""
|
|
162
|
+
Gx = torch.tensor([0.3, 0.3, 0.3], requires_grad=True)
|
|
163
|
+
Gy = torch.tensor([0.4, 0.4, 0.4], requires_grad=True)
|
|
164
|
+
v = torch.tensor([0.5, 0.1, -0.2]) # v[0] = rho exactly: tangency
|
|
165
|
+
wp, _ = _solve(Gx, Gy, v)
|
|
166
|
+
g = torch.autograd.grad(wp.sum(), [Gx, Gy])
|
|
167
|
+
assert all(torch.isfinite(x).all() for x in g), "tangency poisoned the batch"
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
# ------------------------------------------------------- end to end
|
|
171
|
+
|
|
172
|
+
def _ff_model():
|
|
173
|
+
geom = HEDMGeometry(
|
|
174
|
+
Lsd=1_000_000.0, y_BC=1024.0, z_BC=1024.0, px=200.0,
|
|
175
|
+
omega_start=0.0, omega_step=0.25, n_frames=1440,
|
|
176
|
+
n_pixels_y=2048, n_pixels_z=2048, min_eta=6.0, wavelength=0.295,
|
|
177
|
+
)
|
|
178
|
+
a, wl = 2.87, 0.295
|
|
179
|
+
hkls_int = torch.tensor([[1, 1, 0], [2, 0, 0], [2, 1, 1], [2, 2, 0]])
|
|
180
|
+
hkls_cart = (torch.eye(3) / a @ hkls_int.double().T).T
|
|
181
|
+
thetas = torch.asin(wl / (2.0 * (1.0 / torch.norm(hkls_cart, dim=-1))))
|
|
182
|
+
return HEDMForwardModel(hkls=hkls_cart, thetas=thetas, geometry=geom,
|
|
183
|
+
hkls_int=hkls_int.double(),
|
|
184
|
+
device=torch.device("cpu"))
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
def test_model_forward_is_finite_and_differentiable():
|
|
188
|
+
model = _ff_model()
|
|
189
|
+
eul = torch.tensor([[0.4, 0.6, 0.8]], requires_grad=True)
|
|
190
|
+
pos = torch.zeros(1, 3)
|
|
191
|
+
out = model(eul.unsqueeze(0), pos.unsqueeze(0))
|
|
192
|
+
assert torch.isfinite(out.omega).all()
|
|
193
|
+
assert torch.isfinite(out.eta).all()
|
|
194
|
+
g = torch.autograd.grad(
|
|
195
|
+
(out.omega * out.valid).sum(), eul, allow_unused=True)[0]
|
|
196
|
+
assert g is not None and torch.isfinite(g).all()
|
|
197
|
+
|
|
198
|
+
|
|
199
|
+
def test_model_omega_within_180():
|
|
200
|
+
model = _ff_model()
|
|
201
|
+
rng = np.random.default_rng(0)
|
|
202
|
+
eul = torch.tensor(rng.uniform(0, 2 * np.pi, size=(64, 3)))
|
|
203
|
+
out = model(eul.unsqueeze(0), torch.zeros(64, 3).unsqueeze(0))
|
|
204
|
+
om = out.omega[out.valid > 0.5]
|
|
205
|
+
assert float(om.abs().max()) <= 180.0 + 1e-9
|
|
@@ -0,0 +1,92 @@
|
|
|
1
|
+
"""A single-grain fit with too few matches must not look like a converged one.
|
|
2
|
+
|
|
3
|
+
Both early-outs in ``optimize_single_grain``'s closure used to return
|
|
4
|
+
``torch.tensor(1e6, requires_grad=True)`` -- a FRESH LEAF, disconnected from
|
|
5
|
+
``opt_euler``, on which ``.backward()`` was never called. So ``opt_euler.grad``
|
|
6
|
+
stayed None, L-BFGS took no step, and the function returned
|
|
7
|
+
``loss_history = [1e6, 1e6, 1e6]`` with the parameters moved exactly 0.0 and no
|
|
8
|
+
error raised. A constant sentinel is the worst possible signal: it reads as a
|
|
9
|
+
fit that plateaued.
|
|
10
|
+
|
|
11
|
+
Nothing can be restored -- with no matched spots there are no observations and
|
|
12
|
+
so no gradient. The contract (spec_autograd_classB_classC.md, C2) is that the
|
|
13
|
+
failure is legible instead.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import math
|
|
19
|
+
|
|
20
|
+
import numpy as np
|
|
21
|
+
import pytest
|
|
22
|
+
import torch
|
|
23
|
+
|
|
24
|
+
from midas_diffract import NoMatchError, optimize_single_grain
|
|
25
|
+
from midas_diffract.forward import HEDMForwardModel, HEDMGeometry
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def _model():
|
|
29
|
+
geom = HEDMGeometry(
|
|
30
|
+
Lsd=1_000_000.0, y_BC=1024.0, z_BC=1024.0, px=200.0,
|
|
31
|
+
omega_start=0.0, omega_step=0.25, n_frames=1440,
|
|
32
|
+
n_pixels_y=2048, n_pixels_z=2048, min_eta=6.0, wavelength=0.295,
|
|
33
|
+
)
|
|
34
|
+
a, wl = 2.87, 0.295
|
|
35
|
+
hkls_int = torch.tensor([[1, 1, 0], [2, 0, 0], [2, 1, 1], [2, 2, 0]],
|
|
36
|
+
dtype=torch.float32)
|
|
37
|
+
hkls_cart = (torch.eye(3) / a @ hkls_int.T).T
|
|
38
|
+
thetas = torch.asin(torch.tensor(wl) / (2.0 * (1.0 / torch.norm(hkls_cart, dim=-1))))
|
|
39
|
+
return HEDMForwardModel(hkls=hkls_cart, thetas=thetas, geometry=geom,
|
|
40
|
+
hkls_int=hkls_int, device=torch.device("cpu"))
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def _run_with_unmatchable_observations():
|
|
44
|
+
model = _model()
|
|
45
|
+
init_euler = torch.tensor([0.4, 0.6, 0.8])
|
|
46
|
+
init_lattice = torch.tensor([2.87, 2.87, 2.87, 90.0, 90.0, 90.0])
|
|
47
|
+
# Observations nowhere near any prediction, and a tolerance too tight for
|
|
48
|
+
# anything to associate.
|
|
49
|
+
obs = torch.tensor([[3.0, 3.0, 3.0], [3.1, 3.1, 3.1]])
|
|
50
|
+
return optimize_single_grain(
|
|
51
|
+
model, observed_spots=obs, init_euler=init_euler,
|
|
52
|
+
init_lattice=init_lattice, position=torch.zeros(3),
|
|
53
|
+
max_match_distance=1e-12, min_matches=50,
|
|
54
|
+
phase1_steps=2, phase2_steps=2, phase3_steps=2,
|
|
55
|
+
)
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def test_failure_is_reported_not_disguised_as_a_plateau():
|
|
59
|
+
r = _run_with_unmatchable_observations()
|
|
60
|
+
assert r["success"] is False
|
|
61
|
+
assert "failure_reason" in r
|
|
62
|
+
assert all(math.isinf(x) for x in r["loss_history"]), (
|
|
63
|
+
f"loss_history={r['loss_history']} -- a finite constant is "
|
|
64
|
+
f"indistinguishable from a fit that stopped improving"
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def test_failed_fit_never_beats_a_real_one_on_loss():
|
|
69
|
+
r = _run_with_unmatchable_observations()
|
|
70
|
+
assert min(r["loss_history"]) == float("inf")
|
|
71
|
+
assert not any(x <= 1e6 for x in r["loss_history"])
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def test_parameters_are_returned_as_the_caller_s_own_seed():
|
|
75
|
+
"""Not a defect -- there is nothing to move toward. Pinned so it is
|
|
76
|
+
explicit that unmoved parameters come labelled as a failure."""
|
|
77
|
+
init_euler = torch.tensor([0.4, 0.6, 0.8])
|
|
78
|
+
r = _run_with_unmatchable_observations()
|
|
79
|
+
assert torch.allclose(r["euler_rad"], init_euler, atol=1e-12)
|
|
80
|
+
assert r["success"] is False
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def test_n_matched_is_reported():
|
|
84
|
+
r = _run_with_unmatchable_observations()
|
|
85
|
+
assert "n_matched" in r
|
|
86
|
+
assert r["n_matched"] < 50
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def test_no_match_error_is_public_and_carries_the_count():
|
|
90
|
+
assert issubclass(NoMatchError, RuntimeError)
|
|
91
|
+
e = NoMatchError("x", n_matched=3)
|
|
92
|
+
assert e.n_matched == 3
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|