midas-diffract 0.6.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.6.0 → midas_diffract-0.8.0}/PKG-INFO +1 -1
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/midas_diffract/__init__.py +3 -2
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/midas_diffract/forward.py +242 -70
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/midas_diffract/optimize.py +84 -21
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/midas_diffract.egg-info/PKG-INFO +1 -1
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/midas_diffract.egg-info/SOURCES.txt +3 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/pyproject.toml +1 -1
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/tests/test_distortion_layer.py +1 -1
- midas_diffract-0.8.0/tests/test_omega_box_filter.py +352 -0
- 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.6.0 → midas_diffract-0.8.0}/LICENSE +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/README.md +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/midas_diffract/hkls.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/midas_diffract/losses.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/midas_diffract/simulate_panel_zarrs.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/midas_diffract.egg-info/dependency_links.txt +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/midas_diffract.egg-info/requires.txt +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/midas_diffract.egg-info/top_level.txt +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/setup.cfg +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/tests/test_c_comparison.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/tests/test_forward.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/tests/test_hkls.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/tests/test_losses.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/tests/test_multi_detector.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/tests/test_strain_tensor.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.8.0}/tests/test_tilts.py +0 -0
- {midas_diffract-0.6.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
|
]
|
|
@@ -117,6 +117,26 @@ class HEDMGeometry:
|
|
|
117
117
|
# spot is valid if it lands on at least
|
|
118
118
|
# one detector; output gains a det_id
|
|
119
119
|
# field naming which one.
|
|
120
|
+
# Paired OmegaRange / BoxSize acceptance filter (the ``OmegaRange`` and
|
|
121
|
+
# ``BoxSize`` paramfile keys). Ports the KeepSpot gate in
|
|
122
|
+
# NF_HEDM/src/CalcDiffractionSpots.c:225-236 -- a spot is kept only if
|
|
123
|
+
# SOME index i satisfies BOTH omega in (omega_ranges[i]) AND the nominal
|
|
124
|
+
# ring position (yl, zl) inside box_sizes[i]; a spot that fails every
|
|
125
|
+
# pair is dropped from the theoretical spot list entirely, i.e. it is
|
|
126
|
+
# excluded from BOTH numerator and denominator of the overlap fraction.
|
|
127
|
+
#
|
|
128
|
+
# Entry i of ``omega_ranges`` pairs with entry i of ``box_sizes``, so the
|
|
129
|
+
# two lists must have equal length. Both default to None => the filter is
|
|
130
|
+
# OFF and every existing FF/pf/NF caller is bit-unchanged.
|
|
131
|
+
#
|
|
132
|
+
# Units: omega_ranges in degrees; box_sizes in MICROMETRES, as
|
|
133
|
+
# (y_min, y_max, z_min, z_max) on the detector plane relative to the beam
|
|
134
|
+
# centre, evaluated at the FIRST distance (Lsd[0]) with no grain
|
|
135
|
+
# displacement, no tilt and no distortion -- exactly the C's
|
|
136
|
+
# CalcSpotPosition(RingRadius = Lsd[0]*tan(2*theta), eta). Bounds are
|
|
137
|
+
# STRICT (a spot exactly on an edge is rejected), matching the C.
|
|
138
|
+
omega_ranges: "list[tuple[float, float]] | None" = None
|
|
139
|
+
box_sizes: "list[tuple[float, float, float, float]] | None" = None
|
|
120
140
|
|
|
121
141
|
@property
|
|
122
142
|
def n_distances(self) -> int:
|
|
@@ -204,6 +224,90 @@ class SpotDescriptors:
|
|
|
204
224
|
det_id: Optional[torch.Tensor] = None # (..., K, M) int64 panel index (panel mode only)
|
|
205
225
|
|
|
206
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
|
+
|
|
207
311
|
# ---------------------------------------------------------------------------
|
|
208
312
|
# Forward model
|
|
209
313
|
# ---------------------------------------------------------------------------
|
|
@@ -399,6 +503,38 @@ class HEDMForwardModel(nn.Module):
|
|
|
399
503
|
)
|
|
400
504
|
self._has_wedge = abs(float(geometry.wedge)) > 0.0
|
|
401
505
|
|
|
506
|
+
# Paired OmegaRange / BoxSize acceptance filter. OFF unless BOTH
|
|
507
|
+
# lists are supplied and non-empty (see HEDMGeometry.box_sizes).
|
|
508
|
+
om_ranges = getattr(geometry, "omega_ranges", None) or []
|
|
509
|
+
bx_sizes = getattr(geometry, "box_sizes", None) or []
|
|
510
|
+
if bool(om_ranges) != bool(bx_sizes):
|
|
511
|
+
raise ValueError(
|
|
512
|
+
"omega_ranges and box_sizes are a paired filter: supply both "
|
|
513
|
+
f"or neither (got {len(om_ranges)} omega_ranges and "
|
|
514
|
+
f"{len(bx_sizes)} box_sizes)."
|
|
515
|
+
)
|
|
516
|
+
if len(om_ranges) != len(bx_sizes):
|
|
517
|
+
raise ValueError(
|
|
518
|
+
f"omega_ranges (len {len(om_ranges)}) must match box_sizes "
|
|
519
|
+
f"(len {len(bx_sizes)}); entry i of one pairs with entry i "
|
|
520
|
+
"of the other."
|
|
521
|
+
)
|
|
522
|
+
self._has_omega_box = bool(bx_sizes)
|
|
523
|
+
self.register_buffer(
|
|
524
|
+
"_omega_ranges",
|
|
525
|
+
torch.tensor(
|
|
526
|
+
[[float(a), float(b)] for a, b in om_ranges],
|
|
527
|
+
dtype=torch.float64, device=device,
|
|
528
|
+
).reshape(-1, 2),
|
|
529
|
+
)
|
|
530
|
+
self.register_buffer(
|
|
531
|
+
"_box_sizes",
|
|
532
|
+
torch.tensor(
|
|
533
|
+
[[float(v) for v in bs] for bs in bx_sizes],
|
|
534
|
+
dtype=torch.float64, device=device,
|
|
535
|
+
).reshape(-1, 4),
|
|
536
|
+
)
|
|
537
|
+
|
|
402
538
|
# Scan config
|
|
403
539
|
self.scan_config = scan_config
|
|
404
540
|
if scan_config is not None:
|
|
@@ -518,8 +654,21 @@ class HEDMForwardModel(nn.Module):
|
|
|
518
654
|
# ------------------------------------------------------------------
|
|
519
655
|
|
|
520
656
|
def safe_arccos(self, x: torch.Tensor) -> torch.Tensor:
|
|
521
|
-
"""Numerically stable arccos: clamp to [-1+eps, 1-eps].
|
|
522
|
-
|
|
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))
|
|
523
672
|
|
|
524
673
|
# ------------------------------------------------------------------
|
|
525
674
|
# strain_as_voigt (accept full 3x3 tensor OR plain-Voigt 6-vector)
|
|
@@ -833,67 +982,9 @@ class HEDMForwardModel(nn.Module):
|
|
|
833
982
|
v = v_no_wedge + sin_W * Gz_p
|
|
834
983
|
# ---------------------------------------------------------------
|
|
835
984
|
|
|
836
|
-
#
|
|
837
|
-
|
|
838
|
-
|
|
839
|
-
# C uses almostzero=1e-12 for the Gy≈0 branch (see
|
|
840
|
-
# NF_HEDM/src/CalcDiffractionSpots.c:96 and
|
|
841
|
-
# FF_HEDM/src/ForwardSimulationCompressed.c:168). Match exactly.
|
|
842
|
-
almostzero = 1e-12
|
|
843
|
-
x2 = Gx * Gx
|
|
844
|
-
y2 = Gy * Gy
|
|
845
|
-
a = 1.0 + x2 / (y2 + self.epsilon)
|
|
846
|
-
b_coeff = 2.0 * v * Gx / (y2 + self.epsilon)
|
|
847
|
-
c_coeff = v * v / (y2 + self.epsilon) - 1.0
|
|
848
|
-
discriminant = b_coeff * b_coeff - 4.0 * a * c_coeff
|
|
849
|
-
|
|
850
|
-
sqrt_disc = torch.sqrt(torch.abs(discriminant))
|
|
851
|
-
|
|
852
|
-
coswp = (-b_coeff + sqrt_disc) / (2.0 * a)
|
|
853
|
-
coswn = (-b_coeff - sqrt_disc) / (2.0 * a)
|
|
854
|
-
|
|
855
|
-
wap = self.safe_arccos(coswp)
|
|
856
|
-
wan = self.safe_arccos(coswn)
|
|
857
|
-
wbp = -wap
|
|
858
|
-
wbn = -wan
|
|
859
|
-
|
|
860
|
-
# Select correct branch: the one satisfying -Gx*cos(w)+Gy*sin(w)=v
|
|
861
|
-
eqap = -Gx * torch.cos(wap) + Gy * torch.sin(wap)
|
|
862
|
-
eqbp = -Gx * torch.cos(wbp) + Gy * torch.sin(wbp)
|
|
863
|
-
eqan = -Gx * torch.cos(wan) + Gy * torch.sin(wan)
|
|
864
|
-
eqbn = -Gx * torch.cos(wbn) + Gy * torch.sin(wbn)
|
|
865
|
-
|
|
866
|
-
Dap = torch.abs(eqap - v)
|
|
867
|
-
Dbp = torch.abs(eqbp - v)
|
|
868
|
-
Dan = torch.abs(eqan - v)
|
|
869
|
-
Dbn = torch.abs(eqbn - v)
|
|
870
|
-
|
|
871
|
-
all_wp = torch.where(Dap < Dbp, wap, wbp)
|
|
872
|
-
all_wn = torch.where(Dan < Dbn, wan, wbn)
|
|
873
|
-
|
|
874
|
-
# Special case: Gy ~ 0 (C uses almostzero=1e-12)
|
|
875
|
-
# C code (CalcDiffractionSpots.c:97-106):
|
|
876
|
-
# cosome1 = -v / x;
|
|
877
|
-
# if (|cosome1| <= 1) { ome = acos(cosome1); solutions: +ome, -ome }
|
|
878
|
-
gy_zero = torch.abs(Gy) < almostzero
|
|
879
|
-
cosome_special = -v / (Gx + self.epsilon)
|
|
880
|
-
cosome_special_valid = (torch.abs(cosome_special) <= 1.0) & gy_zero & (torch.abs(Gx) > self.epsilon)
|
|
881
|
-
special_w = self.safe_arccos(cosome_special) # positive omega solution
|
|
882
|
-
# Two solutions: +ome and -ome
|
|
883
|
-
special_wp = special_w # positive
|
|
884
|
-
special_wn = -special_w # negative
|
|
885
|
-
|
|
886
|
-
# When |Gy| < almostzero, use the special case; otherwise use the quadratic
|
|
887
|
-
disc_valid = (discriminant >= 0) & (~gy_zero)
|
|
888
|
-
coswp_valid = (coswp >= -1.0) & (coswp <= 1.0)
|
|
889
|
-
coswn_valid = (coswn >= -1.0) & (coswn <= 1.0)
|
|
890
|
-
|
|
891
|
-
omega_p = torch.where(cosome_special_valid, special_wp,
|
|
892
|
-
torch.where(disc_valid & coswp_valid, all_wp,
|
|
893
|
-
torch.zeros_like(all_wp)))
|
|
894
|
-
omega_n = torch.where(cosome_special_valid, special_wn,
|
|
895
|
-
torch.where(disc_valid & coswn_valid, all_wn,
|
|
896
|
-
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
|
+
|
|
897
988
|
|
|
898
989
|
# Concatenate two solutions: (..., 2N, M)
|
|
899
990
|
all_omega = torch.cat([omega_p, omega_n], dim=-2)
|
|
@@ -931,21 +1022,102 @@ class HEDMForwardModel(nn.Module):
|
|
|
931
1022
|
tt = tt.expand_as(omega_p)
|
|
932
1023
|
two_theta = torch.cat([tt, tt], dim=-2)
|
|
933
1024
|
|
|
934
|
-
# Validity mask
|
|
935
|
-
|
|
936
|
-
|
|
937
|
-
# For gy_zero special case, valid only if cosome is in [-1, 1]
|
|
938
|
-
valid_p = valid_p | cosome_special_valid
|
|
939
|
-
valid_n = valid_n | cosome_special_valid
|
|
940
|
-
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()
|
|
941
1028
|
|
|
942
1029
|
# Eta bounds
|
|
943
1030
|
eta_ok = (torch.abs(eta) >= self.min_eta) & \
|
|
944
1031
|
((math.pi - torch.abs(eta)) >= self.min_eta)
|
|
945
1032
|
valid = valid * eta_ok.float()
|
|
946
1033
|
|
|
1034
|
+
# Paired OmegaRange / BoxSize gate (no-op unless both were supplied).
|
|
1035
|
+
if self._has_omega_box:
|
|
1036
|
+
valid = valid * self.omega_box_mask(
|
|
1037
|
+
all_omega, eta, two_theta
|
|
1038
|
+
).to(valid.dtype)
|
|
1039
|
+
|
|
947
1040
|
return all_omega, eta, two_theta, valid
|
|
948
1041
|
|
|
1042
|
+
# ------------------------------------------------------------------
|
|
1043
|
+
# omega_box_mask (port of the KeepSpot gate in CalcDiffrSpots_Furnace)
|
|
1044
|
+
# ------------------------------------------------------------------
|
|
1045
|
+
|
|
1046
|
+
def omega_box_mask(
|
|
1047
|
+
self,
|
|
1048
|
+
omega: torch.Tensor,
|
|
1049
|
+
eta: torch.Tensor,
|
|
1050
|
+
two_theta: torch.Tensor,
|
|
1051
|
+
) -> torch.Tensor:
|
|
1052
|
+
"""Paired ``OmegaRange`` / ``BoxSize`` acceptance mask.
|
|
1053
|
+
|
|
1054
|
+
Direct port of the ``KeepSpot`` gate in
|
|
1055
|
+
``NF_HEDM/src/CalcDiffrSpots_Furnace`` (CalcDiffractionSpots.c:225-236,
|
|
1056
|
+
identical code at MakeDiffrSpots.c:248-259)::
|
|
1057
|
+
|
|
1058
|
+
CalcSpotPosition(RingRadius, eta, &yl, &zl); // C:224
|
|
1059
|
+
for (OmegaRangeNo = 0; OmegaRangeNo < NOmegaRanges; OmegaRangeNo++) {
|
|
1060
|
+
KeepSpot = 0;
|
|
1061
|
+
if ((Omega > OmegaRange[i][0]) && (Omega < OmegaRange[i][1]) &&
|
|
1062
|
+
(yl > BoxSizes[i][0]) && (yl < BoxSizes[i][1]) &&
|
|
1063
|
+
(zl > BoxSizes[i][2]) && (zl < BoxSizes[i][3])) {
|
|
1064
|
+
KeepSpot = 1; break;
|
|
1065
|
+
}
|
|
1066
|
+
}
|
|
1067
|
+
|
|
1068
|
+
with ``RingRadius = Lsd[0] * tan(2*theta)`` (C:209, and the caller
|
|
1069
|
+
passes ``Lsd[0]`` -- SharedFuncsFit.c:830), and
|
|
1070
|
+
``yl = -sin(eta)*RingRadius``, ``zl = cos(eta)*RingRadius``
|
|
1071
|
+
(CalcDiffractionSpots.c:80-85).
|
|
1072
|
+
|
|
1073
|
+
Semantics that matter:
|
|
1074
|
+
|
|
1075
|
+
- ``(yl, zl)`` are the **nominal ring** coordinates in micrometres at
|
|
1076
|
+
the FIRST distance, with **no** grain displacement, **no** tilt and
|
|
1077
|
+
**no** distortion applied. They are not pixel coordinates and carry
|
|
1078
|
+
no beam-centre offset.
|
|
1079
|
+
- Bounds are **strict** on all six comparisons (edge-exact spots are
|
|
1080
|
+
rejected).
|
|
1081
|
+
- Range ``i`` pairs with box ``i``; a spot is kept if **any** pair
|
|
1082
|
+
accepts (``break`` on first hit).
|
|
1083
|
+
- A rejected spot never enters ``TheorSpots``, so ``nTspots`` shrinks
|
|
1084
|
+
and the spot is dropped from **both** the numerator and the
|
|
1085
|
+
denominator of ``CalcFracOverlap`` (SharedFuncsFit.c:645, 648-650) --
|
|
1086
|
+
it is *not* counted as a miss. Folding the mask into ``valid`` here
|
|
1087
|
+
reproduces that, because ``valid`` is the denominator weight in
|
|
1088
|
+
``ObsVolume.hard_fraction`` / ``soft_fraction``.
|
|
1089
|
+
|
|
1090
|
+
Parameters
|
|
1091
|
+
----------
|
|
1092
|
+
omega, eta, two_theta : Tensor
|
|
1093
|
+
Radians, all broadcastable to a common shape.
|
|
1094
|
+
|
|
1095
|
+
Returns
|
|
1096
|
+
-------
|
|
1097
|
+
Tensor (bool), same broadcast shape.
|
|
1098
|
+
"""
|
|
1099
|
+
dtype = omega.dtype
|
|
1100
|
+
# C uses Lsd[0] for the ring radius regardless of nDistances.
|
|
1101
|
+
Lsd0 = self._Lsd_eff.to(dtype).reshape(-1)[0]
|
|
1102
|
+
ring_radius = Lsd0 * torch.tan(two_theta)
|
|
1103
|
+
yl = -torch.sin(eta) * ring_radius
|
|
1104
|
+
zl = torch.cos(eta) * ring_radius
|
|
1105
|
+
omega_deg = omega * self.RAD2DEG
|
|
1106
|
+
|
|
1107
|
+
om = self._omega_ranges.to(dtype)
|
|
1108
|
+
bx = self._box_sizes.to(dtype)
|
|
1109
|
+
accept = torch.zeros(
|
|
1110
|
+
torch.broadcast_shapes(omega_deg.shape, yl.shape, zl.shape),
|
|
1111
|
+
dtype=torch.bool, device=omega.device,
|
|
1112
|
+
)
|
|
1113
|
+
for i in range(om.shape[0]):
|
|
1114
|
+
accept = accept | (
|
|
1115
|
+
(omega_deg > om[i, 0]) & (omega_deg < om[i, 1])
|
|
1116
|
+
& (yl > bx[i, 0]) & (yl < bx[i, 1])
|
|
1117
|
+
& (zl > bx[i, 2]) & (zl < bx[i, 3])
|
|
1118
|
+
)
|
|
1119
|
+
return accept
|
|
1120
|
+
|
|
949
1121
|
# ------------------------------------------------------------------
|
|
950
1122
|
# Tilt rotation matrix (RotationTilts in SharedFuncsFit.c:230-266)
|
|
951
1123
|
# ------------------------------------------------------------------
|
|
@@ -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
|
|
@@ -18,6 +18,9 @@ tests/test_forward.py
|
|
|
18
18
|
tests/test_hkls.py
|
|
19
19
|
tests/test_losses.py
|
|
20
20
|
tests/test_multi_detector.py
|
|
21
|
+
tests/test_omega_box_filter.py
|
|
22
|
+
tests/test_omega_solver_singularities.py
|
|
23
|
+
tests/test_optimize_no_match.py
|
|
21
24
|
tests/test_strain_tensor.py
|
|
22
25
|
tests/test_tilts.py
|
|
23
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,352 @@
|
|
|
1
|
+
"""Tests for the paired ``OmegaRange`` / ``BoxSize`` acceptance gate.
|
|
2
|
+
|
|
3
|
+
Reference C: ``CalcDiffrSpots_Furnace`` in
|
|
4
|
+
``NF_HEDM/src/CalcDiffractionSpots.c:185-248`` (identical code in
|
|
5
|
+
``NF_HEDM/src/MakeDiffrSpots.c:216-259``)::
|
|
6
|
+
|
|
7
|
+
RealType RingRadius = distance * tan(2 * deg2rad * Thetas[indexhkl]); # :209
|
|
8
|
+
CalcSpotPosition(RingRadius, etas[i], &yl, &zl); # :224
|
|
9
|
+
# yl = -(sin(eta) * RingRadius); zl = cos(eta) * RingRadius; # :83-84
|
|
10
|
+
for (OmegaRangeNo = 0; OmegaRangeNo < NOmegaRanges; OmegaRangeNo++) { # :225
|
|
11
|
+
KeepSpot = 0;
|
|
12
|
+
if ((Omega > OmegaRange[i][0]) && (Omega < OmegaRange[i][1]) && # :227-232
|
|
13
|
+
(yl > BoxSizes[i][0]) && (yl < BoxSizes[i][1]) &&
|
|
14
|
+
(zl > BoxSizes[i][2]) && (zl < BoxSizes[i][3])) {
|
|
15
|
+
KeepSpot = 1; break;
|
|
16
|
+
}
|
|
17
|
+
}
|
|
18
|
+
if (KeepSpot == 1) { spots[spotnr*3 + ...] = ...; spotnr++; } # :237-243
|
|
19
|
+
|
|
20
|
+
``distance`` is ``Lsd[0]`` -- the caller passes it explicitly
|
|
21
|
+
(``NF_HEDM/src/SharedFuncsFit.c:830``). A rejected spot is never written into
|
|
22
|
+
``TheorSpots``, so it is dropped from both the numerator and the denominator of
|
|
23
|
+
``CalcFracOverlap`` (``SharedFuncsFit.c:645`` increments ``TotalPixels`` only
|
|
24
|
+
for spots present in the list).
|
|
25
|
+
|
|
26
|
+
Run with:
|
|
27
|
+
cd packages/midas_diffract
|
|
28
|
+
python -m pytest tests/test_omega_box_filter.py -v
|
|
29
|
+
"""
|
|
30
|
+
|
|
31
|
+
import math
|
|
32
|
+
|
|
33
|
+
import pytest
|
|
34
|
+
import torch
|
|
35
|
+
|
|
36
|
+
from midas_diffract.forward import HEDMForwardModel, HEDMGeometry
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
LSD = 1_000_000.0
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _geom(**kw):
|
|
43
|
+
base = dict(
|
|
44
|
+
Lsd=LSD,
|
|
45
|
+
y_BC=1024.0,
|
|
46
|
+
z_BC=1024.0,
|
|
47
|
+
px=200.0,
|
|
48
|
+
omega_start=-180.0,
|
|
49
|
+
omega_step=1.0,
|
|
50
|
+
n_frames=360,
|
|
51
|
+
n_pixels_y=2048,
|
|
52
|
+
n_pixels_z=2048,
|
|
53
|
+
min_eta=6.0,
|
|
54
|
+
wavelength=0.172979,
|
|
55
|
+
flip_y=False, # NF convention
|
|
56
|
+
multi_mode="layered",
|
|
57
|
+
)
|
|
58
|
+
base.update(kw)
|
|
59
|
+
return HEDMGeometry(**base)
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def _model(geom):
|
|
63
|
+
"""Gold (a = 4.08 A) reflections, matching the NF fit-orientation path."""
|
|
64
|
+
hkls_int = torch.tensor(
|
|
65
|
+
[[1, 0, 0], [1, 1, 0], [1, 1, 1], [2, 0, 0], [2, 1, 1]],
|
|
66
|
+
dtype=torch.float64,
|
|
67
|
+
)
|
|
68
|
+
B = torch.eye(3, dtype=torch.float64) / 4.08
|
|
69
|
+
hkls_cart = hkls_int @ B.T
|
|
70
|
+
g = torch.linalg.norm(hkls_cart, dim=-1)
|
|
71
|
+
thetas = torch.asin(g * geom.wavelength / 2.0)
|
|
72
|
+
return HEDMForwardModel(
|
|
73
|
+
hkls=hkls_cart, thetas=thetas, geometry=geom,
|
|
74
|
+
hkls_int=hkls_int, device=torch.device("cpu"),
|
|
75
|
+
)
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
EUL = torch.tensor([[0.3, 0.5, 0.7]], dtype=torch.float64)
|
|
79
|
+
POS = torch.tensor([[0.0, 0.0, 0.0]], dtype=torch.float64)
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def _nominal_ring_coords(model, eul=EUL):
|
|
83
|
+
"""(yl, zl) in um exactly as the C CalcSpotPosition computes them."""
|
|
84
|
+
om, eta, tt, valid = model.calc_bragg_geometry(model.euler2mat(eul))
|
|
85
|
+
R = LSD * torch.tan(tt)
|
|
86
|
+
return -torch.sin(eta) * R, torch.cos(eta) * R, om, valid
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
# ---------------------------------------------------------------------------
|
|
90
|
+
# Default: no BoxSize => filter off, behaviour unchanged
|
|
91
|
+
# ---------------------------------------------------------------------------
|
|
92
|
+
|
|
93
|
+
class TestDefaultOff:
|
|
94
|
+
def test_flag_off_and_buffers_empty(self):
|
|
95
|
+
m = _model(_geom())
|
|
96
|
+
assert m._has_omega_box is False
|
|
97
|
+
assert m._omega_ranges.shape == (0, 2)
|
|
98
|
+
assert m._box_sizes.shape == (0, 4)
|
|
99
|
+
|
|
100
|
+
def test_valid_mask_identical_to_wide_open_box(self):
|
|
101
|
+
"""A box wider than any possible ring must not change ``valid``."""
|
|
102
|
+
m_off = _model(_geom())
|
|
103
|
+
m_open = _model(_geom(
|
|
104
|
+
omega_ranges=[(-180.0, 180.0)],
|
|
105
|
+
box_sizes=[(-1e9, 1e9, -1e9, 1e9)],
|
|
106
|
+
))
|
|
107
|
+
s_off = m_off(EUL, POS)
|
|
108
|
+
s_open = m_open(EUL, POS)
|
|
109
|
+
assert torch.equal(s_off.valid, s_open.valid)
|
|
110
|
+
assert float(s_off.valid.sum()) > 0 # the test is not vacuous
|
|
111
|
+
|
|
112
|
+
def test_ff_geometry_unaffected(self):
|
|
113
|
+
"""FF/pf callers never pass the keys -> byte-identical forward."""
|
|
114
|
+
m = _model(_geom(flip_y=True, multi_mode="panel"))
|
|
115
|
+
assert m._has_omega_box is False
|
|
116
|
+
s = m(EUL, POS)
|
|
117
|
+
assert float(s.valid.sum()) > 0
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
# ---------------------------------------------------------------------------
|
|
121
|
+
# Inside is kept, outside is dropped
|
|
122
|
+
# ---------------------------------------------------------------------------
|
|
123
|
+
|
|
124
|
+
class TestInsideOutside:
|
|
125
|
+
def test_spot_inside_box_is_kept_outside_is_dropped(self):
|
|
126
|
+
m_off = _model(_geom())
|
|
127
|
+
yl, zl, _om, _v = _nominal_ring_coords(m_off)
|
|
128
|
+
s_off = m_off(EUL, POS)
|
|
129
|
+
base = s_off.valid.bool()
|
|
130
|
+
n_base = int(base.sum())
|
|
131
|
+
assert n_base > 2
|
|
132
|
+
|
|
133
|
+
# Cut just above the smallest valid zl -> drop exactly those spots.
|
|
134
|
+
z_cut = float(zl[base].min()) + 1.0
|
|
135
|
+
n_below = int((base & (zl <= z_cut)).sum())
|
|
136
|
+
assert n_below > 0
|
|
137
|
+
|
|
138
|
+
m_box = _model(_geom(
|
|
139
|
+
omega_ranges=[(-180.0, 180.0)],
|
|
140
|
+
box_sizes=[(-1e9, 1e9, z_cut, 1e9)],
|
|
141
|
+
))
|
|
142
|
+
kept = m_box(EUL, POS).valid.bool()
|
|
143
|
+
|
|
144
|
+
assert int(kept.sum()) == n_base - n_below
|
|
145
|
+
# Every kept spot was valid before and is inside the box.
|
|
146
|
+
assert bool((kept <= base).all())
|
|
147
|
+
assert bool((zl[kept] > z_cut).all())
|
|
148
|
+
|
|
149
|
+
def test_box_is_micrometres_at_lsd0_not_pixels(self):
|
|
150
|
+
"""A 0..2048 box read as *pixels* would keep everything; read as um
|
|
151
|
+
(the C convention) it keeps nothing for this geometry."""
|
|
152
|
+
m_off = _model(_geom())
|
|
153
|
+
yl, zl, _om, _v = _nominal_ring_coords(m_off)
|
|
154
|
+
base = m_off(EUL, POS).valid.bool()
|
|
155
|
+
# Sanity: every valid ring here is far outside a +/-2048 um window.
|
|
156
|
+
assert float(zl[base].abs().min()) > 2048.0
|
|
157
|
+
|
|
158
|
+
m_px = _model(_geom(
|
|
159
|
+
omega_ranges=[(-180.0, 180.0)],
|
|
160
|
+
box_sizes=[(-2048.0, 2048.0, -2048.0, 2048.0)],
|
|
161
|
+
))
|
|
162
|
+
assert float(m_px(EUL, POS).valid.sum()) == 0.0
|
|
163
|
+
|
|
164
|
+
def test_y_sign_convention(self):
|
|
165
|
+
"""yl = -sin(eta) * R: an eta = +90 deg spot lands at negative y."""
|
|
166
|
+
m = _model(_geom(
|
|
167
|
+
omega_ranges=[(-180.0, 180.0)],
|
|
168
|
+
box_sizes=[(-1e9, 1e9, -1e9, 1e9)],
|
|
169
|
+
))
|
|
170
|
+
two_theta = torch.tensor([[0.1]], dtype=torch.float64)
|
|
171
|
+
eta = torch.tensor([[math.pi / 2]], dtype=torch.float64)
|
|
172
|
+
omega = torch.zeros_like(eta)
|
|
173
|
+
R = float(LSD * math.tan(0.1))
|
|
174
|
+
# Positive-y half plane rejects it, negative-y half plane accepts it.
|
|
175
|
+
m._box_sizes = torch.tensor([[0.0, 1e9, -1e9, 1e9]], dtype=torch.float64)
|
|
176
|
+
assert not bool(m.omega_box_mask(omega, eta, two_theta).item())
|
|
177
|
+
m._box_sizes = torch.tensor([[-1e9, 0.0, -1e9, 1e9]], dtype=torch.float64)
|
|
178
|
+
assert bool(m.omega_box_mask(omega, eta, two_theta).item())
|
|
179
|
+
assert R > 0 # ring radius is positive, so the sign came from -sin
|
|
180
|
+
|
|
181
|
+
|
|
182
|
+
# ---------------------------------------------------------------------------
|
|
183
|
+
# Strict (exclusive) bounds -- C uses > and <, never >= / <=
|
|
184
|
+
# ---------------------------------------------------------------------------
|
|
185
|
+
|
|
186
|
+
class TestStrictBounds:
|
|
187
|
+
@staticmethod
|
|
188
|
+
def _one_spot(box):
|
|
189
|
+
"""Evaluate the mask for a single spot at eta = 0, omega = 0."""
|
|
190
|
+
m = _model(_geom(
|
|
191
|
+
omega_ranges=[(-180.0, 180.0)], box_sizes=[box],
|
|
192
|
+
))
|
|
193
|
+
two_theta = torch.tensor([[0.1]], dtype=torch.float64)
|
|
194
|
+
eta = torch.zeros_like(two_theta)
|
|
195
|
+
omega = torch.zeros_like(two_theta)
|
|
196
|
+
# zl = cos(0) * Lsd*tan(2theta); reproduce bit-for-bit.
|
|
197
|
+
zl = float(
|
|
198
|
+
torch.cos(eta) * (torch.tensor(LSD, dtype=torch.float64)
|
|
199
|
+
* torch.tan(two_theta))
|
|
200
|
+
)
|
|
201
|
+
return m, omega, eta, two_theta, zl
|
|
202
|
+
|
|
203
|
+
def test_zmin_edge_exact_is_rejected(self):
|
|
204
|
+
m, om, eta, tt, zl = self._one_spot((-1e6, 1e6, 0.0, 1e9))
|
|
205
|
+
m._box_sizes = torch.tensor([[-1e6, 1e6, zl, 1e9]], dtype=torch.float64)
|
|
206
|
+
assert not bool(m.omega_box_mask(om, eta, tt).item())
|
|
207
|
+
|
|
208
|
+
def test_zmax_edge_exact_is_rejected(self):
|
|
209
|
+
m, om, eta, tt, zl = self._one_spot((-1e6, 1e6, 0.0, 1e9))
|
|
210
|
+
m._box_sizes = torch.tensor([[-1e6, 1e6, -1e9, zl]], dtype=torch.float64)
|
|
211
|
+
assert not bool(m.omega_box_mask(om, eta, tt).item())
|
|
212
|
+
|
|
213
|
+
def test_strictly_inside_is_accepted(self):
|
|
214
|
+
m, om, eta, tt, zl = self._one_spot((-1e6, 1e6, 0.0, 1e9))
|
|
215
|
+
m._box_sizes = torch.tensor(
|
|
216
|
+
[[-1e6, 1e6, zl - 1e-6, zl + 1e-6]], dtype=torch.float64,
|
|
217
|
+
)
|
|
218
|
+
assert bool(m.omega_box_mask(om, eta, tt).item())
|
|
219
|
+
|
|
220
|
+
def test_omega_edge_exact_is_rejected(self):
|
|
221
|
+
m, om, eta, tt, zl = self._one_spot((-1e6, 1e6, -1e9, 1e9))
|
|
222
|
+
# omega == 0 exactly; a range starting at 0 must reject it.
|
|
223
|
+
m._omega_ranges = torch.tensor([[0.0, 180.0]], dtype=torch.float64)
|
|
224
|
+
assert not bool(m.omega_box_mask(om, eta, tt).item())
|
|
225
|
+
m._omega_ranges = torch.tensor([[-1e-9, 180.0]], dtype=torch.float64)
|
|
226
|
+
assert bool(m.omega_box_mask(om, eta, tt).item())
|
|
227
|
+
|
|
228
|
+
|
|
229
|
+
# ---------------------------------------------------------------------------
|
|
230
|
+
# Paired any-pair-accepts logic
|
|
231
|
+
# ---------------------------------------------------------------------------
|
|
232
|
+
|
|
233
|
+
class TestPairing:
|
|
234
|
+
def test_any_pair_accepts(self):
|
|
235
|
+
"""Pair 0 accepts on omega only, pair 1 on box only; each spot needs
|
|
236
|
+
BOTH halves of the SAME pair, and either pair is enough."""
|
|
237
|
+
m_off = _model(_geom())
|
|
238
|
+
yl, zl, om, _v = _nominal_ring_coords(m_off)
|
|
239
|
+
base = m_off(EUL, POS).valid.bool()
|
|
240
|
+
om_deg = om * (180.0 / math.pi)
|
|
241
|
+
|
|
242
|
+
# Split the valid spots by omega sign.
|
|
243
|
+
neg = base & (om_deg < 0)
|
|
244
|
+
pos = base & (om_deg > 0)
|
|
245
|
+
assert int(neg.sum()) > 0 and int(pos.sum()) > 0
|
|
246
|
+
|
|
247
|
+
z_hi = float(zl[base].max()) + 1.0
|
|
248
|
+
z_lo = float(zl[base].min()) - 1.0
|
|
249
|
+
|
|
250
|
+
# Pair 0: negative omega, box open. Pair 1: positive omega, box shut.
|
|
251
|
+
m2 = _model(_geom(
|
|
252
|
+
omega_ranges=[(-180.0, 0.0), (0.0, 180.0)],
|
|
253
|
+
box_sizes=[(-1e9, 1e9, z_lo, z_hi), (-1e9, 1e9, 1e8, 1e9)],
|
|
254
|
+
))
|
|
255
|
+
kept = m2(EUL, POS).valid.bool()
|
|
256
|
+
assert torch.equal(kept, neg)
|
|
257
|
+
|
|
258
|
+
# Swapping the boxes swaps which half survives -- the pairing is by
|
|
259
|
+
# index, not a global OR of all omegas with all boxes.
|
|
260
|
+
m3 = _model(_geom(
|
|
261
|
+
omega_ranges=[(-180.0, 0.0), (0.0, 180.0)],
|
|
262
|
+
box_sizes=[(-1e9, 1e9, 1e8, 1e9), (-1e9, 1e9, z_lo, z_hi)],
|
|
263
|
+
))
|
|
264
|
+
assert torch.equal(m3(EUL, POS).valid.bool(), pos)
|
|
265
|
+
|
|
266
|
+
def test_both_pairs_open_is_union(self):
|
|
267
|
+
m_off = _model(_geom())
|
|
268
|
+
yl, zl, om, _v = _nominal_ring_coords(m_off)
|
|
269
|
+
base = m_off(EUL, POS).valid.bool()
|
|
270
|
+
z_hi = float(zl[base].max()) + 1.0
|
|
271
|
+
z_lo = float(zl[base].min()) - 1.0
|
|
272
|
+
m2 = _model(_geom(
|
|
273
|
+
omega_ranges=[(-180.0, 0.0), (0.0, 180.0)],
|
|
274
|
+
box_sizes=[(-1e9, 1e9, z_lo, z_hi), (-1e9, 1e9, z_lo, z_hi)],
|
|
275
|
+
))
|
|
276
|
+
assert torch.equal(m2(EUL, POS).valid.bool(), base)
|
|
277
|
+
|
|
278
|
+
|
|
279
|
+
# ---------------------------------------------------------------------------
|
|
280
|
+
# Arity validation
|
|
281
|
+
# ---------------------------------------------------------------------------
|
|
282
|
+
|
|
283
|
+
class TestValidation:
|
|
284
|
+
def test_box_without_omega_range_raises(self):
|
|
285
|
+
with pytest.raises(ValueError, match="paired filter"):
|
|
286
|
+
_model(_geom(box_sizes=[(-1e6, 1e6, -1e6, 1e6)]))
|
|
287
|
+
|
|
288
|
+
def test_omega_range_without_box_raises(self):
|
|
289
|
+
with pytest.raises(ValueError, match="paired filter"):
|
|
290
|
+
_model(_geom(omega_ranges=[(-180.0, 180.0)]))
|
|
291
|
+
|
|
292
|
+
def test_length_mismatch_raises(self):
|
|
293
|
+
with pytest.raises(ValueError, match="must match box_sizes"):
|
|
294
|
+
_model(_geom(
|
|
295
|
+
omega_ranges=[(-180.0, 180.0)],
|
|
296
|
+
box_sizes=[(-1e6, 1e6, -1e6, 1e6), (-1e6, 1e6, -1e6, 1e6)],
|
|
297
|
+
))
|
|
298
|
+
|
|
299
|
+
|
|
300
|
+
# ---------------------------------------------------------------------------
|
|
301
|
+
# The mask reaches every downstream consumer of ``valid``
|
|
302
|
+
# ---------------------------------------------------------------------------
|
|
303
|
+
|
|
304
|
+
class TestPropagation:
|
|
305
|
+
def _pair(self):
|
|
306
|
+
m_off = _model(_geom())
|
|
307
|
+
yl, zl, _om, _v = _nominal_ring_coords(m_off)
|
|
308
|
+
base = m_off(EUL, POS).valid.bool()
|
|
309
|
+
z_cut = float(zl[base].min()) + 1.0
|
|
310
|
+
m_box = _model(_geom(
|
|
311
|
+
omega_ranges=[(-180.0, 180.0)],
|
|
312
|
+
box_sizes=[(-1e9, 1e9, z_cut, 1e9)],
|
|
313
|
+
))
|
|
314
|
+
return m_off, m_box
|
|
315
|
+
|
|
316
|
+
def test_layer_valid_inherits_in_multi_distance(self):
|
|
317
|
+
hkls_int = torch.tensor([[1, 1, 1], [2, 0, 0]], dtype=torch.float64)
|
|
318
|
+
B = torch.eye(3, dtype=torch.float64) / 4.08
|
|
319
|
+
hkls_cart = hkls_int @ B.T
|
|
320
|
+
g = torch.linalg.norm(hkls_cart, dim=-1)
|
|
321
|
+
thetas = torch.asin(g * 0.172979 / 2.0)
|
|
322
|
+
|
|
323
|
+
def build(**kw):
|
|
324
|
+
geom = _geom(
|
|
325
|
+
Lsd=[LSD, 1.05 * LSD],
|
|
326
|
+
y_BC=[1024.0, 1024.0],
|
|
327
|
+
z_BC=[1024.0, 1024.0],
|
|
328
|
+
**kw,
|
|
329
|
+
)
|
|
330
|
+
return HEDMForwardModel(
|
|
331
|
+
hkls=hkls_cart, thetas=thetas, geometry=geom,
|
|
332
|
+
hkls_int=hkls_int, device=torch.device("cpu"),
|
|
333
|
+
)
|
|
334
|
+
|
|
335
|
+
m_off = build()
|
|
336
|
+
s_off = m_off(EUL, POS)
|
|
337
|
+
assert s_off.layer_valid is not None
|
|
338
|
+
assert float(s_off.valid.sum()) > 0
|
|
339
|
+
|
|
340
|
+
m_box = build(
|
|
341
|
+
omega_ranges=[(-180.0, 180.0)],
|
|
342
|
+
box_sizes=[(-1e9, 1e9, 1e8, 1e9)], # shut
|
|
343
|
+
)
|
|
344
|
+
s_box = m_box(EUL, POS)
|
|
345
|
+
assert float(s_box.valid.sum()) == 0.0
|
|
346
|
+
assert float(s_box.layer_valid.sum()) == 0.0
|
|
347
|
+
|
|
348
|
+
def test_forward_per_grain_agrees_with_forward(self):
|
|
349
|
+
_m_off, m_box = self._pair()
|
|
350
|
+
s_fwd = m_box(EUL, POS)
|
|
351
|
+
s_pg = m_box.forward_per_grain(EUL, POS)
|
|
352
|
+
assert torch.equal(s_fwd.valid.reshape(-1), s_pg.valid.reshape(-1))
|
|
@@ -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
|