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.
Files changed (26) hide show
  1. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/PKG-INFO +1 -1
  2. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/__init__.py +1 -1
  3. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/forward.py +137 -0
  4. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/PKG-INFO +1 -1
  5. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/SOURCES.txt +1 -0
  6. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/pyproject.toml +1 -1
  7. midas_diffract-0.7.0/tests/test_omega_box_filter.py +352 -0
  8. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/LICENSE +0 -0
  9. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/README.md +0 -0
  10. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/hkls.py +0 -0
  11. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/losses.py +0 -0
  12. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/optimize.py +0 -0
  13. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract/simulate_panel_zarrs.py +0 -0
  14. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/dependency_links.txt +0 -0
  15. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/requires.txt +0 -0
  16. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/top_level.txt +0 -0
  17. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/setup.cfg +0 -0
  18. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_c_comparison.py +0 -0
  19. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_distortion_layer.py +0 -0
  20. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_forward.py +0 -0
  21. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_hkls.py +0 -0
  22. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_losses.py +0 -0
  23. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_multi_detector.py +0 -0
  24. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_strain_tensor.py +0 -0
  25. {midas_diffract-0.6.0 → midas_diffract-0.7.0}/tests/test_tilts.py +0 -0
  26. {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.6.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
@@ -28,7 +28,7 @@ Quick start
28
28
  loss.backward()
29
29
  """
30
30
 
31
- __version__ = "0.6.0"
31
+ __version__ = "0.7.0"
32
32
 
33
33
  from .forward import (
34
34
  HEDMForwardModel,
@@ -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.6.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
@@ -18,6 +18,7 @@ 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
21
22
  tests/test_strain_tensor.py
22
23
  tests/test_tilts.py
23
24
  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.6.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