midas-pf-odf 0.1.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 (31) hide show
  1. midas_pf_odf-0.1.0/LICENSE +53 -0
  2. midas_pf_odf-0.1.0/PKG-INFO +94 -0
  3. midas_pf_odf-0.1.0/README.md +70 -0
  4. midas_pf_odf-0.1.0/midas_pf_odf/__init__.py +98 -0
  5. midas_pf_odf-0.1.0/midas_pf_odf/calibrate.py +362 -0
  6. midas_pf_odf-0.1.0/midas_pf_odf/centroid_baseline.py +302 -0
  7. midas_pf_odf-0.1.0/midas_pf_odf/forward.py +321 -0
  8. midas_pf_odf-0.1.0/midas_pf_odf/henningsson_faithful.py +485 -0
  9. midas_pf_odf-0.1.0/midas_pf_odf/inversion.py +922 -0
  10. midas_pf_odf-0.1.0/midas_pf_odf/io.py +1009 -0
  11. midas_pf_odf-0.1.0/midas_pf_odf/multi_grain.py +299 -0
  12. midas_pf_odf-0.1.0/midas_pf_odf/simulate.py +510 -0
  13. midas_pf_odf-0.1.0/midas_pf_odf/validation.py +134 -0
  14. midas_pf_odf-0.1.0/midas_pf_odf.egg-info/PKG-INFO +94 -0
  15. midas_pf_odf-0.1.0/midas_pf_odf.egg-info/SOURCES.txt +29 -0
  16. midas_pf_odf-0.1.0/midas_pf_odf.egg-info/dependency_links.txt +1 -0
  17. midas_pf_odf-0.1.0/midas_pf_odf.egg-info/requires.txt +10 -0
  18. midas_pf_odf-0.1.0/midas_pf_odf.egg-info/top_level.txt +1 -0
  19. midas_pf_odf-0.1.0/pyproject.toml +38 -0
  20. midas_pf_odf-0.1.0/setup.cfg +4 -0
  21. midas_pf_odf-0.1.0/tests/test_anisotropic_spread.py +122 -0
  22. midas_pf_odf-0.1.0/tests/test_calibrate.py +136 -0
  23. midas_pf_odf-0.1.0/tests/test_centroid_vs_peakshape.py +140 -0
  24. midas_pf_odf-0.1.0/tests/test_chunked_splatter.py +100 -0
  25. midas_pf_odf-0.1.0/tests/test_io.py +403 -0
  26. midas_pf_odf-0.1.0/tests/test_multi_grain.py +154 -0
  27. midas_pf_odf-0.1.0/tests/test_robustness.py +112 -0
  28. midas_pf_odf-0.1.0/tests/test_round_trip_small.py +155 -0
  29. midas_pf_odf-0.1.0/tests/test_saturation_mask.py +114 -0
  30. midas_pf_odf-0.1.0/tests/test_smoke.py +73 -0
  31. midas_pf_odf-0.1.0/tests/test_sparse_smooth_and_spread_init.py +112 -0
