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.
Files changed (26) hide show
  1. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/PKG-INFO +1 -1
  2. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/__init__.py +1 -1
  3. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/forward.py +270 -5
  4. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/simulate_panel_zarrs.py +12 -16
  5. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/PKG-INFO +1 -1
  6. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/SOURCES.txt +1 -0
  7. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/pyproject.toml +1 -1
  8. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_forward.py +84 -0
  9. midas_diffract-0.7.0/tests/test_omega_box_filter.py +352 -0
  10. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/LICENSE +0 -0
  11. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/README.md +0 -0
  12. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/hkls.py +0 -0
  13. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/losses.py +0 -0
  14. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract/optimize.py +0 -0
  15. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/dependency_links.txt +0 -0
  16. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/requires.txt +0 -0
  17. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/midas_diffract.egg-info/top_level.txt +0 -0
  18. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/setup.cfg +0 -0
  19. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_c_comparison.py +0 -0
  20. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_distortion_layer.py +0 -0
  21. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_hkls.py +0 -0
  22. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_losses.py +0 -0
  23. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_multi_detector.py +0 -0
  24. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_strain_tensor.py +0 -0
  25. {midas_diffract-0.4.0 → midas_diffract-0.7.0}/tests/test_tilts.py +0 -0
  26. {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.4.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.4.0"
31
+ __version__ = "0.7.0"
32
32
 
33
33
  from .forward import (
34
34
  HEDMForwardModel,
@@ -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
- G_C = torch.einsum("...nij,...mj->...nmi", orientation_matrices, hkls_cart)
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
- v_no_wedge = v_no_wedge.unsqueeze(-2).expand_as(G_C[..., 0]) # (..., N, M)
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 (broadcast thetas to match 2N dimension)
904
- two_theta_single = 2.0 * thetas.unsqueeze(-2) # (..., 1, M) or (1, M)
905
- two_theta = two_theta_single.expand_as(all_omega)
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
- spots = model(eulers_t, positions_t,
617
- lattice_params=latc, strain=strain_t)
618
-
619
- valid_np = (spots.valid > 0.5).cpu().numpy()
620
- # Strain × orientation diagonal mask
621
- G_strain, Kdim, Mdim = valid_np.shape
622
- if G_strain == n_grains and Kdim == 2 * n_grains:
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
- grain_ids = np.broadcast_to(np.arange(G_strain)[:, None, None],
656
- (G_strain, Kdim, Mdim))
657
- hkl_ids = np.broadcast_to(np.arange(Mdim)[None, None, :],
658
- (G_strain, Kdim, Mdim))
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.4.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.4.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