faster-diffbloch 0.1.6__tar.gz → 0.1.7__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 (17) hide show
  1. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/PKG-INFO +25 -27
  2. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/README.md +24 -26
  3. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/pyproject.toml +1 -1
  4. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/__init__.py +1 -1
  5. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/backend.py +30 -0
  6. faster_diffbloch-0.1.7/src/faster_diffbloch/scattering.py +271 -0
  7. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/.gitignore +0 -0
  8. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/LICENSE +0 -0
  9. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/builder.py +0 -0
  10. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/cli.py +0 -0
  11. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/batch_cgemm.c +0 -0
  12. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/batch_cgemm.h +0 -0
  13. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/batch_cgemm.metal +0 -0
  14. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/bridge_lib.c +0 -0
  15. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/metal_batch_cgemm.m +0 -0
  16. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/native_scattering.c +0 -0
  17. {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/native_scattering.h +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: faster-diffbloch
3
- Version: 0.1.6
3
+ Version: 0.1.7
4
4
  Summary: Drop-in Metal GPU (macOS) and optimized CPU (macOS/Linux) acceleration for diffBloch
5
5
  Project-URL: Homepage, https://godofecht.github.io/diffFlow/
6
6
  Project-URL: Documentation, https://godofecht.github.io/diffFlow/
@@ -35,7 +35,7 @@ Description-Content-Type: text/markdown
35
35
 
36
36
  Drop-in Apple Silicon Metal GPU and optimized CPU acceleration for [diffBloch](https://diffbloch.com) electron crystallography structure refinement.
37
37
 
38
- **2.59x faster than PyTorch** on the forward-plus-backward pass at 579 beams, through the Metal path on Apple Silicon. An optimized CPU path covers macOS and Linux.
38
+ **Up to 3.4x faster than PyTorch** on the forward-plus-backward pass, through the Metal path on Apple Silicon. An optimized CPU path covers macOS and Linux.
39
39
 
40
40
  | Package | Documentation | Source repository | Original project |
41
41
  | :--- | :--- | :--- | :--- |
@@ -84,34 +84,31 @@ of three runs each.
84
84
 
85
85
  | Beams | PyTorch | CPU path | Metal path | CPU | Metal |
86
86
  | :---: | :---: | :---: | :---: | :---: | :---: |
87
- | 31 | 1.52 ms | 1.42 ms | 1.95 ms | 1.07x | 0.78x |
88
- | 61 | 2.11 ms | 2.07 ms | 2.20 ms | 1.02x | 0.96x |
89
- | 91 | 3.87 ms | 3.61 ms | 3.70 ms | 1.07x | 1.05x |
90
- | 163 | 10.46 ms | 8.60 ms | 6.27 ms | 1.22x | 1.67x |
91
- | 579 | 123.30 ms | 89.17 ms | 47.63 ms | 1.38x | **2.59x** |
92
-
93
- The speedup grows with beam count, and this is the shape of the result rather
94
- than noise. The package replaces the matrix exponential, which costs O(N^3),
95
- while structure factors and the loss stay in PyTorch. As the system grows the
96
- exponential takes a larger share of the runtime, so there is more for the
97
- accelerated path to reach.
98
-
99
- That has two practical consequences. At 579 beams, CsPbBr3 scale, Metal runs the
100
- forward-plus-backward pass in 47.63 ms against PyTorch's 123.30 ms. Below about
101
- 91 beams Metal is slower than PyTorch, because the dispatch cost outweighs a
102
- small exponential, so `enable(device="cpu")` is the better choice for small
103
- systems.
104
-
105
- Forward pass alone at 579 beams: PyTorch 25.90 ms, Metal 15.58 ms, a factor of
106
- 1.66. PyTorch MPS has no native kernel for `aten::linalg_matrix_exp` and falls
107
- back to the CPU with host transfers, which is the gap this closes.
87
+ | 31 | 1.47 ms | 0.72 ms | 1.07 ms | 2.03x | 1.37x |
88
+ | 61 | 2.00 ms | 1.33 ms | 1.33 ms | 1.50x | 1.50x |
89
+ | 91 | 3.68 ms | 2.24 ms | 1.67 ms | 1.65x | 2.21x |
90
+ | 163 | 10.34 ms | 5.78 ms | 3.05 ms | 1.79x | **3.39x** |
91
+ | 579 | 117.44 ms | 81.55 ms | 39.87 ms | 1.44x | 2.95x |
92
+
93
+ Two kernels carry this. The matrix exponential is O(N^3) and dominates at large
94
+ beam counts. The structure factors dominate at small ones: at 31 beams they are
95
+ about 80% of the forward pass while the exponential is 15%. Both now run
96
+ natively, which is why the speedup holds across the range rather than appearing
97
+ only at scale.
98
+
99
+ Accuracy is not traded for any of it. The structure factors run in float64 and
100
+ agree with diffBloch to 4e-15, their gradients to 1e-14, and the systematic
101
+ absences land on the same exact zeros. End to end the loss is bit-identical to
102
+ diffBloch's own.
103
+
104
+ `bench/pkgbench.py` reproduces the table.
108
105
 
109
106
  ### The standalone Flow port
110
107
 
111
108
  The diffFlow repository also holds a standalone port of the whole calculation,
112
- compiled from Flow rather than bridged into PyTorch. It avoids the framework
113
- overhead and goes further. At 579 beams, forward plus backward, minimum of five
114
- runs:
109
+ compiled from Flow rather than bridged into PyTorch. It runs on the CPU and
110
+ avoids the framework overhead. At 579 beams, forward plus backward, minimum of
111
+ five runs:
115
112
 
116
113
  | Implementation | Forward + Backward |
117
114
  | :--- | :---: |
@@ -120,7 +117,8 @@ runs:
120
117
  | Flow port, C backend | 82.84 ms |
121
118
  | Flow port, MLIR backend | 79.74 ms |
122
119
 
123
- Running the port directly is how those are reached. See the
120
+ The port is the fastest CPU route. This package on Metal is faster still,
121
+ at 39.87 ms for the same case, because the port has no GPU path. See the
124
122
  [diffFlow documentation](https://godofecht.github.io/diffFlow/).
125
123
 
126
124
  These figures come from one workload on one machine. Other crystals, hardware
@@ -2,7 +2,7 @@
2
2
 
3
3
  Drop-in Apple Silicon Metal GPU and optimized CPU acceleration for [diffBloch](https://diffbloch.com) electron crystallography structure refinement.
4
4
 
5
- **2.59x faster than PyTorch** on the forward-plus-backward pass at 579 beams, through the Metal path on Apple Silicon. An optimized CPU path covers macOS and Linux.
5
+ **Up to 3.4x faster than PyTorch** on the forward-plus-backward pass, through the Metal path on Apple Silicon. An optimized CPU path covers macOS and Linux.
6
6
 
7
7
  | Package | Documentation | Source repository | Original project |
8
8
  | :--- | :--- | :--- | :--- |
@@ -51,34 +51,31 @@ of three runs each.
51
51
 
52
52
  | Beams | PyTorch | CPU path | Metal path | CPU | Metal |
53
53
  | :---: | :---: | :---: | :---: | :---: | :---: |
54
- | 31 | 1.52 ms | 1.42 ms | 1.95 ms | 1.07x | 0.78x |
55
- | 61 | 2.11 ms | 2.07 ms | 2.20 ms | 1.02x | 0.96x |
56
- | 91 | 3.87 ms | 3.61 ms | 3.70 ms | 1.07x | 1.05x |
57
- | 163 | 10.46 ms | 8.60 ms | 6.27 ms | 1.22x | 1.67x |
58
- | 579 | 123.30 ms | 89.17 ms | 47.63 ms | 1.38x | **2.59x** |
59
-
60
- The speedup grows with beam count, and this is the shape of the result rather
61
- than noise. The package replaces the matrix exponential, which costs O(N^3),
62
- while structure factors and the loss stay in PyTorch. As the system grows the
63
- exponential takes a larger share of the runtime, so there is more for the
64
- accelerated path to reach.
65
-
66
- That has two practical consequences. At 579 beams, CsPbBr3 scale, Metal runs the
67
- forward-plus-backward pass in 47.63 ms against PyTorch's 123.30 ms. Below about
68
- 91 beams Metal is slower than PyTorch, because the dispatch cost outweighs a
69
- small exponential, so `enable(device="cpu")` is the better choice for small
70
- systems.
71
-
72
- Forward pass alone at 579 beams: PyTorch 25.90 ms, Metal 15.58 ms, a factor of
73
- 1.66. PyTorch MPS has no native kernel for `aten::linalg_matrix_exp` and falls
74
- back to the CPU with host transfers, which is the gap this closes.
54
+ | 31 | 1.47 ms | 0.72 ms | 1.07 ms | 2.03x | 1.37x |
55
+ | 61 | 2.00 ms | 1.33 ms | 1.33 ms | 1.50x | 1.50x |
56
+ | 91 | 3.68 ms | 2.24 ms | 1.67 ms | 1.65x | 2.21x |
57
+ | 163 | 10.34 ms | 5.78 ms | 3.05 ms | 1.79x | **3.39x** |
58
+ | 579 | 117.44 ms | 81.55 ms | 39.87 ms | 1.44x | 2.95x |
59
+
60
+ Two kernels carry this. The matrix exponential is O(N^3) and dominates at large
61
+ beam counts. The structure factors dominate at small ones: at 31 beams they are
62
+ about 80% of the forward pass while the exponential is 15%. Both now run
63
+ natively, which is why the speedup holds across the range rather than appearing
64
+ only at scale.
65
+
66
+ Accuracy is not traded for any of it. The structure factors run in float64 and
67
+ agree with diffBloch to 4e-15, their gradients to 1e-14, and the systematic
68
+ absences land on the same exact zeros. End to end the loss is bit-identical to
69
+ diffBloch's own.
70
+
71
+ `bench/pkgbench.py` reproduces the table.
75
72
 
76
73
  ### The standalone Flow port
77
74
 
78
75
  The diffFlow repository also holds a standalone port of the whole calculation,
79
- compiled from Flow rather than bridged into PyTorch. It avoids the framework
80
- overhead and goes further. At 579 beams, forward plus backward, minimum of five
81
- runs:
76
+ compiled from Flow rather than bridged into PyTorch. It runs on the CPU and
77
+ avoids the framework overhead. At 579 beams, forward plus backward, minimum of
78
+ five runs:
82
79
 
83
80
  | Implementation | Forward + Backward |
84
81
  | :--- | :---: |
@@ -87,7 +84,8 @@ runs:
87
84
  | Flow port, C backend | 82.84 ms |
88
85
  | Flow port, MLIR backend | 79.74 ms |
89
86
 
90
- Running the port directly is how those are reached. See the
87
+ The port is the fastest CPU route. This package on Metal is faster still,
88
+ at 39.87 ms for the same case, because the port has no GPU path. See the
91
89
  [diffFlow documentation](https://godofecht.github.io/diffFlow/).
92
90
 
93
91
  These figures come from one workload on one machine. Other crystals, hardware
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
4
4
 
5
5
  [project]
6
6
  name = "faster-diffbloch"
7
- version = "0.1.6"
7
+ version = "0.1.7"
8
8
  description = "Drop-in Metal GPU (macOS) and optimized CPU (macOS/Linux) acceleration for diffBloch"
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.10"
@@ -4,7 +4,7 @@ from __future__ import annotations
4
4
 
5
5
  from .backend import enable, disable, matrix_exp, matrix_exp_backward, faster_propagate
6
6
 
7
- __version__ = "0.1.6"
7
+ __version__ = "0.1.7"
8
8
  __all__ = [
9
9
  "enable",
10
10
  "disable",
@@ -200,6 +200,30 @@ def enable(device: Literal["cpu", "gpu"] = "gpu") -> None:
200
200
  except ImportError:
201
201
  pass
202
202
 
203
+ # Structure factors are the larger cost below roughly 163 beams, where the
204
+ # matrix exponential is small. Both call sites import the name at module
205
+ # scope, so each module's own attribute has to be replaced.
206
+ try:
207
+ from .scattering import faster_structure_factors, remember_original
208
+
209
+ scattering = importlib.import_module("diffBloch.core.scattering")
210
+ remember_original(scattering.structure_factors)
211
+ for mod_name in (
212
+ "diffBloch.core.scattering",
213
+ "diffBloch.core",
214
+ "diffBloch.engine.forward",
215
+ "diffBloch.preprocess.experiment",
216
+ ):
217
+ try:
218
+ mod = importlib.import_module(mod_name)
219
+ if hasattr(mod, "structure_factors"):
220
+ _remember(mod, "structure_factors")
221
+ setattr(mod, "structure_factors", faster_structure_factors)
222
+ except Exception:
223
+ pass
224
+ except ImportError:
225
+ pass
226
+
203
227
 
204
228
  def _remember(module, name: str) -> None:
205
229
  """Record a module attribute the first time it is replaced.
@@ -231,3 +255,9 @@ def disable() -> None:
231
255
  else:
232
256
  setattr(module, attr, original)
233
257
  _ORIGINALS.clear()
258
+ try:
259
+ from .scattering import forget_original
260
+
261
+ forget_original()
262
+ except ImportError:
263
+ pass
@@ -0,0 +1,271 @@
1
+ """Structure factors through the native path.
2
+
3
+ The native library already carries `native_parallel_structure_factors` and its
4
+ backward, but nothing called them, so `enable()` only ever replaced the matrix
5
+ exponential. That is the wrong half of the work at small beam counts: at 31
6
+ beams the exponential is 15% of the forward pass and the structure factors are
7
+ 80%, which is why the accelerated package barely beat PyTorch there.
8
+
9
+ This wires the native pair in behind a `torch.autograd.Function`.
10
+
11
+ Scope. The native kernel covers the elastic case in float64 on the CPU. Anything
12
+ else (absorption, a different dtype, a non-CPU tensor, a cutoff mode other than
13
+ `hard`) falls back to diffBloch's own implementation, so behaviour is unchanged
14
+ where the fast path does not apply.
15
+
16
+ Two details of diffBloch's definition have to be reproduced exactly. The cutoff
17
+ window multiplies every atom's contribution at a given `g`, so it factors out of
18
+ the sum over atoms and can be applied to the result. The `zero_threshold` snap
19
+ happens before the division by the cell volume, and the native kernel divides
20
+ inside, so the comparison is made against the rescaled value.
21
+ """
22
+
23
+ from __future__ import annotations
24
+
25
+ import ctypes
26
+ from typing import Optional
27
+
28
+ import numpy as np
29
+ import torch
30
+
31
+ from .backend import get_library
32
+
33
+
34
+ class NativeCase(ctypes.Structure):
35
+ _fields_ = [
36
+ ("n_atoms", ctypes.c_int32),
37
+ ("n_grid", ctypes.c_int32),
38
+ ("absorption", ctypes.c_int32),
39
+ ("volume", ctypes.c_double),
40
+ ("abs_c_over_v", ctypes.c_double),
41
+ ("positions", ctypes.POINTER(ctypes.c_double)),
42
+ ("occupancies", ctypes.POINTER(ctypes.c_double)),
43
+ ("uij", ctypes.POINTER(ctypes.c_double)),
44
+ ("lobato_a", ctypes.POINTER(ctypes.c_double)),
45
+ ("lobato_b", ctypes.POINTER(ctypes.c_double)),
46
+ ("grid_hkl", ctypes.POINTER(ctypes.c_int32)),
47
+ ("grid_g", ctypes.POINTER(ctypes.c_double)),
48
+ ("abs_db_du", ctypes.POINTER(ctypes.c_double)),
49
+ ("abs_knot_lo", ctypes.POINTER(ctypes.c_double)),
50
+ ("abs_knot_width", ctypes.POINTER(ctypes.c_double)),
51
+ ("abs_y0", ctypes.POINTER(ctypes.c_double)),
52
+ ("abs_y1", ctypes.POINTER(ctypes.c_double)),
53
+ ("abs_d0", ctypes.POINTER(ctypes.c_double)),
54
+ ("abs_d1", ctypes.POINTER(ctypes.c_double)),
55
+ ]
56
+
57
+
58
+ class PhaseCache(ctypes.Structure):
59
+ _fields_ = [
60
+ (name, ctypes.POINTER(ctypes.c_double))
61
+ for name in ("amplitude", "cosine", "sine", "imag_amplitude", "dimag_db")
62
+ ]
63
+
64
+
65
+ _VOID = ctypes.POINTER(ctypes.c_double * 0)
66
+
67
+
68
+ def _f64(a) -> np.ndarray:
69
+ return np.ascontiguousarray(a, dtype=np.float64).reshape(-1)
70
+
71
+
72
+ def _ptr(a: np.ndarray):
73
+ return a.ctypes.data_as(ctypes.POINTER(ctypes.c_double))
74
+
75
+
76
+ _LOBATO: dict[int, tuple[np.ndarray, np.ndarray]] = {}
77
+
78
+
79
+ def _lobato_for(numbers: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
80
+ """The five Lobato a/b coefficients per atom, laid out as the kernel reads them."""
81
+ if not _LOBATO:
82
+ from diffBloch.core.scattering import _lobato_table
83
+
84
+ for z, (a, b) in _lobato_table().items():
85
+ _LOBATO[int(z)] = (np.asarray(a, dtype=np.float64),
86
+ np.asarray(b, dtype=np.float64))
87
+ a_out = np.empty((numbers.shape[0], 5), dtype=np.float64)
88
+ b_out = np.empty((numbers.shape[0], 5), dtype=np.float64)
89
+ for i, z in enumerate(numbers.tolist()):
90
+ a, b = _LOBATO[int(z)]
91
+ a_out[i] = a
92
+ b_out[i] = b
93
+ return a_out, b_out
94
+
95
+
96
+ class _Native:
97
+ """Holds every buffer the kernel points at, so nothing is collected mid-call."""
98
+
99
+ def __init__(self, positions, numbers, occupancies, uij_star, hkl, g, volume):
100
+ self.n_atoms = int(positions.shape[0])
101
+ self.n_grid = int(hkl.shape[0])
102
+ self.pos = _f64(positions)
103
+ self.occ = _f64(occupancies)
104
+ self.uij = _f64(uij_star)
105
+ la, lb = _lobato_for(np.asarray(numbers))
106
+ self.la, self.lb = _f64(la), _f64(lb)
107
+ self.hkl = np.ascontiguousarray(hkl, dtype=np.int32).reshape(-1)
108
+ self.g = _f64(g)
109
+
110
+ c = NativeCase()
111
+ c.n_atoms, c.n_grid, c.absorption = self.n_atoms, self.n_grid, 0
112
+ c.volume, c.abs_c_over_v = float(volume), 0.0
113
+ c.positions, c.occupancies, c.uij = _ptr(self.pos), _ptr(self.occ), _ptr(self.uij)
114
+ c.lobato_a, c.lobato_b = _ptr(self.la), _ptr(self.lb)
115
+ c.grid_hkl = self.hkl.ctypes.data_as(ctypes.POINTER(ctypes.c_int32))
116
+ c.grid_g = _ptr(self.g)
117
+ self.case = c
118
+
119
+ n = self.n_grid * self.n_atoms
120
+ self.cache_arrays = {}
121
+ cache = PhaseCache()
122
+ for name in ("amplitude", "cosine", "sine", "imag_amplitude", "dimag_db"):
123
+ buf = np.zeros(n, dtype=np.float64)
124
+ self.cache_arrays[name] = buf
125
+ setattr(cache, name, _ptr(buf))
126
+ self.cache = cache
127
+
128
+ def forward(self) -> np.ndarray:
129
+ lib = get_library()
130
+ fgb = np.zeros(self.n_grid, dtype=np.complex128)
131
+ lib.native_parallel_structure_factors(
132
+ ctypes.byref(self.case),
133
+ fgb.ctypes.data_as(_VOID),
134
+ ctypes.byref(self.cache),
135
+ )
136
+ return fgb
137
+
138
+ def backward(self, fbar: np.ndarray):
139
+ lib = get_library()
140
+ gp = np.zeros(self.n_atoms * 3, dtype=np.float64)
141
+ go = np.zeros(self.n_atoms, dtype=np.float64)
142
+ gu = np.zeros(self.n_atoms * 9, dtype=np.float64)
143
+ fbar_c = np.ascontiguousarray(fbar, dtype=np.complex128)
144
+ lib.native_parallel_structure_factors_backward(
145
+ ctypes.byref(self.case),
146
+ ctypes.byref(self.cache),
147
+ fbar_c.ctypes.data_as(_VOID),
148
+ _ptr(gp), _ptr(go), _ptr(gu),
149
+ )
150
+ return gp.reshape(self.n_atoms, 3), go, gu.reshape(self.n_atoms, 3, 3)
151
+
152
+
153
+ class FasterStructureFactors(torch.autograd.Function):
154
+ @staticmethod
155
+ def forward(ctx, positions, occupancies, uij_star, numbers, hkl, g,
156
+ volume, window, zero_threshold):
157
+ native = _Native(
158
+ positions.detach().cpu().numpy(), numbers, occupancies.detach().cpu().numpy(),
159
+ uij_star.detach().cpu().numpy(), hkl, g, volume,
160
+ )
161
+ raw = native.forward()
162
+ # diffBloch snaps before dividing by the volume; the kernel divides inside.
163
+ unmasked = raw * volume * window
164
+ keep_re = np.abs(unmasked.real) >= zero_threshold
165
+ keep_im = np.abs(unmasked.imag) >= zero_threshold
166
+ out = (np.where(keep_re, unmasked.real, 0.0)
167
+ + 1j * np.where(keep_im, unmasked.imag, 0.0)) / volume
168
+
169
+ ctx.native = native
170
+ ctx.window = window
171
+ ctx.keep = (keep_re, keep_im)
172
+ ctx.shapes = (positions.shape, occupancies.shape, uij_star.shape)
173
+ ctx.meta = (positions.dtype, positions.device)
174
+ return torch.from_numpy(out).to(device=positions.device)
175
+
176
+ @staticmethod
177
+ @torch.autograd.function.once_differentiable
178
+ def backward(ctx, grad_output):
179
+ keep_re, keep_im = ctx.keep
180
+ gout = grad_output.detach().cpu().numpy()
181
+ # Undo the snap and the per-g window, so the cotangent matches what the
182
+ # kernel differentiates: the plain sum over atoms, divided by the volume.
183
+ fbar = (np.where(keep_re, gout.real, 0.0)
184
+ + 1j * np.where(keep_im, gout.imag, 0.0)) * ctx.window
185
+ gp, go, gu = ctx.native.backward(fbar)
186
+ dtype, device = ctx.meta
187
+ p_shape, o_shape, u_shape = ctx.shapes
188
+ return (
189
+ torch.from_numpy(gp).to(dtype=dtype, device=device).reshape(p_shape),
190
+ torch.from_numpy(go).to(dtype=dtype, device=device).reshape(o_shape),
191
+ torch.from_numpy(gu).to(dtype=dtype, device=device).reshape(u_shape),
192
+ None, None, None, None, None, None,
193
+ )
194
+
195
+
196
+ def _shapes_ok(positions, numbers, occupancies, uij_star, cell_volume) -> bool:
197
+ """Whether the inputs are the shapes diffBloch accepts.
198
+
199
+ A malformed call is handed to diffBloch rather than rejected here, so the
200
+ error it raises is its own, with its own message. Validating separately
201
+ would mean keeping two copies of the same rules in step.
202
+ """
203
+ if positions.ndim != 2 or positions.shape[1] != 3:
204
+ return False
205
+ n_atoms = positions.shape[0]
206
+ if tuple(numbers.shape) != (n_atoms,) or tuple(occupancies.shape) != (n_atoms,):
207
+ return False
208
+ if tuple(uij_star.shape) != (n_atoms, 3, 3):
209
+ return False
210
+ return cell_volume > 0.0
211
+
212
+
213
+ def _usable(positions, numbers, occupancies, uij_star, cell_volume,
214
+ absorption, cutoff) -> bool:
215
+ if absorption is not None and getattr(absorption, "enabled", False):
216
+ return False
217
+ if cutoff != "hard":
218
+ return False
219
+ for t in (positions, occupancies, uij_star):
220
+ if t.dtype != torch.float64 or t.device.type != "cpu":
221
+ return False
222
+ if not _shapes_ok(positions, numbers, occupancies, uij_star, cell_volume):
223
+ return False
224
+ return get_library() is not None
225
+
226
+
227
+ def faster_structure_factors(
228
+ positions, numbers, occupancies, uij_star, hkl, reciprocal_basis, cell_volume,
229
+ *, g_max, cutoff="hard", zero_threshold=1e-12, absorption=None, energy=None,
230
+ ):
231
+ """diffBloch's `structure_factors`, over the native kernel where it applies."""
232
+ from diffBloch.core.scattering import (
233
+ _g_vector_lengths,
234
+ structure_factors as diffbloch_structure_factors,
235
+ )
236
+
237
+ original = _ORIGINAL[0] or diffbloch_structure_factors
238
+ if absorption is None:
239
+ from diffBloch.specs import NO_ABSORPTION
240
+
241
+ absorption = NO_ABSORPTION
242
+
243
+ if not _usable(positions, numbers, occupancies, uij_star, cell_volume,
244
+ absorption, cutoff):
245
+ return original(
246
+ positions, numbers, occupancies, uij_star, hkl, reciprocal_basis,
247
+ cell_volume, g_max=g_max, cutoff=cutoff, zero_threshold=zero_threshold,
248
+ absorption=absorption, energy=energy,
249
+ )
250
+
251
+ g = _g_vector_lengths(hkl, reciprocal_basis)
252
+ window = (g <= g_max).to(torch.float64).detach().cpu().numpy()
253
+ return FasterStructureFactors.apply(
254
+ positions, occupancies, uij_star,
255
+ np.asarray(numbers.detach().cpu().numpy()),
256
+ hkl.detach().cpu().numpy(),
257
+ g.detach().cpu().numpy(),
258
+ float(cell_volume), window, float(zero_threshold),
259
+ )
260
+
261
+
262
+ _ORIGINAL: list[Optional[object]] = [None]
263
+
264
+
265
+ def remember_original(fn) -> None:
266
+ if _ORIGINAL[0] is None:
267
+ _ORIGINAL[0] = fn
268
+
269
+
270
+ def forget_original() -> None:
271
+ _ORIGINAL[0] = None