@@ -0,0 +1,53 @@
1
+ Copyright (c) 2012, UChicago Argonne, LLC
2
+
3
+ All Rights Reserved
4
+
5
+ MIDAS Microstructural Imaging using Diffraction Analysis Software
6
+
7
+ Materials Physics and Engineering
8
+ Computational X-ray Science
9
+ Advanced Photon Source
10
+ Argonne National Laboratory
11
+
12
+ Contributing Authors:
13
+ Hemant Sharma (hsharma@anl.gov)
14
+
15
+ OPEN SOURCE LICENSE
16
+
17
+ Redistribution and use in source and binary forms, with or without
18
+ modification, are permitted provided that the following conditions are met:
19
+
20
+ 1. Redistributions of source code must retain the above copyright notice,
21
+ this list of conditions and the following disclaimer. Software changes,
22
+ modifications, or derivative works, should be noted with comments and
23
+ the author and organization's name.
24
+
25
+ 2. Redistributions in binary form must reproduce the above copyright notice,
26
+ this list of conditions and the following disclaimer in the documentation
27
+ and/or other materials provided with the distribution.
28
+
29
+ 3. Neither the names of UChicago Argonne, LLC or the Department of Energy
30
+ nor the names of its contributors may be used to endorse or promote
31
+ products derived from this software without specific prior written
32
+ permission.
33
+
34
+ 4. The software and the end-user documentation included with the
35
+ redistribution, if any, must include the following acknowledgment:
36
+
37
+ "This product includes software produced by UChicago Argonne, LLC
38
+ under Contract No. DE-AC02-06CH11357 with the Department of Energy."
39
+
40
+ ****************************************************************************
41
+
42
+ DISCLAIMER
43
+
44
+ THE SOFTWARE IS SUPPLIED "AS IS" WITHOUT WARRANTY OF ANY KIND.
45
+
46
+ Neither the United States GOVERNMENT, nor the United States Department
47
+ of Energy, NOR UChicago Argonne, LLC, nor any of their employees, makes
48
+ any warranty, express or implied, or assumes any legal liability or
49
+ responsibility for the accuracy, completeness, or usefulness of any
50
+ information, data, apparatus, product, or process disclosed, or
51
+ represents that its use would not infringe privately owned rights.
52
+
53
+ ****************************************************************************
@@ -0,0 +1,94 @@
1
+ Metadata-Version: 2.4
2
+ Name: midas-pf-odf
3
+ Version: 0.1.0
4
+ Summary: Joint per-grain peak-shape inversion of pf-HEDM data; per-voxel orientation/strain (and Phase 2: per-voxel ODF)
5
+ Author-email: Hemant Sharma <hsharma@anl.gov>
6
+ Classifier: Development Status :: 2 - Pre-Alpha
7
+ Classifier: Intended Audience :: Science/Research
8
+ Classifier: Programming Language :: Python :: 3
9
+ Classifier: Operating System :: OS Independent
10
+ Classifier: Topic :: Scientific/Engineering :: Physics
11
+ Requires-Python: >=3.9
12
+ Description-Content-Type: text/markdown
13
+ License-File: LICENSE
14
+ Requires-Dist: numpy>=1.20
15
+ Requires-Dist: torch>=2.0
16
+ Requires-Dist: h5py>=3.0
17
+ Requires-Dist: midas-diffract>=0.1
18
+ Requires-Dist: midas-stress>=0.1
19
+ Requires-Dist: midas-grain-odf>=0.1
20
+ Provides-Extra: dev
21
+ Requires-Dist: pytest>=7.0; extra == "dev"
22
+ Requires-Dist: matplotlib>=3.5; extra == "dev"
23
+ Dynamic: license-file
24
+
25
+ # midas-pf-odf
26
+
27
+ Joint per-grain peak-shape inversion of pf-HEDM (point-focused HEDM, also
28
+ known as scanning 3DXRD) data.
29
+
30
+ **Phase 1 (current):** for each grain, fit all of its voxels'
31
+ `(R_V, ε_V)` jointly to image-MSE on measured 3D peak patches, given
32
+ voxel→grain assignment as input. Differentiable PyTorch forward;
33
+ voxel-summed splat per (spot, scan); closed-form per-spot intensity
34
+ scale `c_s*`.
35
+
36
+ **Phase 2 (later):** per-voxel ODF (sub-grain mosaic / GND density) on
37
+ the same forward.
38
+
39
+ ## Notebooks
40
+
41
+ Worked-example Jupyter notebooks live in `notebooks/`. They are **not shipped with `pip install`** — get them by cloning the [MIDAS repository](https://github.com/marinerhemant/MIDAS/tree/master/packages/midas_pf_odf/notebooks).
42
+
43
+ ## Status
44
+
45
+ Pre-alpha; private. See [dev/RESTART.md](dev/RESTART.md) for the active
46
+ worklog and resume point.
47
+
48
+ ## Where this sits relative to the literature
49
+
50
+ Henningsson et al. 2020 (J. Appl. Cryst.) introduced PCR / ASR — joint
51
+ multi-voxel intragranular strain reconstruction from scanning 3DXRD,
52
+ using **peak centroids only**. They explicitly call out peak-shape
53
+ inversion as the next direction (untapped at time of writing). This
54
+ package takes that step:
55
+
56
+ - Centroid → peak-shape (image-MSE on `(F, P, P)` patches).
57
+ - Hand-derived Jacobian → autograd / PyTorch.
58
+ - Single (R_V, ε_V) per voxel → per-voxel ODF (Phase 2).
59
+
60
+ The forward model is `midas_diffract.HEDMForwardModel`; reusable
61
+ infrastructure (sparse splatter, ODF parameterizations, held-out CV
62
+ selector, sparse-tile chunking) comes from `midas_grain_odf`.
63
+
64
+ ## Quickstart
65
+
66
+ ```bash
67
+ cd packages/midas_pf_odf
68
+ pip install -e .[dev]
69
+ python -m pytest tests/ -xvs
70
+ ```
71
+
72
+ ## Layout
73
+
74
+ ```
75
+ midas_pf_odf/ — public library
76
+ simulate.py — synthetic plant (orientation+strain gradients, multi-grain, noise)
77
+ forward.py — joint per-grain forward (soft beam gate, voxel-summed splat)
78
+ inversion.py — joint inversion driver (Adam / L-BFGS, identifiability knobs)
79
+ validation.py — per-voxel RMSE vs plant, held-out R²
80
+ io.py — I/O for real data (deferred)
81
+ dev/ — implementation plan, paper, notebooks, worklog
82
+ tests/ — synthetic round-trip tests
83
+ ```
84
+
85
+ ## Decisions locked
86
+
87
+ 1. **Identifiability knob:** project ε to mean-zero per grain (default,
88
+ recommended) or per-voxel-lattice over-parameterized (toggle).
89
+ 2. **Smoothness regularizer:** off by default (λ_smooth = 0).
90
+ 3. **Synthetic source:** torch-native simulator inside this package.
91
+ 4. **Outlier voxels:** soft-mask by held-out R² in the loss.
92
+ 5. **Multi-grain parallelism:** one grain × one GPU each via parsl
93
+ (deferred until workflow integration).
94
+ 6. **Output schema:** parallel HDF5; centroid path untouched (deferred).
@@ -0,0 +1,70 @@
1
+ # midas-pf-odf
2
+
3
+ Joint per-grain peak-shape inversion of pf-HEDM (point-focused HEDM, also
4
+ known as scanning 3DXRD) data.
5
+
6
+ **Phase 1 (current):** for each grain, fit all of its voxels'
7
+ `(R_V, ε_V)` jointly to image-MSE on measured 3D peak patches, given
8
+ voxel→grain assignment as input. Differentiable PyTorch forward;
9
+ voxel-summed splat per (spot, scan); closed-form per-spot intensity
10
+ scale `c_s*`.
11
+
12
+ **Phase 2 (later):** per-voxel ODF (sub-grain mosaic / GND density) on
13
+ the same forward.
14
+
15
+ ## Notebooks
16
+
17
+ Worked-example Jupyter notebooks live in `notebooks/`. They are **not shipped with `pip install`** — get them by cloning the [MIDAS repository](https://github.com/marinerhemant/MIDAS/tree/master/packages/midas_pf_odf/notebooks).
18
+
19
+ ## Status
20
+
21
+ Pre-alpha; private. See [dev/RESTART.md](dev/RESTART.md) for the active
22
+ worklog and resume point.
23
+
24
+ ## Where this sits relative to the literature
25
+
26
+ Henningsson et al. 2020 (J. Appl. Cryst.) introduced PCR / ASR — joint
27
+ multi-voxel intragranular strain reconstruction from scanning 3DXRD,
28
+ using **peak centroids only**. They explicitly call out peak-shape
29
+ inversion as the next direction (untapped at time of writing). This
30
+ package takes that step:
31
+
32
+ - Centroid → peak-shape (image-MSE on `(F, P, P)` patches).
33
+ - Hand-derived Jacobian → autograd / PyTorch.
34
+ - Single (R_V, ε_V) per voxel → per-voxel ODF (Phase 2).
35
+
36
+ The forward model is `midas_diffract.HEDMForwardModel`; reusable
37
+ infrastructure (sparse splatter, ODF parameterizations, held-out CV
38
+ selector, sparse-tile chunking) comes from `midas_grain_odf`.
39
+
40
+ ## Quickstart
41
+
42
+ ```bash
43
+ cd packages/midas_pf_odf
44
+ pip install -e .[dev]
45
+ python -m pytest tests/ -xvs
46
+ ```
47
+
48
+ ## Layout
49
+
50
+ ```
51
+ midas_pf_odf/ — public library
52
+ simulate.py — synthetic plant (orientation+strain gradients, multi-grain, noise)
53
+ forward.py — joint per-grain forward (soft beam gate, voxel-summed splat)
54
+ inversion.py — joint inversion driver (Adam / L-BFGS, identifiability knobs)
55
+ validation.py — per-voxel RMSE vs plant, held-out R²
56
+ io.py — I/O for real data (deferred)
57
+ dev/ — implementation plan, paper, notebooks, worklog
58
+ tests/ — synthetic round-trip tests
59
+ ```
60
+
61
+ ## Decisions locked
62
+
63
+ 1. **Identifiability knob:** project ε to mean-zero per grain (default,
64
+ recommended) or per-voxel-lattice over-parameterized (toggle).
65
+ 2. **Smoothness regularizer:** off by default (λ_smooth = 0).
66
+ 3. **Synthetic source:** torch-native simulator inside this package.
67
+ 4. **Outlier voxels:** soft-mask by held-out R² in the loss.
68
+ 5. **Multi-grain parallelism:** one grain × one GPU each via parsl
69
+ (deferred until workflow integration).
70
+ 6. **Output schema:** parallel HDF5; centroid path untouched (deferred).
@@ -0,0 +1,98 @@
1
+ """midas-pf-odf — joint per-grain peak-shape inversion of pf-HEDM data.
2
+
3
+ Phase 1: per-voxel (R, ε) from peak shapes. Phase 2: per-voxel ODF.
4
+ """
5
+
6
+ __version__ = "0.1.0"
7
+
8
+ from midas_pf_odf.simulate import (
9
+ SinglePhaseGrainPlant,
10
+ plant_single_grain,
11
+ simulate_grain_patches,
12
+ )
13
+ from midas_pf_odf.forward import (
14
+ joint_grain_forward,
15
+ soft_beam_gate,
16
+ )
17
+ from midas_pf_odf.inversion import (
18
+ neighbor_edges_from_grid_ij,
19
+ fit_grain_peakshape,
20
+ GrainPeakFitResult,
21
+ IdentifiabilityMode,
22
+ )
23
+ from midas_pf_odf.validation import (
24
+ recovery_metrics,
25
+ holdout_score,
26
+ )
27
+ from midas_pf_odf.multi_grain import (
28
+ MultiGrainPlant,
29
+ plant_multi_grain,
30
+ split_into_grains,
31
+ simulate_multi_grain,
32
+ fit_multi_grain,
33
+ )
34
+ from midas_pf_odf.centroid_baseline import (
35
+ fit_grain_centroid_baseline,
36
+ measured_centroids_from_patches,
37
+ predicted_centroids,
38
+ )
39
+ from midas_pf_odf.calibrate import (
40
+ RawFrameCalibration,
41
+ calibrate_raw_frame_geometry,
42
+ layer_model_factory,
43
+ measure_patch_offsets,
44
+ )
45
+ from midas_pf_odf.io import (
46
+ PFGrainDataset,
47
+ distortion_from_paramstest,
48
+ saturation_threshold_from_paramstest,
49
+ load_pf_grain,
50
+ build_model_from_paramstest,
51
+ geometry_from_paramstest,
52
+ parse_paramstest,
53
+ assemble_grain_patch_data,
54
+ crop_patches_from_frames,
55
+ build_model_from_zarr,
56
+ geometry_from_zarr,
57
+ read_zarr_params,
58
+ ZarrFrameSource,
59
+ )
60
+
61
+ __all__ = [
62
+ "SinglePhaseGrainPlant",
63
+ "plant_single_grain",
64
+ "simulate_grain_patches",
65
+ "joint_grain_forward",
66
+ "soft_beam_gate",
67
+ "fit_grain_peakshape",
68
+ "GrainPeakFitResult",
69
+ "IdentifiabilityMode",
70
+ "neighbor_edges_from_grid_ij",
71
+ "recovery_metrics",
72
+ "holdout_score",
73
+ "MultiGrainPlant",
74
+ "plant_multi_grain",
75
+ "split_into_grains",
76
+ "simulate_multi_grain",
77
+ "fit_multi_grain",
78
+ "fit_grain_centroid_baseline",
79
+ "measured_centroids_from_patches",
80
+ "predicted_centroids",
81
+ "PFGrainDataset",
82
+ "load_pf_grain",
83
+ "build_model_from_paramstest",
84
+ "geometry_from_paramstest",
85
+ "parse_paramstest",
86
+ "assemble_grain_patch_data",
87
+ "crop_patches_from_frames",
88
+ "build_model_from_zarr",
89
+ "geometry_from_zarr",
90
+ "read_zarr_params",
91
+ "ZarrFrameSource",
92
+ "distortion_from_paramstest",
93
+ "saturation_threshold_from_paramstest",
94
+ "RawFrameCalibration",
95
+ "calibrate_raw_frame_geometry",
96
+ "layer_model_factory",
97
+ "measure_patch_offsets",
98
+ ]
@@ -0,0 +1,362 @@
1
+ """Raw-frame geometry mini-calibration (P1-6).
2
+
3
+ Empirical-Jacobian damped Gauss-Newton: bump each geometry parameter in
4
+ the ACTUAL forward model (convention-proof — no analytic derivative can
5
+ disagree with the model), fit the parameter steps to the measured anchor
6
+ residuals with a robust (MAD-rejecting) LSQ, iterate, verify collapse.
7
+
8
+ On the SOH wt316 campaign this collapsed anchor residuals from 7.06 to
9
+ 1.4 px RMS and cross-validated between independent loads (tx/BC/tz agree;
10
+ ty/Lsd/ω₀ scatter with 2-ring data). Any raw-frame consumer of a promoted
11
+ ``paramstest.txt`` needs it: the powder calibration behind paramstest is
12
+ blind to tx (~0.27° ≈ 3-4e3 µε fake strain) and its convention chain
13
+ (ideal↔raw, flips, ω origin) is easy to get subtly wrong — this
14
+ calibration absorbs all of it into the model actually being used.
15
+
16
+ Typical use::
17
+
18
+ factory = layer_model_factory(layer_dir, ring_numbers=[1, 2, 3, 4],
19
+ n_pixels_y=2880, n_pixels_z=2880,
20
+ n_frames=1440, apply_distortion=True,
21
+ device=dev)
22
+ ds = load_pf_grain(layer_dir, grain_id, ..., model=factory())
23
+ cal = calibrate_raw_frame_geometry(ds, cached_patches,
24
+ model_factory=factory)
25
+ model = factory(geom_mod=cal.calibrated, omega_start=cal.calibrated["om0"])
26
+ """
27
+
28
+ from __future__ import annotations
29
+
30
+ import copy
31
+ import math
32
+ from dataclasses import dataclass, field
33
+ from pathlib import Path
34
+ from typing import Callable, Dict, Optional, Sequence, Tuple
35
+
36
+ import numpy as np
37
+ import torch
38
+
39
+ from midas_diffract.forward import HEDMForwardModel
40
+
41
+ __all__ = [
42
+ "RawFrameCalibration",
43
+ "calibrate_raw_frame_geometry",
44
+ "layer_model_factory",
45
+ "measure_patch_offsets",
46
+ ]
47
+
48
+ # Default parameter set + bump sizes (the empirical-derivative step per
49
+ # parameter; also the natural scale of one GN step component).
50
+ DEFAULT_BUMPS: Dict[str, float] = {
51
+ "y_BC": 1.0, # px
52
+ "z_BC": 1.0, # px
53
+ "tx": 0.05, # deg
54
+ "ty": 0.05, # deg
55
+ "tz": 0.05, # deg
56
+ "Lsd": 500.0, # um
57
+ "om0": 0.05, # deg
58
+ }
59
+
60
+
61
+ @dataclass
62
+ class RawFrameCalibration:
63
+ """Result of :func:`calibrate_raw_frame_geometry`."""
64
+
65
+ calibrated: Dict[str, float] # fitted parameter values
66
+ original: Dict[str, float] # starting values
67
+ rms_before_px: float # sqrt(mean(dy^2+dz^2)), first iter
68
+ rms_after_px: float # same, after the final iteration
69
+ rms_frames_after: float # RMS omega-frame residual after
70
+ n_spots_used: int # spots with a measured centroid
71
+ per_iter_rms_px: list = field(default_factory=list)
72
+
73
+ @property
74
+ def delta(self) -> Dict[str, float]:
75
+ return {k: self.calibrated[k] - self.original[k]
76
+ for k in self.calibrated}
77
+
78
+
79
+ def measure_patch_offsets(
80
+ measured: torch.Tensor,
81
+ *,
82
+ signal_threshold_frac: float = 0.2,
83
+ ) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
84
+ """Per-spot measured centroid offsets from a cached patch tensor.
85
+
86
+ Parameters
87
+ ----------
88
+ measured : (S, Sigma, F, P, P) tensor
89
+ The cached measured patches (``assemble_grain_patch_data`` layout).
90
+ signal_threshold_frac : float
91
+ A (spot, scan) cell counts as signal when its summed intensity
92
+ exceeds this fraction of the median non-zero cell total.
93
+
94
+ Returns
95
+ -------
96
+ (dy, dz, df, ok) : per-spot median centroid offsets (px, px, frames)
97
+ relative to the patch centre, NaN where no scan had signal;
98
+ ``ok`` = finite mask.
99
+ """
100
+ meas = measured.detach().to(torch.float64).cpu()
101
+ S, _Sigma, F, P, _P2 = meas.shape
102
+ c = P // 2
103
+ ps = meas.sum(dim=2) # (S, Sigma, P, P)
104
+ tot = ps.sum(dim=(-1, -2)) # (S, Sigma)
105
+ nz = tot[tot > 0]
106
+ thr = float(nz.median()) * signal_threshold_frac if nz.numel() else 0.0
107
+ yy, zz = torch.meshgrid(
108
+ torch.arange(P, dtype=torch.float64),
109
+ torch.arange(P, dtype=torch.float64), indexing="ij",
110
+ )
111
+ cy = ((ps * yy).sum(dim=(-1, -2)) / (tot + 1e-12)).numpy() - c
112
+ cz = ((ps * zz).sum(dim=(-1, -2)) / (tot + 1e-12)).numpy() - c
113
+ fidx = torch.arange(F, dtype=torch.float64)
114
+ pf = meas.sum(dim=(-1, -2)) # (S, Sigma, F)
115
+ cf = ((pf * fidx).sum(dim=-1) / (pf.sum(dim=-1) + 1e-12)).numpy() - F // 2
116
+ has = (tot.numpy() > thr)
117
+
118
+ dy = np.full(S, np.nan)
119
+ dz = np.full(S, np.nan)
120
+ df = np.full(S, np.nan)
121
+ for s in range(S):
122
+ g = np.nonzero(has[s])[0]
123
+ if g.size >= 1:
124
+ dy[s] = np.median(cy[s, g])
125
+ dz[s] = np.median(cz[s, g])
126
+ df[s] = np.median(cf[s, g])
127
+ ok = np.isfinite(dy)
128
+ return dy, dz, df, ok
129
+
130
+
131
+ def _anchors_with_model(ds, model: HEDMForwardModel, dtype, device):
132
+ from . import io as _io
133
+ d2 = copy.copy(ds)
134
+ d2.model = model
135
+ ay, az, af, _valid, _obs, S = _io._forward_anchors(d2, dtype, device)
136
+ return (ay.detach().cpu().numpy(), az.detach().cpu().numpy(),
137
+ af.detach().cpu().numpy(), S)
138
+
139
+
140
+ def calibrate_raw_frame_geometry(
141
+ ds,
142
+ measured: torch.Tensor,
143
+ *,
144
+ model_factory: Callable[..., HEDMForwardModel],
145
+ anchor_model: Optional[HEDMForwardModel] = None,
146
+ params_to_fit: Sequence[str] = tuple(DEFAULT_BUMPS),
147
+ bumps: Optional[Dict[str, float]] = None,
148
+ n_iters: int = 3,
149
+ frame_weight: float = 4.0,
150
+ damping: float = 1e-2,
151
+ max_step_bumps: float = 40.0,
152
+ mad_reject: float = 6.0,
153
+ signal_threshold_frac: float = 0.2,
154
+ dtype: torch.dtype = torch.float64,
155
+ device: torch.device | str = "cpu",
156
+ verbose: bool = True,
157
+ ) -> RawFrameCalibration:
158
+ """Fit (y_BC, z_BC, tx, ty, tz, Lsd, ω₀) to measured patch centroids.
159
+
160
+ Parameters
161
+ ----------
162
+ ds : PFGrainDataset
163
+ The grain whose cached patches anchor the calibration.
164
+ measured : (S, Sigma, F, P, P) tensor
165
+ Cached measured patches, in the anchor layout of ``anchor_model``.
166
+ model_factory : callable
167
+ ``model_factory(geom_mod: dict | None = None, omega_start:
168
+ float | None = None) -> HEDMForwardModel``. ``geom_mod`` maps
169
+ parameter names (see :data:`DEFAULT_BUMPS`, minus ``om0`` which is
170
+ passed as ``omega_start``) to absolute values. Must build the model
171
+ the CALIBRATION should converge (e.g. distortion ON for raw-frame
172
+ work); :func:`layer_model_factory` builds one from a layer dir.
173
+ anchor_model : optional
174
+ The model the cache was ASSEMBLED with (defaults to ``ds.model``).
175
+ Absolute measured positions are ``anchors(anchor_model) + offsets``.
176
+
177
+ Notes
178
+ -----
179
+ Derivatives are empirical: each parameter is bumped by ``bumps[name]``
180
+ in a rebuilt model and the anchor shift is the Jacobian column — no
181
+ convention (flips, ideal↔raw, ω origin) can silently disagree with
182
+ the model. The LSQ is Levenberg-damped and MAD-rejecting; steps are
183
+ clipped to ``max_step_bumps`` bump units.
184
+ """
185
+ bumps = dict(DEFAULT_BUMPS if bumps is None else bumps)
186
+ names = [n for n in params_to_fit if n in bumps]
187
+ if not names:
188
+ raise ValueError(f"params_to_fit {params_to_fit!r} matches no known "
189
+ f"parameter (choose from {sorted(DEFAULT_BUMPS)})")
190
+
191
+ base = anchor_model if anchor_model is not None else ds.model
192
+ ay0, az0, af0, S = _anchors_with_model(ds, base, dtype, device)
193
+
194
+ dy, dz, df, ok = measure_patch_offsets(
195
+ measured, signal_threshold_frac=signal_threshold_frac)
196
+ if int(ok.sum()) < len(names):
197
+ raise ValueError(
198
+ f"only {int(ok.sum())} spots have a measured centroid — fewer "
199
+ f"than the {len(names)} parameters being fit."
200
+ )
201
+ abs_y = ay0 + dy
202
+ abs_z = az0 + dz
203
+ abs_f = af0 + df
204
+
205
+ def _getf(model, name, fallback=0.0):
206
+ v = getattr(model, name, fallback)
207
+ try:
208
+ v = v.detach().cpu().item()
209
+ except AttributeError:
210
+ pass
211
+ if isinstance(v, (list, tuple)):
212
+ v = v[0]
213
+ return float(v)
214
+
215
+ fit0 = model_factory()
216
+ state: Dict[str, float] = {}
217
+ for n in names:
218
+ if n == "om0":
219
+ state[n] = _getf(fit0, "omega_start")
220
+ else:
221
+ state[n] = _getf(fit0, n)
222
+ original = dict(state)
223
+
224
+ def model_at(st: Dict[str, float]) -> HEDMForwardModel:
225
+ gm = {k: v for k, v in st.items() if k != "om0"}
226
+ om0 = st.get("om0")
227
+ return model_factory(geom_mod=gm, omega_start=om0)
228
+
229
+ rms_before = None
230
+ per_iter = []
231
+ for it in range(n_iters):
232
+ m_it = model_at(state)
233
+ ayi, azi, afi, _ = _anchors_with_model(ds, m_it, dtype, device)
234
+ ry = (abs_y - ayi)[ok]
235
+ rz = (abs_z - azi)[ok]
236
+ rf = (abs_f - afi)[ok]
237
+ rms = math.sqrt(np.nanmean(ry ** 2 + rz ** 2))
238
+ per_iter.append(rms)
239
+ if rms_before is None:
240
+ rms_before = rms
241
+ if verbose:
242
+ print(f"[mini-calib iter {it}] RMS(dy,dz)={rms:.2f}px "
243
+ f"RMS(df)={math.sqrt(np.nanmean(rf ** 2)):.2f}fr",
244
+ flush=True)
245
+ Jy, Jz, Jf = [], [], []
246
+ for n in names:
247
+ st2 = dict(state)
248
+ st2[n] = state[n] + bumps[n]
249
+ ayp, azp, afp, _ = _anchors_with_model(
250
+ ds, model_at(st2), dtype, device)
251
+ Jy.append((ayp - ayi)[ok])
252
+ Jz.append((azp - azi)[ok])
253
+ Jf.append((afp - afi)[ok])
254
+ Jy = np.stack(Jy, 1)
255
+ Jz = np.stack(Jz, 1)
256
+ Jf = np.stack(Jf, 1)
257
+ A = np.vstack([Jy, Jz, frame_weight * Jf])
258
+ b = np.concatenate([ry, rz, frame_weight * rf])
259
+ good = np.isfinite(b) & np.all(np.isfinite(A), axis=1)
260
+ med = np.median(b[good])
261
+ mad = np.median(np.abs(b[good] - med)) + 1e-9
262
+ good &= np.abs(b - med) < mad_reject * 1.4826 * mad
263
+ Ag, bg = A[good], b[good]
264
+ ata = Ag.T @ Ag
265
+ ata += damping * np.eye(len(names)) * np.trace(ata) / len(names)
266
+ x = np.linalg.solve(ata, Ag.T @ bg) # step in bump units
267
+ x = np.clip(x, -max_step_bumps, max_step_bumps)
268
+ for k, n in enumerate(names):
269
+ state[n] += x[k] * bumps[n]
270
+ if verbose:
271
+ print(" step: " + " ".join(
272
+ f"d{n}={x[k] * bumps[n]:+.4f}" for k, n in enumerate(names)),
273
+ flush=True)
274
+
275
+ m_fin = model_at(state)
276
+ ayf, azf, aff, _ = _anchors_with_model(ds, m_fin, dtype, device)
277
+ ryf = (abs_y - ayf)[ok]
278
+ rzf = (abs_z - azf)[ok]
279
+ rff = (abs_f - aff)[ok]
280
+ rms_after = math.sqrt(np.nanmean(ryf ** 2 + rzf ** 2))
281
+ per_iter.append(rms_after)
282
+ return RawFrameCalibration(
283
+ calibrated=state,
284
+ original=original,
285
+ rms_before_px=float(rms_before),
286
+ rms_after_px=float(rms_after),
287
+ rms_frames_after=float(math.sqrt(np.nanmean(rff ** 2))),
288
+ n_spots_used=int(ok.sum()),
289
+ per_iter_rms_px=per_iter,
290
+ )
291
+
292
+
293
+ def layer_model_factory(
294
+ layer_dir: str | Path,
295
+ *,
296
+ ring_numbers: Sequence[int],
297
+ n_pixels_y: int,
298
+ n_pixels_z: int,
299
+ n_frames: Optional[int] = None,
300
+ omega_step: Optional[float] = None,
301
+ apply_distortion: bool = True,
302
+ max_two_theta_deg: Optional[float] = None,
303
+ device: torch.device | str = "cpu",
304
+ dtype: torch.dtype = torch.float64,
305
+ ) -> Callable[..., HEDMForwardModel]:
306
+ """Build a ``model_factory`` for :func:`calibrate_raw_frame_geometry`
307
+ from a MIDAS layer directory (paramstest/hkls/positions).
308
+
309
+ The returned callable accepts ``geom_mod`` (dict of absolute geometry
310
+ values: y_BC, z_BC, tx, ty, tz, Lsd) and ``omega_start`` overrides.
311
+ """
312
+ from midas_fit_grain.driver import _cartesian_B_matrix, _read_hkls_csv
313
+
314
+ from .io import (
315
+ _first_float, geometry_from_paramstest, parse_paramstest,
316
+ scan_config_from_positions,
317
+ )
318
+
319
+ layer = Path(layer_dir)
320
+ params = parse_paramstest(layer / "paramstest.txt")
321
+ lat = tuple(float(x) for x in params["LatticeParameter"][0][:6])
322
+ beam_size = _first_float(params, "BeamSize")
323
+
324
+ def factory(geom_mod: Optional[Dict[str, float]] = None,
325
+ omega_start: Optional[float] = None) -> HEDMForwardModel:
326
+ geom = geometry_from_paramstest(
327
+ params, n_pixels_y=n_pixels_y, n_pixels_z=n_pixels_z,
328
+ n_frames=n_frames, omega_step=omega_step,
329
+ apply_distortion=apply_distortion,
330
+ )
331
+ if geom_mod:
332
+ for k, v in geom_mod.items():
333
+ setattr(geom, k, float(v))
334
+ if omega_start is not None:
335
+ geom.omega_start = float(omega_start)
336
+ mtth = max_two_theta_deg
337
+ if mtth is None:
338
+ if "MaxRingRad" in params:
339
+ mrr = _first_float(params, "MaxRingRad")
340
+ lsd = geom.Lsd[0] if isinstance(geom.Lsd, list) else geom.Lsd
341
+ mtth = 2.0 * math.degrees(math.atan(mrr / lsd))
342
+ else:
343
+ mtth = 180.0
344
+ hkls_int, thetas_deg, _ring_nr = _read_hkls_csv(
345
+ layer / "hkls.csv", [int(r) for r in ring_numbers], mtth)
346
+ B = _cartesian_B_matrix(lat)
347
+ cart = (B @ hkls_int.astype(np.float64).T).T
348
+ scan_cfg = scan_config_from_positions(
349
+ layer / "positions.csv", beam_size, dtype=dtype)
350
+ model = HEDMForwardModel(
351
+ torch.from_numpy(cart),
352
+ torch.from_numpy(np.asarray(thetas_deg) * math.pi / 180.0),
353
+ geom,
354
+ hkls_int=torch.from_numpy(hkls_int.astype(np.float64)),
355
+ scan_config=scan_cfg,
356
+ device=device,
357
+ )
358
+ # Raw-frame convention: tilts are part of the raw prediction.
359
+ model.apply_tilts = True
360
+ return model
361
+
362
+ return factory