midas-diffract 0.6.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.6.0 → midas_diffract-0.7.0}/PKG-INFO +1 -1
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/__init__.py +1 -1
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/forward.py +137 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/PKG-INFO +1 -1
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/SOURCES.txt +1 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/pyproject.toml +1 -1
- midas_diffract-0.7.0/tests/test_omega_box_filter.py +352 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/LICENSE +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/README.md +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/hkls.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/losses.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/optimize.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/simulate_panel_zarrs.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/dependency_links.txt +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/requires.txt +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/top_level.txt +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/setup.cfg +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_c_comparison.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_distortion_layer.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_forward.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_hkls.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_losses.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_multi_detector.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_strain_tensor.py +0 -0
- {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_tilts.py +0 -0
- {midas_diffract-0.6.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
|
|
@@ -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:
|
|
@@ -399,6 +419,38 @@ class HEDMForwardModel(nn.Module):
|
|
|
399
419
|
)
|
|
400
420
|
self._has_wedge = abs(float(geometry.wedge)) > 0.0
|
|
401
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
|
+
|
|
402
454
|
# Scan config
|
|
403
455
|
self.scan_config = scan_config
|
|
404
456
|
if scan_config is not None:
|
|
@@ -944,8 +996,93 @@ class HEDMForwardModel(nn.Module):
|
|
|
944
996
|
((math.pi - torch.abs(eta)) >= self.min_eta)
|
|
945
997
|
valid = valid * eta_ok.float()
|
|
946
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
|
+
|
|
947
1005
|
return all_omega, eta, two_theta, valid
|
|
948
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
|
+
|
|
949
1086
|
# ------------------------------------------------------------------
|
|
950
1087
|
# Tilt rotation matrix (RotationTilts in SharedFuncsFit.c:230-266)
|
|
951
1088
|
# ------------------------------------------------------------------
|
|
@@ -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"
|
|
@@ -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
|
|
File without changes
|
|
File without changes
|