midas-diffract 0.4.0__tar.gz → 0.7.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.4.0 → midas_diffract-0.7.0}/PKG-INFO +1 -1
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/__init__.py +1 -1
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/forward.py +270 -5
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/simulate_panel_zarrs.py +12 -16
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/PKG-INFO +1 -1
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/SOURCES.txt +1 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/pyproject.toml +1 -1
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_forward.py +84 -0
- midas_diffract-0.7.0/tests/test_omega_box_filter.py +352 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/LICENSE +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/README.md +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/hkls.py +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/losses.py +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/optimize.py +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/dependency_links.txt +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/requires.txt +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/top_level.txt +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/setup.cfg +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_c_comparison.py +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_distortion_layer.py +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_hkls.py +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_losses.py +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_multi_detector.py +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_strain_tensor.py +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_tilts.py +0 -0
- {midas_diffract-0.4.0 → midas_diffract-0.7.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.7.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
|
|
@@ -27,6 +27,7 @@ Reference C code:
|
|
|
27
27
|
"""
|
|
28
28
|
|
|
29
29
|
import math
|
|
30
|
+
import warnings
|
|
30
31
|
from dataclasses import dataclass, field
|
|
31
32
|
from typing import Optional, Tuple
|
|
32
33
|
|
|
@@ -116,6 +117,26 @@ class HEDMGeometry:
|
|
|
116
117
|
# spot is valid if it lands on at least
|
|
117
118
|
# one detector; output gains a det_id
|
|
118
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
|
|
119
140
|
|
|
120
141
|
@property
|
|
121
142
|
def n_distances(self) -> int:
|
|
@@ -398,6 +419,38 @@ class HEDMForwardModel(nn.Module):
|
|
|
398
419
|
)
|
|
399
420
|
self._has_wedge = abs(float(geometry.wedge)) > 0.0
|
|
400
421
|
|
|
422
|
+
# Paired OmegaRange / BoxSize acceptance filter. OFF unless BOTH
|
|
423
|
+
# lists are supplied and non-empty (see HEDMGeometry.box_sizes).
|
|
424
|
+
om_ranges = getattr(geometry, "omega_ranges", None) or []
|
|
425
|
+
bx_sizes = getattr(geometry, "box_sizes", None) or []
|
|
426
|
+
if bool(om_ranges) != bool(bx_sizes):
|
|
427
|
+
raise ValueError(
|
|
428
|
+
"omega_ranges and box_sizes are a paired filter: supply both "
|
|
429
|
+
f"or neither (got {len(om_ranges)} omega_ranges and "
|
|
430
|
+
f"{len(bx_sizes)} box_sizes)."
|
|
431
|
+
)
|
|
432
|
+
if len(om_ranges) != len(bx_sizes):
|
|
433
|
+
raise ValueError(
|
|
434
|
+
f"omega_ranges (len {len(om_ranges)}) must match box_sizes "
|
|
435
|
+
f"(len {len(bx_sizes)}); entry i of one pairs with entry i "
|
|
436
|
+
"of the other."
|
|
437
|
+
)
|
|
438
|
+
self._has_omega_box = bool(bx_sizes)
|
|
439
|
+
self.register_buffer(
|
|
440
|
+
"_omega_ranges",
|
|
441
|
+
torch.tensor(
|
|
442
|
+
[[float(a), float(b)] for a, b in om_ranges],
|
|
443
|
+
dtype=torch.float64, device=device,
|
|
444
|
+
).reshape(-1, 2),
|
|
445
|
+
)
|
|
446
|
+
self.register_buffer(
|
|
447
|
+
"_box_sizes",
|
|
448
|
+
torch.tensor(
|
|
449
|
+
[[float(v) for v in bs] for bs in bx_sizes],
|
|
450
|
+
dtype=torch.float64, device=device,
|
|
451
|
+
).reshape(-1, 4),
|
|
452
|
+
)
|
|
453
|
+
|
|
401
454
|
# Scan config
|
|
402
455
|
self.scan_config = scan_config
|
|
403
456
|
if scan_config is not None:
|
|
@@ -738,6 +791,7 @@ class HEDMForwardModel(nn.Module):
|
|
|
738
791
|
orientation_matrices: torch.Tensor,
|
|
739
792
|
hkls_cart: Optional[torch.Tensor] = None,
|
|
740
793
|
thetas: Optional[torch.Tensor] = None,
|
|
794
|
+
per_grain: bool = False,
|
|
741
795
|
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
|
742
796
|
"""Core Bragg geometry: orientations + G-vectors -> angles.
|
|
743
797
|
|
|
@@ -779,14 +833,32 @@ class HEDMForwardModel(nn.Module):
|
|
|
779
833
|
# batch, (b) per-voxel hkls_cart shape (..., M, 3) for strained
|
|
780
834
|
# rendering. Both flow through the same einsum via leading-dim
|
|
781
835
|
# broadcasting on the second arg.
|
|
782
|
-
|
|
836
|
+
#
|
|
837
|
+
# per_grain=True: ELEMENT-WISE pairing of grain i's orientation with
|
|
838
|
+
# grain i's strained hkls -- NO orientation x strain cross-product.
|
|
839
|
+
# Requires orientation_matrices (N,3,3); hkls_cart (N,M,3) per-grain or
|
|
840
|
+
# (M,3) shared. Used by forward_per_grain() for the O(N*M) fast path.
|
|
841
|
+
if per_grain:
|
|
842
|
+
if hkls_cart.dim() == 2: # (M,3) shared lattice
|
|
843
|
+
G_C = torch.einsum("nij,mj->nmi", orientation_matrices, hkls_cart)
|
|
844
|
+
else: # (N,M,3) per-grain strain
|
|
845
|
+
G_C = torch.einsum("nij,nmj->nmi", orientation_matrices, hkls_cart)
|
|
846
|
+
else:
|
|
847
|
+
G_C = torch.einsum("...nij,...mj->...nmi", orientation_matrices, hkls_cart)
|
|
783
848
|
|
|
784
849
|
# v = sin(theta)*|G| -- C precomputes Gs from the UNROTATED G-vector norm
|
|
785
850
|
# (rotation preserves norm in exact arithmetic but not in float64).
|
|
786
851
|
# Match C: use |hkls_cart| (pre-rotation), not |R @ hkls_cart|.
|
|
787
852
|
len_hkl = torch.norm(hkls_cart, dim=-1) # (M,) or (..., M)
|
|
788
853
|
v_no_wedge = torch.sin(thetas) * len_hkl # (M,) or (..., M)
|
|
789
|
-
|
|
854
|
+
# Broadcast v to G_C's (..., N, M) grid. Layout-agnostic: in the
|
|
855
|
+
# cross-product layout v gains an orientation axis at -2; in the
|
|
856
|
+
# per_grain / shared layouts it already matches. Numerically identical
|
|
857
|
+
# to the former ``unsqueeze(-2).expand_as`` for those existing paths.
|
|
858
|
+
gc0 = G_C[..., 0]
|
|
859
|
+
while v_no_wedge.dim() < gc0.dim():
|
|
860
|
+
v_no_wedge = v_no_wedge.unsqueeze(-2)
|
|
861
|
+
v_no_wedge = v_no_wedge.expand_as(gc0)
|
|
790
862
|
|
|
791
863
|
# ---- Wedge: rigorous geometric formulation -------------------
|
|
792
864
|
# The rotation axis tilts from z to n_hat = (sin W, 0, cos W).
|
|
@@ -900,9 +972,16 @@ class HEDMForwardModel(nn.Module):
|
|
|
900
972
|
eta = self.safe_arccos(Gz_lab / r_yz)
|
|
901
973
|
eta = -torch.sign(Gy_lab) * eta
|
|
902
974
|
|
|
903
|
-
# 2*theta
|
|
904
|
-
|
|
905
|
-
|
|
975
|
+
# 2*theta -- broadcast to the single-branch (pre-cat) grid, then double
|
|
976
|
+
# along the grain axis to match all_omega's 2N. Layout-agnostic (handles
|
|
977
|
+
# the per_grain layout where the 2N axis is grain-doubled) and
|
|
978
|
+
# numerically identical to the former expand for the cross-product path
|
|
979
|
+
# (thetas does not depend on the orientation axis).
|
|
980
|
+
tt = 2.0 * thetas
|
|
981
|
+
while tt.dim() < omega_p.dim():
|
|
982
|
+
tt = tt.unsqueeze(-2)
|
|
983
|
+
tt = tt.expand_as(omega_p)
|
|
984
|
+
two_theta = torch.cat([tt, tt], dim=-2)
|
|
906
985
|
|
|
907
986
|
# Validity mask
|
|
908
987
|
valid_p = disc_valid & coswp_valid
|
|
@@ -917,8 +996,93 @@ class HEDMForwardModel(nn.Module):
|
|
|
917
996
|
((math.pi - torch.abs(eta)) >= self.min_eta)
|
|
918
997
|
valid = valid * eta_ok.float()
|
|
919
998
|
|
|
999
|
+
# Paired OmegaRange / BoxSize gate (no-op unless both were supplied).
|
|
1000
|
+
if self._has_omega_box:
|
|
1001
|
+
valid = valid * self.omega_box_mask(
|
|
1002
|
+
all_omega, eta, two_theta
|
|
1003
|
+
).to(valid.dtype)
|
|
1004
|
+
|
|
920
1005
|
return all_omega, eta, two_theta, valid
|
|
921
1006
|
|
|
1007
|
+
# ------------------------------------------------------------------
|
|
1008
|
+
# omega_box_mask (port of the KeepSpot gate in CalcDiffrSpots_Furnace)
|
|
1009
|
+
# ------------------------------------------------------------------
|
|
1010
|
+
|
|
1011
|
+
def omega_box_mask(
|
|
1012
|
+
self,
|
|
1013
|
+
omega: torch.Tensor,
|
|
1014
|
+
eta: torch.Tensor,
|
|
1015
|
+
two_theta: torch.Tensor,
|
|
1016
|
+
) -> torch.Tensor:
|
|
1017
|
+
"""Paired ``OmegaRange`` / ``BoxSize`` acceptance mask.
|
|
1018
|
+
|
|
1019
|
+
Direct port of the ``KeepSpot`` gate in
|
|
1020
|
+
``NF_HEDM/src/CalcDiffrSpots_Furnace`` (CalcDiffractionSpots.c:225-236,
|
|
1021
|
+
identical code at MakeDiffrSpots.c:248-259)::
|
|
1022
|
+
|
|
1023
|
+
CalcSpotPosition(RingRadius, eta, &yl, &zl); // C:224
|
|
1024
|
+
for (OmegaRangeNo = 0; OmegaRangeNo < NOmegaRanges; OmegaRangeNo++) {
|
|
1025
|
+
KeepSpot = 0;
|
|
1026
|
+
if ((Omega > OmegaRange[i][0]) && (Omega < OmegaRange[i][1]) &&
|
|
1027
|
+
(yl > BoxSizes[i][0]) && (yl < BoxSizes[i][1]) &&
|
|
1028
|
+
(zl > BoxSizes[i][2]) && (zl < BoxSizes[i][3])) {
|
|
1029
|
+
KeepSpot = 1; break;
|
|
1030
|
+
}
|
|
1031
|
+
}
|
|
1032
|
+
|
|
1033
|
+
with ``RingRadius = Lsd[0] * tan(2*theta)`` (C:209, and the caller
|
|
1034
|
+
passes ``Lsd[0]`` -- SharedFuncsFit.c:830), and
|
|
1035
|
+
``yl = -sin(eta)*RingRadius``, ``zl = cos(eta)*RingRadius``
|
|
1036
|
+
(CalcDiffractionSpots.c:80-85).
|
|
1037
|
+
|
|
1038
|
+
Semantics that matter:
|
|
1039
|
+
|
|
1040
|
+
- ``(yl, zl)`` are the **nominal ring** coordinates in micrometres at
|
|
1041
|
+
the FIRST distance, with **no** grain displacement, **no** tilt and
|
|
1042
|
+
**no** distortion applied. They are not pixel coordinates and carry
|
|
1043
|
+
no beam-centre offset.
|
|
1044
|
+
- Bounds are **strict** on all six comparisons (edge-exact spots are
|
|
1045
|
+
rejected).
|
|
1046
|
+
- Range ``i`` pairs with box ``i``; a spot is kept if **any** pair
|
|
1047
|
+
accepts (``break`` on first hit).
|
|
1048
|
+
- A rejected spot never enters ``TheorSpots``, so ``nTspots`` shrinks
|
|
1049
|
+
and the spot is dropped from **both** the numerator and the
|
|
1050
|
+
denominator of ``CalcFracOverlap`` (SharedFuncsFit.c:645, 648-650) --
|
|
1051
|
+
it is *not* counted as a miss. Folding the mask into ``valid`` here
|
|
1052
|
+
reproduces that, because ``valid`` is the denominator weight in
|
|
1053
|
+
``ObsVolume.hard_fraction`` / ``soft_fraction``.
|
|
1054
|
+
|
|
1055
|
+
Parameters
|
|
1056
|
+
----------
|
|
1057
|
+
omega, eta, two_theta : Tensor
|
|
1058
|
+
Radians, all broadcastable to a common shape.
|
|
1059
|
+
|
|
1060
|
+
Returns
|
|
1061
|
+
-------
|
|
1062
|
+
Tensor (bool), same broadcast shape.
|
|
1063
|
+
"""
|
|
1064
|
+
dtype = omega.dtype
|
|
1065
|
+
# C uses Lsd[0] for the ring radius regardless of nDistances.
|
|
1066
|
+
Lsd0 = self._Lsd_eff.to(dtype).reshape(-1)[0]
|
|
1067
|
+
ring_radius = Lsd0 * torch.tan(two_theta)
|
|
1068
|
+
yl = -torch.sin(eta) * ring_radius
|
|
1069
|
+
zl = torch.cos(eta) * ring_radius
|
|
1070
|
+
omega_deg = omega * self.RAD2DEG
|
|
1071
|
+
|
|
1072
|
+
om = self._omega_ranges.to(dtype)
|
|
1073
|
+
bx = self._box_sizes.to(dtype)
|
|
1074
|
+
accept = torch.zeros(
|
|
1075
|
+
torch.broadcast_shapes(omega_deg.shape, yl.shape, zl.shape),
|
|
1076
|
+
dtype=torch.bool, device=omega.device,
|
|
1077
|
+
)
|
|
1078
|
+
for i in range(om.shape[0]):
|
|
1079
|
+
accept = accept | (
|
|
1080
|
+
(omega_deg > om[i, 0]) & (omega_deg < om[i, 1])
|
|
1081
|
+
& (yl > bx[i, 0]) & (yl < bx[i, 1])
|
|
1082
|
+
& (zl > bx[i, 2]) & (zl < bx[i, 3])
|
|
1083
|
+
)
|
|
1084
|
+
return accept
|
|
1085
|
+
|
|
922
1086
|
# ------------------------------------------------------------------
|
|
923
1087
|
# Tilt rotation matrix (RotationTilts in SharedFuncsFit.c:230-266)
|
|
924
1088
|
# ------------------------------------------------------------------
|
|
@@ -1354,6 +1518,11 @@ class HEDMForwardModel(nn.Module):
|
|
|
1354
1518
|
if positions.shape[-1] == 2:
|
|
1355
1519
|
positions = F.pad(positions, (0, 1), value=0.0)
|
|
1356
1520
|
|
|
1521
|
+
# Footgun guard: per-grain lattice/strain + N>1 orientations forms an
|
|
1522
|
+
# N x N orientation x strain cross-product (output (N, 2N, M)); callers
|
|
1523
|
+
# simulating a fixed polycrystal almost always want only the diagonal.
|
|
1524
|
+
self._warn_if_cross_product(euler_angles, lattice_params, strain)
|
|
1525
|
+
|
|
1357
1526
|
# 1. Orientation matrices
|
|
1358
1527
|
orientation_matrices = self.euler2mat(euler_angles)
|
|
1359
1528
|
|
|
@@ -1382,6 +1551,102 @@ class HEDMForwardModel(nn.Module):
|
|
|
1382
1551
|
|
|
1383
1552
|
return spots
|
|
1384
1553
|
|
|
1554
|
+
@staticmethod
|
|
1555
|
+
def _warn_if_cross_product(euler_angles, lattice_params, strain):
|
|
1556
|
+
"""Warn when forward() would form an orientation x strain cross-product.
|
|
1557
|
+
|
|
1558
|
+
Fires only when there are N>1 orientations AND lattice_params/strain
|
|
1559
|
+
carry a matching per-grain axis -- the case where the (N, 2N, M) output
|
|
1560
|
+
is an N x N cross-product and the caller likely wanted the diagonal.
|
|
1561
|
+
Shared lattice/strain (no grain axis) is the correct (2N, M) path and
|
|
1562
|
+
does not warn.
|
|
1563
|
+
"""
|
|
1564
|
+
n_orient = euler_angles.shape[-2] if euler_angles.dim() >= 2 else 1
|
|
1565
|
+
if n_orient <= 1:
|
|
1566
|
+
return
|
|
1567
|
+
|
|
1568
|
+
def _has_grain_axis(t):
|
|
1569
|
+
if t is None:
|
|
1570
|
+
return False
|
|
1571
|
+
if t.shape[-1] == 6 and t.dim() >= 2: # Voigt lattice or strain
|
|
1572
|
+
return t.shape[-2] == n_orient
|
|
1573
|
+
if tuple(t.shape[-2:]) == (3, 3) and t.dim() >= 3: # full-tensor strain
|
|
1574
|
+
return t.shape[-3] == n_orient
|
|
1575
|
+
return False
|
|
1576
|
+
|
|
1577
|
+
if _has_grain_axis(lattice_params) or _has_grain_axis(strain):
|
|
1578
|
+
warnings.warn(
|
|
1579
|
+
f"forward() called with {n_orient} orientations and per-grain "
|
|
1580
|
+
"lattice_params/strain forms an orientation x strain "
|
|
1581
|
+
f"cross-product (output shape ({n_orient}, {2 * n_orient}, M)); "
|
|
1582
|
+
"only the diagonal [i, i] and [i, i+N] is physical. For a fixed "
|
|
1583
|
+
"polycrystal use forward_per_grain() (element-wise, O(N*M)), "
|
|
1584
|
+
"or index the diagonal of this output.",
|
|
1585
|
+
stacklevel=3,
|
|
1586
|
+
)
|
|
1587
|
+
|
|
1588
|
+
def forward_per_grain(
|
|
1589
|
+
self,
|
|
1590
|
+
euler_angles: torch.Tensor,
|
|
1591
|
+
positions: torch.Tensor,
|
|
1592
|
+
lattice_params: Optional[torch.Tensor] = None,
|
|
1593
|
+
strain: Optional[torch.Tensor] = None,
|
|
1594
|
+
) -> SpotDescriptors:
|
|
1595
|
+
"""Element-wise per-grain forward simulation -- the fast path.
|
|
1596
|
+
|
|
1597
|
+
Grain ``i`` is simulated with orientation ``i``, lattice/strain ``i``
|
|
1598
|
+
and position ``i``, WITHOUT the orientation x strain cross-product that
|
|
1599
|
+
:meth:`forward` forms when both are per-grain. The output has leading
|
|
1600
|
+
shape ``(2N, M)`` (the two omega branches doubled along the grain axis),
|
|
1601
|
+
which is exactly the diagonal of :meth:`forward`'s ``(N, 2N, M)`` output
|
|
1602
|
+
-- so gradient and gradient-free callers agree bit-for-bit.
|
|
1603
|
+
|
|
1604
|
+
Cost is O(N*M) rather than O(N^2 * M), matching the algorithm of the C
|
|
1605
|
+
reference ``ForwardSimulationCompressed.c``. Fully differentiable and
|
|
1606
|
+
device-portable; for pure forward simulation wrap the call in
|
|
1607
|
+
``torch.inference_mode()``.
|
|
1608
|
+
|
|
1609
|
+
Parameters
|
|
1610
|
+
----------
|
|
1611
|
+
euler_angles : Tensor (N, 3)
|
|
1612
|
+
Bunge ZXZ Euler angles (radians), one per grain. No leading batch.
|
|
1613
|
+
positions : Tensor (N, 3) or (N, 2)
|
|
1614
|
+
Real-space grain positions (micrometers).
|
|
1615
|
+
lattice_params : Tensor (6,) or (N, 6), optional
|
|
1616
|
+
Shared or per-grain lattice [a,b,c,alpha,beta,gamma] (Ang/deg).
|
|
1617
|
+
strain : Tensor (6,), (N, 6), (3, 3), or (N, 3, 3), optional
|
|
1618
|
+
Shared or per-grain crystal-frame strain (plain-Voigt or full 3x3).
|
|
1619
|
+
|
|
1620
|
+
Returns
|
|
1621
|
+
-------
|
|
1622
|
+
SpotDescriptors with leading shape ``(2N, M)``: grain ``i``'s two omega
|
|
1623
|
+
solutions live at axis-(-2) indices ``i`` and ``i + N``.
|
|
1624
|
+
"""
|
|
1625
|
+
if positions.shape[-1] == 2:
|
|
1626
|
+
positions = F.pad(positions, (0, 1), value=0.0)
|
|
1627
|
+
|
|
1628
|
+
orientation_matrices = self.euler2mat(euler_angles) # (N, 3, 3)
|
|
1629
|
+
|
|
1630
|
+
hkls_cart = None
|
|
1631
|
+
thetas = None
|
|
1632
|
+
if lattice_params is not None:
|
|
1633
|
+
hkls_cart, thetas = self.correct_hkls_latc(lattice_params, strain=strain)
|
|
1634
|
+
elif strain is not None:
|
|
1635
|
+
raise ValueError(
|
|
1636
|
+
"strain was supplied but lattice_params is None; strain "
|
|
1637
|
+
"requires a reference lattice to apply (I + eps)^{-1} @ B0."
|
|
1638
|
+
)
|
|
1639
|
+
|
|
1640
|
+
omega, eta, two_theta, valid = self.calc_bragg_geometry(
|
|
1641
|
+
orientation_matrices, hkls_cart, thetas, per_grain=True
|
|
1642
|
+
)
|
|
1643
|
+
spots = self.project_to_detector(omega, eta, two_theta, positions, valid)
|
|
1644
|
+
|
|
1645
|
+
if self.scan_config is not None:
|
|
1646
|
+
spots = self.filter_by_scan(spots, positions)
|
|
1647
|
+
|
|
1648
|
+
return spots
|
|
1649
|
+
|
|
1385
1650
|
# ------------------------------------------------------------------
|
|
1386
1651
|
# filter_by_scan (beam proximity for pf-HEDM)
|
|
1387
1652
|
# ------------------------------------------------------------------
|
|
@@ -613,18 +613,13 @@ def simulate_panel_zarrs(
|
|
|
613
613
|
f"BC=({panel.y_bc:.1f},{panel.z_bc:.1f}), "
|
|
614
614
|
f"tx={panel.tx:.2f}° ty={panel.ty:.2f}° tz={panel.tz:.2f}°")
|
|
615
615
|
with torch.no_grad():
|
|
616
|
-
|
|
617
|
-
|
|
618
|
-
|
|
619
|
-
|
|
620
|
-
|
|
621
|
-
|
|
622
|
-
|
|
623
|
-
diag_mask = np.zeros((G_strain, Kdim, Mdim), dtype=bool)
|
|
624
|
-
for gi in range(G_strain):
|
|
625
|
-
diag_mask[gi, gi, :] = True
|
|
626
|
-
diag_mask[gi, gi + n_grains, :] = True
|
|
627
|
-
valid_np = valid_np & diag_mask
|
|
616
|
+
# Per-grain fast path: O(N*M), no orientation x strain cross-product.
|
|
617
|
+
# Output is (2N, M) -- already the per-grain diagonal, no mask needed.
|
|
618
|
+
spots = model.forward_per_grain(eulers_t, positions_t,
|
|
619
|
+
lattice_params=latc, strain=strain_t)
|
|
620
|
+
|
|
621
|
+
valid_np = (spots.valid > 0.5).cpu().numpy() # (2N, M)
|
|
622
|
+
Kdim, Mdim = valid_np.shape # Kdim == 2 * n_grains
|
|
628
623
|
|
|
629
624
|
# First-wins: drop spots that an earlier panel already took
|
|
630
625
|
# (avoids double-counting in the rare overlap).
|
|
@@ -652,10 +647,11 @@ def simulate_panel_zarrs(
|
|
|
652
647
|
y_pix = spots.y_pixel.cpu().numpy()
|
|
653
648
|
z_pix = spots.z_pixel.cpu().numpy()
|
|
654
649
|
frame_nr = spots.frame_nr.cpu().numpy()
|
|
655
|
-
|
|
656
|
-
|
|
657
|
-
|
|
658
|
-
|
|
650
|
+
# (2N, M): row k holds grain (k % n_grains); rows [0:N]/[N:2N] are the
|
|
651
|
+
# two omega branches.
|
|
652
|
+
grain_ids = np.broadcast_to((np.arange(Kdim) % n_grains)[:, None],
|
|
653
|
+
(Kdim, Mdim))
|
|
654
|
+
hkl_ids = np.broadcast_to(np.arange(Mdim)[None, :], (Kdim, Mdim))
|
|
659
655
|
|
|
660
656
|
rec = {
|
|
661
657
|
"grain_id": grain_ids.reshape(-1)[flat_idx].astype(np.int32),
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: midas-diffract
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.7.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
|
|
@@ -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.7.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"
|
|
@@ -1201,5 +1201,89 @@ class TestStrainTensorInput:
|
|
|
1201
1201
|
torch.testing.assert_close(sp_v.omega, sp_S.omega)
|
|
1202
1202
|
|
|
1203
1203
|
|
|
1204
|
+
# ===================================================================
|
|
1205
|
+
# Test: forward_per_grain (fast path) == diagonal of forward()
|
|
1206
|
+
# ===================================================================
|
|
1207
|
+
|
|
1208
|
+
class TestForwardPerGrain:
|
|
1209
|
+
"""The O(N*M) per-grain path must equal the diagonal of forward()'s
|
|
1210
|
+
O(N^2*M) orientation x strain cross-product, bit-for-bit, so the
|
|
1211
|
+
gradient and gradient-free callers agree."""
|
|
1212
|
+
|
|
1213
|
+
def _setup(self, nf_geometry, device, N=6, strained=True):
|
|
1214
|
+
model, _, _ = make_model_with_cubic_iron(nf_geometry, device)
|
|
1215
|
+
torch.manual_seed(7)
|
|
1216
|
+
euler = torch.rand(N, 3, dtype=torch.float64) * 2 * math.pi
|
|
1217
|
+
pos = torch.arange(N * 3, dtype=torch.float64).reshape(N, 3) * 5.0
|
|
1218
|
+
latc = torch.tensor([2.87, 2.87, 2.87, 90., 90., 90.],
|
|
1219
|
+
dtype=torch.float64).expand(N, 6)
|
|
1220
|
+
if strained:
|
|
1221
|
+
s = torch.randn(N, 3, 3, dtype=torch.float64) * 1e-3
|
|
1222
|
+
strain = (s + s.transpose(-1, -2)) / 2
|
|
1223
|
+
else:
|
|
1224
|
+
strain = None
|
|
1225
|
+
return model, euler, pos, latc, strain, N
|
|
1226
|
+
|
|
1227
|
+
def test_per_grain_equals_forward_diagonal_strained(self, nf_geometry, device):
|
|
1228
|
+
# Compare per-grain (no global sort, which would reorder near-degenerate
|
|
1229
|
+
# spots): forward() diagonal rows [gi,gi] & [gi,gi+N] vs
|
|
1230
|
+
# forward_per_grain() rows [gi] & [gi+N]. Element-wise einsum differs
|
|
1231
|
+
# from the cross-product einsum only by fp64 reduction order (~1e-12).
|
|
1232
|
+
model, euler, pos, latc, strain, N = self._setup(nf_geometry, device)
|
|
1233
|
+
spf = model(euler, pos, lattice_params=latc, strain=strain) # (N,2N,M)
|
|
1234
|
+
spg = model.forward_per_grain(euler, pos, lattice_params=latc, strain=strain) # (2N,M)
|
|
1235
|
+
|
|
1236
|
+
vf = spf.valid > 0.5
|
|
1237
|
+
vg = spg.valid > 0.5
|
|
1238
|
+
max_diff = 0.0
|
|
1239
|
+
for fld in ("two_theta", "eta", "omega", "y_pixel", "z_pixel", "frame_nr"):
|
|
1240
|
+
F = getattr(spf, fld)
|
|
1241
|
+
G = getattr(spg, fld)
|
|
1242
|
+
for gi in range(N):
|
|
1243
|
+
for kf, kg in ((gi, gi), (gi + N, gi + N)):
|
|
1244
|
+
# validity must agree exactly on the diagonal
|
|
1245
|
+
assert torch.equal(vf[gi, kf], vg[kg])
|
|
1246
|
+
max_diff = max(max_diff, float((F[gi, kf] - G[kg]).abs().max()))
|
|
1247
|
+
assert max_diff < 1e-9, f"per-grain vs forward-diagonal max diff {max_diff:.2e}"
|
|
1248
|
+
|
|
1249
|
+
def test_per_grain_shape_is_2N(self, nf_geometry, device):
|
|
1250
|
+
model, euler, pos, latc, strain, N = self._setup(nf_geometry, device)
|
|
1251
|
+
sp = model.forward_per_grain(euler, pos, lattice_params=latc, strain=strain)
|
|
1252
|
+
assert sp.valid.shape[0] == 2 * N
|
|
1253
|
+
|
|
1254
|
+
def test_per_grain_shared_lattice_no_strain(self, nf_geometry, device):
|
|
1255
|
+
# nominal (no lattice) path must also work and avoid cross-product
|
|
1256
|
+
model, euler, pos, latc, _, N = self._setup(nf_geometry, device, strained=False)
|
|
1257
|
+
sp = model.forward_per_grain(euler, pos)
|
|
1258
|
+
assert sp.valid.shape[0] == 2 * N
|
|
1259
|
+
|
|
1260
|
+
def test_per_grain_differentiable(self, nf_geometry, device):
|
|
1261
|
+
model, euler, pos, latc, strain, N = self._setup(nf_geometry, device)
|
|
1262
|
+
euler = euler.clone().requires_grad_(True)
|
|
1263
|
+
sp = model.forward_per_grain(euler, pos, lattice_params=latc, strain=strain)
|
|
1264
|
+
# omega/eta/pixels depend on orientation (two_theta does not)
|
|
1265
|
+
(sp.omega * sp.valid).sum().backward()
|
|
1266
|
+
assert euler.grad is not None and torch.all(torch.isfinite(euler.grad))
|
|
1267
|
+
|
|
1268
|
+
def test_per_grain_fp32(self, nf_geometry, device):
|
|
1269
|
+
model, euler, pos, latc, strain, N = self._setup(nf_geometry, device)
|
|
1270
|
+
sp = model.forward_per_grain(euler.float(), pos.float(),
|
|
1271
|
+
lattice_params=latc.float(), strain=strain.float())
|
|
1272
|
+
assert sp.valid.dtype == torch.float32 or sp.valid.dtype == torch.float64
|
|
1273
|
+
|
|
1274
|
+
def test_forward_warns_on_cross_product(self, nf_geometry, device):
|
|
1275
|
+
model, euler, pos, latc, strain, N = self._setup(nf_geometry, device)
|
|
1276
|
+
with pytest.warns(UserWarning, match="cross-product"):
|
|
1277
|
+
model(euler, pos, lattice_params=latc, strain=strain)
|
|
1278
|
+
|
|
1279
|
+
def test_forward_no_warn_shared_lattice(self, nf_geometry, device):
|
|
1280
|
+
import warnings as _w
|
|
1281
|
+
model, euler, pos, _, _, N = self._setup(nf_geometry, device, strained=False)
|
|
1282
|
+
shared = torch.tensor([2.87, 2.87, 2.87, 90., 90., 90.], dtype=torch.float64)
|
|
1283
|
+
with _w.catch_warnings():
|
|
1284
|
+
_w.simplefilter("error") # any cross-product warning becomes an error
|
|
1285
|
+
model(euler, pos, lattice_params=shared) # shared -> (2N,M), no warn
|
|
1286
|
+
|
|
1287
|
+
|
|
1204
1288
|
if __name__ == "__main__":
|
|
1205
1289
|
pytest.main([__file__, "-v"])
|
|
@@ -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))
|
|
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
|